feat(v1): 架构与算法优化二轮——共享 LLM 客户端去重、语义缓存 O(1) 提升/淘汰、分类器热路径清理

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