test: 新增测试文件并更新现有测试适配命名与架构变更

- 新增: test_knowledge_base_backend.py, test_skills_middleware.py
- 更新: 所有现有测试适配 db_id→kb_id, get_subagents_from_names→slugs 等接口变更
This commit is contained in:
Wenjie Zhang 2026-05-21 19:32:36 +08:00
parent 99c590d003
commit 708ff29515
25 changed files with 689 additions and 265 deletions

View File

@ -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",

View File

@ -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)

View File

@ -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)

View File

@ -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,
)

View File

@ -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)

View File

@ -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}")

View File

@ -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 == []

View File

@ -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,

View 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 == []

View File

@ -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/"]

View File

@ -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": []}

View File

@ -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,

View File

@ -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,

View File

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

View File

@ -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:

View 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

View File

@ -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 == []

View File

@ -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

View File

@ -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"),
)

View File

@ -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) == []

View File

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

View File

@ -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(),
)

View File

@ -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"

View File

@ -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": {},

View File

@ -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