218 lines
8.9 KiB
Python
218 lines
8.9 KiB
Python
"""主路由器:协调 缓存 -> 分类 -> 专家 -> Judge -> 大模型回退 的完整链路。
|
||
|
||
流程(对齐实现方案):
|
||
1. 检查缓存(L1 精确 / L2 语义)
|
||
2. 低置信度查询直接走大模型(should_fallback)
|
||
3. 分类器输出领域 + 难度
|
||
4. 选择专家模型生成
|
||
5. Judge 评估质量
|
||
6. 质量不达标 -> 升级大模型
|
||
7. 记录指标、写缓存、返回结果
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from typing import Any, Dict, Optional
|
||
|
||
from .cache import RouterCache
|
||
from .classifier import BaseClassifier, build_classifier
|
||
from .config import load_config
|
||
from .experts import Expert, build_expert_pool
|
||
from .fallback import FallbackProvider, build_fallback
|
||
from .judge import BaseJudge, build_judge
|
||
from .models import Classification, ExpertResponse, RouterResult, now_ms
|
||
from .stats import Stats
|
||
|
||
|
||
class Router:
|
||
def __init__(
|
||
self,
|
||
classifier: BaseClassifier,
|
||
experts: Dict[str, Expert],
|
||
judge: BaseJudge,
|
||
fallback: FallbackProvider,
|
||
cache: Optional[RouterCache] = None,
|
||
stats: Optional[Stats] = None,
|
||
config: Optional[Dict[str, Any]] = None,
|
||
):
|
||
self.classifier = classifier
|
||
self.experts = experts
|
||
self.judge = judge
|
||
self.fallback = fallback
|
||
self.cache = cache or RouterCache()
|
||
self.stats = stats or Stats()
|
||
cfg = config or {}
|
||
rcfg = cfg.get("router", {})
|
||
self.low_confidence_threshold = rcfg.get("low_confidence_threshold", 0.60)
|
||
self.judge_fallback_threshold = rcfg.get("judge_fallback_threshold", 0.70)
|
||
self.cache_enabled = cfg.get("cache", {}).get("enabled", True)
|
||
|
||
# ---------------------------------------------------------------
|
||
async def route(self, query: str) -> RouterResult:
|
||
start = now_ms()
|
||
route: list = []
|
||
|
||
# ---- Step 1: 缓存 ----
|
||
if self.cache_enabled:
|
||
hit = self.cache.get(query)
|
||
if hit is not None:
|
||
level, cached = hit
|
||
latency = now_ms() - start
|
||
result = RouterResult(
|
||
query=query,
|
||
response=cached.get("response", ""),
|
||
domain=cached.get("domain", "general"),
|
||
difficulty=cached.get("difficulty", "medium"),
|
||
confidence=cached.get("confidence", 0.0),
|
||
upgraded=False,
|
||
quality_score=cached.get("quality_score", 0.0),
|
||
model_used=cached.get("model_used", ""),
|
||
route=["cache:" + level],
|
||
latency_ms=latency,
|
||
cache_hit=True,
|
||
cache_level=level,
|
||
cost_est=0.0,
|
||
)
|
||
self.stats.record(latency, result.domain, result.difficulty, False, True, level, 0.0, result.model_used)
|
||
return result
|
||
route.append("cache:miss")
|
||
|
||
# ---- Step 2: 分类 ----
|
||
classification = self.classifier.classify(query)
|
||
route.append(f"classify:{classification.domain}@{classification.confidence:.2f}/{classification.difficulty}")
|
||
|
||
# 低置信度 -> 直接走大模型
|
||
if self.classifier.should_fallback(classification, self.low_confidence_threshold):
|
||
route.append("direct_fallback")
|
||
fb = await self._call_fallback(query)
|
||
latency = now_ms() - start
|
||
result = self._finalize(query, classification, fb, quality_score=0.0,
|
||
upgraded=True, route=route, latency_ms=latency,
|
||
model_used=fb.model_used, cost_est=fb.cost_est)
|
||
self._record(result, latency)
|
||
return result
|
||
|
||
# ---- Step 3: 选择专家 ----
|
||
domain = classification.domain
|
||
expert = self.experts.get(domain)
|
||
if expert is None:
|
||
expert = self.experts.get("general")
|
||
route.append("expert:fallback-to-general")
|
||
else:
|
||
route.append(f"expert:{expert.name}")
|
||
|
||
# ---- Step 4: 生成 ----
|
||
try:
|
||
expert_resp = await expert.generate(query, classification.difficulty)
|
||
except Exception as e:
|
||
self.stats.record_error()
|
||
route.append(f"expert_error:{type(e).__name__}")
|
||
fb = await self._call_fallback(query)
|
||
latency = now_ms() - start
|
||
result = self._finalize(query, classification, fb, quality_score=0.0,
|
||
upgraded=True, route=route, latency_ms=latency,
|
||
model_used=fb.model_used, cost_est=fb.cost_est,
|
||
error=str(e))
|
||
self._record(result, latency)
|
||
return result
|
||
|
||
# ---- Step 5: Judge 评估 ----
|
||
try:
|
||
evaluation = await self.judge.evaluate(query, expert_resp.text, domain)
|
||
except Exception:
|
||
evaluation = None
|
||
route.append("judge_error")
|
||
|
||
quality_score = evaluation.overall_score if evaluation else 0.0
|
||
route.append(f"judge:{quality_score:.2f}")
|
||
|
||
upgraded = False
|
||
final_resp = expert_resp
|
||
if evaluation is not None and evaluation.needs_fallback:
|
||
route.append("upgrade")
|
||
final_resp = await self._call_fallback(query)
|
||
upgraded = True
|
||
|
||
latency = now_ms() - start
|
||
result = self._finalize(query, classification, final_resp, quality_score=quality_score,
|
||
upgraded=upgraded, route=route, latency_ms=latency,
|
||
model_used=final_resp.model_used, cost_est=final_resp.cost_est)
|
||
self._record(result, latency)
|
||
|
||
# 未升级的结果写缓存
|
||
if self.cache_enabled and not upgraded and result.response:
|
||
self.cache.put(query, result.to_dict())
|
||
|
||
return result
|
||
|
||
# ---------------------------------------------------------------
|
||
async def _call_fallback(self, query: str) -> ExpertResponse:
|
||
try:
|
||
return await self.fallback.generate(query)
|
||
except Exception as e:
|
||
# 回退也失败:返回错误占位响应
|
||
return ExpertResponse(
|
||
text=f"[系统错误] 专家与大模型回退均失败:{type(e).__name__}: {e}",
|
||
model_used=f"error:{self.fallback.name}",
|
||
cost_est=0.0,
|
||
)
|
||
|
||
@staticmethod
|
||
def _finalize(query: str, classification: Classification, resp: ExpertResponse,
|
||
quality_score: float, upgraded: bool, route: list,
|
||
latency_ms: float, model_used: str, cost_est: float,
|
||
error: Optional[str] = None) -> RouterResult:
|
||
return RouterResult(
|
||
query=query,
|
||
response=resp.text,
|
||
domain=classification.domain,
|
||
difficulty=classification.difficulty,
|
||
confidence=classification.confidence,
|
||
upgraded=upgraded,
|
||
quality_score=quality_score,
|
||
model_used=model_used,
|
||
route=route,
|
||
latency_ms=latency_ms,
|
||
cache_hit=False,
|
||
cost_est=cost_est,
|
||
error=error,
|
||
)
|
||
|
||
def _record(self, result: RouterResult, latency_ms: float):
|
||
self.stats.record(
|
||
latency_ms,
|
||
result.domain,
|
||
result.difficulty,
|
||
result.upgraded,
|
||
result.cache_hit,
|
||
result.cache_level,
|
||
result.cost_est,
|
||
result.model_used,
|
||
)
|
||
|
||
# ---------------------------------------------------------------
|
||
def health(self) -> Dict[str, Any]:
|
||
return {
|
||
"status": "ok",
|
||
"domains": list(self.experts.keys()),
|
||
"classifier": type(self.classifier).__name__,
|
||
"judge": type(self.judge).__name__,
|
||
"fallback": type(self.fallback).__name__,
|
||
}
|
||
|
||
|
||
def build_router(config_path: Optional[str] = None) -> Router:
|
||
"""从配置构建完整 Router(默认 mock 全链路,零依赖可跑)。"""
|
||
config = load_config(config_path)
|
||
classifier = build_classifier(config.get("classifier", {}))
|
||
experts = build_expert_pool(config.get("experts", {}), config.get("domains", []))
|
||
judge = build_judge(config.get("judge", {}), config.get("router", {}).get("judge_fallback_threshold", 0.70))
|
||
fallback = build_fallback(config.get("fallback", {}))
|
||
cache_cfg = config.get("cache", {})
|
||
cache = RouterCache(
|
||
semantic_enabled=cache_cfg.get("semantic_enabled", True),
|
||
similarity_threshold=cache_cfg.get("similarity_threshold", 0.88),
|
||
promote_frequency=cache_cfg.get("promote_frequency", 5),
|
||
)
|
||
stats = Stats()
|
||
return Router(classifier, experts, judge, fallback, cache, stats, config)
|
||
|