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),
) )
+6 -2
View File
@@ -25,6 +25,10 @@ _MEDIUM_MARKERS = [
"translate", "fix", "explain", "describe", "list", "translate", "fix", "explain", "describe", "list",
] ]
# 预编译正则(模块级一次,避免每次调用走 re 内部缓存查找)
_RE_CODE_EXPR = re.compile(r"\b(def|class|function|import)\b")
_RE_ARITH_EXPR = re.compile(r"[0-9]+\s*[+\-*/^=]\s*[0-9xya-z]")
def estimate_difficulty(query: str) -> Tuple[str, float]: def estimate_difficulty(query: str) -> Tuple[str, float]:
"""返回 (difficulty, score)score 属于 [0,1]。""" """返回 (difficulty, score)score 属于 [0,1]。"""
@@ -39,9 +43,9 @@ def estimate_difficulty(query: str) -> Tuple[str, float]:
score += min(0.30, medium_hits * 0.08) # 一般指令贡献 score += min(0.30, medium_hits * 0.08) # 一般指令贡献
# 额外:代码 / 数学表达式(多步骤信号) # 额外:代码 / 数学表达式(多步骤信号)
if "```" in query or re.search(r"\b(def|class|function|import)\b", q): if "```" in query or _RE_CODE_EXPR.search(q):
score += 0.15 score += 0.15
if re.search(r"[0-9]+\s*[+\-*/^=]\s*[0-9xya-z]", q): if _RE_ARITH_EXPR.search(q):
score += 0.15 score += 0.15
score = max(0.0, min(1.0, score)) score = max(0.0, min(1.0, score))
+25 -34
View File
@@ -10,8 +10,11 @@ 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 .config import get_api_key
from .llm_client import OpenAICompatClient
from .models import ExpertResponse from .models import ExpertResponse
# 按模型规模粗估的相对成本($ / 1M output tokens,近似) # 按模型规模粗估的相对成本($ / 1M output tokens,近似)
@@ -38,8 +41,8 @@ def _cost_for(model_name: str, default: str = "1b") -> float:
return MODEL_COST_EST.get(default, 0.1) return MODEL_COST_EST.get(default, 0.1)
def extract_content_terms(query: str) -> List[str]: @lru_cache(maxsize=2048)
"""抽取查询中的"内容词"(中文词/英文单词),用于 Judge 覆盖度与 Mock 回显。""" def _extract_content_terms_cached(query: str) -> Tuple[str, ...]:
q = query.lower() q = query.lower()
terms: List[str] = [] terms: List[str] = []
# 英文单词(>=2 字符) # 英文单词(>=2 字符)
@@ -50,7 +53,16 @@ def extract_content_terms(query: str) -> List[str]:
cn = re.findall(r"[\u4e00-\u9fff]{2,8}", q) cn = re.findall(r"[\u4e00-\u9fff]{2,8}", q)
for c in cn: for c in cn:
terms.append(c) terms.append(c)
return terms return tuple(terms)
def extract_content_terms(query: str) -> Tuple[str, ...]:
"""抽取查询中的"内容词"(中文词/英文单词),用于 Judge 覆盖度与 Mock 回显。
纯函数,但 Router 主链路中专家生成与 Judge 评估会对同一查询各调一次,
故以 lru_cache 去重;返回不可变 tuple(调用方只读)。
"""
return _extract_content_terms_cached(query)
_STOPWORDS_EN = { _STOPWORDS_EN = {
@@ -202,33 +214,17 @@ class APIExpert(Expert):
self.name = name self.name = name
self.domain = domain self.domain = domain
self.model = model self.model = model
self.base_url = base_url.rstrip("/") self._client = OpenAICompatClient(base_url, api_key, timeout=60.0)
self.api_key = api_key
self._client = None
def _get_client(self):
if self._client is None:
import httpx
self._client = httpx.AsyncClient(timeout=60.0)
return self._client
async def generate(self, query: str, difficulty: str) -> ExpertResponse: async def generate(self, query: str, difficulty: str) -> ExpertResponse:
client = self._get_client() data = await self._client.chat(
resp = await client.post( self.model,
f"{self.base_url}/chat/completions", [{"role": "user", "content": query}],
headers={"Authorization": f"Bearer {self.api_key}"}, temperature=0.2,
json={ max_tokens=1024,
"model": self.model,
"messages": [{"role": "user", "content": query}],
"temperature": 0.2,
"max_tokens": 1024,
},
) )
resp.raise_for_status() body = self._client.completion_text(data)
data = resp.json() tokens = self._client.completion_tokens(data, body)
body = data["choices"][0]["message"]["content"]
usage = data.get("usage", {})
tokens = usage.get("completion_tokens", int(len(body) / 2.2))
return ExpertResponse( return ExpertResponse(
text=body, text=body,
model_used=self.model, model_used=self.model,
@@ -249,18 +245,13 @@ def build_expert(domain: str, cfg: Dict) -> Expert:
return HFExpert(name, domain, model) return HFExpert(name, domain, model)
if etype == "api": if etype == "api":
base_url = cfg.get("base_url", "https://api.deepseek.com/v1") base_url = cfg.get("base_url", "https://api.deepseek.com/v1")
api_key = cfg.get("api_key") or _env(cfg.get("api_key_env", "")) api_key = cfg.get("api_key") or get_api_key(cfg)
if not api_key: if not api_key:
raise RuntimeError(f"APIExpert({domain}) 缺少 API Keyenv: {cfg.get('api_key_env')}") raise RuntimeError(f"APIExpert({domain}) 缺少 API Keyenv: {cfg.get('api_key_env')}")
return APIExpert(name, domain, model, base_url, api_key) return APIExpert(name, domain, model, base_url, api_key)
raise ValueError(f"未知专家后端类型: {etype}(支持 mock | hf | api") raise ValueError(f"未知专家后端类型: {etype}(支持 mock | hf | api")
def _env(name: str) -> Optional[str]:
import os
return os.environ.get(name) if name else None
def build_expert_pool(experts_cfg: Dict[str, Dict], domains: List[str]) -> Dict[str, Expert]: def build_expert_pool(experts_cfg: Dict[str, Dict], domains: List[str]) -> Dict[str, Expert]:
"""构建完整专家池。""" """构建完整专家池。"""
pool: Dict[str, Expert] = {} pool: Dict[str, Expert] = {}
+12 -27
View File
@@ -2,8 +2,10 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
from typing import Dict, Optional from typing import Dict
from .config import get_api_key
from .llm_client import OpenAICompatClient
from .models import ExpertResponse from .models import ExpertResponse
@@ -45,34 +47,18 @@ class APIFallback(FallbackProvider):
def __init__(self, model: str, base_url: str, api_key: str): def __init__(self, model: str, base_url: str, api_key: str):
self.model = model self.model = model
self.base_url = base_url.rstrip("/")
self.api_key = api_key
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):
if self._client is None:
import httpx
self._client = httpx.AsyncClient(timeout=90.0)
return self._client
async def generate(self, query: str) -> ExpertResponse: async def generate(self, query: str) -> ExpertResponse:
client = self._get_client() data = await self._client.chat(
resp = await client.post( self.model,
f"{self.base_url}/chat/completions", [{"role": "user", "content": query}],
headers={"Authorization": f"Bearer {self.api_key}"}, temperature=0.3,
json={ max_tokens=2048,
"model": self.model,
"messages": [{"role": "user", "content": query}],
"temperature": 0.3,
"max_tokens": 2048,
},
) )
resp.raise_for_status() body = self._client.completion_text(data)
data = resp.json() tokens = self._client.completion_tokens(data, body)
body = data["choices"][0]["message"]["content"]
usage = data.get("usage", {})
tokens = usage.get("completion_tokens", int(len(body) / 2.2))
return ExpertResponse( return ExpertResponse(
text=body, text=body,
model_used=self.model, model_used=self.model,
@@ -89,8 +75,7 @@ def build_fallback(cfg: Dict) -> FallbackProvider:
if ftype == "mock": if ftype == "mock":
return MockFallback(model=model) return MockFallback(model=model)
if ftype == "api": if ftype == "api":
import os api_key = cfg.get("api_key") or get_api_key(cfg)
api_key = cfg.get("api_key") or os.environ.get(cfg.get("api_key_env", "API_KEY"))
if not api_key: if not api_key:
raise RuntimeError( raise RuntimeError(
f"APIFallback 缺少 API Key:请设置环境变量 {cfg.get('api_key_env')} 或配置 api_key" f"APIFallback 缺少 API Key:请设置环境变量 {cfg.get('api_key_env')} 或配置 api_key"
+13 -19
View File
@@ -66,7 +66,8 @@ class RuleJudge(BaseJudge):
# 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 一次,避免逐词重复复制长响应
hit = sum(1 for t in terms if t in lowered)
coverage = hit / len(terms) coverage = hit / len(terms)
scores["coverage"] = coverage scores["coverage"] = coverage
if coverage < 0.4: if coverage < 0.4:
@@ -120,33 +121,26 @@ class LLMJudge(BaseJudge):
"""可选:LLM-as-JudgeAPI 后端)。""" """可选:LLM-as-JudgeAPI 后端)。"""
def __init__(self, model: str, base_url: str, api_key: str, fallback_threshold: float = 0.70): def __init__(self, model: str, base_url: str, api_key: str, fallback_threshold: float = 0.70):
from .llm_client import OpenAICompatClient
self.model = model self.model = model
self.base_url = base_url.rstrip("/")
self.api_key = api_key
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:
import httpx
self._client = httpx.AsyncClient(timeout=60.0)
return self._client
async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation: async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation:
client = self._get_client()
prompt = ( prompt = (
f"你是质量评审员。评估以下回答对查询的满足程度,输出 0-1 分(相关性/正确性/完整性)。\n" f"你是质量评审员。评估以下回答对查询的满足程度,输出 0-1 分(相关性/正确性/完整性)。\n"
f"查询: {query}\n领域: {domain}\n回答: {response[:2000]}\n" f"查询: {query}\n领域: {domain}\n回答: {response[:2000]}\n"
f"只输出一个 0 到 1 之间的数字。" f"只输出一个 0 到 1 之间的数字。"
) )
resp = await client.post( data = await self._client.chat(
f"{self.base_url}/chat/completions", self.model,
headers={"Authorization": f"Bearer {self.api_key}"}, [{"role": "user", "content": prompt}],
json={"model": self.model, "messages": [{"role": "user", "content": prompt}]}, temperature=None,
max_tokens=None,
) )
resp.raise_for_status()
try: try:
score = float(resp.json()["choices"][0]["message"]["content"].strip()) score = float(self._client.completion_text(data).strip())
score = max(0.0, min(1.0, score)) score = max(0.0, min(1.0, score))
except Exception: except Exception:
score = 0.5 score = 0.5
@@ -163,8 +157,8 @@ def build_judge(cfg: Dict, fallback_threshold: float = 0.70) -> BaseJudge:
if jtype == "rule": if jtype == "rule":
return RuleJudge(fallback_threshold=fallback_threshold) return RuleJudge(fallback_threshold=fallback_threshold)
if jtype == "llm": if jtype == "llm":
import os from .config import get_api_key
api_key = cfg.get("api_key") or os.environ.get(cfg.get("api_key_env", "API_KEY")) api_key = cfg.get("api_key") or get_api_key(cfg)
return LLMJudge( return LLMJudge(
cfg.get("model", "deepseek-chat"), cfg.get("model", "deepseek-chat"),
cfg.get("base_url", "https://api.deepseek.com/v1"), cfg.get("base_url", "https://api.deepseek.com/v1"),
+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"]