"""本系统 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), }