ForcePilot/test/api/test_knowledge_router.py

394 lines
15 KiB
Python
Raw Normal View History

"""
Integration tests for knowledge router and mindmap router endpoints.
"""
from __future__ import annotations
import uuid
import pytest
pytestmark = [pytest.mark.asyncio, pytest.mark.integration]
def _assert_forbidden_response(response):
"""验证 403 禁止访问响应的格式"""
assert response.status_code == 403
payload = response.json()
assert "detail" in payload
assert isinstance(payload["detail"], str)
async def test_admin_can_manage_knowledge_databases(test_client, admin_headers, knowledge_database):
db_id = knowledge_database["db_id"]
list_response = await test_client.get("/api/knowledge/databases", headers=admin_headers)
assert list_response.status_code == 200, list_response.text
databases = list_response.json().get("databases", [])
assert any(entry["db_id"] == db_id for entry in databases)
get_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
assert get_response.status_code == 200, get_response.text
assert get_response.json()["db_id"] == db_id
update_response = await test_client.put(
f"/api/knowledge/databases/{db_id}",
json={"name": knowledge_database["name"], "description": "Updated by pytest"},
headers=admin_headers,
)
assert update_response.status_code == 200, update_response.text
assert update_response.json()["database"]["description"] == "Updated by pytest"
async def test_create_database_with_chunk_preset(test_client, admin_headers):
db_name = f"pytest_chunk_preset_{uuid.uuid4().hex[:6]}"
payload = {
"database_name": db_name,
"description": "Chunk preset create test",
"embed_model_name": "siliconflow/BAAI/bge-m3",
"kb_type": "milvus",
"additional_params": {"chunk_preset_id": "book"},
}
create_response = await test_client.post("/api/knowledge/databases", json=payload, headers=admin_headers)
assert create_response.status_code == 200, create_response.text
db_id = create_response.json()["db_id"]
info_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
assert info_response.status_code == 200, info_response.text
assert info_response.json()["additional_params"]["chunk_preset_id"] == "book"
delete_response = await test_client.delete(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
assert delete_response.status_code == 200, delete_response.text
async def test_update_database_additional_params_merge_keeps_chunk_preset(
test_client, admin_headers, knowledge_database
):
db_id = knowledge_database["db_id"]
first_update = await test_client.put(
f"/api/knowledge/databases/{db_id}",
json={
"name": knowledge_database["name"],
"description": "update with chunk preset",
"additional_params": {"chunk_preset_id": "qa"},
},
headers=admin_headers,
)
assert first_update.status_code == 200, first_update.text
second_update = await test_client.put(
f"/api/knowledge/databases/{db_id}",
json={
"name": knowledge_database["name"],
"description": "update without additional params",
},
headers=admin_headers,
)
assert second_update.status_code == 200, second_update.text
info_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
assert info_response.status_code == 200, info_response.text
assert info_response.json()["additional_params"]["chunk_preset_id"] == "qa"
async def test_knowledge_routes_enforce_permissions(test_client, standard_user, knowledge_database):
db_id = knowledge_database["db_id"]
forbidden_create = await test_client.post(
"/api/knowledge/databases",
json={
"database_name": "unauthorized_db",
"description": "Should not succeed",
"embed_model_name": "siliconflow/BAAI/bge-m3",
},
headers=standard_user["headers"],
)
_assert_forbidden_response(forbidden_create)
forbidden_list = await test_client.get("/api/knowledge/databases", headers=standard_user["headers"])
_assert_forbidden_response(forbidden_list)
forbidden_get = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=standard_user["headers"])
_assert_forbidden_response(forbidden_get)
async def test_admin_can_create_vector_db_with_reranker(test_client, admin_headers):
"""测试创建向量库并配置 reranker 参数(通过 query_params.options
注意数据库清理由 conftest.py 中的 session fixture 自动处理
"""
db_name = f"pytest_rerank_{uuid.uuid4().hex[:6]}"
payload = {
"database_name": db_name,
"description": "Vector DB with reranker",
"embed_model_name": "siliconflow/BAAI/bge-m3",
"kb_type": "milvus",
"additional_params": {},
}
create_response = await test_client.post("/api/knowledge/databases", json=payload, headers=admin_headers)
assert create_response.status_code == 200, create_response.text
db_payload = create_response.json()
db_id = db_payload["db_id"]
# 获取查询参数配置
params_response = await test_client.get(f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers)
assert params_response.status_code == 200, params_response.text
params_payload = params_response.json()
options = params_payload.get("params", {}).get("options", [])
option_keys = {option.get("key") for option in options}
# 验证新的参数名称
assert "final_top_k" in option_keys
assert "use_reranker" in option_keys
assert "recall_top_k" in option_keys
assert "reranker_model" in option_keys
# 验证参数配置
final_top_k_option = next((opt for opt in options if opt.get("key") == "final_top_k"), None)
assert final_top_k_option is not None
assert final_top_k_option.get("default") == 10
use_reranker_option = next((opt for opt in options if opt.get("key") == "use_reranker"), None)
assert use_reranker_option is not None
assert use_reranker_option.get("default") is False
# 保存查询参数(模拟前端配置)
update_params = {
"final_top_k": 5,
"use_reranker": True,
"recall_top_k": 20,
}
update_response = await test_client.put(
f"/api/knowledge/databases/{db_id}/query-params", json=update_params, headers=admin_headers
)
assert update_response.status_code == 200, update_response.text
# 再次获取参数,验证保存成功
params_response2 = await test_client.get(f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers)
assert params_response2.status_code == 200, params_response2.text
params_payload2 = params_response2.json()
options2 = params_payload2.get("params", {}).get("options", [])
# 验证保存的值
final_top_k_option2 = next((opt for opt in options2 if opt.get("key") == "final_top_k"), None)
assert final_top_k_option2 is not None
assert final_top_k_option2.get("default") == 5 # 保存的值
use_reranker_option2 = next((opt for opt in options2 if opt.get("key") == "use_reranker"), None)
assert use_reranker_option2 is not None
assert use_reranker_option2.get("default") is True # 保存的值
# =============================================================================
# === Mindmap Router Tests ===
# =============================================================================
async def test_get_databases_overview(test_client, admin_headers, knowledge_database):
"""测试获取所有知识库概览"""
response = await test_client.get("/api/mindmap/databases", headers=admin_headers)
assert response.status_code == 200, response.text
payload = response.json()
assert payload["message"] == "success"
assert "databases" in payload
assert "total" in payload
# 验证知识库在列表中
db_ids = [db["db_id"] for db in payload["databases"]]
assert knowledge_database["db_id"] in db_ids
async def test_get_database_files(test_client, admin_headers, knowledge_database):
"""测试获取知识库文件列表"""
db_id = knowledge_database["db_id"]
response = await test_client.get(f"/api/mindmap/databases/{db_id}/files", headers=admin_headers)
assert response.status_code == 200, response.text
payload = response.json()
assert payload["message"] == "success"
assert payload["db_id"] == db_id
assert "files" in payload
assert "total" in payload
assert payload["db_name"] == knowledge_database["name"]
async def test_get_database_files_not_found(test_client, admin_headers):
"""测试获取不存在的知识库文件列表"""
response = await test_client.get("/api/mindmap/databases/nonexistent_db_id/files", headers=admin_headers)
assert response.status_code == 404
async def test_generate_mindmap_empty_files(test_client, admin_headers, knowledge_database):
"""测试空文件列表生成思维导图"""
db_id = knowledge_database["db_id"]
response = await test_client.post(
"/api/mindmap/generate",
json={"db_id": db_id, "file_ids": [], "user_prompt": ""},
headers=admin_headers,
)
# 空文件应该返回400错误
assert response.status_code == 400
assert "中没有文件" in response.json()["detail"]
async def test_get_database_mindmap_not_exists(test_client, admin_headers, knowledge_database):
"""测试获取不存在的思维导图"""
db_id = knowledge_database["db_id"]
response = await test_client.get(f"/api/mindmap/database/{db_id}", headers=admin_headers)
assert response.status_code == 200, response.text
payload = response.json()
assert payload["db_id"] == db_id
assert payload["mindmap"] is None # 尚未生成思维导图
async def test_generate_and_get_mindmap(test_client, admin_headers, knowledge_database):
"""测试生成并获取思维导图
注意此测试需要知识库中有文件才能完整测试核心功能
由于没有前置的文件上传 fixture测试会先验证空文件场景预期400
然后使用 xfail 标记等待后续完善
"""
db_id = knowledge_database["db_id"]
# 空文件场景 - 预期返回400错误
generate_response = await test_client.post(
"/api/mindmap/generate",
json={"db_id": db_id, "file_ids": [], "user_prompt": ""},
headers=admin_headers,
)
assert generate_response.status_code == 400
assert "中没有文件" in generate_response.json()["detail"]
# 标记此测试需要文件上传支持才能完整执行
pytest.skip("需要先上传文件才能完整测试思维导图生成功能")
# =============================================================================
# === Knowledge Router Additional Tests ===
# =============================================================================
async def test_get_accessible_databases(test_client, admin_headers, knowledge_database):
"""测试获取可访问的知识库列表"""
response = await test_client.get("/api/knowledge/databases/accessible", headers=admin_headers)
assert response.status_code == 200, response.text
payload = response.json()
assert "databases" in payload
# 验证知识库在列表中
db_ids = [db["db_id"] for db in payload["databases"]]
assert knowledge_database["db_id"] in db_ids
async def test_get_knowledge_base_types(test_client, admin_headers):
"""测试获取支持的知识库类型"""
response = await test_client.get("/api/knowledge/types", headers=admin_headers)
assert response.status_code == 200, response.text
payload = response.json()
assert payload["message"] == "success"
assert "kb_types" in payload
async def test_get_knowledge_base_statistics(test_client, admin_headers):
"""测试获取知识库统计信息"""
response = await test_client.get("/api/knowledge/stats", headers=admin_headers)
assert response.status_code == 200, response.text
payload = response.json()
assert payload["message"] == "success"
assert "stats" in payload
async def test_get_supported_file_types(test_client, admin_headers):
"""测试获取支持的文件类型"""
response = await test_client.get("/api/knowledge/files/supported-types", headers=admin_headers)
assert response.status_code == 200, response.text
payload = response.json()
assert payload["message"] == "success"
assert "file_types" in payload
assert isinstance(payload["file_types"], list)
async def test_duplicate_database_name(test_client, admin_headers, knowledge_database):
"""测试重复创建同名知识库"""
db_name = knowledge_database["name"]
response = await test_client.post(
"/api/knowledge/databases",
json={
"database_name": db_name,
"description": "Duplicate name test",
"embed_model_name": "siliconflow/BAAI/bge-m3",
"kb_type": "lightrag",
"additional_params": {},
},
headers=admin_headers,
)
assert response.status_code == 409
assert "已存在" in response.json()["detail"]
async def test_create_milvus_knowledge_base(test_client, admin_headers):
"""测试创建 Milvus 知识库
注意数据库清理由 conftest.py 中的 session fixture 自动处理
"""
db_name = f"pytest_milvus_{uuid.uuid4().hex[:6]}"
payload = {
"database_name": db_name,
"description": "Pytest Milvus knowledge base",
"embed_model_name": "siliconflow/BAAI/bge-m3",
"kb_type": "milvus",
"additional_params": {},
}
create_response = await test_client.post("/api/knowledge/databases", json=payload, headers=admin_headers)
assert create_response.status_code == 200, create_response.text
db_payload = create_response.json()
assert db_payload["kb_type"] == "milvus"
async def test_sample_questions_endpoints(test_client, admin_headers, knowledge_database):
"""测试示例问题接口空文件时预期返回400"""
db_id = knowledge_database["db_id"]
# 获取示例问题(空知识库应该返回空列表)
get_response = await test_client.get(f"/api/knowledge/databases/{db_id}/sample-questions", headers=admin_headers)
assert get_response.status_code == 200, get_response.text
get_payload = get_response.json()
assert get_payload["db_id"] == db_id
assert "questions" in get_payload
assert get_payload["count"] == 0 # 空知识库没有问题
# 生成示例问题空知识库应该返回400
generate_response = await test_client.post(
f"/api/knowledge/databases/{db_id}/sample-questions",
json={"count": 5},
headers=admin_headers,
)
assert generate_response.status_code == 400
assert "中没有文件" in generate_response.json()["detail"]
async def test_mindmap_permissions(test_client, standard_user, knowledge_database):
"""测试思维导图接口的权限控制"""
db_id = knowledge_database["db_id"]
# 普通用户应该无法访问
forbidden_list = await test_client.get("/api/mindmap/databases", headers=standard_user["headers"])
_assert_forbidden_response(forbidden_list)
forbidden_files = await test_client.get(f"/api/mindmap/databases/{db_id}/files", headers=standard_user["headers"])
_assert_forbidden_response(forbidden_files)
forbidden_generate = await test_client.post(
"/api/mindmap/generate",
json={"db_id": db_id, "file_ids": []},
headers=standard_user["headers"],
)
_assert_forbidden_response(forbidden_generate)