test: 新增测试文件并更新现有测试适配命名与架构变更
- 新增: test_knowledge_base_backend.py, test_skills_middleware.py - 更新: 所有现有测试适配 db_id→kb_id, get_subagents_from_names→slugs 等接口变更
This commit is contained in:
parent
99c590d003
commit
708ff29515
@ -11,12 +11,12 @@ import pytest
|
||||
pytestmark = [pytest.mark.asyncio, pytest.mark.integration]
|
||||
|
||||
|
||||
async def _upload_test_dataset(test_client, admin_headers: dict[str, str], db_id: str) -> tuple[str, str]:
|
||||
async def _upload_test_dataset(test_client, admin_headers: dict[str, str], slug: str) -> tuple[str, str]:
|
||||
dataset_name = f"pytest_dataset_{uuid.uuid4().hex[:8]}"
|
||||
line = '{"query":"什么是单元测试?","gold_answer":"用于验证代码行为的自动化测试"}\n'
|
||||
|
||||
response = await test_client.post(
|
||||
f"/api/evaluation/databases/{db_id}/datasets/upload",
|
||||
f"/api/evaluation/databases/{slug}/datasets/upload",
|
||||
data={"name": dataset_name, "description": "pytest dataset for download"},
|
||||
files={"file": ("pytest_dataset.jsonl", line.encode("utf-8"), "application/x-ndjson")},
|
||||
headers=admin_headers,
|
||||
@ -39,7 +39,7 @@ async def test_download_dataset_requires_admin(test_client, standard_user):
|
||||
|
||||
|
||||
async def test_admin_can_download_dataset(test_client, admin_headers, knowledge_database):
|
||||
dataset_id, expected_line = await _upload_test_dataset(test_client, admin_headers, knowledge_database["db_id"])
|
||||
dataset_id, expected_line = await _upload_test_dataset(test_client, admin_headers, knowledge_database["slug"])
|
||||
|
||||
response = await test_client.get(
|
||||
f"/api/evaluation/datasets/{dataset_id}/download",
|
||||
|
||||
@ -102,26 +102,26 @@ async def _create_test_database(test_client, admin_headers, share_config=None):
|
||||
return response.json()
|
||||
|
||||
|
||||
async def _accessible_db_ids(test_client, headers):
|
||||
async def _accessible_slugs(test_client, headers):
|
||||
response = await test_client.get("/api/knowledge/databases/accessible", headers=headers)
|
||||
assert response.status_code == 200, response.text
|
||||
return {item["db_id"] for item in response.json().get("databases", [])}
|
||||
return {item["slug"] for item in response.json().get("databases", [])}
|
||||
|
||||
|
||||
async def test_admin_can_manage_knowledge_databases(test_client, admin_headers, knowledge_database):
|
||||
db_id = knowledge_database["db_id"]
|
||||
slug = knowledge_database["slug"]
|
||||
|
||||
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)
|
||||
assert any(entry["slug"] == slug for entry in databases)
|
||||
|
||||
get_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
get_response = await test_client.get(f"/api/knowledge/databases/{slug}", headers=admin_headers)
|
||||
assert get_response.status_code == 200, get_response.text
|
||||
assert get_response.json()["db_id"] == db_id
|
||||
assert get_response.json()["slug"] == slug
|
||||
|
||||
update_response = await test_client.put(
|
||||
f"/api/knowledge/databases/{db_id}",
|
||||
f"/api/knowledge/databases/{slug}",
|
||||
json={"name": knowledge_database["name"], "description": "Updated by pytest"},
|
||||
headers=admin_headers,
|
||||
)
|
||||
@ -141,23 +141,23 @@ async def test_create_database_with_chunk_preset(test_client, admin_headers):
|
||||
|
||||
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"]
|
||||
slug = create_response.json()["slug"]
|
||||
|
||||
info_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
info_response = await test_client.get(f"/api/knowledge/databases/{slug}", 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)
|
||||
delete_response = await test_client.delete(f"/api/knowledge/databases/{slug}", 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"]
|
||||
slug = knowledge_database["slug"]
|
||||
|
||||
first_update = await test_client.put(
|
||||
f"/api/knowledge/databases/{db_id}",
|
||||
f"/api/knowledge/databases/{slug}",
|
||||
json={
|
||||
"name": knowledge_database["name"],
|
||||
"description": "update with chunk preset",
|
||||
@ -168,7 +168,7 @@ async def test_update_database_additional_params_merge_keeps_chunk_preset(
|
||||
assert first_update.status_code == 200, first_update.text
|
||||
|
||||
second_update = await test_client.put(
|
||||
f"/api/knowledge/databases/{db_id}",
|
||||
f"/api/knowledge/databases/{slug}",
|
||||
json={
|
||||
"name": knowledge_database["name"],
|
||||
"description": "update without additional params",
|
||||
@ -177,13 +177,13 @@ async def test_update_database_additional_params_merge_keeps_chunk_preset(
|
||||
)
|
||||
assert second_update.status_code == 200, second_update.text
|
||||
|
||||
info_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
info_response = await test_client.get(f"/api/knowledge/databases/{slug}", 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"]
|
||||
slug = knowledge_database["slug"]
|
||||
|
||||
forbidden_create = await test_client.post(
|
||||
"/api/knowledge/databases",
|
||||
@ -199,7 +199,7 @@ async def test_knowledge_routes_enforce_permissions(test_client, standard_user,
|
||||
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"])
|
||||
forbidden_get = await test_client.get(f"/api/knowledge/databases/{slug}", headers=standard_user["headers"])
|
||||
_assert_forbidden_response(forbidden_get)
|
||||
|
||||
|
||||
@ -221,10 +221,10 @@ async def test_admin_can_create_vector_db_with_reranker(test_client, admin_heade
|
||||
assert create_response.status_code == 200, create_response.text
|
||||
|
||||
db_payload = create_response.json()
|
||||
db_id = db_payload["db_id"]
|
||||
slug = db_payload["slug"]
|
||||
|
||||
# 获取查询参数配置
|
||||
params_response = await test_client.get(f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers)
|
||||
params_response = await test_client.get(f"/api/knowledge/databases/{slug}/query-params", headers=admin_headers)
|
||||
assert params_response.status_code == 200, params_response.text
|
||||
|
||||
params_payload = params_response.json()
|
||||
@ -253,12 +253,12 @@ async def test_admin_can_create_vector_db_with_reranker(test_client, admin_heade
|
||||
"recall_top_k": 20,
|
||||
}
|
||||
update_response = await test_client.put(
|
||||
f"/api/knowledge/databases/{db_id}/query-params", json=update_params, headers=admin_headers
|
||||
f"/api/knowledge/databases/{slug}/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)
|
||||
params_response2 = await test_client.get(f"/api/knowledge/databases/{slug}/query-params", headers=admin_headers)
|
||||
assert params_response2.status_code == 200, params_response2.text
|
||||
|
||||
params_payload2 = params_response2.json()
|
||||
@ -290,11 +290,11 @@ async def test_create_dify_database_success(test_client, admin_headers):
|
||||
create_response = await test_client.post("/api/knowledge/databases", json=payload, headers=admin_headers)
|
||||
assert create_response.status_code == 200, create_response.text
|
||||
created_payload = create_response.json()
|
||||
db_id = created_payload["db_id"]
|
||||
slug = created_payload["slug"]
|
||||
assert created_payload["embedding_model_spec"] is None
|
||||
assert "chunk_preset_id" not in created_payload["metadata"]
|
||||
|
||||
info_response = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
info_response = await test_client.get(f"/api/knowledge/databases/{slug}", headers=admin_headers)
|
||||
assert info_response.status_code == 200, info_response.text
|
||||
additional_params = info_response.json()["additional_params"]
|
||||
assert additional_params["dify_api_url"] == "https://api.dify.ai/v1"
|
||||
@ -350,16 +350,16 @@ async def test_dify_query_params_and_documents_readonly(test_client, admin_heade
|
||||
|
||||
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"]
|
||||
slug = create_response.json()["slug"]
|
||||
|
||||
params_response = await test_client.get(f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers)
|
||||
params_response = await test_client.get(f"/api/knowledge/databases/{slug}/query-params", headers=admin_headers)
|
||||
assert params_response.status_code == 200, params_response.text
|
||||
options = params_response.json().get("params", {}).get("options", [])
|
||||
option_keys = {item.get("key") for item in options}
|
||||
assert option_keys == {"search_mode", "final_top_k", "score_threshold_enabled", "similarity_threshold"}
|
||||
|
||||
add_response = await test_client.post(
|
||||
f"/api/knowledge/databases/{db_id}/documents",
|
||||
f"/api/knowledge/databases/{slug}/documents",
|
||||
json={"items": ["/tmp/demo.txt"], "params": {"content_type": "file"}},
|
||||
headers=admin_headers,
|
||||
)
|
||||
@ -367,7 +367,7 @@ async def test_dify_query_params_and_documents_readonly(test_client, admin_heade
|
||||
assert "只支持检索" in add_response.json()["detail"]
|
||||
|
||||
parse_response = await test_client.post(
|
||||
f"/api/knowledge/databases/{db_id}/documents/parse",
|
||||
f"/api/knowledge/databases/{slug}/documents/parse",
|
||||
json=["file_id_1"],
|
||||
headers=admin_headers,
|
||||
)
|
||||
@ -375,7 +375,7 @@ async def test_dify_query_params_and_documents_readonly(test_client, admin_heade
|
||||
assert "只支持检索" in parse_response.json()["detail"]
|
||||
|
||||
index_response = await test_client.post(
|
||||
f"/api/knowledge/databases/{db_id}/documents/index",
|
||||
f"/api/knowledge/databases/{slug}/documents/index",
|
||||
json={"file_ids": ["file_id_1"], "params": {}},
|
||||
headers=admin_headers,
|
||||
)
|
||||
@ -398,18 +398,18 @@ async def test_get_databases_overview(test_client, admin_headers, knowledge_data
|
||||
assert "total" in payload
|
||||
|
||||
# 验证知识库在列表中
|
||||
db_ids = [db["db_id"] for db in payload["databases"]]
|
||||
assert knowledge_database["db_id"] in db_ids
|
||||
slugs = [db["slug"] for db in payload["databases"]]
|
||||
assert knowledge_database["slug"] in slugs
|
||||
|
||||
|
||||
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)
|
||||
slug = knowledge_database["slug"]
|
||||
response = await test_client.get(f"/api/mindmap/databases/{slug}/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 payload["slug"] == slug
|
||||
assert "files" in payload
|
||||
assert "total" in payload
|
||||
assert payload["db_name"] == knowledge_database["name"]
|
||||
@ -417,16 +417,16 @@ async def test_get_database_files(test_client, admin_headers, knowledge_database
|
||||
|
||||
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)
|
||||
response = await test_client.get("/api/mindmap/databases/nonexistent_slug/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"]
|
||||
slug = knowledge_database["slug"]
|
||||
response = await test_client.post(
|
||||
"/api/mindmap/generate",
|
||||
json={"db_id": db_id, "file_ids": [], "user_prompt": ""},
|
||||
json={"slug": slug, "file_ids": [], "user_prompt": ""},
|
||||
headers=admin_headers,
|
||||
)
|
||||
# 空文件应该返回400错误
|
||||
@ -436,11 +436,11 @@ async def test_generate_mindmap_empty_files(test_client, admin_headers, knowledg
|
||||
|
||||
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)
|
||||
slug = knowledge_database["slug"]
|
||||
response = await test_client.get(f"/api/mindmap/database/{slug}", headers=admin_headers)
|
||||
assert response.status_code == 200, response.text
|
||||
payload = response.json()
|
||||
assert payload["db_id"] == db_id
|
||||
assert payload["slug"] == slug
|
||||
assert payload["mindmap"] is None # 尚未生成思维导图
|
||||
|
||||
|
||||
@ -451,12 +451,12 @@ async def test_generate_and_get_mindmap(test_client, admin_headers, knowledge_da
|
||||
由于没有前置的文件上传 fixture,测试会先验证空文件场景(预期400),
|
||||
然后使用 xfail 标记等待后续完善。
|
||||
"""
|
||||
db_id = knowledge_database["db_id"]
|
||||
slug = knowledge_database["slug"]
|
||||
|
||||
# 空文件场景 - 预期返回400错误
|
||||
generate_response = await test_client.post(
|
||||
"/api/mindmap/generate",
|
||||
json={"db_id": db_id, "file_ids": [], "user_prompt": ""},
|
||||
json={"slug": slug, "file_ids": [], "user_prompt": ""},
|
||||
headers=admin_headers,
|
||||
)
|
||||
assert generate_response.status_code == 400
|
||||
@ -479,17 +479,17 @@ async def test_get_accessible_databases(test_client, admin_headers, knowledge_da
|
||||
assert "databases" in payload
|
||||
|
||||
# 验证知识库在列表中
|
||||
db_ids = [db["db_id"] for db in payload["databases"]]
|
||||
assert knowledge_database["db_id"] in db_ids
|
||||
slugs = [db["slug"] for db in payload["databases"]]
|
||||
assert knowledge_database["slug"] in slugs
|
||||
|
||||
|
||||
async def test_create_database_defaults_to_global_share_config(test_client, admin_headers):
|
||||
database = await _create_test_database(test_client, admin_headers)
|
||||
db_id = database["db_id"]
|
||||
slug = database["slug"]
|
||||
try:
|
||||
assert database["share_config"] == {"access_level": "global", "department_ids": [], "user_uids": []}
|
||||
finally:
|
||||
await test_client.delete(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
await test_client.delete(f"/api/knowledge/databases/{slug}", headers=admin_headers)
|
||||
|
||||
|
||||
async def test_department_share_config_filters_accessible_databases(test_client, admin_headers):
|
||||
@ -511,11 +511,11 @@ async def test_department_share_config_filters_accessible_databases(test_client,
|
||||
assert saved_config["access_level"] == "department"
|
||||
assert department_a["id"] in saved_config["department_ids"]
|
||||
|
||||
assert database["db_id"] in await _accessible_db_ids(test_client, user_a["headers"])
|
||||
assert database["db_id"] not in await _accessible_db_ids(test_client, user_b["headers"])
|
||||
assert database["slug"] in await _accessible_slugs(test_client, user_a["headers"])
|
||||
assert database["slug"] not in await _accessible_slugs(test_client, user_b["headers"])
|
||||
finally:
|
||||
if database:
|
||||
await test_client.delete(f"/api/knowledge/databases/{database['db_id']}", headers=admin_headers)
|
||||
await test_client.delete(f"/api/knowledge/databases/{database['slug']}", headers=admin_headers)
|
||||
if user_a:
|
||||
await _delete_user_by_id(test_client, admin_headers, user_a["user"]["id"])
|
||||
if user_b:
|
||||
@ -543,11 +543,11 @@ async def test_user_share_config_filters_accessible_databases(test_client, admin
|
||||
assert saved_config["access_level"] == "user"
|
||||
assert user_a["user"]["uid"] in saved_config["user_uids"]
|
||||
|
||||
assert database["db_id"] in await _accessible_db_ids(test_client, user_a["headers"])
|
||||
assert database["db_id"] not in await _accessible_db_ids(test_client, user_b["headers"])
|
||||
assert database["slug"] in await _accessible_slugs(test_client, user_a["headers"])
|
||||
assert database["slug"] not in await _accessible_slugs(test_client, user_b["headers"])
|
||||
finally:
|
||||
if database:
|
||||
await test_client.delete(f"/api/knowledge/databases/{database['db_id']}", headers=admin_headers)
|
||||
await test_client.delete(f"/api/knowledge/databases/{database['slug']}", headers=admin_headers)
|
||||
if user_a:
|
||||
await _delete_user_by_id(test_client, admin_headers, user_a["user"]["id"])
|
||||
if user_b:
|
||||
@ -703,19 +703,19 @@ async def test_create_milvus_knowledge_base(test_client, admin_headers):
|
||||
|
||||
async def test_sample_questions_endpoints(test_client, admin_headers, knowledge_database):
|
||||
"""测试示例问题接口(空文件时预期返回400)"""
|
||||
db_id = knowledge_database["db_id"]
|
||||
slug = knowledge_database["slug"]
|
||||
|
||||
# 获取示例问题(空知识库应该返回空列表)
|
||||
get_response = await test_client.get(f"/api/knowledge/databases/{db_id}/sample-questions", headers=admin_headers)
|
||||
get_response = await test_client.get(f"/api/knowledge/databases/{slug}/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 get_payload["slug"] == slug
|
||||
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",
|
||||
f"/api/knowledge/databases/{slug}/sample-questions",
|
||||
json={"count": 5},
|
||||
headers=admin_headers,
|
||||
)
|
||||
@ -725,18 +725,18 @@ async def test_sample_questions_endpoints(test_client, admin_headers, knowledge_
|
||||
|
||||
async def test_mindmap_permissions(test_client, standard_user, knowledge_database):
|
||||
"""测试思维导图接口的权限控制"""
|
||||
db_id = knowledge_database["db_id"]
|
||||
slug = knowledge_database["slug"]
|
||||
|
||||
# 普通用户应该无法访问
|
||||
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"])
|
||||
forbidden_files = await test_client.get(f"/api/mindmap/databases/{slug}/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": []},
|
||||
json={"slug": slug, "file_ids": []},
|
||||
headers=standard_user["headers"],
|
||||
)
|
||||
_assert_forbidden_response(forbidden_generate)
|
||||
|
||||
@ -73,11 +73,11 @@ async def test_enqueue_document_creates_task(
|
||||
headers=admin_headers,
|
||||
)
|
||||
assert create_response.status_code == 200, create_response.text
|
||||
db_id = create_response.json()["db_id"]
|
||||
slug = create_response.json()["slug"]
|
||||
|
||||
try:
|
||||
enqueue_response = await test_client.post(
|
||||
f"/api/knowledge/databases/{db_id}/documents",
|
||||
f"/api/knowledge/databases/{slug}/documents",
|
||||
json={
|
||||
"items": [],
|
||||
"params": {"content_type": "file"},
|
||||
@ -119,4 +119,4 @@ async def test_enqueue_document_creates_task(
|
||||
else:
|
||||
pytest.fail("Task did not reach a terminal status within timeout window")
|
||||
finally:
|
||||
await test_client.delete(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
await test_client.delete(f"/api/knowledge/databases/{slug}", headers=admin_headers)
|
||||
|
||||
@ -23,11 +23,11 @@ async def _create_dify_database(test_client, admin_headers) -> str:
|
||||
headers=admin_headers,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
return response.json()["db_id"]
|
||||
return response.json()["slug"]
|
||||
|
||||
|
||||
async def _delete_database(test_client, admin_headers, db_id: str) -> None:
|
||||
await test_client.delete(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
async def _delete_database(test_client, admin_headers, slug: str) -> None:
|
||||
await test_client.delete(f"/api/knowledge/databases/{slug}", headers=admin_headers)
|
||||
|
||||
|
||||
async def test_graph_routes_require_auth(test_client):
|
||||
@ -49,16 +49,16 @@ async def test_get_graphs_list_only_returns_milvus(test_client, admin_headers, k
|
||||
assert isinstance(payload["data"], list)
|
||||
assert payload["data"]
|
||||
assert all(graph["type"] == "milvus" for graph in payload["data"])
|
||||
assert any(graph["id"] == knowledge_database["db_id"] for graph in payload["data"])
|
||||
assert any(graph["id"] == knowledge_database["slug"] for graph in payload["data"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["/api/graph/subgraph", "/api/graph/stats", "/api/graph/labels"])
|
||||
async def test_graph_endpoints_reject_non_milvus_types(test_client, admin_headers, path):
|
||||
db_id = await _create_dify_database(test_client, admin_headers)
|
||||
slug = await _create_dify_database(test_client, admin_headers)
|
||||
try:
|
||||
response = await test_client.get(path, params={"db_id": db_id}, headers=admin_headers)
|
||||
response = await test_client.get(path, params={"slug": slug}, headers=admin_headers)
|
||||
finally:
|
||||
await _delete_database(test_client, admin_headers, db_id)
|
||||
await _delete_database(test_client, admin_headers, slug)
|
||||
|
||||
assert response.status_code == 404
|
||||
assert "only supports Milvus" in response.text
|
||||
@ -67,7 +67,7 @@ async def test_graph_endpoints_reject_non_milvus_types(test_client, admin_header
|
||||
async def test_milvus_subgraph_endpoint(test_client, admin_headers, knowledge_database):
|
||||
response = await test_client.get(
|
||||
"/api/graph/subgraph",
|
||||
params={"db_id": knowledge_database["db_id"], "node_label": "*", "max_nodes": 10},
|
||||
params={"slug": knowledge_database["slug"], "node_label": "*", "max_nodes": 10},
|
||||
headers=admin_headers,
|
||||
)
|
||||
|
||||
@ -81,7 +81,7 @@ async def test_milvus_subgraph_endpoint(test_client, admin_headers, knowledge_da
|
||||
async def test_milvus_stats_endpoint(test_client, admin_headers, knowledge_database):
|
||||
response = await test_client.get(
|
||||
"/api/graph/stats",
|
||||
params={"db_id": knowledge_database["db_id"]},
|
||||
params={"slug": knowledge_database["slug"]},
|
||||
headers=admin_headers,
|
||||
)
|
||||
|
||||
@ -96,7 +96,7 @@ async def test_milvus_stats_endpoint(test_client, admin_headers, knowledge_datab
|
||||
async def test_milvus_labels_endpoint(test_client, admin_headers, knowledge_database):
|
||||
response = await test_client.get(
|
||||
"/api/graph/labels",
|
||||
params={"db_id": knowledge_database["db_id"]},
|
||||
params={"slug": knowledge_database["slug"]},
|
||||
headers=admin_headers,
|
||||
)
|
||||
|
||||
|
||||
@ -105,7 +105,7 @@ async def test_viewer_tree_root_does_not_require_sandbox_listing(test_client, st
|
||||
service_module = importlib.import_module("yuxi.services.viewer_filesystem_service")
|
||||
|
||||
async def _fake_resolve_viewer_state(**kwargs):
|
||||
return _FailingSandbox(), _EmptyBackend(), _EmptyBackend(), []
|
||||
return _FailingSandbox(), _EmptyBackend(), []
|
||||
|
||||
monkeypatch.setattr(service_module, "_resolve_viewer_state", _fake_resolve_viewer_state)
|
||||
|
||||
@ -141,7 +141,7 @@ async def test_viewer_tree_user_data_uses_local_thread_directory(test_client, st
|
||||
service_module = importlib.import_module("yuxi.services.viewer_filesystem_service")
|
||||
|
||||
async def _fake_resolve_viewer_state(**kwargs):
|
||||
return _FailingSandbox(), _EmptyBackend(), _EmptyBackend(), []
|
||||
return _FailingSandbox(), _EmptyBackend(), []
|
||||
|
||||
monkeypatch.setattr(service_module, "_resolve_viewer_state", _fake_resolve_viewer_state)
|
||||
|
||||
|
||||
@ -146,15 +146,15 @@ def cleanup_test_knowledge_databases():
|
||||
prefixes = ("pytest_", "py_test")
|
||||
for entry in databases:
|
||||
name = entry.get("name") or ""
|
||||
db_id = entry.get("db_id")
|
||||
if not db_id or not isinstance(name, str) or not name.startswith(prefixes):
|
||||
slug = entry.get("slug")
|
||||
if not slug or not isinstance(name, str) or not name.startswith(prefixes):
|
||||
continue
|
||||
try:
|
||||
delete_response = await client.delete(f"/api/knowledge/databases/{db_id}", headers=headers)
|
||||
delete_response = await client.delete(f"/api/knowledge/databases/{slug}", headers=headers)
|
||||
if delete_response.status_code not in (200, 404):
|
||||
print(f"Warning: Failed to cleanup knowledge database {db_id}: {delete_response.text}")
|
||||
print(f"Warning: Failed to cleanup knowledge database {slug}: {delete_response.text}")
|
||||
except Exception as exc:
|
||||
print(f"Warning: Exception during cleanup of {db_id}: {exc}")
|
||||
print(f"Warning: Exception during cleanup of {slug}: {exc}")
|
||||
|
||||
try:
|
||||
anyio.run(run_cleanup)
|
||||
@ -275,7 +275,7 @@ async def knowledge_database(
|
||||
unique_id = uuid.uuid4().hex
|
||||
timestamp = int(time.time() * 1000000)
|
||||
db_name = f"pytest_kb_{timestamp}_{unique_id}"
|
||||
db_id = None
|
||||
slug = None
|
||||
|
||||
try:
|
||||
create_response = await test_client.post(
|
||||
@ -292,7 +292,7 @@ async def knowledge_database(
|
||||
|
||||
if create_response.status_code == 200:
|
||||
db_payload = create_response.json()
|
||||
db_id = db_payload["db_id"]
|
||||
slug = db_payload["slug"]
|
||||
elif create_response.status_code == 409:
|
||||
error_detail = create_response.json().get("detail", "")
|
||||
pytest.fail(f"Knowledge database name conflict: {error_detail}. Please clean up old test databases first.")
|
||||
@ -301,13 +301,13 @@ async def knowledge_database(
|
||||
f"Failed to create knowledge database (status={create_response.status_code}): {create_response.text}"
|
||||
)
|
||||
|
||||
yield db_payload if db_id else {"db_id": db_id, "name": db_name}
|
||||
yield db_payload if slug else {"slug": slug, "name": db_name}
|
||||
|
||||
finally:
|
||||
if db_id:
|
||||
if slug:
|
||||
try:
|
||||
delete_response = await test_client.delete(f"/api/knowledge/databases/{db_id}", headers=admin_headers)
|
||||
delete_response = await test_client.delete(f"/api/knowledge/databases/{slug}", headers=admin_headers)
|
||||
if delete_response.status_code != 200:
|
||||
print(f"Warning: Failed to cleanup knowledge database {db_id}: {delete_response.text}")
|
||||
print(f"Warning: Failed to cleanup knowledge database {slug}: {delete_response.text}")
|
||||
except Exception as exc:
|
||||
print(f"Warning: Exception during cleanup of {db_id}: {exc}")
|
||||
print(f"Warning: Exception during cleanup of {slug}: {exc}")
|
||||
|
||||
@ -80,27 +80,49 @@ def test_filter_config_by_role_keeps_admin_context_values_for_admin():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normalize_agent_context_config_expands_null_and_filters_explicit_lists(monkeypatch):
|
||||
async def fake_get_databases_by_user(_user):
|
||||
return {"databases": [{"db_id": "kb-a"}, {"db_id": "kb-b"}]}
|
||||
|
||||
async def fake_get_enabled_mcp_server_names(db=None):
|
||||
del db
|
||||
return ["mcp-a", "mcp-b"]
|
||||
|
||||
async def fake_list_skill_slugs(_db):
|
||||
return ["skill-a", "skill-b"]
|
||||
|
||||
async def fake_get_enabled_subagent_names(_db=None):
|
||||
return ["research-agent", "critique-agent"]
|
||||
async def test_resolve_agent_resource_options_empty_fields_loads_nothing(monkeypatch):
|
||||
async def fail_if_loaded(*_args, **_kwargs):
|
||||
raise AssertionError("empty resource_fields should not load resources")
|
||||
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.agents.toolkits",
|
||||
"yuxi.knowledge",
|
||||
types.SimpleNamespace(knowledge_base=types.SimpleNamespace(get_databases_by_user=fail_if_loaded)),
|
||||
)
|
||||
|
||||
assert await context_module.resolve_agent_resource_options(set(), db=object(), user=object()) == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normalize_agent_context_config_expands_null_and_filters_explicit_lists(monkeypatch):
|
||||
async def fake_get_databases_by_user(_user):
|
||||
return {"databases": [{"kb_id": "kb-a"}, {"kb_id": "kb-b"}]}
|
||||
|
||||
async def fake_get_all_mcp_servers(_db):
|
||||
return [
|
||||
types.SimpleNamespace(slug="mcp-a", name="MCP A", description="", enabled=True),
|
||||
types.SimpleNamespace(slug="mcp-b", name="MCP B", description="", enabled=True),
|
||||
]
|
||||
|
||||
async def fake_list_skills(_db):
|
||||
return [
|
||||
types.SimpleNamespace(slug="skill-a", name="Skill A", description=""),
|
||||
types.SimpleNamespace(slug="skill-b", name="Skill B", description=""),
|
||||
]
|
||||
|
||||
async def fake_get_all_subagents(_db=None):
|
||||
return [
|
||||
{"slug": "research-agent", "name": "Research", "description": "", "enabled": True},
|
||||
{"slug": "critique-agent", "name": "Critique", "description": "", "enabled": True},
|
||||
]
|
||||
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.tool_service",
|
||||
types.SimpleNamespace(
|
||||
get_all_tool_instances=lambda: [
|
||||
types.SimpleNamespace(name="ask_user_question"),
|
||||
types.SimpleNamespace(name="tavily_search"),
|
||||
get_tool_metadata=lambda: [
|
||||
{"slug": "ask_user_question", "name": "Ask User", "description": ""},
|
||||
{"slug": "tavily_search", "name": "Tavily", "description": ""},
|
||||
]
|
||||
),
|
||||
)
|
||||
@ -112,17 +134,17 @@ async def test_normalize_agent_context_config_expands_null_and_filters_explicit_
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.mcp_service",
|
||||
types.SimpleNamespace(get_enabled_mcp_server_names=fake_get_enabled_mcp_server_names),
|
||||
types.SimpleNamespace(get_all_mcp_servers=fake_get_all_mcp_servers),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.skill_service",
|
||||
types.SimpleNamespace(list_skill_slugs=fake_list_skill_slugs),
|
||||
types.SimpleNamespace(list_skills=fake_list_skills),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.subagent_service",
|
||||
types.SimpleNamespace(get_enabled_subagent_names=fake_get_enabled_subagent_names),
|
||||
types.SimpleNamespace(get_all_subagents=fake_get_all_subagents),
|
||||
)
|
||||
|
||||
normalized = await normalize_agent_context_config(
|
||||
@ -144,3 +166,170 @@ async def test_normalize_agent_context_config_expands_null_and_filters_explicit_
|
||||
assert normalized["skills"] == []
|
||||
assert normalized["subagents"] == ["research-agent"]
|
||||
assert "summary_threshold" not in normalized
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_agent_runtime_context_filters_resources_and_derives_runtime_scope(monkeypatch):
|
||||
async def fake_get_databases_by_user(_user):
|
||||
return {"databases": [{"kb_id": "kb-a"}, {"kb_id": "kb-b"}]}
|
||||
|
||||
async def fake_get_all_mcp_servers(_db):
|
||||
return [types.SimpleNamespace(slug="mcp-a", name="MCP A", description="", enabled=True)]
|
||||
|
||||
async def fake_list_skills(_db):
|
||||
return [
|
||||
types.SimpleNamespace(slug="skill-a", name="Skill A", description=""),
|
||||
types.SimpleNamespace(slug="skill-b", name="Skill B", description=""),
|
||||
]
|
||||
|
||||
async def fake_get_all_subagents(_db=None):
|
||||
return [{"slug": "sub-a", "name": "Sub A", "description": "", "enabled": True}]
|
||||
|
||||
async def fake_resolve_visible_knowledge_bases(context):
|
||||
assert context.knowledges == ["kb-a"]
|
||||
context._visible_knowledge_bases = [{"slug": "kb-a", "name": "Docs A"}]
|
||||
return context._visible_knowledge_bases
|
||||
|
||||
async def fake_resolve_runtime_skills_for_context(context, *, db=None):
|
||||
del db
|
||||
assert context.skills == ["skill-a"]
|
||||
return {
|
||||
"context_skills": ["skill-a"],
|
||||
"prompt_skills": ["skill-a", "skill-b"],
|
||||
"readable_skills": ["skill-a", "skill-b"],
|
||||
}
|
||||
|
||||
class FakeSessionContext:
|
||||
async def __aenter__(self):
|
||||
return object()
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return None
|
||||
|
||||
class FakeUserRepository:
|
||||
async def get_by_uid_with_db(self, _db, uid):
|
||||
assert uid == "u1"
|
||||
return types.SimpleNamespace(role="user", uid="u1")
|
||||
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.agents.backends.knowledge_base_backend",
|
||||
types.SimpleNamespace(resolve_visible_knowledge_bases_for_context=fake_resolve_visible_knowledge_bases),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.agents.middlewares.skills_middleware",
|
||||
types.SimpleNamespace(resolve_runtime_skills_for_context=fake_resolve_runtime_skills_for_context),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.repositories.user_repository",
|
||||
types.SimpleNamespace(UserRepository=FakeUserRepository),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.storage.postgres.manager",
|
||||
types.SimpleNamespace(pg_manager=types.SimpleNamespace(get_async_session_context=lambda: FakeSessionContext())),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.tool_service",
|
||||
types.SimpleNamespace(
|
||||
get_tool_metadata=lambda: [{"slug": "ask_user_question", "name": "Ask User", "description": ""}]
|
||||
),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.knowledge",
|
||||
types.SimpleNamespace(knowledge_base=types.SimpleNamespace(get_databases_by_user=fake_get_databases_by_user)),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.mcp_service",
|
||||
types.SimpleNamespace(get_all_mcp_servers=fake_get_all_mcp_servers),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.skill_service",
|
||||
types.SimpleNamespace(list_skills=fake_list_skills),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.services.subagent_service",
|
||||
types.SimpleNamespace(get_all_subagents=fake_get_all_subagents),
|
||||
)
|
||||
|
||||
context = BaseContext(
|
||||
uid="u1",
|
||||
tools=["ask_user_question", "missing"],
|
||||
knowledges=["kb-a", "missing"],
|
||||
mcps=None,
|
||||
skills=["skill-a", "missing"],
|
||||
subagents=[],
|
||||
)
|
||||
|
||||
prepared = await context_module.prepare_agent_runtime_context(context)
|
||||
|
||||
assert prepared.tools == ["ask_user_question"]
|
||||
assert prepared.knowledges == ["kb-a"]
|
||||
assert prepared.mcps == ["mcp-a"]
|
||||
assert prepared.skills == ["skill-a"]
|
||||
assert prepared.subagents == []
|
||||
assert prepared._visible_knowledge_bases == [{"slug": "kb-a", "name": "Docs A"}]
|
||||
assert prepared._prompt_skills == ["skill-a", "skill-b"]
|
||||
assert prepared._readable_skills == ["skill-a", "skill-b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_agent_runtime_context_clears_resources_for_missing_user(monkeypatch):
|
||||
class FakeSessionContext:
|
||||
async def __aenter__(self):
|
||||
return object()
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return None
|
||||
|
||||
class FakeUserRepository:
|
||||
async def get_by_uid_with_db(self, _db, _uid):
|
||||
return None
|
||||
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.agents.backends.knowledge_base_backend",
|
||||
types.SimpleNamespace(resolve_visible_knowledge_bases_for_context=lambda _context: None),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.agents.middlewares.skills_middleware",
|
||||
types.SimpleNamespace(resolve_runtime_skills_for_context=lambda _context, db=None: None),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.repositories.user_repository",
|
||||
types.SimpleNamespace(UserRepository=FakeUserRepository),
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"yuxi.storage.postgres.manager",
|
||||
types.SimpleNamespace(pg_manager=types.SimpleNamespace(get_async_session_context=lambda: FakeSessionContext())),
|
||||
)
|
||||
|
||||
context = BaseContext(
|
||||
uid="missing",
|
||||
tools=["tool"],
|
||||
knowledges=["kb"],
|
||||
mcps=["mcp"],
|
||||
skills=["skill"],
|
||||
subagents=["agent"],
|
||||
)
|
||||
|
||||
prepared = await context_module.prepare_agent_runtime_context(context)
|
||||
|
||||
assert prepared.tools == []
|
||||
assert prepared.knowledges == []
|
||||
assert prepared.mcps == []
|
||||
assert prepared.skills == []
|
||||
assert prepared.subagents == []
|
||||
assert prepared._visible_knowledge_bases == []
|
||||
assert prepared._prompt_skills == []
|
||||
assert prepared._readable_skills == []
|
||||
|
||||
@ -350,7 +350,7 @@ async def test_install_skill_git_with_skill_names_passes_admin_check(mock_pg):
|
||||
# Mock enabling skill in config
|
||||
mock_enable.return_value = True
|
||||
|
||||
with patch("yuxi.services.skill_service.sync_thread_visible_skills"):
|
||||
with patch("yuxi.services.skill_service.sync_thread_readable_skills"):
|
||||
result = await _install_skill_func(
|
||||
source="owner/repo",
|
||||
skill_names=["test-skill"],
|
||||
@ -387,7 +387,7 @@ async def test_install_skill_sandbox_success(mock_pg):
|
||||
) as mock_enable:
|
||||
mock_enable.return_value = True
|
||||
|
||||
with patch("yuxi.services.skill_service.sync_thread_visible_skills"):
|
||||
with patch("yuxi.services.skill_service.sync_thread_readable_skills"):
|
||||
result = await _install_skill_func(
|
||||
source="/home/gem/user-data/workspace/my-skill",
|
||||
runtime=runtime,
|
||||
@ -476,7 +476,7 @@ async def test_install_skill_partial_config_failure(mock_pg):
|
||||
# Simulate config persistence failure
|
||||
mock_enable.return_value = False
|
||||
|
||||
with patch("yuxi.services.skill_service.sync_thread_visible_skills"):
|
||||
with patch("yuxi.services.skill_service.sync_thread_readable_skills"):
|
||||
result = await _install_skill_func(
|
||||
source="/home/gem/user-data/workspace/my-skill",
|
||||
runtime=runtime,
|
||||
@ -514,7 +514,7 @@ async def test_install_skill_slug_warning_for_renamed(mock_pg):
|
||||
) as mock_enable:
|
||||
mock_enable.return_value = True
|
||||
|
||||
with patch("yuxi.services.skill_service.sync_thread_visible_skills"):
|
||||
with patch("yuxi.services.skill_service.sync_thread_readable_skills"):
|
||||
result = await _install_skill_func(
|
||||
source="/home/gem/user-data/workspace/my-skill",
|
||||
runtime=runtime,
|
||||
|
||||
41
backend/test/unit/backends/test_knowledge_base_backend.py
Normal file
41
backend/test/unit/backends/test_knowledge_base_backend.py
Normal file
@ -0,0 +1,41 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
import yuxi.agents.backends.knowledge_base_backend as knowledge_base_backend
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_visible_knowledge_bases_filters_by_slug(monkeypatch):
|
||||
async def fake_get_databases_by_uid(_uid):
|
||||
return {
|
||||
"databases": [
|
||||
{"slug": "kb-a", "name": "Same Name"},
|
||||
{"slug": "kb-b", "name": "Same Name"},
|
||||
]
|
||||
}
|
||||
|
||||
monkeypatch.setattr(knowledge_base_backend.knowledge_base, "get_databases_by_uid", fake_get_databases_by_uid)
|
||||
|
||||
context = SimpleNamespace(uid="u1", knowledges=["kb-b"])
|
||||
|
||||
databases = await knowledge_base_backend.resolve_visible_knowledge_bases_for_context(context)
|
||||
|
||||
assert databases == [{"slug": "kb-b", "name": "Same Name"}]
|
||||
assert context._visible_knowledge_bases == databases
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_visible_knowledge_bases_requires_slug(monkeypatch):
|
||||
async def fake_get_databases_by_uid(_uid):
|
||||
return {"databases": [{"id": "legacy-id", "name": "Legacy"}]}
|
||||
|
||||
monkeypatch.setattr(knowledge_base_backend.knowledge_base, "get_databases_by_uid", fake_get_databases_by_uid)
|
||||
|
||||
context = SimpleNamespace(uid="u1", knowledges=["legacy-id"])
|
||||
|
||||
databases = await knowledge_base_backend.resolve_visible_knowledge_bases_for_context(context)
|
||||
|
||||
assert databases == []
|
||||
@ -16,6 +16,7 @@ def _runtime(
|
||||
thread_id: str | None = "thread-1",
|
||||
uid: str | None = "user-1",
|
||||
skills: list[str] | None = None,
|
||||
readable_skills: list[str] | None = None,
|
||||
visible_kbs: list[dict] | None = None,
|
||||
):
|
||||
configurable = {"thread_id": thread_id, "uid": uid} if thread_id and uid else {}
|
||||
@ -23,21 +24,22 @@ def _runtime(
|
||||
config={"configurable": configurable},
|
||||
context=SimpleNamespace(
|
||||
skills=skills or [],
|
||||
_readable_skills=readable_skills,
|
||||
_visible_knowledge_bases=visible_kbs or [],
|
||||
uid=uid,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_create_agent_composite_backend_uses_provisioner_default(monkeypatch):
|
||||
def test_create_agent_composite_backend_uses_prepared_readable_skills(monkeypatch):
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
|
||||
backend = create_agent_composite_backend(
|
||||
_runtime(skills=["reporter"], visible_kbs=[{"db_id": "db-1", "name": "Docs"}])
|
||||
_runtime(readable_skills=["reporter"], visible_kbs=[{"slug": "db-1", "name": "Docs"}])
|
||||
)
|
||||
|
||||
assert isinstance(backend.default, ProvisionerSandboxBackend)
|
||||
assert backend.default._visible_skills == ["reporter"]
|
||||
assert backend.default._readable_skills == ["reporter"]
|
||||
assert "/skills/" in backend.routes
|
||||
assert "/home/gem/kbs/" not in backend.routes
|
||||
|
||||
@ -47,6 +49,14 @@ def test_create_agent_composite_backend_requires_thread_id():
|
||||
create_agent_composite_backend(_runtime(thread_id=None))
|
||||
|
||||
|
||||
def test_create_agent_composite_backend_ignores_unprepared_context_skills(monkeypatch):
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
|
||||
backend = create_agent_composite_backend(_runtime(skills=["configured"], readable_skills=None))
|
||||
|
||||
assert backend.default._readable_skills == []
|
||||
|
||||
|
||||
def test_skills_middleware_extracts_slug_for_new_paths() -> None:
|
||||
middleware = SkillsMiddleware()
|
||||
assert middleware.skills_sources_for_prompt == ["/home/gem/skills/"]
|
||||
|
||||
@ -102,19 +102,19 @@ async def test_milvus_graph_service_configure_persists_updated_concurrency():
|
||||
)
|
||||
|
||||
class Repo:
|
||||
async def get_by_id(self, db_id):
|
||||
async def get_by_kb_id(self, kb_id):
|
||||
return kb
|
||||
|
||||
async def update(self, db_id, data):
|
||||
async def update(self, kb_id, data):
|
||||
kb.additional_params = data["additional_params"]
|
||||
return kb
|
||||
|
||||
chunk_repo = SimpleNamespace(
|
||||
count_by_db_id=AsyncMock(return_value=0),
|
||||
count_graph_pending_by_db_id=AsyncMock(return_value=0),
|
||||
count_graph_indexed_by_db_id=AsyncMock(return_value=0),
|
||||
count_by_kb_id=AsyncMock(return_value=0),
|
||||
count_graph_pending_by_kb_id=AsyncMock(return_value=0),
|
||||
count_graph_indexed_by_kb_id=AsyncMock(return_value=0),
|
||||
)
|
||||
graph_repo = SimpleNamespace(count_by_db_id=AsyncMock(return_value=(3, 2)))
|
||||
graph_repo = SimpleNamespace(count_by_kb_id=AsyncMock(return_value=(3, 2)))
|
||||
service = MilvusGraphService(kb_repo=Repo(), chunk_repo=chunk_repo, graph_repo=graph_repo)
|
||||
|
||||
await service.configure(
|
||||
@ -142,7 +142,7 @@ def test_milvus_graph_service_writes_chunk_entity_and_relation():
|
||||
chunk = SimpleNamespace(
|
||||
chunk_id="chunk_1",
|
||||
file_id="file_1",
|
||||
db_id="kb_test",
|
||||
kb_id="kb_test",
|
||||
chunk_index=1,
|
||||
content="张三任职于公司",
|
||||
start_char_pos=0,
|
||||
@ -182,25 +182,25 @@ def test_milvus_graph_service_writes_chunk_entity_and_relation():
|
||||
assert entity_call.kwargs["attributes"] == '[{"text": "工程师", "label": "Occupation"}]'
|
||||
|
||||
|
||||
def test_milvus_graph_service_query_nodes_empty_db_id():
|
||||
def test_milvus_graph_service_query_nodes_empty_kb_id():
|
||||
service = MilvusGraphService()
|
||||
import asyncio
|
||||
|
||||
result = asyncio.get_event_loop().run_until_complete(service.query_nodes(db_id=None, keyword="test"))
|
||||
result = asyncio.get_event_loop().run_until_complete(service.query_nodes(kb_id=None, keyword="test"))
|
||||
assert result == {"nodes": [], "edges": []}
|
||||
|
||||
|
||||
def test_milvus_graph_service_get_labels_empty_db_id():
|
||||
def test_milvus_graph_service_get_labels_empty_kb_id():
|
||||
service = MilvusGraphService()
|
||||
import asyncio
|
||||
|
||||
result = asyncio.get_event_loop().run_until_complete(service.get_labels(db_id=None))
|
||||
result = asyncio.get_event_loop().run_until_complete(service.get_labels(kb_id=None))
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_milvus_graph_service_get_stats_empty_db_id():
|
||||
def test_milvus_graph_service_get_stats_empty_kb_id():
|
||||
service = MilvusGraphService()
|
||||
import asyncio
|
||||
|
||||
result = asyncio.get_event_loop().run_until_complete(service.get_stats(db_id=None))
|
||||
result = asyncio.get_event_loop().run_until_complete(service.get_stats(kb_id=None))
|
||||
assert result == {"total_nodes": 0, "total_edges": 0, "entity_types": []}
|
||||
|
||||
@ -20,11 +20,11 @@ from yuxi.knowledge.eval.benchmark_generation import (
|
||||
|
||||
class FakeKnowledgeBase:
|
||||
files_meta = {
|
||||
"file_a": {"database_id": "db_1"},
|
||||
"file_b": {"database_id": "db_2"},
|
||||
"file_a": {"kb_id": "db_1"},
|
||||
"file_b": {"kb_id": "db_2"},
|
||||
}
|
||||
|
||||
async def get_file_content(self, db_id, fid):
|
||||
async def get_file_content(self, kb_id, fid):
|
||||
return {
|
||||
"lines": [
|
||||
{"id": f"{fid}_chunk", "content": "内容", "chunk_order_index": 0},
|
||||
@ -33,21 +33,21 @@ class FakeKnowledgeBase:
|
||||
|
||||
|
||||
class FakeGenerationKnowledgeBase:
|
||||
files_meta = {"file_a": {"database_id": "db_1"}}
|
||||
files_meta = {"file_a": {"kb_id": "db_1"}}
|
||||
|
||||
def __init__(self, query_results=None):
|
||||
self.query_results = query_results or []
|
||||
self.query_calls = []
|
||||
|
||||
async def get_file_content(self, db_id, fid):
|
||||
async def get_file_content(self, kb_id, fid):
|
||||
return {
|
||||
"lines": [
|
||||
{"id": "anchor_chunk", "content": "anchor content", "chunk_order_index": 0},
|
||||
]
|
||||
}
|
||||
|
||||
async def aquery(self, query_text, db_id, **kwargs):
|
||||
self.query_calls.append({"query_text": query_text, "db_id": db_id, **kwargs})
|
||||
async def aquery(self, query_text, kb_id, **kwargs):
|
||||
self.query_calls.append({"query_text": query_text, "kb_id": kb_id, **kwargs})
|
||||
return self.query_results
|
||||
|
||||
|
||||
@ -64,7 +64,7 @@ class FakeLlm:
|
||||
|
||||
|
||||
class NoQueryKnowledgeBase(FakeGenerationKnowledgeBase):
|
||||
async def aquery(self, query_text, db_id, **kwargs):
|
||||
async def aquery(self, query_text, kb_id, **kwargs):
|
||||
raise AssertionError("neighbors_count=1 时不应调用 aquery")
|
||||
|
||||
|
||||
@ -89,7 +89,7 @@ class TrackingLlm:
|
||||
|
||||
|
||||
class FakeGraphGenerationKnowledgeBase(FakeGenerationKnowledgeBase):
|
||||
async def get_file_content(self, db_id, fid):
|
||||
async def get_file_content(self, kb_id, fid):
|
||||
return {
|
||||
"lines": [
|
||||
{
|
||||
@ -139,7 +139,7 @@ def test_build_benchmark_generation_prompt_contains_required_schema():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collect_kb_chunks_filters_database_id():
|
||||
async def test_collect_kb_chunks_filters_kb_id():
|
||||
chunks = await collect_kb_chunks(FakeKnowledgeBase(), "db_1")
|
||||
|
||||
assert chunks == [
|
||||
@ -165,7 +165,7 @@ async def test_iter_generated_benchmark_items_with_one_chunk_does_not_query(monk
|
||||
item
|
||||
async for item in iter_generated_benchmark_items(
|
||||
kb_instance=NoQueryKnowledgeBase(),
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
count=1,
|
||||
neighbors_count=1,
|
||||
llm_model_spec="test-provider:test-model",
|
||||
@ -193,7 +193,7 @@ async def test_select_neighbor_chunks_by_kb_query_filters_anchor():
|
||||
|
||||
chunks = await select_neighbor_chunks_by_kb_query(
|
||||
kb_instance=kb,
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
anchor_chunk={"id": "anchor_chunk", "content": "anchor content", "file_id": "file_a", "chunk_index": 0},
|
||||
neighbors_count=1,
|
||||
)
|
||||
@ -202,7 +202,7 @@ async def test_select_neighbor_chunks_by_kb_query_filters_anchor():
|
||||
assert kb.query_calls == [
|
||||
{
|
||||
"query_text": "anchor content",
|
||||
"db_id": "db_1",
|
||||
"kb_id": "db_1",
|
||||
"search_mode": "vector",
|
||||
"final_top_k": 4,
|
||||
"use_reranker": False,
|
||||
@ -215,7 +215,7 @@ async def test_select_neighbor_chunks_by_kb_query_filters_anchor():
|
||||
async def test_select_graph_enhanced_chunks_expands_by_ppr_with_anchor_bias(monkeypatch):
|
||||
calls = []
|
||||
|
||||
async def fake_rank(self, db_id, seed_weights, *, max_nodes, top_k, damping):
|
||||
async def fake_rank(self, kb_id, seed_weights, *, max_nodes, top_k, damping):
|
||||
calls.append(dict(seed_weights))
|
||||
if len(calls) == 1:
|
||||
return [("anchor", 0.9), ("neighbor_1", 0.8)]
|
||||
@ -232,7 +232,7 @@ async def test_select_graph_enhanced_chunks_expands_by_ppr_with_anchor_bias(monk
|
||||
}
|
||||
|
||||
chunks = await select_graph_enhanced_chunks(
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
anchor_chunk=chunks_by_id["anchor"],
|
||||
chunks_by_id=chunks_by_id,
|
||||
context_count=3,
|
||||
@ -247,7 +247,7 @@ async def test_select_graph_enhanced_chunks_expands_by_ppr_with_anchor_bias(monk
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_iter_generated_benchmark_items_graph_mode_uses_graph_indexed_anchor(monkeypatch):
|
||||
async def fake_rank(self, db_id, seed_weights, *, max_nodes, top_k, damping):
|
||||
async def fake_rank(self, kb_id, seed_weights, *, max_nodes, top_k, damping):
|
||||
assert seed_weights["anchor_entity"] == 1.0
|
||||
return [("graph_anchor", 0.9), ("graph_neighbor", 0.8)]
|
||||
|
||||
@ -263,7 +263,7 @@ async def test_iter_generated_benchmark_items_graph_mode_uses_graph_indexed_anch
|
||||
item
|
||||
async for item in iter_generated_benchmark_items(
|
||||
kb_instance=kb,
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
count=1,
|
||||
neighbors_count=2,
|
||||
llm_model_spec="test-provider:test-model",
|
||||
@ -295,7 +295,7 @@ async def test_iter_generated_benchmark_items_uses_query_neighbor(monkeypatch):
|
||||
item
|
||||
async for item in iter_generated_benchmark_items(
|
||||
kb_instance=kb,
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
count=1,
|
||||
neighbors_count=2,
|
||||
llm_model_spec="test-provider:test-model",
|
||||
@ -317,7 +317,7 @@ async def test_iter_generated_benchmark_items_falls_back_to_anchor_when_query_em
|
||||
item
|
||||
async for item in iter_generated_benchmark_items(
|
||||
kb_instance=FakeGenerationKnowledgeBase(query_results=[]),
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
count=1,
|
||||
neighbors_count=2,
|
||||
llm_model_spec="test-provider:test-model",
|
||||
@ -337,7 +337,7 @@ async def test_iter_generated_benchmark_items_respects_concurrency_count(monkeyp
|
||||
item
|
||||
async for item in iter_generated_benchmark_items(
|
||||
kb_instance=NoQueryKnowledgeBase(),
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
count=4,
|
||||
neighbors_count=1,
|
||||
concurrency_count=2,
|
||||
@ -358,7 +358,7 @@ async def test_iter_generated_benchmark_items_returns_at_most_count(monkeypatch)
|
||||
item
|
||||
async for item in iter_generated_benchmark_items(
|
||||
kb_instance=NoQueryKnowledgeBase(),
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
count=3,
|
||||
neighbors_count=1,
|
||||
concurrency_count=10,
|
||||
@ -378,7 +378,7 @@ async def test_iter_generated_benchmark_items_stops_at_max_attempts(monkeypatch)
|
||||
item
|
||||
async for item in iter_generated_benchmark_items(
|
||||
kb_instance=NoQueryKnowledgeBase(),
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
count=2,
|
||||
neighbors_count=1,
|
||||
concurrency_count=10,
|
||||
|
||||
@ -22,7 +22,7 @@ class FakeChunkRepository:
|
||||
def __init__(self, indexed_count):
|
||||
self.indexed_count = indexed_count
|
||||
|
||||
async def count_graph_indexed_by_db_id(self, db_id):
|
||||
async def count_graph_indexed_by_kb_id(self, kb_id):
|
||||
return self.indexed_count
|
||||
|
||||
|
||||
@ -37,7 +37,7 @@ async def test_generate_dataset_saves_generation_params(monkeypatch):
|
||||
service.chunk_repo = FakeChunkRepository(indexed_count=1)
|
||||
|
||||
result = await service.generate_dataset(
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
name="dataset",
|
||||
description="desc",
|
||||
count=2,
|
||||
@ -65,7 +65,7 @@ async def test_generate_dataset_rejects_graph_mode_without_indexed_chunks():
|
||||
|
||||
with pytest.raises(ValueError, match="尚未完成图索引"):
|
||||
await service.generate_dataset(
|
||||
db_id="db_1",
|
||||
kb_id="db_1",
|
||||
name="dataset",
|
||||
description="desc",
|
||||
count=2,
|
||||
|
||||
@ -67,7 +67,7 @@ class TestAddFileRecordSizeFallback:
|
||||
def kb_type(self):
|
||||
return "test"
|
||||
|
||||
async def _create_kb_instance(self, db_id, config):
|
||||
async def _create_kb_instance(self, kb_id, config):
|
||||
pass
|
||||
|
||||
async def _initialize_kb_instance(self, instance):
|
||||
@ -76,34 +76,34 @@ class TestAddFileRecordSizeFallback:
|
||||
async def _persist_file(self, file_id):
|
||||
pass
|
||||
|
||||
async def _persist_kb(self, db_id):
|
||||
async def _persist_kb(self, kb_id):
|
||||
pass
|
||||
|
||||
async def _save_metadata(self):
|
||||
pass
|
||||
|
||||
async def index_file(self, db_id, file_id, operator_id=None):
|
||||
async def index_file(self, kb_id, file_id, operator_id=None):
|
||||
return {}
|
||||
|
||||
async def aquery(self, query_text, db_id, **kwargs):
|
||||
async def aquery(self, query_text, kb_id, **kwargs):
|
||||
return []
|
||||
|
||||
async def delete_file(self, db_id, file_id):
|
||||
async def delete_file(self, kb_id, file_id):
|
||||
pass
|
||||
|
||||
async def get_file_basic_info(self, db_id, file_id):
|
||||
async def get_file_basic_info(self, kb_id, file_id):
|
||||
return {}
|
||||
|
||||
async def get_file_content(self, db_id, file_id):
|
||||
async def get_file_content(self, kb_id, file_id):
|
||||
return {}
|
||||
|
||||
async def get_file_info(self, db_id, file_id):
|
||||
async def get_file_info(self, kb_id, file_id):
|
||||
return {}
|
||||
|
||||
async def update_content(self, db_id, file_ids, params=None):
|
||||
async def update_content(self, kb_id, file_ids, params=None):
|
||||
return []
|
||||
|
||||
async def get_query_params_config(self, db_id, **kwargs):
|
||||
async def get_query_params_config(self, kb_id, **kwargs):
|
||||
return {"type": "test", "options": []}
|
||||
|
||||
kb = TestKB(work_dir=work_dir)
|
||||
@ -177,34 +177,34 @@ async def test_save_metadata_does_not_overwrite_existing_kb_config():
|
||||
def kb_type(self):
|
||||
return "test"
|
||||
|
||||
async def _create_kb_instance(self, db_id, config):
|
||||
async def _create_kb_instance(self, kb_id, config):
|
||||
pass
|
||||
|
||||
async def _initialize_kb_instance(self, instance):
|
||||
pass
|
||||
|
||||
async def index_file(self, db_id, file_id, operator_id=None):
|
||||
async def index_file(self, kb_id, file_id, operator_id=None):
|
||||
return {}
|
||||
|
||||
async def update_content(self, db_id, file_ids, params=None):
|
||||
async def update_content(self, kb_id, file_ids, params=None):
|
||||
return []
|
||||
|
||||
async def aquery(self, query_text, db_id, **kwargs):
|
||||
async def aquery(self, query_text, kb_id, **kwargs):
|
||||
return []
|
||||
|
||||
def get_query_params_config(self, db_id, **kwargs):
|
||||
def get_query_params_config(self, kb_id, **kwargs):
|
||||
return {"type": "test", "options": []}
|
||||
|
||||
async def delete_file(self, db_id, file_id):
|
||||
async def delete_file(self, kb_id, file_id):
|
||||
pass
|
||||
|
||||
async def get_file_basic_info(self, db_id, file_id):
|
||||
async def get_file_basic_info(self, kb_id, file_id):
|
||||
return {}
|
||||
|
||||
async def get_file_content(self, db_id, file_id):
|
||||
async def get_file_content(self, kb_id, file_id):
|
||||
return {}
|
||||
|
||||
async def get_file_info(self, db_id, file_id):
|
||||
async def get_file_info(self, kb_id, file_id):
|
||||
return {}
|
||||
|
||||
class ExistingKbRepo:
|
||||
@ -212,14 +212,14 @@ async def test_save_metadata_does_not_overwrite_existing_kb_config():
|
||||
self.created = []
|
||||
self.updated = []
|
||||
|
||||
async def get_by_id(self, db_id):
|
||||
return SimpleNamespace(db_id=db_id)
|
||||
async def get_by_kb_id(self, kb_id):
|
||||
return SimpleNamespace(kb_id=kb_id)
|
||||
|
||||
async def create(self, payload):
|
||||
self.created.append(payload)
|
||||
|
||||
async def update(self, db_id, data):
|
||||
self.updated.append((db_id, data))
|
||||
async def update(self, kb_id, data):
|
||||
self.updated.append((kb_id, data))
|
||||
|
||||
kb_repo = ExistingKbRepo()
|
||||
kb = TestKB(work_dir="/tmp/test_kb")
|
||||
|
||||
@ -8,34 +8,34 @@ class FakeKnowledgeBase(KnowledgeBase):
|
||||
def kb_type(self) -> str:
|
||||
return "fake"
|
||||
|
||||
async def _create_kb_instance(self, db_id: str, config: dict):
|
||||
async def _create_kb_instance(self, slug: str, config: dict):
|
||||
return None
|
||||
|
||||
async def _initialize_kb_instance(self, instance) -> None:
|
||||
pass
|
||||
|
||||
async def index_file(self, db_id: str, file_id: str, operator_id: str | None = None) -> dict:
|
||||
async def index_file(self, slug: str, file_id: str, operator_id: str | None = None) -> dict:
|
||||
return {}
|
||||
|
||||
async def update_content(self, db_id: str, file_ids: list[str], params: dict | None = None) -> list[dict]:
|
||||
async def update_content(self, slug: str, file_ids: list[str], params: dict | None = None) -> list[dict]:
|
||||
return []
|
||||
|
||||
async def aquery(self, query_text: str, db_id: str, **kwargs) -> list[dict]:
|
||||
async def aquery(self, query_text: str, slug: str, **kwargs) -> list[dict]:
|
||||
return []
|
||||
|
||||
def get_query_params_config(self, db_id: str, **kwargs) -> dict:
|
||||
def get_query_params_config(self, slug: str, **kwargs) -> dict:
|
||||
return {"options": []}
|
||||
|
||||
async def delete_file(self, db_id: str, file_id: str) -> None:
|
||||
async def delete_file(self, slug: str, file_id: str) -> None:
|
||||
pass
|
||||
|
||||
async def get_file_basic_info(self, db_id: str, file_id: str) -> dict:
|
||||
async def get_file_basic_info(self, slug: str, file_id: str) -> dict:
|
||||
return {}
|
||||
|
||||
async def get_file_content(self, db_id: str, file_id: str) -> dict:
|
||||
async def get_file_content(self, slug: str, file_id: str) -> dict:
|
||||
return {}
|
||||
|
||||
async def get_file_info(self, db_id: str, file_id: str) -> dict:
|
||||
async def get_file_info(self, slug: str, file_id: str) -> dict:
|
||||
return {}
|
||||
|
||||
async def _save_metadata(self) -> None:
|
||||
|
||||
133
backend/test/unit/middlewares/test_skills_middleware.py
Normal file
133
backend/test/unit/middlewares/test_skills_middleware.py
Normal file
@ -0,0 +1,133 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langgraph.types import Command
|
||||
|
||||
import yuxi.agents.middlewares.skills_middleware as skills_middleware
|
||||
from yuxi.agents.middlewares.skills_middleware import SkillsMiddleware, resolve_runtime_skills_for_context
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_runtime_skills_derives_prompt_and_readable_closure(monkeypatch):
|
||||
async def fake_get_dependency_map(db=None):
|
||||
del db
|
||||
return {
|
||||
"alpha": {"tools": [], "mcps": [], "skills": ["beta"]},
|
||||
"beta": {"tools": [], "mcps": [], "skills": []},
|
||||
}
|
||||
|
||||
monkeypatch.setattr(skills_middleware, "get_dependency_map", fake_get_dependency_map)
|
||||
|
||||
context = SimpleNamespace(skills=["alpha", "missing"])
|
||||
|
||||
scope = await resolve_runtime_skills_for_context(context)
|
||||
|
||||
assert scope == {
|
||||
"context_skills": ["alpha"],
|
||||
"prompt_skills": ["alpha", "beta"],
|
||||
"readable_skills": ["alpha", "beta"],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skills_prompt_uses_prepared_prompt_skills(monkeypatch):
|
||||
async def fake_get_prompt_metadata(db=None):
|
||||
del db
|
||||
return {
|
||||
"alpha": {
|
||||
"name": "Alpha",
|
||||
"description": "alpha desc",
|
||||
"path": "/home/gem/skills/alpha/SKILL.md",
|
||||
},
|
||||
"configured-only": {
|
||||
"name": "Configured Only",
|
||||
"description": "should not appear",
|
||||
"path": "/home/gem/skills/configured-only/SKILL.md",
|
||||
},
|
||||
}
|
||||
|
||||
monkeypatch.setattr(skills_middleware, "get_prompt_metadata", fake_get_prompt_metadata)
|
||||
|
||||
context = SimpleNamespace(
|
||||
system_prompt="base",
|
||||
skills=["configured-only"],
|
||||
_prompt_skills=["alpha"],
|
||||
)
|
||||
|
||||
await SkillsMiddleware().abefore_agent({}, SimpleNamespace(context=context))
|
||||
|
||||
assert "base" in context.system_prompt
|
||||
assert "Alpha" in context.system_prompt
|
||||
assert "Configured Only" not in context.system_prompt
|
||||
assert getattr(context, "_skills_prompt_injected") is True
|
||||
assert not hasattr(context, "_visible_skills")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_awrap_model_call_mounts_dependencies_only_for_readable_activated_skills(monkeypatch):
|
||||
async def fake_get_dependency_map(db=None):
|
||||
del db
|
||||
return {
|
||||
"alpha": {"tools": ["tool-a"], "mcps": [], "skills": []},
|
||||
"beta": {"tools": ["tool-b"], "mcps": [], "skills": []},
|
||||
}
|
||||
|
||||
monkeypatch.setattr(skills_middleware, "get_dependency_map", fake_get_dependency_map)
|
||||
monkeypatch.setattr(
|
||||
skills_middleware,
|
||||
"get_all_tool_instances",
|
||||
lambda: [SimpleNamespace(name="tool-a"), SimpleNamespace(name="tool-b")],
|
||||
)
|
||||
|
||||
class FakeRequest:
|
||||
def __init__(self, tools=None):
|
||||
self.runtime = SimpleNamespace(context=SimpleNamespace(_readable_skills=["alpha"], mcps=[]))
|
||||
self.state = {"activated_skills": ["alpha", "beta"]}
|
||||
self.tools = tools or []
|
||||
|
||||
def override(self, *, tools):
|
||||
new_request = FakeRequest(tools=tools)
|
||||
new_request.runtime = self.runtime
|
||||
new_request.state = self.state
|
||||
return new_request
|
||||
|
||||
captured = {}
|
||||
|
||||
async def handler(request):
|
||||
captured["tools"] = [tool.name for tool in request.tools]
|
||||
return "ok"
|
||||
|
||||
result = await SkillsMiddleware().awrap_model_call(FakeRequest(), handler)
|
||||
|
||||
assert result == "ok"
|
||||
assert captured["tools"] == ["tool-a"]
|
||||
|
||||
|
||||
def test_read_file_activates_only_readable_skill() -> None:
|
||||
middleware = SkillsMiddleware()
|
||||
result = ToolMessage(content="ok", tool_call_id="tool-1", name="read_file")
|
||||
request = SimpleNamespace(
|
||||
runtime=SimpleNamespace(context=SimpleNamespace(_readable_skills=["alpha"])),
|
||||
tool_call={"name": "read_file", "args": {"file_path": "/home/gem/skills/alpha/SKILL.md"}},
|
||||
)
|
||||
|
||||
updated = middleware._process_tool_call_result(result, request)
|
||||
|
||||
assert isinstance(updated, Command)
|
||||
assert updated.update["activated_skills"] == ["alpha"]
|
||||
|
||||
|
||||
def test_read_file_denies_skill_outside_readable_scope() -> None:
|
||||
middleware = SkillsMiddleware()
|
||||
result = ToolMessage(content="ok", tool_call_id="tool-1", name="read_file")
|
||||
request = SimpleNamespace(
|
||||
runtime=SimpleNamespace(context=SimpleNamespace(_readable_skills=["alpha"])),
|
||||
tool_call={"name": "read_file", "args": {"file_path": "/home/gem/skills/beta/SKILL.md"}},
|
||||
)
|
||||
|
||||
updated = middleware._process_tool_call_result(result, request)
|
||||
|
||||
assert updated is result
|
||||
@ -79,8 +79,8 @@ def test_dify_validation_rejects_missing_or_invalid_params():
|
||||
@pytest.mark.asyncio
|
||||
async def test_dify_kb_aquery_maps_records(monkeypatch, tmp_path):
|
||||
kb = DifyKB(str(tmp_path))
|
||||
db_id = "kb_test_dify"
|
||||
kb.databases_meta[db_id] = {
|
||||
slug = "kb_test_dify"
|
||||
kb.databases_meta[slug] = {
|
||||
"name": "dify-kb",
|
||||
"description": "test",
|
||||
"kb_type": "dify",
|
||||
@ -118,7 +118,7 @@ async def test_dify_kb_aquery_maps_records(monkeypatch, tmp_path):
|
||||
lambda **kwargs: _FakeAsyncClient(response_payload=payload, **kwargs),
|
||||
)
|
||||
|
||||
result = await kb.aquery("hello", db_id)
|
||||
result = await kb.aquery("hello", slug)
|
||||
assert len(result) == 1
|
||||
assert result[0]["content"] == "hello world"
|
||||
assert result[0]["score"] == 0.98
|
||||
@ -131,8 +131,8 @@ async def test_dify_kb_aquery_maps_records(monkeypatch, tmp_path):
|
||||
@pytest.mark.asyncio
|
||||
async def test_dify_kb_aquery_error_returns_empty(monkeypatch, tmp_path):
|
||||
kb = DifyKB(str(tmp_path))
|
||||
db_id = "kb_test_dify_error"
|
||||
kb.databases_meta[db_id] = {
|
||||
slug = "kb_test_dify_error"
|
||||
kb.databases_meta[slug] = {
|
||||
"name": "dify-kb",
|
||||
"description": "test",
|
||||
"kb_type": "dify",
|
||||
@ -149,5 +149,5 @@ async def test_dify_kb_aquery_error_returns_empty(monkeypatch, tmp_path):
|
||||
lambda **kwargs: _FakeAsyncClient(raises=RuntimeError("boom"), **kwargs),
|
||||
)
|
||||
|
||||
result = await kb.aquery("hello", db_id)
|
||||
result = await kb.aquery("hello", slug)
|
||||
assert result == []
|
||||
|
||||
@ -37,11 +37,11 @@ class FakeCollection:
|
||||
def make_kb(collection: FakeCollection) -> MilvusKB:
|
||||
kb = MilvusKB.__new__(MilvusKB)
|
||||
kb.databases_meta = {"db": {"embedding_model_spec": "test-provider:test-embedding"}}
|
||||
kb.files_meta = {"file-1": {"filename": "demo.md", "database_id": "db"}}
|
||||
kb._get_query_params = lambda db_id: {}
|
||||
kb.files_meta = {"file-1": {"filename": "demo.md", "kb_id": "db"}}
|
||||
kb._get_query_params = lambda kb_id: {}
|
||||
kb._get_embedding_function = lambda embedding_model_spec, **kwargs: lambda texts: [[0.1, 0.2] for _ in texts]
|
||||
|
||||
async def get_collection(db_id: str):
|
||||
async def get_collection(kb_id: str):
|
||||
return collection
|
||||
|
||||
kb._get_milvus_collection = get_collection
|
||||
|
||||
@ -14,17 +14,17 @@ async def test_import_workspace_files_uploads_workspace_file_to_minio(tmp_path,
|
||||
source = tmp_path / "note.md"
|
||||
source.write_text("# workspace note\n", encoding="utf-8")
|
||||
|
||||
async def fake_ensure_database_not_dify(db_id: str, operation: str) -> None:
|
||||
assert db_id == "db_1"
|
||||
async def fake_ensure_database_supports_documents(slug: str, operation: str) -> None:
|
||||
assert slug == "db_1"
|
||||
assert "文档添加" in operation
|
||||
|
||||
async def fake_file_existed_in_db(db_id: str, content_hash: str) -> bool:
|
||||
assert db_id == "db_1"
|
||||
async def fake_file_existed_in_db(slug: str, content_hash: str) -> bool:
|
||||
assert slug == "db_1"
|
||||
assert content_hash
|
||||
return False
|
||||
|
||||
async def fake_get_same_name_files(db_id: str, filename: str) -> list:
|
||||
assert db_id == "db_1"
|
||||
async def fake_get_same_name_files(slug: str, filename: str) -> list:
|
||||
assert slug == "db_1"
|
||||
assert filename == "note.md"
|
||||
return []
|
||||
|
||||
@ -34,14 +34,18 @@ async def test_import_workspace_files_uploads_workspace_file_to_minio(tmp_path,
|
||||
assert data == b"# workspace note\n"
|
||||
return f"http://minio/{bucket_name}/{file_name}"
|
||||
|
||||
monkeypatch.setattr(knowledge_router, "_ensure_database_not_dify", fake_ensure_database_not_dify)
|
||||
monkeypatch.setattr(
|
||||
knowledge_router,
|
||||
"_ensure_database_supports_documents",
|
||||
fake_ensure_database_supports_documents,
|
||||
)
|
||||
monkeypatch.setattr(knowledge_router, "resolve_workspace_file_path", lambda **_kwargs: source)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "file_existed_in_db", fake_file_existed_in_db)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "get_same_name_files", fake_get_same_name_files)
|
||||
monkeypatch.setattr(knowledge_router, "aupload_file_to_minio", fake_upload)
|
||||
|
||||
result = await knowledge_router.import_workspace_files(
|
||||
knowledge_router.WorkspaceImportRequest(db_id="db_1", paths=["/note.md"]),
|
||||
knowledge_router.WorkspaceImportRequest(slug="db_1", paths=["/note.md"]),
|
||||
current_user=SimpleNamespace(id="user_1"),
|
||||
)
|
||||
|
||||
@ -58,18 +62,22 @@ async def test_import_workspace_files_uploads_workspace_file_to_minio(tmp_path,
|
||||
|
||||
|
||||
async def test_import_workspace_files_rejects_directory(tmp_path, monkeypatch):
|
||||
async def fake_ensure_database_not_dify(db_id: str, operation: str) -> None:
|
||||
async def fake_ensure_database_supports_documents(slug: str, operation: str) -> None:
|
||||
return None
|
||||
|
||||
def fake_resolve_workspace_file_path(**_kwargs):
|
||||
raise HTTPException(status_code=400, detail="当前路径不是文件: /folder")
|
||||
|
||||
monkeypatch.setattr(knowledge_router, "_ensure_database_not_dify", fake_ensure_database_not_dify)
|
||||
monkeypatch.setattr(
|
||||
knowledge_router,
|
||||
"_ensure_database_supports_documents",
|
||||
fake_ensure_database_supports_documents,
|
||||
)
|
||||
monkeypatch.setattr(knowledge_router, "resolve_workspace_file_path", fake_resolve_workspace_file_path)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await knowledge_router.import_workspace_files(
|
||||
knowledge_router.WorkspaceImportRequest(db_id="db_1", paths=["/folder"]),
|
||||
knowledge_router.WorkspaceImportRequest(slug="db_1", paths=["/folder"]),
|
||||
current_user=SimpleNamespace(id="user_1"),
|
||||
)
|
||||
|
||||
|
||||
@ -48,6 +48,7 @@ def test_list_subagents_returns_data(monkeypatch):
|
||||
async def fake_get_all_subagents(_db):
|
||||
return [
|
||||
{
|
||||
"slug": "research-agent",
|
||||
"name": "research-agent",
|
||||
"description": "Test research agent",
|
||||
"system_prompt": "You are a researcher",
|
||||
@ -78,6 +79,7 @@ def test_get_single_subagent(monkeypatch):
|
||||
async def fake_get_subagent(name, db=None):
|
||||
if name == "research-agent":
|
||||
return {
|
||||
"slug": "research-agent",
|
||||
"name": "research-agent",
|
||||
"description": "Test research agent",
|
||||
"system_prompt": "You are a researcher",
|
||||
@ -122,6 +124,7 @@ def test_create_subagent(monkeypatch):
|
||||
captured["data"] = data
|
||||
captured["created_by"] = created_by
|
||||
return {
|
||||
"slug": data["slug"],
|
||||
"name": data["name"],
|
||||
"description": data["description"],
|
||||
"system_prompt": data["system_prompt"],
|
||||
@ -142,7 +145,8 @@ def test_create_subagent(monkeypatch):
|
||||
resp = client.post(
|
||||
"/api/system/subagents",
|
||||
json={
|
||||
"name": "my-agent",
|
||||
"slug": "my-agent",
|
||||
"name": "My Agent",
|
||||
"description": "My custom agent",
|
||||
"system_prompt": "You are a helpful assistant",
|
||||
"tools": ["tool_a", "tool_b"],
|
||||
@ -152,7 +156,8 @@ def test_create_subagent(monkeypatch):
|
||||
assert resp.status_code == 200, resp.text
|
||||
payload = resp.json()
|
||||
assert payload["success"] is True
|
||||
assert captured["data"]["name"] == "my-agent"
|
||||
assert captured["data"]["slug"] == "my-agent"
|
||||
assert captured["data"]["name"] == "My Agent"
|
||||
assert captured["created_by"] == "admin"
|
||||
|
||||
|
||||
@ -161,7 +166,10 @@ def test_create_subagent_duplicate_returns_409(monkeypatch):
|
||||
raise IntegrityError(
|
||||
"duplicate",
|
||||
{},
|
||||
Exception('duplicate key value violates unique constraint "subagents_pkey"'),
|
||||
Exception(
|
||||
'duplicate key value violates unique constraint "subagents_slug_key" '
|
||||
"Detail: Key (slug)=(my-agent) already exists."
|
||||
),
|
||||
)
|
||||
|
||||
monkeypatch.setattr("server.routers.subagent_router.service.create_subagent", fake_create_subagent)
|
||||
@ -171,7 +179,8 @@ def test_create_subagent_duplicate_returns_409(monkeypatch):
|
||||
resp = client.post(
|
||||
"/api/system/subagents",
|
||||
json={
|
||||
"name": "my-agent",
|
||||
"slug": "my-agent",
|
||||
"name": "My Agent",
|
||||
"description": "My custom agent",
|
||||
"system_prompt": "You are a helpful assistant",
|
||||
"tools": [],
|
||||
@ -189,6 +198,7 @@ def test_update_subagent(monkeypatch):
|
||||
captured["data"] = data
|
||||
captured["updated_by"] = updated_by
|
||||
return {
|
||||
"slug": name,
|
||||
"name": name,
|
||||
"description": data.get("description", "Updated description"),
|
||||
"system_prompt": data.get("system_prompt", "Updated prompt"),
|
||||
@ -271,6 +281,7 @@ def test_update_subagent_status(monkeypatch):
|
||||
captured["enabled"] = enabled
|
||||
captured["updated_by"] = updated_by
|
||||
return {
|
||||
"slug": name,
|
||||
"name": name,
|
||||
"description": "Test",
|
||||
"system_prompt": "Prompt",
|
||||
@ -322,6 +333,7 @@ class TestSubAgentRepository:
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalars.return_value.all.return_value = [
|
||||
SubAgent(
|
||||
slug="test-agent",
|
||||
name="test-agent",
|
||||
description="Test agent",
|
||||
system_prompt="You are a test",
|
||||
@ -344,12 +356,13 @@ class TestSubAgentRepository:
|
||||
mock_db.execute.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_by_name_found(self):
|
||||
async def test_get_by_slug_found(self):
|
||||
from yuxi.repositories.subagent_repository import SubAgentRepository
|
||||
|
||||
mock_db = AsyncMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none.return_value = SubAgent(
|
||||
slug="test-agent",
|
||||
name="test-agent",
|
||||
description="Test agent",
|
||||
system_prompt="You are a test",
|
||||
@ -364,13 +377,13 @@ class TestSubAgentRepository:
|
||||
mock_db.execute.return_value = mock_result
|
||||
|
||||
repo = SubAgentRepository(mock_db)
|
||||
result = await repo.get_by_name("test-agent")
|
||||
result = await repo.get_by_slug("test-agent")
|
||||
|
||||
assert result is not None
|
||||
assert result.name == "test-agent"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_by_name_not_found(self):
|
||||
async def test_get_by_slug_not_found(self):
|
||||
from yuxi.repositories.subagent_repository import SubAgentRepository
|
||||
|
||||
mock_db = AsyncMock()
|
||||
@ -379,7 +392,7 @@ class TestSubAgentRepository:
|
||||
mock_db.execute.return_value = mock_result
|
||||
|
||||
repo = SubAgentRepository(mock_db)
|
||||
result = await repo.get_by_name("nonexistent")
|
||||
result = await repo.get_by_slug("nonexistent")
|
||||
|
||||
assert result is None
|
||||
|
||||
@ -390,6 +403,7 @@ class TestSubAgentRepository:
|
||||
mock_db = AsyncMock()
|
||||
repo = SubAgentRepository(mock_db)
|
||||
item = SubAgent(
|
||||
slug="test-agent",
|
||||
name="test-agent",
|
||||
description="Test agent",
|
||||
system_prompt="You are a test",
|
||||
@ -404,6 +418,7 @@ class TestSubAgentRepository:
|
||||
|
||||
await repo.update(
|
||||
item,
|
||||
name=None,
|
||||
description=None,
|
||||
system_prompt=None,
|
||||
tools=None,
|
||||
@ -431,7 +446,7 @@ class TestSubAgentService:
|
||||
def __init__(self, session):
|
||||
pass
|
||||
|
||||
async def get_by_name(self, name):
|
||||
async def get_by_slug(self, name):
|
||||
return None
|
||||
|
||||
async def create(self, **kwargs):
|
||||
@ -458,15 +473,16 @@ class TestSubAgentService:
|
||||
await service_module.init_builtin_subagents()
|
||||
|
||||
assert len(created_agents) == 2
|
||||
agent_names = [a["name"] for a in created_agents]
|
||||
assert "research-agent" in agent_names
|
||||
assert "critique-agent" in agent_names
|
||||
agent_slugs = [a["slug"] for a in created_agents]
|
||||
assert "research-agent" in agent_slugs
|
||||
assert "critique-agent" in agent_slugs
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_subagent_specs_returns_list(self, monkeypatch):
|
||||
from yuxi.services import subagent_service as service_module
|
||||
|
||||
mock_spec = {
|
||||
"slug": "test-agent",
|
||||
"name": "test-agent",
|
||||
"description": "Test",
|
||||
"system_prompt": "You are a test",
|
||||
@ -501,6 +517,7 @@ class TestSubAgentService:
|
||||
|
||||
service_module._subagent_specs_cache = [
|
||||
{
|
||||
"slug": "test-agent",
|
||||
"name": "test-agent",
|
||||
"description": "Test",
|
||||
"system_prompt": "You are a test",
|
||||
@ -525,6 +542,7 @@ class TestSubAgentModel:
|
||||
def test_to_dict(self):
|
||||
now = utc_now_naive()
|
||||
agent = SubAgent(
|
||||
slug="test-agent",
|
||||
name="test-agent",
|
||||
description="Test agent",
|
||||
system_prompt="You are a test",
|
||||
@ -549,6 +567,7 @@ class TestSubAgentModel:
|
||||
|
||||
def test_to_subagent_spec(self):
|
||||
agent = SubAgent(
|
||||
slug="test-agent",
|
||||
name="test-agent",
|
||||
description="Test agent",
|
||||
system_prompt="You are a test",
|
||||
@ -571,6 +590,7 @@ class TestSubAgentModel:
|
||||
|
||||
def test_to_subagent_spec_no_model(self):
|
||||
agent = SubAgent(
|
||||
slug="test-agent",
|
||||
name="test-agent",
|
||||
description="Test agent",
|
||||
system_prompt="You are a test",
|
||||
@ -590,18 +610,20 @@ class TestSubAgentModel:
|
||||
|
||||
class TestDeepAgentSubagentSelection:
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_subagents_from_names_filters_and_resolves_tools(self, monkeypatch):
|
||||
async def test_get_subagents_from_slugs_filters_and_resolves_tools(self, monkeypatch):
|
||||
from yuxi.services import subagent_service as service_module
|
||||
|
||||
async def fake_get_specs(_db=None):
|
||||
return [
|
||||
{
|
||||
"slug": "research-agent",
|
||||
"name": "research-agent",
|
||||
"description": "r",
|
||||
"system_prompt": "s",
|
||||
"tools": ["tool_a"],
|
||||
},
|
||||
{
|
||||
"slug": "critique-agent",
|
||||
"name": "critique-agent",
|
||||
"description": "c",
|
||||
"system_prompt": "s",
|
||||
@ -615,18 +637,19 @@ class TestDeepAgentSubagentSelection:
|
||||
monkeypatch.setattr(service_module, "get_subagent_specs", fake_get_specs)
|
||||
monkeypatch.setattr("yuxi.agents.toolkits.get_all_tool_instances", lambda: [mock_tool])
|
||||
|
||||
resolved_specs = await service_module.get_subagents_from_names(["research-agent", "missing-agent"])
|
||||
resolved_specs = await service_module.get_subagents_from_slugs(["research-agent", "missing-agent"])
|
||||
|
||||
assert [item["name"] for item in resolved_specs] == ["research-agent"]
|
||||
assert resolved_specs[0]["tools"] == [mock_tool]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_subagents_from_names_none_selects_none(self, monkeypatch):
|
||||
async def test_get_subagents_from_slugs_none_selects_none(self, monkeypatch):
|
||||
from yuxi.services import subagent_service as service_module
|
||||
|
||||
async def fake_get_specs(_db=None):
|
||||
return [
|
||||
{
|
||||
"slug": "research-agent",
|
||||
"name": "research-agent",
|
||||
"description": "r",
|
||||
"system_prompt": "s",
|
||||
@ -636,4 +659,4 @@ class TestDeepAgentSubagentSelection:
|
||||
|
||||
monkeypatch.setattr(service_module, "get_subagent_specs", fake_get_specs)
|
||||
|
||||
assert await service_module.get_subagents_from_names(None) == []
|
||||
assert await service_module.get_subagents_from_slugs(None) == []
|
||||
|
||||
@ -9,6 +9,10 @@ from langchain.messages import AIMessageChunk, HumanMessage
|
||||
from yuxi.services import chat_service as svc
|
||||
|
||||
|
||||
async def _fake_normalize_agent_context_config(context, **_kwargs):
|
||||
return dict(context or {})
|
||||
|
||||
|
||||
class _FakeAgentConfigRepo:
|
||||
def __init__(self, _db):
|
||||
pass
|
||||
@ -75,6 +79,8 @@ async def test_stream_agent_chat_passes_langfuse_callbacks_and_persists_trace_in
|
||||
calls: dict[str, object] = {}
|
||||
|
||||
class FakeAgent:
|
||||
context_schema = None
|
||||
|
||||
async def stream_messages(self, messages, input_context=None, **kwargs):
|
||||
calls["stream_messages"] = messages
|
||||
calls["stream_input_context"] = input_context
|
||||
@ -111,6 +117,7 @@ async def test_stream_agent_chat_passes_langfuse_callbacks_and_persists_trace_in
|
||||
|
||||
monkeypatch.setattr(svc.agent_manager, "get_agent", lambda agent_id: FakeAgent())
|
||||
monkeypatch.setattr(svc, "get_agent_config_by_id", fake_get_agent_config_by_id)
|
||||
monkeypatch.setattr(svc, "normalize_agent_context_config", _fake_normalize_agent_context_config)
|
||||
monkeypatch.setattr(svc, "ConversationRepository", _FakeConvRepo)
|
||||
monkeypatch.setattr(svc, "AgentConfigRepository", _FakeAgentConfigRepo)
|
||||
monkeypatch.setattr(svc, "save_messages_from_langgraph_state", fake_save_messages_from_langgraph_state)
|
||||
@ -144,7 +151,7 @@ async def test_stream_agent_chat_passes_langfuse_callbacks_and_persists_trace_in
|
||||
thread_id="thread-1",
|
||||
meta={"request_id": "req-1"},
|
||||
image_content=None,
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", department_id="dept-1"),
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", role="user", department_id="dept-1"),
|
||||
db=object(),
|
||||
):
|
||||
chunks.append(json.loads(chunk.decode("utf-8")))
|
||||
@ -171,6 +178,8 @@ async def test_stream_agent_chat_emits_realtime_agent_state_from_values(monkeypa
|
||||
return SimpleNamespace(values={"todos": [{"content": "done", "status": "completed"}]})
|
||||
|
||||
class FakeAgent:
|
||||
context_schema = None
|
||||
|
||||
async def stream_messages_with_state(self, messages, input_context=None, **kwargs):
|
||||
yield "values", {"messages": [], "todos": [{"content": "step 1", "status": "pending"}]}
|
||||
yield "values", {"messages": [], "todos": [{"content": "step 1", "status": "in_progress"}]}
|
||||
@ -202,6 +211,7 @@ async def test_stream_agent_chat_emits_realtime_agent_state_from_values(monkeypa
|
||||
|
||||
monkeypatch.setattr(svc.agent_manager, "get_agent", lambda agent_id: FakeAgent())
|
||||
monkeypatch.setattr(svc, "get_agent_config_by_id", fake_get_agent_config_by_id)
|
||||
monkeypatch.setattr(svc, "normalize_agent_context_config", _fake_normalize_agent_context_config)
|
||||
monkeypatch.setattr(svc, "ConversationRepository", _FakeConvRepo)
|
||||
monkeypatch.setattr(svc, "AgentConfigRepository", _FakeAgentConfigRepo)
|
||||
monkeypatch.setattr(svc, "save_messages_from_langgraph_state", fake_save_messages_from_langgraph_state)
|
||||
@ -223,7 +233,7 @@ async def test_stream_agent_chat_emits_realtime_agent_state_from_values(monkeypa
|
||||
thread_id="thread-1",
|
||||
meta={"request_id": "req-1"},
|
||||
image_content=None,
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", department_id="dept-1"),
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", role="user", department_id="dept-1"),
|
||||
db=object(),
|
||||
):
|
||||
chunks.append(json.loads(chunk.decode("utf-8")))
|
||||
|
||||
@ -12,6 +12,10 @@ def _empty_agents_prompt(_thread_id: str, _uid: str) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
async def _fake_normalize_agent_context_config(context, **_kwargs):
|
||||
return dict(context or {})
|
||||
|
||||
|
||||
class _FakeAgentConfigRepo:
|
||||
def __init__(self, _db):
|
||||
pass
|
||||
@ -83,6 +87,8 @@ async def test_agent_chat_uses_invoke_messages_and_persists_langgraph_state(monk
|
||||
return SimpleNamespace(values={"messages": [AIMessage(content="Hi from graph")], "todos": ["todo-1"]})
|
||||
|
||||
class FakeAgent:
|
||||
context_schema = None
|
||||
|
||||
async def invoke_messages(self, messages, input_context=None, **kwargs):
|
||||
calls["invoke_messages"] = messages
|
||||
calls["invoke_input_context"] = input_context
|
||||
@ -131,6 +137,7 @@ async def test_agent_chat_uses_invoke_messages_and_persists_langgraph_state(monk
|
||||
|
||||
monkeypatch.setattr(svc.agent_manager, "get_agent", lambda agent_id: FakeAgent())
|
||||
monkeypatch.setattr(svc, "get_agent_config_by_id", fake_get_agent_config_by_id)
|
||||
monkeypatch.setattr(svc, "normalize_agent_context_config", _fake_normalize_agent_context_config)
|
||||
monkeypatch.setattr(svc, "ConversationRepository", _FakeConvRepo)
|
||||
monkeypatch.setattr(svc, "AgentConfigRepository", _FakeAgentConfigRepo)
|
||||
monkeypatch.setattr(svc, "save_messages_from_langgraph_state", fake_save_messages_from_langgraph_state)
|
||||
@ -142,7 +149,7 @@ async def test_agent_chat_uses_invoke_messages_and_persists_langgraph_state(monk
|
||||
thread_id="thread-1",
|
||||
meta={"request_id": "req-1"},
|
||||
image_content=None,
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", department_id="dept-1"),
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", role="user", department_id="dept-1"),
|
||||
db=object(),
|
||||
)
|
||||
|
||||
@ -185,6 +192,8 @@ async def test_agent_chat_sync_returns_finished_even_when_state_has_interrupt(mo
|
||||
)
|
||||
|
||||
class FakeAgent:
|
||||
context_schema = None
|
||||
|
||||
async def invoke_messages(self, messages, input_context=None, **kwargs):
|
||||
return {"messages": [messages[0], AIMessage(content="Need input later")]}
|
||||
|
||||
@ -214,6 +223,7 @@ async def test_agent_chat_sync_returns_finished_even_when_state_has_interrupt(mo
|
||||
|
||||
monkeypatch.setattr(svc.agent_manager, "get_agent", lambda agent_id: FakeAgent())
|
||||
monkeypatch.setattr(svc, "get_agent_config_by_id", fake_get_agent_config_by_id)
|
||||
monkeypatch.setattr(svc, "normalize_agent_context_config", _fake_normalize_agent_context_config)
|
||||
monkeypatch.setattr(svc, "ConversationRepository", _FakeConvRepo)
|
||||
monkeypatch.setattr(svc, "AgentConfigRepository", _FakeAgentConfigRepo)
|
||||
monkeypatch.setattr(svc, "save_messages_from_langgraph_state", fake_save_messages_from_langgraph_state)
|
||||
@ -225,7 +235,7 @@ async def test_agent_chat_sync_returns_finished_even_when_state_has_interrupt(mo
|
||||
thread_id="thread-2",
|
||||
meta={"request_id": "req-2"},
|
||||
image_content=None,
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", department_id="dept-1"),
|
||||
current_user=SimpleNamespace(id=1, uid="user-1", role="user", department_id="dept-1"),
|
||||
db=object(),
|
||||
)
|
||||
|
||||
|
||||
@ -44,18 +44,18 @@ def test_is_valid_skill_slug():
|
||||
assert svc.is_valid_skill_slug("") is False
|
||||
|
||||
|
||||
def test_sync_thread_visible_skills_none_keeps_no_skills(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
def test_sync_thread_readable_skills_none_keeps_no_skills(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(svc.sys_config, "save_dir", str(tmp_path))
|
||||
skills_root = tmp_path / "skills"
|
||||
(skills_root / "alpha").mkdir(parents=True, exist_ok=True)
|
||||
(skills_root / "alpha" / "SKILL.md").write_text("alpha", encoding="utf-8")
|
||||
|
||||
thread_root = svc.sync_thread_visible_skills("thread_1", None)
|
||||
thread_root = svc.sync_thread_readable_skills("thread_1", None)
|
||||
|
||||
assert list(thread_root.iterdir()) == []
|
||||
|
||||
|
||||
def test_sync_thread_visible_skills_only_keeps_selected(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
def test_sync_thread_readable_skills_only_keeps_selected(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(svc.sys_config, "save_dir", str(tmp_path))
|
||||
skills_root = tmp_path / "skills"
|
||||
(skills_root / "alpha").mkdir(parents=True, exist_ok=True)
|
||||
@ -63,7 +63,7 @@ def test_sync_thread_visible_skills_only_keeps_selected(tmp_path: Path, monkeypa
|
||||
(skills_root / "beta").mkdir(parents=True, exist_ok=True)
|
||||
(skills_root / "beta" / "SKILL.md").write_text("beta", encoding="utf-8")
|
||||
|
||||
thread_root = svc.sync_thread_visible_skills("thread_1", ["alpha", "missing", "alpha"])
|
||||
thread_root = svc.sync_thread_readable_skills("thread_1", ["alpha", "missing", "alpha"])
|
||||
|
||||
assert thread_root == tmp_path / "threads" / "thread_1" / "skills"
|
||||
assert sorted(path.name for path in thread_root.iterdir()) == ["alpha"]
|
||||
@ -71,7 +71,7 @@ def test_sync_thread_visible_skills_only_keeps_selected(tmp_path: Path, monkeypa
|
||||
assert not (thread_root / "alpha").is_symlink()
|
||||
assert (thread_root / "alpha" / "SKILL.md").read_text(encoding="utf-8") == "alpha"
|
||||
|
||||
svc.sync_thread_visible_skills("thread_1", ["beta"])
|
||||
svc.sync_thread_readable_skills("thread_1", ["beta"])
|
||||
assert sorted(path.name for path in thread_root.iterdir()) == ["beta"]
|
||||
assert (thread_root / "beta" / "SKILL.md").read_text(encoding="utf-8") == "beta"
|
||||
|
||||
@ -81,17 +81,17 @@ async def test_get_skill_dependency_options(monkeypatch: pytest.MonkeyPatch):
|
||||
# Mock get_tool_metadata to return tool list
|
||||
def fake_get_tool_metadata(category=None):
|
||||
return [
|
||||
{"id": "calculator", "name": "Calculator"},
|
||||
{"id": "search", "name": "Search"},
|
||||
{"slug": "calculator", "name": "Calculator"},
|
||||
{"slug": "search", "name": "Search"},
|
||||
]
|
||||
|
||||
monkeypatch.setattr(tool_service, "get_tool_metadata", fake_get_tool_metadata)
|
||||
|
||||
async def fake_get_enabled_mcp_server_names(db=None):
|
||||
async def fake_get_enabled_mcp_server_slugs(db=None):
|
||||
del db
|
||||
return ["mcp-a", "mcp-b"]
|
||||
|
||||
monkeypatch.setattr(svc, "get_enabled_mcp_server_names", fake_get_enabled_mcp_server_names)
|
||||
monkeypatch.setattr(svc, "get_enabled_mcp_server_slugs", fake_get_enabled_mcp_server_slugs)
|
||||
|
||||
async def fake_list_skill_slugs(_db):
|
||||
return ["alpha", "beta"]
|
||||
@ -99,7 +99,7 @@ async def test_get_skill_dependency_options(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(svc, "list_skill_slugs", fake_list_skill_slugs)
|
||||
|
||||
result = await svc.get_skill_dependency_options(None)
|
||||
assert result["tools"] == [{"id": "calculator", "name": "Calculator"}, {"id": "search", "name": "Search"}]
|
||||
assert result["tools"] == [{"slug": "calculator", "name": "Calculator"}, {"slug": "search", "name": "Search"}]
|
||||
assert result["mcps"] == ["mcp-a", "mcp-b"]
|
||||
assert result["skills"] == ["alpha", "beta"]
|
||||
|
||||
@ -317,15 +317,15 @@ async def test_update_skill_dependencies(monkeypatch: pytest.MonkeyPatch):
|
||||
|
||||
# Mock get_tool_metadata to return tool list
|
||||
def fake_get_tool_metadata(category=None):
|
||||
return [{"id": "calculator", "name": "Calculator"}]
|
||||
return [{"slug": "calculator", "name": "Calculator"}]
|
||||
|
||||
monkeypatch.setattr(tool_service, "get_tool_metadata", fake_get_tool_metadata)
|
||||
|
||||
async def fake_get_enabled_mcp_server_names(db=None):
|
||||
async def fake_get_enabled_mcp_server_slugs(db=None):
|
||||
del db
|
||||
return ["mcp-a"]
|
||||
|
||||
monkeypatch.setattr(svc, "get_enabled_mcp_server_names", fake_get_enabled_mcp_server_names)
|
||||
monkeypatch.setattr(svc, "get_enabled_mcp_server_slugs", fake_get_enabled_mcp_server_slugs)
|
||||
|
||||
async def fake_get_skill_or_raise(_db, slug: str):
|
||||
assert slug == "alpha"
|
||||
|
||||
@ -34,7 +34,7 @@ def test_get_tool_metadata_includes_config_guide(monkeypatch):
|
||||
|
||||
assert result == [
|
||||
{
|
||||
"id": "demo_tool",
|
||||
"slug": "demo_tool",
|
||||
"name": "演示工具",
|
||||
"description": "demo description",
|
||||
"metadata": {},
|
||||
|
||||
@ -85,7 +85,7 @@ def _patch_retrievers(monkeypatch, *, kb_type: str = "milvus", retriever=None):
|
||||
|
||||
async def _fake_visible_kbs(runtime):
|
||||
del runtime
|
||||
return [{"db_id": "db-1", "name": "FAQ"}]
|
||||
return [{"kb_id": "db-1", "name": "FAQ"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@ -108,11 +108,11 @@ async def test_query_kb_returns_search_schema_without_sandbox_paths(monkeypatch)
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime)
|
||||
result = await _run_query_kb(kb_id="db-1", query_text="auth", runtime=runtime)
|
||||
|
||||
assert result["resource_id"] == "db-1"
|
||||
assert result["kb_id"] == "db-1"
|
||||
assert result["results"][0]["id"] == "file-1:1"
|
||||
assert result["results"][0]["resource_id"] == "db-1"
|
||||
assert result["results"][0]["kb_id"] == "db-1"
|
||||
assert result["results"][0]["file_id"] == "file-1"
|
||||
assert result["results"][0]["content"] == "auth guide"
|
||||
assert result["results"][0]["metadata"]["source"] == "auth-guide.pdf"
|
||||
@ -140,14 +140,14 @@ async def test_query_kb_allows_dify_knowledge_base(monkeypatch) -> None:
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime)
|
||||
result = await _run_query_kb(kb_id="db-1", query_text="auth", runtime=runtime)
|
||||
|
||||
assert result == {
|
||||
"resource_id": "db-1",
|
||||
"kb_id": "db-1",
|
||||
"results": [
|
||||
{
|
||||
"id": "dify-segment-1",
|
||||
"resource_id": "db-1",
|
||||
"kb_id": "db-1",
|
||||
"file_id": "dify-doc-1",
|
||||
"content": "auth guide",
|
||||
"metadata": {
|
||||
@ -171,7 +171,7 @@ async def test_query_kb_returns_plain_result_without_path_injection(monkeypatch)
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime)
|
||||
result = await _run_query_kb(kb_id="db-1", query_text="auth", runtime=runtime)
|
||||
|
||||
assert result == "Milvus context"
|
||||
|
||||
@ -193,11 +193,11 @@ async def test_query_kb_maps_full_doc_id_and_chunk_metadata(monkeypatch) -> None
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_query_kb(resource_id="db-1", query_text="auth", runtime=runtime)
|
||||
result = await _run_query_kb(kb_id="db-1", query_text="auth", runtime=runtime)
|
||||
|
||||
assert result["results"][0] == {
|
||||
"id": "chunk-1",
|
||||
"resource_id": "db-1",
|
||||
"kb_id": "db-1",
|
||||
"file_id": "file-1",
|
||||
"content": "auth guide",
|
||||
"metadata": {"chunk_index": 3},
|
||||
@ -210,7 +210,7 @@ async def test_find_kb_document_returns_context_windows(monkeypatch) -> None:
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
async def _fake_find_file_content(
|
||||
db_id: str,
|
||||
kb_id: str,
|
||||
file_id: str,
|
||||
patterns: list[str],
|
||||
*,
|
||||
@ -219,7 +219,7 @@ async def test_find_kb_document_returns_context_windows(monkeypatch) -> None:
|
||||
max_windows: int = 5,
|
||||
window_size: int = 80,
|
||||
):
|
||||
assert db_id == "db-1"
|
||||
assert kb_id == "db-1"
|
||||
assert file_id == "file-1"
|
||||
assert patterns == ["token"]
|
||||
assert use_regex is False
|
||||
@ -244,14 +244,14 @@ async def test_find_kb_document_returns_context_windows(monkeypatch) -> None:
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_find_kb_document(
|
||||
resource_id="db-1",
|
||||
kb_id="db-1",
|
||||
file_id="file-1",
|
||||
patterns=["token"],
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"resource_id": "db-1",
|
||||
"kb_id": "db-1",
|
||||
"file_id": "file-1",
|
||||
"semantic": False,
|
||||
"match_mode": "keyword",
|
||||
@ -274,7 +274,7 @@ async def test_find_kb_document_rejects_dify(monkeypatch) -> None:
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_find_kb_document(
|
||||
resource_id="db-1",
|
||||
kb_id="db-1",
|
||||
file_id="file-1",
|
||||
patterns=["token"],
|
||||
runtime=runtime,
|
||||
@ -290,17 +290,17 @@ async def test_open_kb_document_reads_markdown_content_by_default_window(monkeyp
|
||||
_patch_retrievers(monkeypatch)
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 1800):
|
||||
assert db_id == "db-1"
|
||||
async def _fake_open_file_content(kb_id: str, file_id: str, offset: int = 0, limit: int = 1800):
|
||||
assert kb_id == "db-1"
|
||||
assert file_id == "file-1"
|
||||
return _build_test_window("\n".join(lines), offset=offset, limit=limit)
|
||||
|
||||
monkeypatch.setattr(tools.knowledge_base, "open_file_content", _fake_open_file_content)
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_open_kb_document(resource_id="db-1", file_id="file-1", runtime=runtime)
|
||||
result = await _run_open_kb_document(kb_id="db-1", file_id="file-1", runtime=runtime)
|
||||
|
||||
assert result["resource_id"] == "db-1"
|
||||
assert result["kb_id"] == "db-1"
|
||||
assert result["file_id"] == "file-1"
|
||||
assert result["start_line"] == 1
|
||||
assert result["end_line"] == 1800
|
||||
@ -320,8 +320,8 @@ async def test_open_kb_document_prefers_line_over_offset(monkeypatch) -> None:
|
||||
_patch_retrievers(monkeypatch)
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 1800):
|
||||
assert db_id == "db-1"
|
||||
async def _fake_open_file_content(kb_id: str, file_id: str, offset: int = 0, limit: int = 1800):
|
||||
assert kb_id == "db-1"
|
||||
assert file_id == "file-1"
|
||||
return _build_test_window("\n".join(lines), offset=offset, limit=limit)
|
||||
|
||||
@ -329,7 +329,7 @@ async def test_open_kb_document_prefers_line_over_offset(monkeypatch) -> None:
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_open_kb_document(
|
||||
resource_id="db-1",
|
||||
kb_id="db-1",
|
||||
file_id="file-1",
|
||||
line=801,
|
||||
offset=0,
|
||||
@ -350,12 +350,12 @@ async def test_open_kb_document_prefers_line_over_offset(monkeypatch) -> None:
|
||||
async def test_open_kb_document_rejects_invisible_resource(monkeypatch) -> None:
|
||||
async def _fake_visible_kbs(runtime):
|
||||
del runtime
|
||||
return [{"db_id": "db-2", "name": "FAQ"}]
|
||||
return [{"kb_id": "db-2", "name": "FAQ"}]
|
||||
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_open_kb_document(resource_id="db-1", file_id="file-1", runtime=runtime)
|
||||
result = await _run_open_kb_document(kb_id="db-1", file_id="file-1", runtime=runtime)
|
||||
|
||||
assert "不存在或当前会话未启用" in result
|
||||
|
||||
@ -365,13 +365,13 @@ async def test_open_kb_document_requires_markdown_content(monkeypatch) -> None:
|
||||
_patch_retrievers(monkeypatch)
|
||||
monkeypatch.setattr(tools, "_resolve_visible_knowledge_bases_for_query", _fake_visible_kbs)
|
||||
|
||||
async def _fake_open_file_content(db_id: str, file_id: str, offset: int = 0, limit: int = 1800):
|
||||
del db_id, file_id, offset, limit
|
||||
async def _fake_open_file_content(kb_id: str, file_id: str, offset: int = 0, limit: int = 1800):
|
||||
del kb_id, file_id, offset, limit
|
||||
raise Exception("文件 file-1 没有解析后的 Markdown 内容")
|
||||
|
||||
monkeypatch.setattr(tools.knowledge_base, "open_file_content", _fake_open_file_content)
|
||||
|
||||
runtime = SimpleNamespace(context=SimpleNamespace())
|
||||
result = await _run_open_kb_document(resource_id="db-1", file_id="file-1", runtime=runtime)
|
||||
result = await _run_open_kb_document(kb_id="db-1", file_id="file-1", runtime=runtime)
|
||||
|
||||
assert "没有解析后的 Markdown 内容" in result
|
||||
|
||||
Loading…
Reference in New Issue
Block a user