ForcePilot/test/api/test_knowledge_router.py
Wenjie Zhang a1812f2e97 feat(知识库): 重构知识库元数据存储为PostgreSQL数据库
- 新增知识库、知识文件、评估基准、评估结果等数据模型
- 实现知识库相关Repository类提供CRUD操作
- 修改知识库实现类使用PostgreSQL存储元数据
- 添加PostgreSQL数据库管理器和初始化逻辑
- 更新相关路由和工具类适配新的存储方式
- 添加数据库迁移脚本和测试工具
- 在docker-compose中配置PostgreSQL服务
2026-01-24 10:56:31 +08:00

341 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
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_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)