- 入库历史遗漏源码/测试: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
127 lines
5.1 KiB
Python
127 lines
5.1 KiB
Python
"""本系统 L0 路由器 → RouterArena BaseRouter 适配器。
|
||
|
||
核心职责:
|
||
- 接收 RouterArena 的 prompt(不含 ground truth)
|
||
- 调用本系统 Router.route() 得到 (domain, confidence, quality_score)
|
||
- 按 `DOMAIN_TO_MODEL_SLOT` 映射到 config.models 中的目标模型名
|
||
- 兜底:confidence<0.60 或 quality<0.70 → 升到 reasoning-mid
|
||
|
||
合规说明(依据 RouterArena README L34-37 + base_router.py L139-167):
|
||
- 路由阶段不接触 ground truth answer
|
||
- 不在 RouterArena 数据上训练/微调
|
||
- 8 领域映射表是手工设计,不基于 RouterArena 标签学习
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import sys
|
||
from pathlib import Path
|
||
from typing import Any, Dict
|
||
|
||
# 把项目根加入 path,以便 import router_system
|
||
_PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
||
if str(_PROJECT_ROOT) not in sys.path:
|
||
sys.path.insert(0, str(_PROJECT_ROOT))
|
||
|
||
from router_system.router import Router, build_router # noqa: E402
|
||
|
||
from .base_router import BaseRouter # noqa: E402
|
||
|
||
|
||
# L0 域 → RouterArena 候选模型槽位 的映射
|
||
# 依据 00_integration_plan.md §3.2:
|
||
# - code → gpt-4o-mini(代码强项)
|
||
# - math/legal/medical/finance → claude-3-haiku-20240307(推理严谨性)
|
||
# - life/general → gemini-2.0-flash-001(实用+快速)
|
||
# - education → deepseek-chat(中文/教育)
|
||
DOMAIN_TO_MODEL_SLOT: Dict[str, str] = {
|
||
"code": "gpt-4o-mini",
|
||
"math": "claude-3-haiku-20240307",
|
||
"legal": "claude-3-haiku-20240307",
|
||
"medical": "claude-3-haiku-20240307",
|
||
"finance": "claude-3-haiku-20240307",
|
||
"life": "gemini-2.0-flash-001",
|
||
"education": "deepseek-chat",
|
||
"general": "gemini-2.0-flash-001",
|
||
}
|
||
|
||
# 兜底:低置信度或低质 → 升到更稳的模型
|
||
ESCALATION_MODEL_SLOT = "mistral-medium"
|
||
|
||
# 阈值:与本系统 config.yaml 默认对齐
|
||
LOW_CONFIDENCE_THRESHOLD = 0.60
|
||
JUDGE_FALLBACK_THRESHOLD = 0.70
|
||
|
||
|
||
class ESExpertRouter(BaseRouter):
|
||
"""把本系统 L0 Router 适配为 RouterArena BaseRouter。
|
||
|
||
重要:_get_prediction 必须只基于 query(无 ground truth),
|
||
返回值必须在 config.models 中。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
router_name: str,
|
||
config_path: str = "",
|
||
underlying_router: Router = None,
|
||
):
|
||
super().__init__(router_name, config_path=config_path or None)
|
||
# 复用本系统 L0 Router(默认配置;如已 build 过可注入)
|
||
self._router = underlying_router or build_router()
|
||
|
||
def _get_prediction(self, query: str) -> str:
|
||
"""依据 L0 路由决策返回目标模型名。"""
|
||
# 注意:这里我们用同步方式跑异步 router.route()
|
||
# BaseRouter 的 _get_prediction 是同步签名;RouterArena 的 generate_prediction_file
|
||
# 默认单线程顺序调用 8400 条,async 包装完全兼容
|
||
import asyncio
|
||
try:
|
||
loop = asyncio.get_event_loop()
|
||
if loop.is_running():
|
||
# 已经在 async 上下文(不太可能,但兜底)
|
||
return self._get_prediction_sync(query)
|
||
return loop.run_until_complete(self._route_one(query))
|
||
except RuntimeError:
|
||
# 没有 event loop,临时建一个
|
||
return asyncio.run(self._route_one(query))
|
||
|
||
async def _route_one(self, query: str) -> str:
|
||
result = await self._router.route(query)
|
||
return self._decide(result.domain, result.confidence, result.quality_score)
|
||
|
||
def _get_prediction_sync(self, query: str) -> str:
|
||
"""event loop 已运行时的兜底(按当前 router 的同步视图返回默认值)。"""
|
||
# 我们没有同步入口,但 generate_prediction_file 是顺序同步调用 _get_prediction,
|
||
# 不会与 async 上下文冲突,所以此分支极少触发;保守返回 generalist-fast
|
||
return DOMAIN_TO_MODEL_SLOT["general"]
|
||
|
||
def _decide(self, domain: str, confidence: float, quality_score: float) -> str:
|
||
"""路由决策:按阈值升档。"""
|
||
if domain not in DOMAIN_TO_MODEL_SLOT:
|
||
domain = "general"
|
||
# 兜底升级
|
||
if confidence < LOW_CONFIDENCE_THRESHOLD or quality_score < JUDGE_FALLBACK_THRESHOLD:
|
||
return ESCALATION_MODEL_SLOT
|
||
return DOMAIN_TO_MODEL_SLOT[domain]
|
||
|
||
def diagnostics(self, query: str) -> Dict[str, Any]:
|
||
"""返回完整路由诊断(用于科研报告,不影响 RouterArena 协议)。"""
|
||
import asyncio
|
||
return asyncio.run(self._diagnose_one(query))
|
||
|
||
async def _diagnose_one(self, query: str) -> Dict[str, Any]:
|
||
result = await self._router.route(query)
|
||
return {
|
||
"query": query,
|
||
"domain": result.domain,
|
||
"subdomain": result.subdomain,
|
||
"subdomain2": result.subdomain2,
|
||
"difficulty": result.difficulty,
|
||
"confidence": result.confidence,
|
||
"quality_score": result.quality_score,
|
||
"upgraded": result.upgraded,
|
||
"model_used": result.model_used,
|
||
"route": result.route,
|
||
"selected_slot": self._decide(result.domain, result.confidence, result.quality_score),
|
||
}
|