chore: T-P-1 工作区收敛——并行会话成果与历史未入库文件整理入库
- 入库历史遗漏源码/测试:router_system 9 模块(agent/executors/inference/knowledge/ memory/planner/skills/trace)、tests 11 个测试文件、config/knowledge 领域知识 - 入库根目录方案文档(v2/v3/可行性×2)、references 文献(arxiv 14-18/cnki_open/ 参考文献清单)、research 论文素材(routerarena/paper/中文文献 PDF) - 前端构建产物刷新(新 hash);webapp 误写文档删除 - gitignore 增补:deepseek-harness、research/_refs、.mimosa/.zcode、网关日志/pid、 临时调试脚本、tests/e2e/node_modules、AI代理功能开发/prefix - 基线确认:318 passed
This commit is contained in:
@@ -0,0 +1,63 @@
|
||||
"""黑板 / DAG(WorkingMemory / TaskGraph)单元测试。"""
|
||||
from router_system.memory import TaskGraph, TaskNode, WorkingMemory
|
||||
|
||||
|
||||
def _graph():
|
||||
g = TaskGraph()
|
||||
g.add_node(TaskNode(id="a", kind="analyze", domain="code", query="q"))
|
||||
g.add_node(TaskNode(id="b", kind="design", domain="code", query="q", deps=["a"]))
|
||||
g.add_node(TaskNode(id="c", kind="implement", domain="code", query="q", deps=["b"]))
|
||||
return g
|
||||
|
||||
|
||||
def test_topo_order_respects_deps():
|
||||
g = _graph()
|
||||
order = [n.id for n in g.topo_order()]
|
||||
assert order.index("a") < order.index("b") < order.index("c")
|
||||
|
||||
|
||||
def test_cycle_fallback_no_crash():
|
||||
g = TaskGraph()
|
||||
g.add_node(TaskNode(id="x", kind="solve", domain="m", query="q", deps=["y"]))
|
||||
g.add_node(TaskNode(id="y", kind="solve", domain="m", query="q", deps=["x"]))
|
||||
order = [n.id for n in g.topo_order()]
|
||||
assert set(order) == {"x", "y"}
|
||||
|
||||
|
||||
def test_failed_and_all_done():
|
||||
g = _graph()
|
||||
assert g.all_done() is False
|
||||
for n in g.nodes():
|
||||
n.status = "done"
|
||||
assert g.all_done() is True
|
||||
g.get("a").status = "failed"
|
||||
assert [n.id for n in g.failed()] == ["a"]
|
||||
|
||||
|
||||
def test_merge_order():
|
||||
mem = WorkingMemory()
|
||||
mem.write_section("s1", "第一部分")
|
||||
mem.write_section("s2", "第二部分")
|
||||
assert mem.merge(["s1", "s2"]) == "第一部分\n\n第二部分"
|
||||
# 指定顺序可调换
|
||||
assert mem.merge(["s2", "s1"]) == "第二部分\n\n第一部分"
|
||||
# 无顺序 → 写入顺序
|
||||
assert mem.merge() == "第一部分\n\n第二部分"
|
||||
|
||||
|
||||
def test_overwrite_fact_records_trace():
|
||||
mem = WorkingMemory()
|
||||
mem.write_fact("k", "v1", rule_id="r1")
|
||||
mem.write_fact("k", "v2", rule_id="r2")
|
||||
assert mem.get_fact("k") == "v2"
|
||||
assert any(t.startswith("overwrite:k@r2") for t in mem.trace)
|
||||
|
||||
|
||||
def test_duplicate_node_id_rejected():
|
||||
g = TaskGraph()
|
||||
g.add_node(TaskNode(id="a", kind="solve", domain="m", query="q"))
|
||||
try:
|
||||
g.add_node(TaskNode(id="a", kind="solve", domain="m", query="q"))
|
||||
assert False, "重复节点 id 应报错"
|
||||
except ValueError:
|
||||
pass
|
||||
@@ -0,0 +1,90 @@
|
||||
"""两级路由(domain_group)测试:用户指定大领域 → 组内路由模型 → 组内执行。"""
|
||||
import pytest
|
||||
|
||||
from router_system.router import build_router
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def router():
|
||||
return build_router()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_group_tech(router):
|
||||
# 指定 tech 组:跳过 8 领域统一分类器,组内路由识别 code
|
||||
res = await router.route("用 Python 写一个快速排序函数", domain_group="tech")
|
||||
assert res.domain == "code"
|
||||
assert res.domain_group == "tech"
|
||||
assert any("group:tech@explicit" in s for s in res.route)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_group_professional(router):
|
||||
res = await router.route("加班费怎么计算", domain_group="professional")
|
||||
assert res.domain == "legal"
|
||||
assert res.domain_group == "professional"
|
||||
assert res.subdomain == "labor"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_group_restricts_domain(router):
|
||||
# 在 tech 组内问法律问题:组内分类器只认识 code/math → 低置信走最后处理者
|
||||
res = await router.route("劳动合同违约怎么办", domain_group="tech")
|
||||
assert res.upgraded is True
|
||||
assert "direct_fallback" in res.route
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_group_detection(router):
|
||||
# 不指定 group:统一分类器识别 → 自动映射大领域组
|
||||
res = await router.route("基金定投的收益率怎么计算")
|
||||
assert res.domain == "finance"
|
||||
assert res.domain_group == "professional"
|
||||
assert any("group:professional@auto" in s for s in res.route)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_group_tech(router):
|
||||
res = await router.route("求 ∫ x^2 dx 从 0 到 1 的定积分")
|
||||
assert res.domain == "math"
|
||||
assert res.domain_group == "tech"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_group_raises(router):
|
||||
with pytest.raises(ValueError):
|
||||
await router.route("任意查询", domain_group="不存在的组")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_router_is_smaller(router):
|
||||
# 组内分类器只认识组内领域:体积/匹配范围 ≈ 全量 1/4
|
||||
assert len(router._group_classifiers["tech"].domains) == 2
|
||||
assert len(router._group_classifiers["professional"].domains) == 3
|
||||
assert len(router.classifier.domains) == 8
|
||||
assert len(router.domain_groups) == 4
|
||||
|
||||
|
||||
def test_health_shows_groups(router):
|
||||
h = router.health()
|
||||
assert "domain_groups" in h
|
||||
assert "tech" in h["domain_groups"]
|
||||
assert h["domain_groups"]["tech"] == ["code", "math"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_with_subdomain2(router):
|
||||
res = await router.route("房贷利率是 LPR 加多少", domain_group="professional")
|
||||
assert res.domain == "finance"
|
||||
assert res.subdomain == "loan"
|
||||
assert res.subdomain2 == "loan"
|
||||
assert "⚠" in res.response # 金融风险免责
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_cache_roundtrip(router):
|
||||
q = "加班费怎么计算"
|
||||
r1 = await router.route(q, domain_group="professional")
|
||||
r2 = await router.route(q, domain_group="professional") # 缓存命中
|
||||
assert r2.cache_hit is True
|
||||
assert r2.domain_group == r1.domain_group == "professional"
|
||||
@@ -0,0 +1,167 @@
|
||||
"""新领域(finance/life/education)与子领域(subdomain)测试。"""
|
||||
import pytest
|
||||
|
||||
from router_system.classifier import RuleClassifier
|
||||
from router_system.knowledge import KnowledgeBase
|
||||
from router_system.models import Classification
|
||||
from router_system.planner import Planner
|
||||
from router_system.router import build_router
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 分类
|
||||
# ---------------------------------------------------------------
|
||||
def test_finance_classification():
|
||||
clf = RuleClassifier()
|
||||
for q in ["基金定投的收益率怎么计算", "房贷利率是多少", "信用卡逾期了怎么办"]:
|
||||
r = clf.classify(q)
|
||||
assert r.domain == "finance", f"{q} -> {r.domain}"
|
||||
|
||||
|
||||
def test_life_classification():
|
||||
clf = RuleClassifier()
|
||||
for q in ["日本旅行攻略", "健身增肌计划", "家常菜谱推荐", "宠物驱虫怎么做"]:
|
||||
r = clf.classify(q)
|
||||
assert r.domain == "life", f"{q} -> {r.domain}"
|
||||
|
||||
|
||||
def test_education_classification():
|
||||
clf = RuleClassifier()
|
||||
for q in ["考研英语怎么备考", "高效学习方法", "面试技巧有哪些"]:
|
||||
r = clf.classify(q)
|
||||
assert r.domain == "education", f"{q} -> {r.domain}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 知识库加载
|
||||
# ---------------------------------------------------------------
|
||||
def test_new_domain_rules_loaded():
|
||||
kb = KnowledgeBase()
|
||||
for d in ("finance", "life", "education"):
|
||||
hits = kb.match("测试", domain=d) # 触发加载检查(domain 存在即可)
|
||||
assert kb.rules_count() >= 60, f"规则总数应 60+,当前 {kb.rules_count()}"
|
||||
# 新领域任务模板
|
||||
assert kb.task_template("finance-advice") is not None
|
||||
assert kb.task_template("life-guide") is not None
|
||||
assert kb.task_template("edu-guide") is not None
|
||||
# 新领域事实表
|
||||
assert len(kb.facts("finance")) >= 5
|
||||
assert len(kb.facts("life")) >= 5
|
||||
assert len(kb.facts("education")) >= 5
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 拆解
|
||||
# ---------------------------------------------------------------
|
||||
def _plan(query: str, domain: str, difficulty: str = "medium"):
|
||||
c = Classification(domain=domain, confidence=0.9, difficulty=difficulty)
|
||||
return Planner(KnowledgeBase()).plan(query, c)
|
||||
|
||||
|
||||
def test_finance_split_four():
|
||||
g = _plan("基金定投的收益率怎么计算", "finance", "medium")
|
||||
ids = [n.id for n in g.nodes()]
|
||||
assert ids == ["facts", "retrieve", "conclude", "disclaimer"]
|
||||
|
||||
|
||||
def test_life_split_three():
|
||||
g = _plan("日本旅行攻略", "life", "medium")
|
||||
kinds = [n.kind for n in g.nodes()]
|
||||
assert kinds == ["analyze", "advise", "conclude"]
|
||||
|
||||
|
||||
def test_edu_split_three():
|
||||
g = _plan("考研英语怎么备考", "education", "medium")
|
||||
kinds = [n.kind for n in g.nodes()]
|
||||
assert kinds == ["analyze", "design", "conclude"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# subdomain / subdomain2(三级子领域)端到端
|
||||
# ---------------------------------------------------------------
|
||||
@pytest.mark.asyncio
|
||||
async def test_subdomain_devops():
|
||||
r = build_router()
|
||||
res = await r.route("git 回滚代码怎么操作")
|
||||
assert res.subdomain == "devops"
|
||||
assert res.subdomain2 == "git"
|
||||
assert any("subdomain:devops" in s for s in res.route)
|
||||
assert any("subdomain2:git" in s for s in res.route)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subdomain_labor():
|
||||
r = build_router()
|
||||
res = await r.route("加班费怎么计算")
|
||||
assert res.domain == "legal"
|
||||
assert res.subdomain == "labor"
|
||||
assert res.subdomain2 == "labor"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subdomain_investing():
|
||||
r = build_router()
|
||||
res = await r.route("基金定投的收益率怎么计算")
|
||||
assert res.domain == "finance"
|
||||
assert res.subdomain == "investing"
|
||||
# finance 强制拆解 + 风险免责
|
||||
assert any("plan:multi[4]" in s for s in res.route)
|
||||
assert "⚠" in res.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_finance_facts_in_response():
|
||||
r = build_router()
|
||||
res = await r.route("信用卡逾期了怎么办")
|
||||
assert res.domain == "finance"
|
||||
assert "征信" in res.response or "逾期" in res.response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 三级子领域(subdomain2)
|
||||
# ---------------------------------------------------------------
|
||||
@pytest.mark.asyncio
|
||||
async def test_subdomain2_sorting():
|
||||
r = build_router()
|
||||
res = await r.route("用 Python 写一个快速排序函数")
|
||||
assert res.domain == "code"
|
||||
assert res.subdomain == "algorithm"
|
||||
assert res.subdomain2 == "sorting"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subdomain2_first_aid():
|
||||
r = build_router()
|
||||
res = await r.route("烫伤后怎么处理")
|
||||
assert res.domain == "medical"
|
||||
assert res.subdomain == "firstaid"
|
||||
assert res.subdomain2 == "first-aid"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subdomain2_loan():
|
||||
r = build_router()
|
||||
res = await r.route("房贷利率是 LPR 加多少")
|
||||
assert res.domain == "finance"
|
||||
assert res.subdomain == "loan"
|
||||
assert res.subdomain2 == "loan"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subdomain2_cache_roundtrip():
|
||||
r = build_router()
|
||||
q = "信用卡逾期了怎么办"
|
||||
r1 = await r.route(q)
|
||||
r2 = await r.route(q) # 缓存命中
|
||||
assert r2.cache_hit is True
|
||||
assert r2.subdomain == r1.subdomain
|
||||
assert r2.subdomain2 == r1.subdomain2
|
||||
assert r2.subdomain2 == "credit"
|
||||
|
||||
|
||||
def test_rule_subdomain2_from_map():
|
||||
kb = KnowledgeBase()
|
||||
r = kb.rule("code-sort")
|
||||
assert r.subdomain2 == "sorting"
|
||||
r2 = kb.rule("medical-firstaid")
|
||||
assert r2.subdomain2 == "first-aid"
|
||||
@@ -0,0 +1,68 @@
|
||||
"""前向链推理机(InferenceEngine)单元测试。"""
|
||||
from router_system.inference import InferenceEngine, render_template
|
||||
from router_system.knowledge import KnowledgeBase, Rule
|
||||
from router_system.memory import WorkingMemory
|
||||
|
||||
|
||||
def _engine():
|
||||
return InferenceEngine(KnowledgeBase())
|
||||
|
||||
|
||||
def test_forward_chain_fires_rules():
|
||||
eng = _engine()
|
||||
mem = WorkingMemory()
|
||||
eng.initialize("求解方程 x^2 - 5x + 6 = 0", "math", "medium", 0.9, mem)
|
||||
fired = eng.run("求解方程 x^2 - 5x + 6 = 0", "math", mem)
|
||||
assert fired, "数学查询应触发规则"
|
||||
# 优先级最高的 math-equation 最先触发
|
||||
assert fired[0] == "math-equation"
|
||||
# 轨迹包含规则记录
|
||||
assert any(t.startswith("rule:") for t in mem.trace)
|
||||
|
||||
|
||||
def test_no_match_terminates():
|
||||
eng = _engine()
|
||||
mem = WorkingMemory()
|
||||
eng.initialize("今天天气怎么样", "general", "medium", 0.5, mem)
|
||||
fired = eng.run("今天天气怎么样", "general", mem)
|
||||
assert fired == []
|
||||
|
||||
|
||||
def test_max_steps_limit():
|
||||
eng = _engine()
|
||||
mem = WorkingMemory()
|
||||
eng.initialize("求解方程 x^2 - 5x + 6 = 0", "math", "medium", 0.9, mem)
|
||||
fired = eng.run("求解方程 x^2 - 5x + 6 = 0", "math", mem, max_steps=1)
|
||||
assert len(fired) == 1
|
||||
|
||||
|
||||
def test_blackboard_facts_written():
|
||||
eng = _engine()
|
||||
mem = WorkingMemory()
|
||||
eng.initialize("求解方程 x^2 - 5x + 6 = 0", "math", "hard", 0.95, mem)
|
||||
assert mem.get_fact("domain") == "math"
|
||||
assert mem.get_fact("difficulty") == "hard"
|
||||
assert "init:" in mem.trace[0]
|
||||
|
||||
|
||||
def test_rule_output_renders_section():
|
||||
kb = KnowledgeBase()
|
||||
# 注入一条带 output 模板的规则
|
||||
kb._rules["test-rule"] = Rule(
|
||||
id="test-rule", domain="general", priority=10,
|
||||
patterns=["测试渲染"], output="领域={facts.domain} 查询={query}",
|
||||
)
|
||||
eng = InferenceEngine(kb)
|
||||
mem = WorkingMemory()
|
||||
eng.initialize("测试渲染一下", "general", "easy", 0.6, mem)
|
||||
fired = eng.run("测试渲染一下", "general", mem)
|
||||
assert "test-rule" in fired
|
||||
section = mem.section("test-rule")
|
||||
assert section is not None
|
||||
assert "领域=general" in section
|
||||
assert "查询=测试渲染一下" in section
|
||||
|
||||
|
||||
def test_render_missing_placeholder():
|
||||
out = render_template("你好 {query} {facts.缺失字段}", "世界", {})
|
||||
assert "你好 世界 [未提供]" == out
|
||||
@@ -0,0 +1,59 @@
|
||||
"""知识库(KnowledgeBase)单元测试。"""
|
||||
from router_system.knowledge import KnowledgeBase, Rule
|
||||
|
||||
|
||||
def test_builtin_rules_loaded():
|
||||
kb = KnowledgeBase()
|
||||
# 内置规则 + 文件规则合并后应至少有 10+ 条
|
||||
assert kb.rules_count() >= 10
|
||||
# 文件规则按 id 覆盖内置(code-sort 存在且可命中)
|
||||
r = kb.rule("code-sort")
|
||||
assert r is not None
|
||||
assert r.template == "code-implement"
|
||||
|
||||
|
||||
def test_match_priority_order():
|
||||
kb = KnowledgeBase()
|
||||
hits = kb.match("用 Python 写一个快速排序函数", domain="code")
|
||||
assert hits, "应命中 code 领域规则"
|
||||
# 优先级降序:排序规则(90) 应在通用实现规则(50) 之前
|
||||
assert hits[0].id == "code-sort"
|
||||
priorities = [h.priority for h in hits]
|
||||
assert priorities == sorted(priorities, reverse=True)
|
||||
|
||||
|
||||
def test_match_cross_domain():
|
||||
kb = KnowledgeBase()
|
||||
# 同一文本在不同领域过滤下命中不同规则
|
||||
code_hits = kb.match("帮我调试报错 TypeError", domain="code")
|
||||
assert code_hits and code_hits[0].id == "code-debug"
|
||||
# 领域过滤:不指定 domain 时全领域匹配
|
||||
all_hits = kb.match("帮我调试报错 TypeError")
|
||||
assert len(all_hits) >= len(code_hits)
|
||||
|
||||
|
||||
def test_task_template():
|
||||
kb = KnowledgeBase()
|
||||
tpl = kb.task_template("code-implement")
|
||||
assert tpl is not None
|
||||
steps = tpl["steps"]
|
||||
assert len(steps) == 4
|
||||
ids = [s["id"] for s in steps]
|
||||
assert ids == ["analyze", "design", "implement", "verify"]
|
||||
# 依赖关系存在
|
||||
assert steps[1]["deps"] == ["analyze"]
|
||||
|
||||
|
||||
def test_facts():
|
||||
kb = KnowledgeBase()
|
||||
legal_facts = kb.facts("legal")
|
||||
assert any("竞业" in f["statement"] for f in legal_facts)
|
||||
medical_facts = kb.facts("medical")
|
||||
assert any("高血压" in f["statement"] for f in medical_facts)
|
||||
|
||||
|
||||
def test_rule_matches_case_insensitive():
|
||||
r = Rule(id="t", domain="code", patterns=["Sort", "Python"])
|
||||
assert r.matches("用 python 写一个 sort 算法")
|
||||
assert r.matches("PYTHON 快速排序")
|
||||
assert not r.matches("java 写个函数")
|
||||
@@ -0,0 +1,130 @@
|
||||
"""知识库扩充验证测试(规则/模板/事实表/output 规则端到端)。"""
|
||||
import pytest
|
||||
|
||||
from router_system.classifier import RuleClassifier
|
||||
from router_system.knowledge import KnowledgeBase
|
||||
from router_system.models import Classification
|
||||
from router_system.planner import Planner
|
||||
|
||||
|
||||
def _kb():
|
||||
return KnowledgeBase()
|
||||
|
||||
|
||||
def _plan(query: str, domain: str, difficulty: str = "medium"):
|
||||
c = Classification(domain=domain, confidence=0.9, difficulty=difficulty)
|
||||
return Planner(_kb()).plan(query, c)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 规则 / 模板 / 事实表规模
|
||||
# ---------------------------------------------------------------
|
||||
def test_expanded_rules_count():
|
||||
kb = _kb()
|
||||
assert kb.rules_count() >= 40, f"知识库应扩充至 40+ 条规则,当前 {kb.rules_count()}"
|
||||
|
||||
|
||||
def test_new_task_templates_exist():
|
||||
kb = _kb()
|
||||
required = ["code-algorithm", "code-refactor", "code-explain", "code-test",
|
||||
"math-proof", "math-optimize", "medical-firstaid", "general-writing"]
|
||||
for tid in required:
|
||||
assert kb.task_template(tid) is not None, f"缺少任务模板 {tid}"
|
||||
|
||||
|
||||
def test_expanded_facts():
|
||||
kb = _kb()
|
||||
legal = kb.facts("legal")
|
||||
medical = kb.facts("medical")
|
||||
assert len(legal) >= 10
|
||||
assert len(medical) >= 10
|
||||
# 新事实条目存在
|
||||
legal_text = " ".join(f["statement"] for f in legal)
|
||||
assert "150%" in legal_text # 加班费
|
||||
assert "押金" in legal_text # 租房押金
|
||||
assert "无理由" in legal_text # 消费者退货
|
||||
assert "继承" in legal_text # 继承顺序
|
||||
medical_text = " ".join(f["statement"] for f in medical)
|
||||
assert "烫伤" in medical_text # 烫伤急救
|
||||
assert "抗生素" in medical_text # 抗生素
|
||||
assert "失眠" in medical_text # 失眠
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 新拆解行为
|
||||
# ---------------------------------------------------------------
|
||||
def test_algorithm_split_five():
|
||||
g = _plan("用动态规划实现背包问题", "code", "hard")
|
||||
ids = [n.id for n in g.nodes()]
|
||||
assert ids == ["analyze", "design", "implement", "complexity", "verify"]
|
||||
|
||||
|
||||
def test_refactor_split_three():
|
||||
g = _plan("这段代码重复太多,帮我重构", "code", "hard")
|
||||
ids = [n.id for n in g.nodes()]
|
||||
assert ids == ["analyze", "refactor", "verify"]
|
||||
|
||||
|
||||
def test_code_explain_split():
|
||||
g = _plan("帮我解释这段代码什么意思", "code", "medium")
|
||||
kinds = [n.kind for n in g.nodes()]
|
||||
assert kinds == ["analyze", "explain", "verify"]
|
||||
|
||||
|
||||
def test_testcase_split():
|
||||
g = _plan("给这个函数写单元测试", "code", "medium")
|
||||
kinds = [n.kind for n in g.nodes()]
|
||||
assert "testcase" in kinds
|
||||
|
||||
|
||||
def test_math_proof_and_optimize():
|
||||
g = _plan("证明勾股定理", "math", "hard")
|
||||
assert len(g) == 3
|
||||
g2 = _plan("求函数 f(x)=x^2 的最小值", "math", "medium")
|
||||
kinds = [n.kind for n in g2.nodes()]
|
||||
assert "optimize" in kinds
|
||||
|
||||
|
||||
def test_firstaid_split():
|
||||
g = _plan("烫伤后怎么处理", "medical", "medium")
|
||||
kinds = [n.kind for n in g.nodes()]
|
||||
assert kinds == ["analyze", "advise", "disclaimer"]
|
||||
|
||||
|
||||
def test_writing_split():
|
||||
g = _plan("写一封请假邮件", "general", "medium")
|
||||
kinds = [n.kind for n in g.nodes()]
|
||||
assert kinds == ["analyze", "draft", "polish"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# output 知识规则端到端
|
||||
# ---------------------------------------------------------------
|
||||
@pytest.mark.asyncio
|
||||
async def test_git_knowledge_in_response():
|
||||
from router_system.router import build_router
|
||||
r = build_router()
|
||||
res = await r.route("git 回滚代码怎么操作")
|
||||
assert "git 核心概念" in res.response
|
||||
assert "git revert" in res.response
|
||||
# 推理轨迹记录规则触发
|
||||
assert any("code-git-knowledge" in s for s in res.route)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_docker_knowledge_in_response():
|
||||
from router_system.router import build_router
|
||||
r = build_router()
|
||||
res = await r.route("docker 部署一个服务")
|
||||
assert "镜像 vs 容器" in res.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_medical_hypertension_full_flow():
|
||||
from router_system.router import build_router
|
||||
r = build_router()
|
||||
res = await r.route("高血压患者可以吃哪些降压药,副作用是什么")
|
||||
assert res.domain == "medical"
|
||||
assert any("plan:multi[3]" in s for s in res.route)
|
||||
assert "⚠" in res.response
|
||||
assert res.upgraded is False
|
||||
@@ -0,0 +1,68 @@
|
||||
"""NodeExecutor 接口抽象测试(T1:整体项目部分拆解·先行实现)。"""
|
||||
import pytest
|
||||
|
||||
from router_system.executors import (ModelNodeExecutor, RuleNodeExecutor,
|
||||
build_node_executor)
|
||||
from router_system.experts import build_expert_pool
|
||||
from router_system.knowledge import KnowledgeBase
|
||||
from router_system.memory import TaskNode, WorkingMemory
|
||||
from router_system.router import build_router
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rule_node_executor_deterministic():
|
||||
ex = RuleNodeExecutor(kb=KnowledgeBase())
|
||||
node = TaskNode(id="t1", kind="solve", domain="math", query="求解方程 x^2-5x+6=0")
|
||||
mem = WorkingMemory()
|
||||
r1 = await ex.execute(node, "math", "medium", mem)
|
||||
r2 = await ex.execute(node, "math", "medium", mem)
|
||||
assert r1.text == r2.text # 确定性
|
||||
assert "求解" in r1.text
|
||||
assert r1.model_used.startswith("rule:")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_node_executor_uses_expert_pool():
|
||||
experts = build_expert_pool({}, ["code", "math", "legal", "medical",
|
||||
"finance", "life", "education", "general"])
|
||||
ex = ModelNodeExecutor(experts)
|
||||
node = TaskNode(id="t1", kind="implement", domain="code", query="写一个函数")
|
||||
mem = WorkingMemory()
|
||||
resp = await ex.execute(node, "code", "medium", mem)
|
||||
assert resp.text # mock 专家输出
|
||||
assert "mock" in resp.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_node_executor_fallback_to_general():
|
||||
# 未知领域 → general 专家兜底
|
||||
experts = {"general": build_expert_pool({}, ["general"])["general"]}
|
||||
ex = ModelNodeExecutor(experts)
|
||||
node = TaskNode(id="t1", kind="explain", domain="unknown", query="测试")
|
||||
resp = await ex.execute(node, "unknown", "easy", WorkingMemory())
|
||||
assert resp.text
|
||||
|
||||
|
||||
def test_factory_rule():
|
||||
ex = build_node_executor("rule", kb=KnowledgeBase())
|
||||
assert isinstance(ex, RuleNodeExecutor)
|
||||
|
||||
|
||||
def test_factory_model():
|
||||
experts = build_expert_pool({}, ["code"])
|
||||
ex = build_node_executor("api", experts=experts)
|
||||
assert isinstance(ex, ModelNodeExecutor)
|
||||
|
||||
|
||||
def test_factory_invalid():
|
||||
with pytest.raises(ValueError):
|
||||
build_node_executor("bad-backend")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_uses_node_executor():
|
||||
r = build_router()
|
||||
assert r.node_executor.name == "rule"
|
||||
res = await r.route("求解方程 x^2 - 5x + 6 = 0")
|
||||
assert res.domain == "math"
|
||||
assert res.response
|
||||
@@ -0,0 +1,63 @@
|
||||
"""Planner(任务拆解)单元测试。"""
|
||||
from router_system.classifier import RuleClassifier
|
||||
from router_system.knowledge import KnowledgeBase
|
||||
from router_system.planner import Planner
|
||||
|
||||
|
||||
def _plan(query: str):
|
||||
clf = RuleClassifier()
|
||||
c = clf.classify(query)
|
||||
return Planner(KnowledgeBase()).plan(query, c), c
|
||||
|
||||
|
||||
def test_simple_task_no_split():
|
||||
graph, c = _plan("2 + 2 等于多少")
|
||||
# easy 难度 → 单节点不拆
|
||||
assert c.domain == "math"
|
||||
assert c.difficulty == "easy"
|
||||
assert len(graph) == 1
|
||||
n = graph.nodes()[0]
|
||||
assert n.id == "solve"
|
||||
assert n.kind == "solve"
|
||||
|
||||
|
||||
def test_code_task_split_four():
|
||||
graph, c = _plan("用 Python 写一个快速排序函数,并解释时间复杂度")
|
||||
assert c.domain == "code"
|
||||
assert len(graph) == 4
|
||||
ids = [n.id for n in graph.nodes()]
|
||||
assert ids == ["analyze", "design", "implement", "verify"]
|
||||
|
||||
|
||||
def test_math_task_split_three():
|
||||
graph, c = _plan("求解方程 x^2 - 5x + 6 = 0")
|
||||
assert c.domain == "math"
|
||||
assert len(graph) == 3
|
||||
ids = [n.id for n in graph.nodes()]
|
||||
assert ids == ["conditions", "solve", "verify"]
|
||||
|
||||
|
||||
def test_deps_wired():
|
||||
graph, _ = _plan("用 Python 写一个快速排序函数,并解释时间复杂度")
|
||||
nodes = {n.id: n for n in graph.nodes()}
|
||||
assert nodes["design"].deps == ["analyze"]
|
||||
assert nodes["implement"].deps == ["design"]
|
||||
assert nodes["verify"].deps == ["implement"]
|
||||
|
||||
|
||||
def test_no_template_falls_back_to_single():
|
||||
graph, c = _plan("你好呀")
|
||||
# general 无模板命中("你好"不在任何 pattern)→ 单节点
|
||||
assert c.domain == "general"
|
||||
assert len(graph) == 1
|
||||
n = graph.nodes()[0]
|
||||
assert n.kind == "explain"
|
||||
|
||||
|
||||
def test_explain_plan_trace():
|
||||
graph, _ = _plan("求解方程 x^2 - 5x + 6 = 0")
|
||||
trace = Planner(KnowledgeBase()).explain_plan(graph)
|
||||
assert trace and "plan:multi[3]" in trace[0]
|
||||
single_graph, _ = _plan("1 + 1 = ?")
|
||||
st = Planner(KnowledgeBase()).explain_plan(single_graph)
|
||||
assert "plan:single" in st[0]
|
||||
@@ -0,0 +1,112 @@
|
||||
"""专家系统内核端到端测试(L0 模式:零参数全链路)。"""
|
||||
import pytest
|
||||
|
||||
from router_system.knowledge import KnowledgeBase
|
||||
from router_system.router import Router, build_router
|
||||
from router_system.experts import build_expert_pool
|
||||
from router_system.judge import build_judge
|
||||
from router_system.fallback import build_fallback
|
||||
from router_system.cache import RouterCache
|
||||
from router_system.planner import Planner
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def es_router():
|
||||
return build_router()
|
||||
|
||||
|
||||
def _hybrid_router(expert_backend: str = "api"):
|
||||
"""构造 expert_backend=api 的路由(专家池为 mock,验证切换路径)。"""
|
||||
config = {"router": {"judge_fallback_threshold": 0.70},
|
||||
"execution": {"expert_backend": expert_backend}}
|
||||
kb = KnowledgeBase()
|
||||
classifier = build_router().classifier
|
||||
experts = build_expert_pool({}, ["code", "math", "legal", "medical", "general"])
|
||||
judge = build_judge({"type": "rule"}, 0.70, kb=kb)
|
||||
fallback = build_fallback({"type": "mock"})
|
||||
return Router(classifier, experts, judge, fallback, RouterCache(),
|
||||
None, config, kb=kb, planner=Planner(kb))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hybrid_backend_uses_expert_pool():
|
||||
r = _hybrid_router("api")
|
||||
res = await r.route("用 Python 写一个快速排序函数,并解释时间复杂度")
|
||||
# 子任务由专家池执行(mock 专家),而非规则执行器
|
||||
assert res.response
|
||||
assert any(s.startswith("analyze:") for s in res.route) or "plan:multi" in " ".join(res.route)
|
||||
# 专家池 mock 输出包含"mock"标识
|
||||
assert "mock" in res.response or "关键点" in res.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dag_route_code(es_router):
|
||||
r = await es_router.route("用 Python 写一个快速排序函数,并解释时间复杂度")
|
||||
assert r.domain == "code"
|
||||
assert r.response
|
||||
# 拆解轨迹:plan:multi[4] 与 4 个节点执行
|
||||
assert any("plan:multi[4]" in step for step in r.route)
|
||||
node_steps = [s for s in r.route if s.startswith("analyze:") or s.startswith("design:")
|
||||
or s.startswith("implement:") or s.startswith("verify:")]
|
||||
assert len(node_steps) == 4
|
||||
# 合并响应包含各步骤产出
|
||||
assert "【code 分析】" in r.response
|
||||
assert "【code 自检】" in r.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dag_route_math(es_router):
|
||||
r = await es_router.route("求解方程 x^2 - 5x + 6 = 0")
|
||||
assert r.domain == "math"
|
||||
assert any("plan:multi[3]" in step for step in r.route)
|
||||
assert "【math 求解】" in r.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_node_easy(es_router):
|
||||
# code/easy 高置信 → 单节点不拆
|
||||
r = await es_router.route("python 排序")
|
||||
assert r.upgraded is False
|
||||
assert any("plan:single" in step for step in r.route)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deterministic_output(es_router):
|
||||
q = "用 Python 写一个快速排序函数,并解释时间复杂度"
|
||||
r1 = await es_router.route(q)
|
||||
r2 = await es_router.route(q)
|
||||
# 第二次缓存命中(响应一致)
|
||||
assert r1.response == r2.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legal_retrieves_facts(es_router):
|
||||
r = await es_router.route("劳动合同到期不续签,公司需要支付经济补偿吗")
|
||||
assert r.domain == "legal"
|
||||
# retrieve 步骤命中知识库事实(竞业/经济补偿类)
|
||||
assert "知识检索" in r.response or "经济补偿" in r.response
|
||||
# 免责提示存在
|
||||
assert "⚠" in r.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_medical_disclaimer(es_router):
|
||||
r = await es_router.route("高血压患者日常饮食需要注意什么")
|
||||
assert r.domain == "medical"
|
||||
assert "⚠" in r.response
|
||||
assert any("warning:disclaimer" in s or s.startswith("warning:") for s in r.route)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_low_confidence_direct_fallback(es_router):
|
||||
r = await es_router.route("今天天气怎么样")
|
||||
assert r.upgraded is True
|
||||
assert "direct_fallback" in r.route
|
||||
|
||||
|
||||
def test_health_es(es_router):
|
||||
h = es_router.health()
|
||||
assert h["status"] == "ok"
|
||||
assert h["execution_mode"] == "rule"
|
||||
assert h["rules"] > 0
|
||||
assert "planner" in h
|
||||
@@ -0,0 +1,158 @@
|
||||
"""RouterArena 适配器最小测试。
|
||||
|
||||
覆盖:
|
||||
- BaseRouter 接口契约:返回模型名必须在 config.models
|
||||
- 路由决策:code→gpt-4o-mini、math→claude-3-haiku、低置信度→mistral-medium
|
||||
- 预测文件格式:包含 RouterArena 必填字段
|
||||
- Arena Score 公式与官方一致
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
# 让 tests/ 目录能找到项目根
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
from research.routerarena.adapter import ( # noqa: E402
|
||||
DOMAIN_TO_MODEL_SLOT,
|
||||
ESCALATION_MODEL_SLOT,
|
||||
ESExpertRouter,
|
||||
LOW_CONFIDENCE_THRESHOLD,
|
||||
)
|
||||
from research.routerarena.base_router import BaseRouter # noqa: E402
|
||||
from research.routerarena.local_runner import ( # noqa: E402
|
||||
MODEL_PRICING,
|
||||
build_mock_dataset,
|
||||
compute_arena_score,
|
||||
)
|
||||
|
||||
|
||||
CONFIG_PATH = str(
|
||||
Path(__file__).resolve().parent.parent
|
||||
/ "research"
|
||||
/ "routerarena"
|
||||
/ "config"
|
||||
/ "es-expert.json"
|
||||
)
|
||||
|
||||
|
||||
def test_base_router_loads_config():
|
||||
"""BaseRouter 应能加载 config 并提取 models 列表。"""
|
||||
r = ESExpertRouter(router_name="es-expert", config_path=CONFIG_PATH)
|
||||
assert "gpt-4o-mini" in r.models
|
||||
assert "claude-3-haiku-20240307" in r.models
|
||||
assert "mistral-medium" in r.models
|
||||
assert len(r.models) == 5
|
||||
|
||||
|
||||
def test_prediction_in_models():
|
||||
"""所有 get_prediction 返回值必须在 config.models 中。"""
|
||||
r = ESExpertRouter(router_name="es-expert", config_path=CONFIG_PATH)
|
||||
samples = [
|
||||
"用 Python 写一个快速排序函数",
|
||||
"Solve x^2 = 4",
|
||||
"Is a non-compete clause enforceable?",
|
||||
"高血压患者日常饮食",
|
||||
"How to calculate ROI on a fund",
|
||||
"Why is the sky blue",
|
||||
]
|
||||
for q in samples:
|
||||
m = r.get_prediction(q)
|
||||
assert m in r.models, f"{q} -> {m} not in {r.models}"
|
||||
|
||||
|
||||
def test_domain_to_slot_mapping_evidence():
|
||||
"""映射表与 00_integration_plan.md §3.2 严格一致。"""
|
||||
assert DOMAIN_TO_MODEL_SLOT["code"] == "gpt-4o-mini"
|
||||
assert DOMAIN_TO_MODEL_SLOT["math"] == "claude-3-haiku-20240307"
|
||||
assert DOMAIN_TO_MODEL_SLOT["legal"] == "claude-3-haiku-20240307"
|
||||
assert DOMAIN_TO_MODEL_SLOT["medical"] == "claude-3-haiku-20240307"
|
||||
assert DOMAIN_TO_MODEL_SLOT["finance"] == "claude-3-haiku-20240307"
|
||||
assert DOMAIN_TO_MODEL_SLOT["life"] == "gemini-2.0-flash-001"
|
||||
assert DOMAIN_TO_MODEL_SLOT["education"] == "deepseek-chat"
|
||||
assert DOMAIN_TO_MODEL_SLOT["general"] == "gemini-2.0-flash-001"
|
||||
|
||||
|
||||
def test_escalation_logic():
|
||||
"""_decide 在低置信度或低质量时升级到 ESCALATION_MODEL_SLOT。"""
|
||||
r = ESExpertRouter(router_name="es-expert", config_path=CONFIG_PATH)
|
||||
# 高置信度 + 高质量 → 主映射
|
||||
assert r._decide("code", 0.95, 0.97) == "gpt-4o-mini"
|
||||
# 低置信度 → 升级
|
||||
assert r._decide("code", 0.50, 0.97) == ESCALATION_MODEL_SLOT
|
||||
# 低质量 → 升级
|
||||
assert r._decide("code", 0.95, 0.50) == ESCALATION_MODEL_SLOT
|
||||
# 未知 domain → general
|
||||
assert r._decide("unknown_xyz", 0.95, 0.97) == "gemini-2.0-flash-001"
|
||||
|
||||
|
||||
def test_compute_arena_score_matches_formula():
|
||||
"""与 RouterArena run.py L65-87 公式逐项核对(0-1 标度)。"""
|
||||
# 已知:Hybrid Router $0.04/1K, 71.38% acc, leaderboard arena=72.08
|
||||
# 官方公式 raw 范围是 [0, 1],leaderboard 显示 ×100 标度
|
||||
s = compute_arena_score(0.04, 0.7138, beta=0.1, c_max=200, c_min=0.0044)
|
||||
# raw 期望:~0.7208(×100 = leaderboard 的 72.08)
|
||||
assert 0.715 < s < 0.725, f"expected ~0.7208, got {s:.4f}"
|
||||
# 同时验证 leaderboard 标度对齐
|
||||
assert 71.5 < s * 100 < 72.5, f"×100 期望 ~72.08, got {s*100:.4f}"
|
||||
|
||||
|
||||
def test_compute_arena_score_clamping():
|
||||
"""成本超 c_max/c_min 应被 clamp。"""
|
||||
s1 = compute_arena_score(0.005, 0.7)
|
||||
s2 = compute_arena_score(0.01, 0.7)
|
||||
# 越便宜分越高(在范围内)
|
||||
assert s1 > s2 > 0
|
||||
# cost = c_max 时 C=0, S=0(公式定义)
|
||||
s_max = compute_arena_score(200, 0.7)
|
||||
assert s_max == 0.0
|
||||
# cost 被 clamp 到 c_min
|
||||
s_under = compute_arena_score(0.001, 0.7)
|
||||
s_at_min = compute_arena_score(0.0044, 0.7)
|
||||
assert abs(s_under - s_at_min) < 1e-9
|
||||
|
||||
|
||||
def test_mock_dataset_protocol():
|
||||
"""mock 数据集符合 RouterArena 协议字段。"""
|
||||
data = build_mock_dataset()
|
||||
assert len(data) >= 80 # 接近 sub_10 的 809
|
||||
for entry in data:
|
||||
assert "global index" in entry
|
||||
assert "prompt" in entry
|
||||
assert "prompt_formatted" in entry
|
||||
assert "domain" in entry # 仅本地诊断字段
|
||||
# prompt 非空
|
||||
assert len(entry["prompt"]) > 0
|
||||
|
||||
|
||||
def test_prediction_file_schema():
|
||||
"""预测文件应包含 RouterArena 必填字段。"""
|
||||
from research.routerarena.local_runner import run_local
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
summary = run_local(
|
||||
source="mock",
|
||||
router_name="es-expert",
|
||||
config_path=CONFIG_PATH,
|
||||
do_mock_inference=True,
|
||||
output_dir=tmp,
|
||||
)
|
||||
# summary 含关键字段
|
||||
assert summary["n_queries"] >= 80
|
||||
assert "routing_distribution" in summary
|
||||
assert "mock_accuracy" in summary
|
||||
assert "arena_score_mock" in summary
|
||||
# 预测文件存在 + 格式合规
|
||||
with open(summary["prediction_file"], "r", encoding="utf-8") as f:
|
||||
preds = json.load(f)
|
||||
for p in preds:
|
||||
assert "global index" in p
|
||||
assert "prompt" in p
|
||||
assert "prediction" in p
|
||||
assert "for_optimality" in p
|
||||
assert p["prediction"] in MODEL_PRICING or p["prediction"] is None
|
||||
@@ -0,0 +1,164 @@
|
||||
"""Skill 技能体系 + RouteAgent 测试(T12:Agent-Skill 路由器·先行实现)。"""
|
||||
import pytest
|
||||
|
||||
from router_system.agent import RouteAgent
|
||||
from router_system.classifier import RuleClassifier
|
||||
from router_system.fallback import build_fallback
|
||||
from router_system.judge import build_judge
|
||||
from router_system.knowledge import KnowledgeBase
|
||||
from router_system.planner import Planner
|
||||
from router_system.skills import (FallbackSkill, JudgeSkill, KBAnswerSkill,
|
||||
KBRetrieveSkill, Skill, SkillContext,
|
||||
SkillRegistry, TemplateSkill,
|
||||
build_skill_registry)
|
||||
|
||||
|
||||
def _registry():
|
||||
kb = KnowledgeBase()
|
||||
judge = build_judge({"type": "rule"}, 0.70, kb=kb)
|
||||
fallback = build_fallback({"type": "mock"})
|
||||
return build_skill_registry(kb=kb, judge=judge, fallback=fallback)
|
||||
|
||||
|
||||
def _agent():
|
||||
kb = KnowledgeBase()
|
||||
clf = RuleClassifier()
|
||||
return RouteAgent(
|
||||
classifier=clf, planner=Planner(kb), kb=kb,
|
||||
judge=build_judge({"type": "rule"}, 0.70, kb=kb),
|
||||
fallback=build_fallback({"type": "mock"}),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Skill 注册表
|
||||
# ---------------------------------------------------------------
|
||||
def test_registry_registers_builtins():
|
||||
reg = _registry()
|
||||
names = reg.list()
|
||||
assert "es.analyze" in names
|
||||
assert "es.solve" in names
|
||||
assert "es.disclaimer" in names
|
||||
assert "kb.retrieve" in names
|
||||
assert "kb.answer" in names
|
||||
assert "judge.evaluate" in names
|
||||
assert "fallback.call" in names
|
||||
assert len(names) >= 20 # 17 模板 + kb×2 + judge + fallback
|
||||
|
||||
|
||||
def test_registry_duplicate_rejected():
|
||||
reg = SkillRegistry()
|
||||
|
||||
class S(Skill):
|
||||
name = "x"
|
||||
|
||||
reg.register(S())
|
||||
with pytest.raises(ValueError):
|
||||
reg.register(S())
|
||||
|
||||
|
||||
def test_registry_unknown_skill():
|
||||
reg = _registry()
|
||||
with pytest.raises(KeyError):
|
||||
import asyncio
|
||||
asyncio.run(reg.execute("no.such.skill", SkillContext(
|
||||
query="q", domain="general", difficulty="easy", memory=None)))
|
||||
|
||||
|
||||
def test_catalog_shape():
|
||||
reg = _registry()
|
||||
catalog = reg.catalog()
|
||||
assert catalog and all({"name", "description", "params"} <= set(c) for c in catalog)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 技能执行
|
||||
# ---------------------------------------------------------------
|
||||
@pytest.mark.asyncio
|
||||
async def test_template_skill_executes():
|
||||
reg = _registry()
|
||||
kb = KnowledgeBase()
|
||||
from router_system.memory import WorkingMemory
|
||||
text = await reg.execute("es.solve", SkillContext(
|
||||
query="求解方程 x^2-5x+6=0", domain="math", difficulty="medium",
|
||||
memory=WorkingMemory(), kb=kb))
|
||||
assert "求解" in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_kb_retrieve_skill():
|
||||
reg = _registry()
|
||||
from router_system.memory import WorkingMemory
|
||||
text = await reg.execute("kb.retrieve", SkillContext(
|
||||
query="加班费怎么计算", domain="legal", difficulty="easy",
|
||||
memory=WorkingMemory(), kb=KnowledgeBase()))
|
||||
assert "150%" in text # 命中加班费知识条目
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_kb_answer_skill():
|
||||
reg = _registry()
|
||||
from router_system.memory import WorkingMemory
|
||||
text = await reg.execute("kb.answer", SkillContext(
|
||||
query="git 回滚代码怎么操作", domain="code", difficulty="easy",
|
||||
memory=WorkingMemory(), kb=KnowledgeBase()))
|
||||
assert "git 核心概念" in text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# RouteAgent 端到端(无用户指定:完全自主分析 + 技能调用)
|
||||
# ---------------------------------------------------------------
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_code_flow_multi_skills():
|
||||
agent = _agent()
|
||||
res = await agent.route("用 Python 写一个快速排序函数,并解释时间复杂度")
|
||||
assert res.domain == "code"
|
||||
assert res.response
|
||||
# 多技能调用链:4 个 es.* 技能
|
||||
skill_steps = [s for s in res.route if s.startswith("skill:es.")]
|
||||
assert len(skill_steps) == 4
|
||||
assert any("es.analyze" in s for s in skill_steps)
|
||||
assert any("es.implement" in s for s in skill_steps)
|
||||
assert res.request_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_legal_uses_retrieve_and_disclaimer():
|
||||
agent = _agent()
|
||||
res = await agent.route("加班费怎么计算")
|
||||
assert res.domain == "legal"
|
||||
# 技能组合:kb.retrieve(法条)+ es.disclaimer(免责)
|
||||
assert any("skill:kb.retrieve" in s for s in res.route)
|
||||
assert any("skill:es.disclaimer" in s for s in res.route)
|
||||
assert "150%" in res.response
|
||||
assert "⚠" in res.response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_low_confidence_fallback():
|
||||
agent = _agent()
|
||||
res = await agent.route("今天天气怎么样")
|
||||
assert res.upgraded is True
|
||||
assert "direct_fallback" in res.route
|
||||
assert "fallback" in res.model_used or "mock" in res.model_used
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_no_user_domain_group():
|
||||
# Agent 模式:不要求 domain_group,完全自主分析
|
||||
agent = _agent()
|
||||
res = await agent.route("基金定投的收益率怎么计算")
|
||||
assert res.domain == "finance"
|
||||
assert res.subdomain == "investing"
|
||||
assert res.domain_group is None # 无用户指定概念
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_trace_and_catalog():
|
||||
agent = _agent()
|
||||
res = await agent.route("求解方程 x^2 - 5x + 6 = 0")
|
||||
trace = agent.trace_store.get(res.request_id)
|
||||
assert trace is not None
|
||||
assert trace["domain"] == "math"
|
||||
catalog = agent.skills_catalog()
|
||||
assert any(s["name"] == "kb.retrieve" for s in catalog)
|
||||
Reference in New Issue
Block a user