From c072a1a23725d9b4ea06343770fc021571f8764b Mon Sep 17 00:00:00 2001 From: tzt <14718231+flying-travel@user.noreply.gitee.com> Date: Fri, 18 Sep 2026 23:27:26 +0800 Subject: [PATCH] =?UTF-8?q?feat(v1):=20=E6=9E=B6=E6=9E=84=E4=B8=8E?= =?UTF-8?q?=E7=AE=97=E6=B3=95=E4=BC=98=E5=8C=96=E4=BA=8C=E8=BD=AE=E2=80=94?= =?UTF-8?q?=E2=80=94=E5=85=B1=E4=BA=AB=20LLM=20=E5=AE=A2=E6=88=B7=E7=AB=AF?= =?UTF-8?q?=E5=8E=BB=E9=87=8D=E3=80=81=E8=AF=AD=E4=B9=89=E7=BC=93=E5=AD=98?= =?UTF-8?q?=20O(1)=20=E6=8F=90=E5=8D=87/=E6=B7=98=E6=B1=B0=E3=80=81?= =?UTF-8?q?=E5=88=86=E7=B1=BB=E5=99=A8=E7=83=AD=E8=B7=AF=E5=BE=84=E6=B8=85?= =?UTF-8?q?=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 router_system/llm_client.py:OpenAICompatClient 统一 experts/judge/fallback 三处复制的懒建 AsyncClient + /chat/completions + choices/usage 解析(~60 行去重); 密钥解析统一走 config.get_api_key(激活原死代码,顺带消除 experts 默认环境名不一致) - 语义缓存 L2:条目容器 list→OrderedDict(提升/淘汰 O(n)→O(1)),按 query 天然去重; n-gram 向量 lru_cache 复用(同一次 miss 的 get/put 免重复分词); A/B:淘汰路径 0.040→0.034s,miss→put 往返 9.41→8.57s(-9%) - RuleJudge 覆盖度:response.lower() 提出逐词循环(原 O(terms×len) 重复复制) - extract_content_terms 纯函数 lru_cache 化(专家与 Judge 对同一查询免重复分词),返回 tuple - RuleClassifier:_score 去掉败者领域白建的命中词 list(胜出后单独收集); 修复 code 规则 ("api",0.7) 重复登记(原命中计 1.4 分) - difficulty:正则模块级预编译 - tests:恢复上一轮引入的乱码中文 docstring;网关测试输入串恢复为可判 code 的中文查询 --- router_system/cache.py | 55 ++-- router_system/classifier.py | 25 +- router_system/difficulty.py | 108 ++++---- router_system/experts.py | 531 ++++++++++++++++++------------------ router_system/fallback.py | 183 ++++++------- router_system/judge.py | 342 ++++++++++++----------- router_system/llm_client.py | 69 +++++ tests/conftest.py | 2 +- tests/test_cache.py | 8 +- tests/test_gateway.py | 4 +- 10 files changed, 696 insertions(+), 631 deletions(-) create mode 100644 router_system/llm_client.py diff --git a/router_system/cache.py b/router_system/cache.py index e4295ac..6a50aad 100644 --- a/router_system/cache.py +++ b/router_system/cache.py @@ -12,11 +12,17 @@ - 每条语义缓存条目在写入时预计算并缓存向量范数,查询时免重复计算(原来每对比较都重算) - 语义查找单遍完成:扫描即跟踪最优条目与命中计数,命中后不再二次线性查找 - 相似度达到 1.0(完全相同查询)时提前终止扫描(余弦相似度上界,不可能更优) +- 语义条目容器为 OrderedDict:提升/淘汰均为 O(1)(原 list pop 是 O(n) 搬移), + 且按 query 天然去重(原实现同一 query 可重复追加,重复条目白白参与扫描) +- n-gram 向量经 lru_cache 复用:同一次未命中请求里 get 与紧随的 put + 不再重复做两遍分词+计数 """ from __future__ import annotations import re +from collections import OrderedDict from dataclasses import dataclass +from functools import lru_cache from typing import Any, Dict, List, Optional, Tuple @@ -52,6 +58,17 @@ def _dot(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float: return sum(v * vec_b.get(k, 0.0) for k, v in vec_a.items()) +@lru_cache(maxsize=1024) +def _embed(text: str) -> Tuple[Dict[str, float], float]: + """n-gram 向量 + 范数(带缓存)。 + + 纯函数且同一次未命中请求会在 get/put 各用一次,缓存可省一遍分词。 + 契约:返回的 dict 仅供只读(内部无任何修改点),调用方不得改写。 + """ + vec = _tf_vector(_ngrams(text)) + return vec, _norm(vec) + + class RouterCache: """L1 精确缓存 + L2 语义缓存。""" @@ -63,7 +80,8 @@ class RouterCache: self.max_exact = max_exact self.max_semantic = max_semantic self._exact: Dict[str, CacheEntry] = {} - self._semantic: List[Tuple[str, CacheEntry]] = [] # (query, entry) + # (query -> entry),插入序即淘汰序;len 兼容原 list 断言 + self._semantic: "OrderedDict[str, CacheEntry]" = OrderedDict() self._sem_vecs: Dict[str, Dict[str, float]] = {} self._sem_norms: Dict[str, float] = {} # 预计算范数,避免查询期重算 self.hits = {"exact": 0, "semantic": 0} @@ -78,57 +96,56 @@ class RouterCache: return ("exact", entry.result) if self.semantic_enabled: - q_vec = _tf_vector(_ngrams(query)) - q_norm = _norm(q_vec) + q_vec, q_norm = _embed(query) best_sim = 0.0 - best_idx = -1 + best_q = "" if q_norm > 0.0: - # 单遍扫描:同时跟踪最优相似度与条目位置 - for i, (q, _e) in enumerate(self._semantic): + # 单遍扫描(插入序):同时跟踪最优条目,sim=1.0 提前终止 + for q, _e in self._semantic.items(): n_q = self._sem_norms.get(q, 0.0) if n_q <= 0.0: continue sim = _dot(q_vec, self._sem_vecs.get(q, {})) / (q_norm * n_q) if sim > best_sim: best_sim = sim - best_idx = i + best_q = q if sim >= 1.0: break # 余弦相似度上界:完全相同查询,提前终止 - if best_idx >= 0 and best_sim >= self.similarity_threshold: - best_q, best_entry = self._semantic[best_idx] + if best_q and best_sim >= self.similarity_threshold: + best_entry = self._semantic[best_q] # 完全相同查询(相似度=1.0)计为 exact 命中 is_exact = best_sim >= 0.999 level = "exact" if is_exact else "semantic" self.hits[level] += 1 - self._bump_semantic(best_idx, best_q, best_entry) + self._bump_semantic(best_q, best_entry) return (level, best_entry.result) self.misses += 1 return None - def _bump_semantic(self, idx: int, query: str, entry: CacheEntry): - """语义命中:累计命中次数,达到阈值提升为精确缓存(O(1),无需二次查找)。""" + def _bump_semantic(self, query: str, entry: CacheEntry): + """语义命中:累计命中次数,达到阈值提升为精确缓存(OrderedDict 删除 O(1))。""" entry.hits += 1 if entry.hits >= self.promote_frequency: self._exact[query] = entry - self._semantic.pop(idx) + del self._semantic[query] self._sem_vecs.pop(query, None) self._sem_norms.pop(query, None) # ---- 写入 ---- def put(self, query: str, result: Dict[str, Any]): - if query in self._exact: - return + if query in self._exact or query in self._semantic: + return # 已缓存(语义区按 query 天然去重,避免重复条目参与扫描) entry = CacheEntry(result=result) if self.semantic_enabled: if len(self._semantic) >= self.max_semantic: - old_q, _ = self._semantic.pop(0) + old_q, _old = self._semantic.popitem(last=False) # O(1) 淘汰最旧 self._sem_vecs.pop(old_q, None) self._sem_norms.pop(old_q, None) - self._semantic.append((query, entry)) - vec = _tf_vector(_ngrams(query)) + self._semantic[query] = entry + vec, norm = _embed(query) # 与 get miss 共享同一份缓存向量 self._sem_vecs[query] = vec - self._sem_norms[query] = _norm(vec) + self._sem_norms[query] = norm else: self._exact[query] = entry if len(self._exact) > self.max_exact: diff --git a/router_system/classifier.py b/router_system/classifier.py index 27db5fa..f275300 100644 --- a/router_system/classifier.py +++ b/router_system/classifier.py @@ -33,7 +33,7 @@ DOMAIN_RULES: Dict[str, List[Tuple[str, float]]] = { ("algorithm", 0.8), ("sort", 0.7), ("array", 0.7), ("regex", 0.8), ("import", 0.7), ("loop", 0.7), ("recursion", 0.8), ("refactor", 0.8), ("deploy", 0.8), ("docker", 0.8), ("kubernetes", 0.8), - ("async", 0.7), ("flask", 0.7), ("django", 0.7), ("api", 0.7), + ("async", 0.7), ("flask", 0.7), ("django", 0.7), ("索引", 0.8), ("优化", 0.7), ("查询", 0.6), ], "math": [ @@ -98,24 +98,29 @@ class RuleClassifier(BaseClassifier): self.confidence_floor = confidence_floor self.rules = DOMAIN_RULES - def _score(self, query: str) -> Tuple[Dict[str, float], Dict[str, List[str]]]: - q = query.lower() + def _score(self, q: str) -> Dict[str, float]: + """对已 lowercase 的查询按领域规则打分,返回有命中的领域分数。 + + 命中词列表只对最终胜出领域有意义(classify 的唯一消费点), + 故不在各领域上白建 list——胜出后由 _matched_rules 单独收集。 + """ scores: Dict[str, float] = {} - matched: Dict[str, List[str]] = {} for domain, rules in self.rules.items(): s = 0.0 - hits = [] for kw, w in rules: if kw in q: s += w - hits.append(kw) if s > 0: scores[domain] = s - matched[domain] = hits - return scores, matched + return scores + + def _matched_rules(self, q: str, domain: str) -> List[str]: + """收集指定领域命中的关键词(保持登记顺序)。""" + return [kw for kw, _w in self.rules.get(domain, ()) if kw in q] def classify(self, query: str) -> Classification: - raw, matched = self._score(query) + q = query.lower() + raw = self._score(q) if not raw: # 完全无命中 -> general,低置信度 diff, ds = estimate_difficulty(query) @@ -150,7 +155,7 @@ class RuleClassifier(BaseClassifier): difficulty=diff, difficulty_score=ds, raw_scores={k: round(v, 3) for k, v in raw.items()}, - matched_rules=matched.get(best_domain, []), + matched_rules=self._matched_rules(q, best_domain), ) diff --git a/router_system/difficulty.py b/router_system/difficulty.py index b8acf6c..31fa5f9 100644 --- a/router_system/difficulty.py +++ b/router_system/difficulty.py @@ -1,52 +1,56 @@ -"""查询难度估计器(启发式,纯标准库)。 - -依据:RouterArena 用 Bloom 分类法把问题分为 easy/medium/hard。 -这里用查询长度、指令动词、数学/推理标记做轻量估计。 -""" -from __future__ import annotations - -import re -from typing import Tuple - -# 触发 hard 的指令动词 / 推理标记 -_HARD_MARKERS = [ - "证明", "推导", "为什么", "如何", "对比", "比较", "分析", "评估", "设计", "优化", - "复杂度", "时间复杂度", "空间复杂度", "原理", "机制", "优缺点", "区别", "推论", "定理", - "proof", "prove", "derive", "explain why", "why", "how", "compare", "contrast", - "analy", "evaluate", "design", "optimize", "refactor", "architect", - "implement", "debug", "review", "plan", "synthesize", - "int", "sum", "sqrt", "lim", "log", "derivative", "integral", -] -# 触发 medium 的标记 -_MEDIUM_MARKERS = [ - "用", "写", "计算", "求解", "生成", "翻译", "总结", "解释", - "注意", "建议", "是否", "实现", "步骤", - "write", "code", "function", "script", "calculate", "solve", "summarize", - "translate", "fix", "explain", "describe", "list", -] - - -def estimate_difficulty(query: str) -> Tuple[str, float]: - """返回 (difficulty, score),score 属于 [0,1]。""" - q = query.lower() - hard_hits = sum(1 for m in _HARD_MARKERS if m in q) - medium_hits = sum(1 for m in _MEDIUM_MARKERS if m in q) - length = len(query) - - score = 0.0 - score += min(0.30, length / 600.0) # 长度贡献 - score += min(0.55, hard_hits * 0.25) # 推理标记贡献 - score += min(0.30, medium_hits * 0.08) # 一般指令贡献 - - # 额外:代码 / 数学表达式(多步骤信号) - if "```" in query or re.search(r"\b(def|class|function|import)\b", q): - score += 0.15 - if re.search(r"[0-9]+\s*[+\-*/^=]\s*[0-9xya-z]", q): - score += 0.15 - - score = max(0.0, min(1.0, score)) - if score >= 0.55: - return "hard", score - if score >= 0.25: - return "medium", score - return "easy", score +"""查询难度估计器(启发式,纯标准库)。 + +依据:RouterArena 用 Bloom 分类法把问题分为 easy/medium/hard。 +这里用查询长度、指令动词、数学/推理标记做轻量估计。 +""" +from __future__ import annotations + +import re +from typing import Tuple + +# 触发 hard 的指令动词 / 推理标记 +_HARD_MARKERS = [ + "证明", "推导", "为什么", "如何", "对比", "比较", "分析", "评估", "设计", "优化", + "复杂度", "时间复杂度", "空间复杂度", "原理", "机制", "优缺点", "区别", "推论", "定理", + "proof", "prove", "derive", "explain why", "why", "how", "compare", "contrast", + "analy", "evaluate", "design", "optimize", "refactor", "architect", + "implement", "debug", "review", "plan", "synthesize", + "int", "sum", "sqrt", "lim", "log", "derivative", "integral", +] +# 触发 medium 的标记 +_MEDIUM_MARKERS = [ + "用", "写", "计算", "求解", "生成", "翻译", "总结", "解释", + "注意", "建议", "是否", "实现", "步骤", + "write", "code", "function", "script", "calculate", "solve", "summarize", + "translate", "fix", "explain", "describe", "list", +] + +# 预编译正则(模块级一次,避免每次调用走 re 内部缓存查找) +_RE_CODE_EXPR = re.compile(r"\b(def|class|function|import)\b") +_RE_ARITH_EXPR = re.compile(r"[0-9]+\s*[+\-*/^=]\s*[0-9xya-z]") + + +def estimate_difficulty(query: str) -> Tuple[str, float]: + """返回 (difficulty, score),score 属于 [0,1]。""" + q = query.lower() + hard_hits = sum(1 for m in _HARD_MARKERS if m in q) + medium_hits = sum(1 for m in _MEDIUM_MARKERS if m in q) + length = len(query) + + score = 0.0 + score += min(0.30, length / 600.0) # 长度贡献 + score += min(0.55, hard_hits * 0.25) # 推理标记贡献 + score += min(0.30, medium_hits * 0.08) # 一般指令贡献 + + # 额外:代码 / 数学表达式(多步骤信号) + if "```" in query or _RE_CODE_EXPR.search(q): + score += 0.15 + if _RE_ARITH_EXPR.search(q): + score += 0.15 + + score = max(0.0, min(1.0, score)) + if score >= 0.55: + return "hard", score + if score >= 0.25: + return "medium", score + return "easy", score diff --git a/router_system/experts.py b/router_system/experts.py index 95dac8f..679287a 100644 --- a/router_system/experts.py +++ b/router_system/experts.py @@ -1,270 +1,261 @@ -"""专家模型池:统一 Expert 接口,支持三种后端。 - -- MockExpert :确定性模板输出(零依赖,离线可跑,便于测试与演示) -- HFExpert :HuggingFace transformers 真实小模型(可选,需 ML 依赖) -- APIExpert :OpenAI 兼容 API(可选,需 API Key,如 DeepSeek) - -成本估计:cost_est 按参数量粗估(美元/百万 token 的近似比例)。 -""" -from __future__ import annotations - -import asyncio -import re -from typing import Dict, List, Optional - -from .models import ExpertResponse - -# 按模型规模粗估的相对成本($ / 1M output tokens,近似) -MODEL_COST_EST = { - "mock": 0.0, - "0.5b": 0.02, - "1b": 0.05, - "1.7b": 0.08, - "3b": 0.12, - "4b": 0.15, - "7b": 0.25, - "70b": 2.50, - "api": 1.00, -} - - -def _cost_for(model_name: str, default: str = "1b") -> float: - mn = model_name.lower() - for key in ("0.5b", "1.7b", "3b", "4b", "7b", "70b"): - if key in mn: - return MODEL_COST_EST[key] - if "api" in mn or mn in ("deepseek-chat", "gpt-4o-mini", "claude"): - return MODEL_COST_EST["api"] - return MODEL_COST_EST.get(default, 0.1) - - -def extract_content_terms(query: str) -> List[str]: - """抽取查询中的"内容词"(中文词/英文单词),用于 Judge 覆盖度与 Mock 回显。""" - q = query.lower() - terms: List[str] = [] - # 英文单词(>=2 字符) - for w in re.findall(r"[a-z][a-z0-9_]{1,}", q): - if w not in _STOPWORDS_EN and w not in terms: - terms.append(w) - # 中文:按 2-4 字窗口切分,保留含中文字符的片段 - cn = re.findall(r"[\u4e00-\u9fff]{2,8}", q) - for c in cn: - terms.append(c) - return terms - - -_STOPWORDS_EN = { - "the", "a", "an", "is", "are", "to", "of", "in", "on", "for", "with", "and", - "or", "do", "does", "can", "could", "would", "should", "please", "me", "my", - "this", "that", "it", "be", "was", "were", "have", "has", "had", "will", - "not", "no", "yes", "i", "you", "he", "she", "we", "they", -} - - -class Expert: - name: str = "expert" - - async def generate(self, query: str, difficulty: str) -> ExpertResponse: - raise NotImplementedError - - -class MockExpert(Expert): - """确定性模板专家:零依赖,离线可跑。 - - 输出会回显查询中的内容词以提高 Judge 覆盖度,并带领域结构, - 使端到端管线(分类 -> 专家 -> Judge -> 缓存)可被稳定测试与演示。 - """ - - def __init__(self, name: str, domain: str, model: str = "mock"): - self.name = name - self.domain = domain - self.model = model - - async def generate(self, query: str, difficulty: str) -> ExpertResponse: - await asyncio.sleep(0.001) # 模拟极短推理延迟 - terms = extract_content_terms(query) - body = self._template(query, terms, difficulty) - # 预估 token 数:中文约 1.5 字符/token,英文约 4 字符/token - tokens = max(8, int(len(body) / 2.2)) - return ExpertResponse( - text=body, - model_used=self.model, - latency_ms=1.0, - tokens=tokens, - cost_est=_cost_for(self.model) * tokens / 1_000_000, - ) - - def _template(self, query: str, terms: List[str], difficulty: str) -> str: - kw = "、".join(terms[:6]) if terms else "该主题" - if self.domain == "code": - return ( - f"(mock 代码专家)针对「{query}」的实现思路如下:\n\n" - f"```python\n" - f"def solve() -> None:\n" - f" # 关键点:{kw}\n" - f" # 1. 明确输入输出约束\n" - f" # 2. 选择合适数据结构\n" - f" # 3. 处理边界条件(空输入、极端值)\n" - f" # 4. 补充单元测试\n" - f" pass\n" - f"```\n\n" - f"复杂度:平均 O(n)。请按上述步骤补充具体实现。" - ) - if self.domain == "math": - return ( - f"(mock 数学专家)求解「{query}」的步骤:\n\n" - f"1. 明确已知条件与目标:{kw}\n" - f"2. 选择合适的方法(代数变形 / 积分 / 归纳等)\n" - f"3. 逐步推导并验证中间结果\n" - f"4. 检查边界与特殊情况\n\n" - f"结论:在标准假设下,结果可化简为闭合形式。完整推导见正式解答。" - ) - if self.domain == "legal": - return ( - f"(mock 法律专家)关于「{query}」的初步法律分析:\n\n" - f"相关要点:{kw}\n" - f"1. 适用法规:请以现行有效法条为准(建议核对最新修订版)\n" - f"2. 合同/合规风险点识别\n" - f"3. 责任划分与救济途径\n\n" - f"⚠️ 提示:以上为一般性分析,不构成正式法律意见,个案请咨询执业律师。" - ) - if self.domain == "medical": - return ( - f"(mock 医学专家)关于「{query}」的科普性说明:\n\n" - f"相关关键词:{kw}\n" - f"1. 常见表现与可能原因\n" - f"2. 一般处理建议与注意事项\n" - f"3. 何时需要就医(警示信号)\n\n" - f"⚠️ 提示:内容仅供健康科普,不能替代医生诊断;如有不适请及时就医。" - ) - return ( - f"(mock 通用专家)关于「{query}」的回答:\n\n" - f"核心要点:{kw}\n" - f"1. 背景与定义\n" - f"2. 主要分类/维度\n" - f"3. 实际应用与注意事项\n\n" - f"如需更深入的分析,可以补充更多上下文。" - ) - - -class HFExpert(Expert): - """可选:HuggingFace 真实小模型(需 requirements-ml.txt)。""" - - def __init__(self, name: str, domain: str, model: str): - self.name = name - self.domain = domain - self.model = model - self._loaded = False - self._model = None - self._tokenizer = None - - def _ensure_loaded(self): - if self._loaded: - return - try: - from transformers import AutoModelForCausalLM, AutoTokenizer - except ImportError as e: - raise RuntimeError("HFExpert 需要安装 ML 依赖:pip install -r requirements-ml.txt") from e - self._tokenizer = AutoTokenizer.from_pretrained(self.model) - self._model = AutoModelForCausalLM.from_pretrained( - self.model, device_map="auto", torch_dtype="auto" - ) - self._loaded = True - - async def generate(self, query: str, difficulty: str) -> ExpertResponse: - self._ensure_loaded() - return await asyncio.to_thread(self._generate_sync, query) - - def _generate_sync(self, query: str) -> ExpertResponse: - messages = [{"role": "user", "content": query}] - text = self._tokenizer.apply_chat_template(messages, tokenize=False) - inputs = self._tokenizer(text, return_tensors="pt").to(self._model.device) - outputs = self._model.generate( - **inputs, - max_new_tokens=512, - temperature=0.2, - do_sample=True, - ) - body = self._tokenizer.decode(outputs[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True) - return ExpertResponse( - text=body, - model_used=self.model, - latency_ms=0.0, - tokens=512, - cost_est=_cost_for(self.model) * 512 / 1_000_000, - ) - - -class APIExpert(Expert): - """可选:OpenAI 兼容 Chat Completions(DeepSeek / OpenAI / 本地 vLLM)。""" - - def __init__(self, name: str, domain: str, model: str, base_url: str, api_key: str): - self.name = name - self.domain = domain - self.model = model - self.base_url = base_url.rstrip("/") - self.api_key = api_key - self._client = None - - def _get_client(self): - if self._client is None: - import httpx - self._client = httpx.AsyncClient(timeout=60.0) - return self._client - - async def generate(self, query: str, difficulty: str) -> ExpertResponse: - client = self._get_client() - resp = await client.post( - f"{self.base_url}/chat/completions", - headers={"Authorization": f"Bearer {self.api_key}"}, - json={ - "model": self.model, - "messages": [{"role": "user", "content": query}], - "temperature": 0.2, - "max_tokens": 1024, - }, - ) - resp.raise_for_status() - data = resp.json() - body = data["choices"][0]["message"]["content"] - usage = data.get("usage", {}) - tokens = usage.get("completion_tokens", int(len(body) / 2.2)) - return ExpertResponse( - text=body, - model_used=self.model, - latency_ms=0.0, - tokens=tokens, - cost_est=_cost_for(self.model) * tokens / 1_000_000, - ) - - -def build_expert(domain: str, cfg: Dict) -> Expert: - """根据配置构建领域专家。cfg 为 experts. 段配置。""" - etype = cfg.get("type", "mock") - model = cfg.get("model", "mock") - name = f"expert-{domain}" - if etype == "mock": - return MockExpert(name, domain, model) - if etype == "hf": - return HFExpert(name, domain, model) - if etype == "api": - base_url = cfg.get("base_url", "https://api.deepseek.com/v1") - api_key = cfg.get("api_key") or _env(cfg.get("api_key_env", "")) - if not api_key: - raise RuntimeError(f"APIExpert({domain}) 缺少 API Key(env: {cfg.get('api_key_env')})") - return APIExpert(name, domain, model, base_url, api_key) - raise ValueError(f"未知专家后端类型: {etype}(支持 mock | hf | api)") - - -def _env(name: str) -> Optional[str]: - import os - return os.environ.get(name) if name else None - - -def build_expert_pool(experts_cfg: Dict[str, Dict], domains: List[str]) -> Dict[str, Expert]: - """构建完整专家池。""" - pool: Dict[str, Expert] = {} - for domain in domains: - cfg = experts_cfg.get(domain, {"type": "mock", "model": "mock"}) - pool[domain] = build_expert(domain, cfg) - return pool +"""专家模型池:统一 Expert 接口,支持三种后端。 + +- MockExpert :确定性模板输出(零依赖,离线可跑,便于测试与演示) +- HFExpert :HuggingFace transformers 真实小模型(可选,需 ML 依赖) +- APIExpert :OpenAI 兼容 API(可选,需 API Key,如 DeepSeek) + +成本估计:cost_est 按参数量粗估(美元/百万 token 的近似比例)。 +""" +from __future__ import annotations + +import asyncio +import re +from functools import lru_cache +from typing import Dict, List, Optional, Tuple + +from .config import get_api_key +from .llm_client import OpenAICompatClient +from .models import ExpertResponse + +# 按模型规模粗估的相对成本($ / 1M output tokens,近似) +MODEL_COST_EST = { + "mock": 0.0, + "0.5b": 0.02, + "1b": 0.05, + "1.7b": 0.08, + "3b": 0.12, + "4b": 0.15, + "7b": 0.25, + "70b": 2.50, + "api": 1.00, +} + + +def _cost_for(model_name: str, default: str = "1b") -> float: + mn = model_name.lower() + for key in ("0.5b", "1.7b", "3b", "4b", "7b", "70b"): + if key in mn: + return MODEL_COST_EST[key] + if "api" in mn or mn in ("deepseek-chat", "gpt-4o-mini", "claude"): + return MODEL_COST_EST["api"] + return MODEL_COST_EST.get(default, 0.1) + + +@lru_cache(maxsize=2048) +def _extract_content_terms_cached(query: str) -> Tuple[str, ...]: + q = query.lower() + terms: List[str] = [] + # 英文单词(>=2 字符) + for w in re.findall(r"[a-z][a-z0-9_]{1,}", q): + if w not in _STOPWORDS_EN and w not in terms: + terms.append(w) + # 中文:按 2-4 字窗口切分,保留含中文字符的片段 + cn = re.findall(r"[\u4e00-\u9fff]{2,8}", q) + for c in cn: + terms.append(c) + return tuple(terms) + + +def extract_content_terms(query: str) -> Tuple[str, ...]: + """抽取查询中的"内容词"(中文词/英文单词),用于 Judge 覆盖度与 Mock 回显。 + + 纯函数,但 Router 主链路中专家生成与 Judge 评估会对同一查询各调一次, + 故以 lru_cache 去重;返回不可变 tuple(调用方只读)。 + """ + return _extract_content_terms_cached(query) + + +_STOPWORDS_EN = { + "the", "a", "an", "is", "are", "to", "of", "in", "on", "for", "with", "and", + "or", "do", "does", "can", "could", "would", "should", "please", "me", "my", + "this", "that", "it", "be", "was", "were", "have", "has", "had", "will", + "not", "no", "yes", "i", "you", "he", "she", "we", "they", +} + + +class Expert: + name: str = "expert" + + async def generate(self, query: str, difficulty: str) -> ExpertResponse: + raise NotImplementedError + + +class MockExpert(Expert): + """确定性模板专家:零依赖,离线可跑。 + + 输出会回显查询中的内容词以提高 Judge 覆盖度,并带领域结构, + 使端到端管线(分类 -> 专家 -> Judge -> 缓存)可被稳定测试与演示。 + """ + + def __init__(self, name: str, domain: str, model: str = "mock"): + self.name = name + self.domain = domain + self.model = model + + async def generate(self, query: str, difficulty: str) -> ExpertResponse: + await asyncio.sleep(0.001) # 模拟极短推理延迟 + terms = extract_content_terms(query) + body = self._template(query, terms, difficulty) + # 预估 token 数:中文约 1.5 字符/token,英文约 4 字符/token + tokens = max(8, int(len(body) / 2.2)) + return ExpertResponse( + text=body, + model_used=self.model, + latency_ms=1.0, + tokens=tokens, + cost_est=_cost_for(self.model) * tokens / 1_000_000, + ) + + def _template(self, query: str, terms: List[str], difficulty: str) -> str: + kw = "、".join(terms[:6]) if terms else "该主题" + if self.domain == "code": + return ( + f"(mock 代码专家)针对「{query}」的实现思路如下:\n\n" + f"```python\n" + f"def solve() -> None:\n" + f" # 关键点:{kw}\n" + f" # 1. 明确输入输出约束\n" + f" # 2. 选择合适数据结构\n" + f" # 3. 处理边界条件(空输入、极端值)\n" + f" # 4. 补充单元测试\n" + f" pass\n" + f"```\n\n" + f"复杂度:平均 O(n)。请按上述步骤补充具体实现。" + ) + if self.domain == "math": + return ( + f"(mock 数学专家)求解「{query}」的步骤:\n\n" + f"1. 明确已知条件与目标:{kw}\n" + f"2. 选择合适的方法(代数变形 / 积分 / 归纳等)\n" + f"3. 逐步推导并验证中间结果\n" + f"4. 检查边界与特殊情况\n\n" + f"结论:在标准假设下,结果可化简为闭合形式。完整推导见正式解答。" + ) + if self.domain == "legal": + return ( + f"(mock 法律专家)关于「{query}」的初步法律分析:\n\n" + f"相关要点:{kw}\n" + f"1. 适用法规:请以现行有效法条为准(建议核对最新修订版)\n" + f"2. 合同/合规风险点识别\n" + f"3. 责任划分与救济途径\n\n" + f"⚠️ 提示:以上为一般性分析,不构成正式法律意见,个案请咨询执业律师。" + ) + if self.domain == "medical": + return ( + f"(mock 医学专家)关于「{query}」的科普性说明:\n\n" + f"相关关键词:{kw}\n" + f"1. 常见表现与可能原因\n" + f"2. 一般处理建议与注意事项\n" + f"3. 何时需要就医(警示信号)\n\n" + f"⚠️ 提示:内容仅供健康科普,不能替代医生诊断;如有不适请及时就医。" + ) + return ( + f"(mock 通用专家)关于「{query}」的回答:\n\n" + f"核心要点:{kw}\n" + f"1. 背景与定义\n" + f"2. 主要分类/维度\n" + f"3. 实际应用与注意事项\n\n" + f"如需更深入的分析,可以补充更多上下文。" + ) + + +class HFExpert(Expert): + """可选:HuggingFace 真实小模型(需 requirements-ml.txt)。""" + + def __init__(self, name: str, domain: str, model: str): + self.name = name + self.domain = domain + self.model = model + self._loaded = False + self._model = None + self._tokenizer = None + + def _ensure_loaded(self): + if self._loaded: + return + try: + from transformers import AutoModelForCausalLM, AutoTokenizer + except ImportError as e: + raise RuntimeError("HFExpert 需要安装 ML 依赖:pip install -r requirements-ml.txt") from e + self._tokenizer = AutoTokenizer.from_pretrained(self.model) + self._model = AutoModelForCausalLM.from_pretrained( + self.model, device_map="auto", torch_dtype="auto" + ) + self._loaded = True + + async def generate(self, query: str, difficulty: str) -> ExpertResponse: + self._ensure_loaded() + return await asyncio.to_thread(self._generate_sync, query) + + def _generate_sync(self, query: str) -> ExpertResponse: + messages = [{"role": "user", "content": query}] + text = self._tokenizer.apply_chat_template(messages, tokenize=False) + inputs = self._tokenizer(text, return_tensors="pt").to(self._model.device) + outputs = self._model.generate( + **inputs, + max_new_tokens=512, + temperature=0.2, + do_sample=True, + ) + body = self._tokenizer.decode(outputs[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True) + return ExpertResponse( + text=body, + model_used=self.model, + latency_ms=0.0, + tokens=512, + cost_est=_cost_for(self.model) * 512 / 1_000_000, + ) + + +class APIExpert(Expert): + """可选:OpenAI 兼容 Chat Completions(DeepSeek / OpenAI / 本地 vLLM)。""" + + def __init__(self, name: str, domain: str, model: str, base_url: str, api_key: str): + self.name = name + self.domain = domain + self.model = model + self._client = OpenAICompatClient(base_url, api_key, timeout=60.0) + + async def generate(self, query: str, difficulty: str) -> ExpertResponse: + data = await self._client.chat( + self.model, + [{"role": "user", "content": query}], + temperature=0.2, + max_tokens=1024, + ) + body = self._client.completion_text(data) + tokens = self._client.completion_tokens(data, body) + return ExpertResponse( + text=body, + model_used=self.model, + latency_ms=0.0, + tokens=tokens, + cost_est=_cost_for(self.model) * tokens / 1_000_000, + ) + + +def build_expert(domain: str, cfg: Dict) -> Expert: + """根据配置构建领域专家。cfg 为 experts. 段配置。""" + etype = cfg.get("type", "mock") + model = cfg.get("model", "mock") + name = f"expert-{domain}" + if etype == "mock": + return MockExpert(name, domain, model) + if etype == "hf": + return HFExpert(name, domain, model) + if etype == "api": + base_url = cfg.get("base_url", "https://api.deepseek.com/v1") + api_key = cfg.get("api_key") or get_api_key(cfg) + if not api_key: + raise RuntimeError(f"APIExpert({domain}) 缺少 API Key(env: {cfg.get('api_key_env')})") + return APIExpert(name, domain, model, base_url, api_key) + raise ValueError(f"未知专家后端类型: {etype}(支持 mock | hf | api)") + + +def build_expert_pool(experts_cfg: Dict[str, Dict], domains: List[str]) -> Dict[str, Expert]: + """构建完整专家池。""" + pool: Dict[str, Expert] = {} + for domain in domains: + cfg = experts_cfg.get(domain, {"type": "mock", "model": "mock"}) + pool[domain] = build_expert(domain, cfg) + return pool diff --git a/router_system/fallback.py b/router_system/fallback.py index f02106c..06c3b2e 100644 --- a/router_system/fallback.py +++ b/router_system/fallback.py @@ -1,99 +1,84 @@ -"""大模型回退层:Mock 与 OpenAI 兼容 API 两种后端。""" -from __future__ import annotations - -import asyncio -from typing import Dict, Optional - -from .models import ExpertResponse - - -class FallbackProvider: - name: str = "fallback" - - async def generate(self, query: str) -> ExpertResponse: - raise NotImplementedError - - -class MockFallback(FallbackProvider): - """确定性 mock 大模型:标识为 fallback,便于测试升级路径。""" - - def __init__(self, model: str = "mock-large"): - self.model = model - self.name = f"fallback-{model}" - - async def generate(self, query: str) -> ExpertResponse: - await asyncio.sleep(0.002) - body = ( - f"(大模型回退)「{query}」\n\n" - "这是一条来自大模型回退路径的完整回答。\n" - "要点:\n" - "1. 对复杂/跨域任务给出综合推理\n" - "2. 补充领域专家未覆盖的上下文\n" - "3. 给出可执行的后续建议\n" - ) - return ExpertResponse( - text=body, - model_used=self.model, - latency_ms=2.0, - tokens=120, - cost_est=2.0 * 120 / 1_000_000, - ) - - -class APIFallback(FallbackProvider): - """OpenAI 兼容大模型 API(如 DeepSeek / OpenAI / 本地 vLLM)。""" - - def __init__(self, model: str, base_url: str, api_key: str): - self.model = model - self.base_url = base_url.rstrip("/") - self.api_key = api_key - self.name = f"fallback-{model}" - self._client = None - - def _get_client(self): - if self._client is None: - import httpx - self._client = httpx.AsyncClient(timeout=90.0) - return self._client - - async def generate(self, query: str) -> ExpertResponse: - client = self._get_client() - resp = await client.post( - f"{self.base_url}/chat/completions", - headers={"Authorization": f"Bearer {self.api_key}"}, - json={ - "model": self.model, - "messages": [{"role": "user", "content": query}], - "temperature": 0.3, - "max_tokens": 2048, - }, - ) - resp.raise_for_status() - data = resp.json() - body = data["choices"][0]["message"]["content"] - usage = data.get("usage", {}) - tokens = usage.get("completion_tokens", int(len(body) / 2.2)) - return ExpertResponse( - text=body, - model_used=self.model, - latency_ms=0.0, - tokens=tokens, - cost_est=2.0 * tokens / 1_000_000, - ) - - -def build_fallback(cfg: Dict) -> FallbackProvider: - """cfg 为 fallback 段配置。""" - ftype = cfg.get("type", "mock") - model = cfg.get("model", "deepseek-chat") - if ftype == "mock": - return MockFallback(model=model) - if ftype == "api": - import os - api_key = cfg.get("api_key") or os.environ.get(cfg.get("api_key_env", "API_KEY")) - if not api_key: - raise RuntimeError( - f"APIFallback 缺少 API Key:请设置环境变量 {cfg.get('api_key_env')} 或配置 api_key" - ) - return APIFallback(model, cfg.get("base_url", "https://api.deepseek.com/v1"), api_key) - raise ValueError(f"未知 fallback 类型: {ftype}(支持 mock | api)") +"""大模型回退层:Mock 与 OpenAI 兼容 API 两种后端。""" +from __future__ import annotations + +import asyncio +from typing import Dict + +from .config import get_api_key +from .llm_client import OpenAICompatClient +from .models import ExpertResponse + + +class FallbackProvider: + name: str = "fallback" + + async def generate(self, query: str) -> ExpertResponse: + raise NotImplementedError + + +class MockFallback(FallbackProvider): + """确定性 mock 大模型:标识为 fallback,便于测试升级路径。""" + + def __init__(self, model: str = "mock-large"): + self.model = model + self.name = f"fallback-{model}" + + async def generate(self, query: str) -> ExpertResponse: + await asyncio.sleep(0.002) + body = ( + f"(大模型回退)「{query}」\n\n" + "这是一条来自大模型回退路径的完整回答。\n" + "要点:\n" + "1. 对复杂/跨域任务给出综合推理\n" + "2. 补充领域专家未覆盖的上下文\n" + "3. 给出可执行的后续建议\n" + ) + return ExpertResponse( + text=body, + model_used=self.model, + latency_ms=2.0, + tokens=120, + cost_est=2.0 * 120 / 1_000_000, + ) + + +class APIFallback(FallbackProvider): + """OpenAI 兼容大模型 API(如 DeepSeek / OpenAI / 本地 vLLM)。""" + + def __init__(self, model: str, base_url: str, api_key: str): + self.model = model + self.name = f"fallback-{model}" + self._client = OpenAICompatClient(base_url, api_key, timeout=90.0) + + async def generate(self, query: str) -> ExpertResponse: + data = await self._client.chat( + self.model, + [{"role": "user", "content": query}], + temperature=0.3, + max_tokens=2048, + ) + body = self._client.completion_text(data) + tokens = self._client.completion_tokens(data, body) + return ExpertResponse( + text=body, + model_used=self.model, + latency_ms=0.0, + tokens=tokens, + cost_est=2.0 * tokens / 1_000_000, + ) + + +def build_fallback(cfg: Dict) -> FallbackProvider: + """cfg 为 fallback 段配置。""" + ftype = cfg.get("type", "mock") + model = cfg.get("model", "deepseek-chat") + if ftype == "mock": + return MockFallback(model=model) + if ftype == "api": + api_key = cfg.get("api_key") or get_api_key(cfg) + if not api_key: + raise RuntimeError( + f"APIFallback 缺少 API Key:请设置环境变量 {cfg.get('api_key_env')} 或配置 api_key" + ) + return APIFallback(model, cfg.get("base_url", "https://api.deepseek.com/v1"), api_key) + raise ValueError(f"未知 fallback 类型: {ftype}(支持 mock | api)") diff --git a/router_system/judge.py b/router_system/judge.py index a0dc2c5..4f00b16 100644 --- a/router_system/judge.py +++ b/router_system/judge.py @@ -1,174 +1,168 @@ -"""质量控制器(Judge):评估专家输出,决定是否升级大模型。 - -- RuleJudge:零依赖启发式(内容覆盖度 / 长度充分性 / 领域格式 / 安全提示), - 稳定可测,适合 MVP 与离线演示。 -- LLMJudge:可选,基于 transformers 小模型或 API 的 LLM-as-Judge。 - -设计对齐实现方案:overall_score < judge_fallback_threshold -> 升级大模型。 -""" -from __future__ import annotations - -from dataclasses import dataclass, field -from typing import Dict, List - -from .experts import extract_content_terms - -# 各领域期望的响应长度范围(字符数) -_EXPECTED_LEN = { - "code": (60, 2000), - "math": (60, 2000), - "legal": (80, 3000), - "medical": (80, 3000), - "general": (40, 2000), -} - -# 领域格式检查:响应应包含的标记 -_DOMAIN_FORMAT_HINTS = { - "code": ["```", "def ", "function", "class "], - "math": ["步骤", "推导", "=", "解", "step"], - "legal": ["⚠", "法律", "意见", "合规", "contract", "law"], - "medical": ["⚠", "就医", "医生", "症状", "诊断", "symptom"], - "general": [], -} - - -@dataclass -class QualityEvaluation: - overall_score: float - scores: Dict[str, float] = field(default_factory=dict) - needs_fallback: bool = False - reasons: List[str] = field(default_factory=list) - - def to_dict(self) -> Dict: - return { - "overall_score": round(self.overall_score, 4), - "scores": {k: round(v, 4) for k, v in self.scores.items()}, - "needs_fallback": self.needs_fallback, - "reasons": self.reasons, - } - - -class BaseJudge: - async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation: - raise NotImplementedError - - -class RuleJudge(BaseJudge): - """启发式质量评估(零依赖)。""" - - def __init__(self, fallback_threshold: float = 0.70): - self.fallback_threshold = fallback_threshold - - async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation: - scores: Dict[str, float] = {} - reasons: List[str] = [] - - # 1) 内容覆盖度:查询中的内容词有多少出现在响应里 - terms = extract_content_terms(query) - if terms: - hit = sum(1 for t in terms if t in response.lower()) - coverage = hit / len(terms) - scores["coverage"] = coverage - if coverage < 0.4: - reasons.append(f"内容覆盖度低 ({coverage:.0%})") - else: - scores["coverage"] = 1.0 - - # 2) 长度充分性 - lo, hi = _EXPECTED_LEN.get(domain, (40, 2000)) - n = len(response) - if n < lo: - scores["length"] = max(0.0, n / lo) - reasons.append(f"响应过短 ({n} 字符)") - elif n > hi: - scores["length"] = 0.8 - reasons.append(f"响应过长 ({n} 字符)") - else: - scores["length"] = 1.0 - - # 3) 领域格式检查 - hints = _DOMAIN_FORMAT_HINTS.get(domain, []) - if hints: - hit_hints = sum(1 for h in hints if h in response) - scores["format"] = min(1.0, 0.4 + 0.2 * hit_hints) - if hit_hints == 0: - reasons.append("缺少领域格式特征") - else: - scores["format"] = 1.0 - - # 4) 安全/免责提示(法律、医疗领域应有警示语) - if domain in ("legal", "medical") and ("⚠" not in response and "提示" not in response): - scores["safety"] = 0.6 - reasons.append("缺少免责提示") - else: - scores["safety"] = 1.0 - - weights = {"coverage": 0.4, "length": 0.2, "format": 0.2, "safety": 0.2} - overall = sum(scores.get(k, 0.0) * w for k, w in weights.items()) - needs = overall < self.fallback_threshold - if needs: - reasons.append("质量分低于阈值,建议升级大模型") - return QualityEvaluation( - overall_score=round(overall, 4), - scores=scores, - needs_fallback=needs, - reasons=reasons, - ) - - -class LLMJudge(BaseJudge): - """可选:LLM-as-Judge(API 后端)。""" - - def __init__(self, model: str, base_url: str, api_key: str, fallback_threshold: float = 0.70): - self.model = model - self.base_url = base_url.rstrip("/") - self.api_key = api_key - self.fallback_threshold = fallback_threshold - self._client = None - - def _get_client(self): - if self._client is None: - import httpx - self._client = httpx.AsyncClient(timeout=60.0) - return self._client - - async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation: - client = self._get_client() - prompt = ( - f"你是质量评审员。评估以下回答对查询的满足程度,输出 0-1 分(相关性/正确性/完整性)。\n" - f"查询: {query}\n领域: {domain}\n回答: {response[:2000]}\n" - f"只输出一个 0 到 1 之间的数字。" - ) - resp = await client.post( - f"{self.base_url}/chat/completions", - headers={"Authorization": f"Bearer {self.api_key}"}, - json={"model": self.model, "messages": [{"role": "user", "content": prompt}]}, - ) - resp.raise_for_status() - try: - score = float(resp.json()["choices"][0]["message"]["content"].strip()) - score = max(0.0, min(1.0, score)) - except Exception: - score = 0.5 - return QualityEvaluation( - overall_score=score, - scores={"llm_judge": score}, - needs_fallback=score < self.fallback_threshold, - ) - - -def build_judge(cfg: Dict, fallback_threshold: float = 0.70) -> BaseJudge: - """cfg 为 judge 段配置。""" - jtype = cfg.get("type", "rule") - if jtype == "rule": - return RuleJudge(fallback_threshold=fallback_threshold) - if jtype == "llm": - import os - api_key = cfg.get("api_key") or os.environ.get(cfg.get("api_key_env", "API_KEY")) - return LLMJudge( - cfg.get("model", "deepseek-chat"), - cfg.get("base_url", "https://api.deepseek.com/v1"), - api_key or "", - fallback_threshold=fallback_threshold, - ) - raise ValueError(f"未知 judge 类型: {jtype}(支持 rule | llm)") +"""质量控制器(Judge):评估专家输出,决定是否升级大模型。 + +- RuleJudge:零依赖启发式(内容覆盖度 / 长度充分性 / 领域格式 / 安全提示), + 稳定可测,适合 MVP 与离线演示。 +- LLMJudge:可选,基于 transformers 小模型或 API 的 LLM-as-Judge。 + +设计对齐实现方案:overall_score < judge_fallback_threshold -> 升级大模型。 +""" +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Dict, List + +from .experts import extract_content_terms + +# 各领域期望的响应长度范围(字符数) +_EXPECTED_LEN = { + "code": (60, 2000), + "math": (60, 2000), + "legal": (80, 3000), + "medical": (80, 3000), + "general": (40, 2000), +} + +# 领域格式检查:响应应包含的标记 +_DOMAIN_FORMAT_HINTS = { + "code": ["```", "def ", "function", "class "], + "math": ["步骤", "推导", "=", "解", "step"], + "legal": ["⚠", "法律", "意见", "合规", "contract", "law"], + "medical": ["⚠", "就医", "医生", "症状", "诊断", "symptom"], + "general": [], +} + + +@dataclass +class QualityEvaluation: + overall_score: float + scores: Dict[str, float] = field(default_factory=dict) + needs_fallback: bool = False + reasons: List[str] = field(default_factory=list) + + def to_dict(self) -> Dict: + return { + "overall_score": round(self.overall_score, 4), + "scores": {k: round(v, 4) for k, v in self.scores.items()}, + "needs_fallback": self.needs_fallback, + "reasons": self.reasons, + } + + +class BaseJudge: + async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation: + raise NotImplementedError + + +class RuleJudge(BaseJudge): + """启发式质量评估(零依赖)。""" + + def __init__(self, fallback_threshold: float = 0.70): + self.fallback_threshold = fallback_threshold + + async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation: + scores: Dict[str, float] = {} + reasons: List[str] = [] + + # 1) 内容覆盖度:查询中的内容词有多少出现在响应里 + terms = extract_content_terms(query) + if terms: + lowered = response.lower() # lowercase 一次,避免逐词重复复制长响应 + hit = sum(1 for t in terms if t in lowered) + coverage = hit / len(terms) + scores["coverage"] = coverage + if coverage < 0.4: + reasons.append(f"内容覆盖度低 ({coverage:.0%})") + else: + scores["coverage"] = 1.0 + + # 2) 长度充分性 + lo, hi = _EXPECTED_LEN.get(domain, (40, 2000)) + n = len(response) + if n < lo: + scores["length"] = max(0.0, n / lo) + reasons.append(f"响应过短 ({n} 字符)") + elif n > hi: + scores["length"] = 0.8 + reasons.append(f"响应过长 ({n} 字符)") + else: + scores["length"] = 1.0 + + # 3) 领域格式检查 + hints = _DOMAIN_FORMAT_HINTS.get(domain, []) + if hints: + hit_hints = sum(1 for h in hints if h in response) + scores["format"] = min(1.0, 0.4 + 0.2 * hit_hints) + if hit_hints == 0: + reasons.append("缺少领域格式特征") + else: + scores["format"] = 1.0 + + # 4) 安全/免责提示(法律、医疗领域应有警示语) + if domain in ("legal", "medical") and ("⚠" not in response and "提示" not in response): + scores["safety"] = 0.6 + reasons.append("缺少免责提示") + else: + scores["safety"] = 1.0 + + weights = {"coverage": 0.4, "length": 0.2, "format": 0.2, "safety": 0.2} + overall = sum(scores.get(k, 0.0) * w for k, w in weights.items()) + needs = overall < self.fallback_threshold + if needs: + reasons.append("质量分低于阈值,建议升级大模型") + return QualityEvaluation( + overall_score=round(overall, 4), + scores=scores, + needs_fallback=needs, + reasons=reasons, + ) + + +class LLMJudge(BaseJudge): + """可选:LLM-as-Judge(API 后端)。""" + + def __init__(self, model: str, base_url: str, api_key: str, fallback_threshold: float = 0.70): + from .llm_client import OpenAICompatClient + self.model = model + self.fallback_threshold = fallback_threshold + # 请求体与旧实现一致:不下发 temperature / max_tokens + self._client = OpenAICompatClient(base_url, api_key, timeout=60.0) + + async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation: + prompt = ( + f"你是质量评审员。评估以下回答对查询的满足程度,输出 0-1 分(相关性/正确性/完整性)。\n" + f"查询: {query}\n领域: {domain}\n回答: {response[:2000]}\n" + f"只输出一个 0 到 1 之间的数字。" + ) + data = await self._client.chat( + self.model, + [{"role": "user", "content": prompt}], + temperature=None, + max_tokens=None, + ) + try: + score = float(self._client.completion_text(data).strip()) + score = max(0.0, min(1.0, score)) + except Exception: + score = 0.5 + return QualityEvaluation( + overall_score=score, + scores={"llm_judge": score}, + needs_fallback=score < self.fallback_threshold, + ) + + +def build_judge(cfg: Dict, fallback_threshold: float = 0.70) -> BaseJudge: + """cfg 为 judge 段配置。""" + jtype = cfg.get("type", "rule") + if jtype == "rule": + return RuleJudge(fallback_threshold=fallback_threshold) + if jtype == "llm": + from .config import get_api_key + api_key = cfg.get("api_key") or get_api_key(cfg) + return LLMJudge( + cfg.get("model", "deepseek-chat"), + cfg.get("base_url", "https://api.deepseek.com/v1"), + api_key or "", + fallback_threshold=fallback_threshold, + ) + raise ValueError(f"未知 judge 类型: {jtype}(支持 rule | llm)") diff --git a/router_system/llm_client.py b/router_system/llm_client.py new file mode 100644 index 0000000..7e9aee3 --- /dev/null +++ b/router_system/llm_client.py @@ -0,0 +1,69 @@ +"""OpenAI 兼容 Chat Completions 共享客户端(架构去重)。 + +此前 experts.APIExpert / judge.LLMJudge / fallback.APIFallback 各自复制一份 +"懒建 AsyncClient + POST /chat/completions + choices/usage 解析"的等价逻辑 +(超时、温度、token 上限各自维护)。本模块将其统一为单一组件: + +- 懒建 httpx.AsyncClient 并长期复用(与仓库内其他客户端一致) +- temperature / max_tokens 传 None 时不出现在请求体(保持各后端原请求形状) +- 密钥解析统一走 config.get_api_key(此前 experts 的默认环境名与另两处不一致) + +依赖说明:httpx 仅 api 类型后端需要,import 延迟到连接创建时, +router_system 核心保持零第三方依赖约束不变。 +""" +from __future__ import annotations + +from typing import Any, Dict, List, Optional + + +class OpenAICompatClient: + """OpenAI 兼容 /chat/completions 客户端(懒建连接、跨调用复用)。""" + + def __init__(self, base_url: str, api_key: str, timeout: float = 60.0): + self.base_url = base_url.rstrip("/") + self.api_key = api_key + self._timeout = timeout + self._client = None + + def _get_client(self): + if self._client is None: + import httpx + self._client = httpx.AsyncClient(timeout=self._timeout) + return self._client + + async def chat( + self, + model: str, + messages: List[Dict[str, str]], + temperature: Optional[float] = 0.2, + max_tokens: Optional[int] = 1024, + ) -> Dict[str, Any]: + """调用 /chat/completions,返回完整响应 JSON(含 usage)。 + + HTTP 非 2xx 时抛 httpx.HTTPStatusError;temperature/max_tokens + 为 None 时对应字段不下发(与旧实现中 LLMJudge 的请求体一致)。 + """ + payload: Dict[str, Any] = {"model": model, "messages": messages} + if temperature is not None: + payload["temperature"] = temperature + if max_tokens is not None: + payload["max_tokens"] = max_tokens + client = self._get_client() + resp = await client.post( + f"{self.base_url}/chat/completions", + headers={"Authorization": f"Bearer {self.api_key}"}, + json=payload, + ) + resp.raise_for_status() + return resp.json() + + @staticmethod + def completion_text(data: Dict[str, Any]) -> str: + """提取首条回复文本。""" + return data["choices"][0]["message"]["content"] + + @staticmethod + def completion_tokens(data: Dict[str, Any], text: str) -> int: + """提取 completion token 数;usage 缺失时按 ~2.2 字符/token 估算。""" + usage = data.get("usage", {}) + return usage.get("completion_tokens", int(len(text) / 2.2)) diff --git a/tests/conftest.py b/tests/conftest.py index 05e476e..eb32dbb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -10,5 +10,5 @@ from router_system.router import build_router @pytest.fixture() def router(): - """????????? mock ?????????????""" + """构建全 mock 链路的 Router 实例(离线可跑,供各测试复用)。""" return build_router() diff --git a/tests/test_cache.py b/tests/test_cache.py index 72b64f0..b373ea9 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -49,9 +49,9 @@ def test_exact_hit(): def test_semantic_hit(): c = RouterCache(semantic_enabled=True, similarity_threshold=0.5) - c.put("?python?????", {"response": "code", "domain": "code"}) - # ?????????? L2 - hit = c.get("?python????????") + c.put("用 python 实现快速排序", {"response": "code", "domain": "code"}) + # 相似但不完全相同的查询应命中 L2 语义缓存 + hit = c.get("用 python 实现快速排序的算法") assert hit is not None assert hit[0] == "semantic" @@ -60,7 +60,7 @@ def test_promote_to_exact(): c = RouterCache(promote_frequency=3) result = {"response": "x", "domain": "general"} c.put("query", result) - # ?????? 3 ? ? ??????? + # 语义命中累计 3 次后提升为精确缓存 for _ in range(3): hit = c.get("query") assert hit is not None diff --git a/tests/test_gateway.py b/tests/test_gateway.py index 39ca51f..18db4f3 100644 --- a/tests/test_gateway.py +++ b/tests/test_gateway.py @@ -1,4 +1,4 @@ -"""???????? fastapi + httpx??""" +"""网关集成测试(需 fastapi + httpx,链路全 mock)。""" import pytest pytest.importorskip("fastapi") @@ -23,7 +23,7 @@ def test_health(client): def test_chat(client): - resp = client.post("/chat", json={"query": "? Python ???????"}) + resp = client.post("/chat", json={"query": "用 Python 写一个快速排序函数"}) assert resp.status_code == 200 data = resp.json() assert data["response"]