test_sample_kb.py 8.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219
  1. """
  2. tests/test_sample_kb.py — 示例知识库种子 + 平台 KB 工具桥接(KB 调用链核查修复)。
  3. 覆盖:
  4. 1. ensure_sample_kb: 建库/登记/索引、幂等、无租户时延迟。
  5. 2. _make_platform_kb_tools: KBSearch 查到挂接库内容;KBList 列出库;
  6. 无挂接时返回 {}(不注入,保留 lambdagent 内置行为)。
  7. 3. compiler 工具合并: subAgents 的 call_* 与宿主注入的自定义工具共存
  8. (原来注入任何 tools 都会让 subAgents 整段跳过)。
  9. """
  10. from __future__ import annotations
  11. import json
  12. import os
  13. import pytest
  14. os.environ.setdefault("AGENTPAAS_DATABASE_URL", "sqlite:///:memory:")
  15. os.environ.setdefault("AGENTPAAS_TESTING", "1")
  16. @pytest.fixture()
  17. def fresh_db(monkeypatch):
  18. """独立内存 DB + 一个活跃租户。"""
  19. from agentpaas.db.models import Database, gen_id, now_utc
  20. import agentpaas.db.session as _session_mod
  21. prev = _session_mod._db
  22. _session_mod._db = Database("sqlite:///:memory:")
  23. db = _session_mod._db
  24. tid = gen_id("tn_")
  25. db.execute(
  26. "INSERT INTO tenants (id, name, plan, status, created_at) "
  27. "VALUES (?, 'test', 'free', 'active', ?)", (tid, now_utc()))
  28. db.commit()
  29. yield db, tid
  30. _session_mod._db = prev
  31. def test_sample_kb_created_and_indexed(fresh_db, tmp_path):
  32. from agentpaas.engine.sample_kb import ensure_sample_kb, SAMPLE_KB_NAME
  33. db, tid = fresh_db
  34. ensure_sample_kb(str(tmp_path))
  35. kb = db.fetchone("SELECT * FROM knowledge_bases WHERE name = ?", (SAMPLE_KB_NAME,))
  36. assert kb, "示例知识库应被创建"
  37. assert kb["tenant_id"] == tid
  38. # 文件已拷贝并登记
  39. files = db.fetchall("SELECT * FROM kb_files WHERE kb_id = ?", (kb["id"],))
  40. assert len(files) >= 4, files
  41. names = {f["file_name"] for f in files}
  42. assert "知识库使用指南.md" in names
  43. # pageindex 已生成且可检索
  44. idx = os.path.join(kb["root_dir"], "rag_page_index.json")
  45. assert os.path.isfile(idx)
  46. entries = json.loads(open(idx, encoding="utf-8").read())
  47. assert len(entries) > 10 # 按 ## 小节切分
  48. from agentpaas.api.v1.knowledge import _page_index_search
  49. hits = _page_index_search(kb["root_dir"], "顺序表 插入 删除", top_k=3)
  50. assert hits and "顺序" in (hits[0].get("text", "") + hits[0].get("summary", ""))
  51. def test_sample_kb_idempotent(fresh_db, tmp_path):
  52. from agentpaas.engine.sample_kb import ensure_sample_kb, SAMPLE_KB_NAME
  53. db, _ = fresh_db
  54. ensure_sample_kb(str(tmp_path))
  55. ensure_sample_kb(str(tmp_path)) # 再跑一次
  56. rows = db.fetchall("SELECT id FROM knowledge_bases WHERE name = ?", (SAMPLE_KB_NAME,))
  57. assert len(rows) == 1, "幂等:不应重复创建"
  58. def test_sample_kb_defers_without_tenant(tmp_path, monkeypatch):
  59. from agentpaas.db.models import Database
  60. import agentpaas.db.session as _session_mod
  61. from agentpaas.engine.sample_kb import ensure_sample_kb, SAMPLE_KB_NAME
  62. prev = _session_mod._db
  63. _session_mod._db = Database("sqlite:///:memory:") # 无租户
  64. try:
  65. ensure_sample_kb(str(tmp_path))
  66. kb = _session_mod._db.fetchone(
  67. "SELECT id FROM knowledge_bases WHERE name = ?", (SAMPLE_KB_NAME,))
  68. assert kb is None, "无租户时应延迟到下次启动"
  69. finally:
  70. _session_mod._db = prev
  71. # ── 平台 KB 工具桥接 ────────────────────────────────────────────────────────
  72. def test_platform_kb_tools_search_and_list(fresh_db, tmp_path):
  73. from agentpaas.engine.sample_kb import ensure_sample_kb, SAMPLE_KB_NAME
  74. from agentpaas.api.v1.agents import _make_platform_kb_tools
  75. db, tid = fresh_db
  76. ensure_sample_kb(str(tmp_path))
  77. kb = db.fetchone("SELECT * FROM knowledge_bases WHERE name = ?", (SAMPLE_KB_NAME,))
  78. tools = _make_platform_kb_tools(db, [kb["id"]], tid, search_mode="pageindex")
  79. assert set(tools) == {"KBSearch", "KBList"}
  80. # KBSearch: JSON 入参与裸字符串都支持
  81. out = tools["KBSearch"]('{"query": "循环队列 队满"}')
  82. assert "知识库检索结果" in out and "队" in out
  83. out2 = tools["KBSearch"]("栈 后进先出")
  84. assert "知识库检索结果" in out2
  85. # 查不到 → NO_MATCH 提示而非空串
  86. miss = tools["KBSearch"]("量子引力 弦论")
  87. assert "NO_MATCH" in miss or "知识库检索结果" in miss
  88. # KBList 列出挂接库与文件数
  89. lst = tools["KBList"]("")
  90. assert SAMPLE_KB_NAME in lst and "文件" in lst
  91. def test_platform_kb_tools_empty_when_no_kb(fresh_db):
  92. from agentpaas.api.v1.agents import _make_platform_kb_tools
  93. db, tid = fresh_db
  94. assert _make_platform_kb_tools(db, [], tid) == {}
  95. def test_platform_kb_tools_tenant_scoped(fresh_db, tmp_path):
  96. """别的租户挂了这个 kb_id 也查不到(audit #30 语义在工具层同样成立)。"""
  97. from agentpaas.engine.sample_kb import ensure_sample_kb, SAMPLE_KB_NAME
  98. from agentpaas.api.v1.agents import _make_platform_kb_tools
  99. db, tid = fresh_db
  100. ensure_sample_kb(str(tmp_path))
  101. kb = db.fetchone("SELECT * FROM knowledge_bases WHERE name = ?", (SAMPLE_KB_NAME,))
  102. tools = _make_platform_kb_tools(db, [kb["id"]], "tn_other_tenant", "pageindex")
  103. out = tools["KBSearch"]("顺序表")
  104. assert "NO_MATCH" in out # 跨租户被静默跳过 → 查无结果
  105. assert SAMPLE_KB_NAME not in tools["KBList"]("")
  106. # ── compiler 工具合并(subAgents 与注入工具共存)────────────────────────────
  107. def test_compiler_merges_injected_tools_with_subagents():
  108. from lambdagent.fromconfig.compiler import build_agent
  109. cfg = {
  110. "name": "orch", "type": "react",
  111. "model": {"name": "ollama/qwen2.5:7b"},
  112. "systemPrompt": "test",
  113. "react": {"maxSteps": 3},
  114. "subAgents": {
  115. "helper": {"inline": {"type": "simple",
  116. "model": {"name": "ollama/qwen2.5:7b"},
  117. "systemPrompt": "sub"},
  118. "tool": "call_helper"},
  119. },
  120. "mcp": {"localTools": ["KBSearch", "terminate"]},
  121. }
  122. sentinel = {"value": None}
  123. def fake_kb_search(x):
  124. sentinel["value"] = x
  125. return "[平台KB] hit"
  126. term = build_agent(cfg, {"tools": {"KBSearch": fake_kb_search}})
  127. assert term is not None
  128. # 注入的 KBSearch 生效(而非 lambdagent 内置实现)
  129. from lambdagent.fromconfig.compiler import _compile_tools # noqa: F401
  130. # 通过重新编译工具表来断言合并语义
  131. tools = {}
  132. from lambdagent.fromconfig import compiler as _c
  133. merged = {** _c._compile_sub_agents(cfg, cfg["subAgents"], {}),
  134. **{"KBSearch": fake_kb_search}}
  135. assert "call_helper" in merged and merged["KBSearch"] is fake_kb_search
  136. def test_page_index_search_chinese_query(fresh_db, tmp_path):
  137. """中文检索回归:原 \\W+ 分词把整句中文当一个 token → 全 0 分乱序。
  138. 现在 CJK bigram 切分,长中文问句应命中正确小节。"""
  139. from agentpaas.engine.sample_kb import ensure_sample_kb, SAMPLE_KB_NAME
  140. from agentpaas.api.v1.knowledge import _page_index_search
  141. db, _ = fresh_db
  142. ensure_sample_kb(str(tmp_path))
  143. kb = db.fetchone("SELECT * FROM knowledge_bases WHERE name = ?", (SAMPLE_KB_NAME,))
  144. hits = _page_index_search(kb["root_dir"], "循环队列怎么判断队满?讲义里怎么说的", 3)
  145. assert hits and hits[0]["score"] > 0, "中文长句必须有非零得分"
  146. assert "栈与队列" in hits[0]["source"], f"应命中第3章, got {hits[0]['source']}"
  147. # 英文 query 不受影响
  148. hits_en = _page_index_search(kb["root_dir"], "ADT stack push pop", 3)
  149. assert hits_en and hits_en[0]["score"] > 0
  150. def test_resolve_tenant_prefers_busiest(fresh_db, tmp_path):
  151. """多租户老库:示例库应归属活跃智能体最多的租户(浏览器会话所在地),
  152. 而非最老租户或 config.json 租户。"""
  153. from agentpaas.db.models import gen_id, now_utc
  154. from agentpaas.engine.sample_kb import ensure_sample_kb, SAMPLE_KB_NAME
  155. db, tid_old = fresh_db
  156. # 第二个租户,拥有 2 个活跃智能体(模拟日常使用租户)
  157. tid_busy = gen_id("tn_")
  158. now = now_utc()
  159. db.execute("INSERT INTO tenants (id, name, plan, status, created_at) "
  160. "VALUES (?, 'busy', 'free', 'active', ?)", (tid_busy, now))
  161. for i in range(2):
  162. db.execute(
  163. "INSERT INTO agents (id, tenant_id, name, current_version, status, "
  164. "created_at, updated_at) VALUES (?, ?, ?, 1, 'active', ?, ?)",
  165. (gen_id("ag_"), tid_busy, f"a{i}", now, now))
  166. db.commit()
  167. ensure_sample_kb(str(tmp_path))
  168. kb = db.fetchone("SELECT tenant_id FROM knowledge_bases WHERE name = ?",
  169. (SAMPLE_KB_NAME,))
  170. assert kb["tenant_id"] == tid_busy, "应归属智能体最多的租户"