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:
+36
-19
@@ -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
@@ -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
@@ -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
@@ -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 Completions(DeepSeek / 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 Completions(DeepSeek / 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 Key(env: {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 Key(env: {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
@@ -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
@@ -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-Judge(API 后端)。"""
|
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
|
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)")␍
|
|
||||||
|
|||||||
@@ -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))
|
||||||
+1
-1
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
Reference in New Issue
Block a user