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(完全相同查询)时提前终止扫描(余弦相似度上界,不可能更优)
|
||||
- 语义条目容器为 OrderedDict:提升/淘汰均为 O(1)(原 list pop 是 O(n) 搬移),
|
||||
且按 query 天然去重(原实现同一 query 可重复追加,重复条目白白参与扫描)
|
||||
- n-gram 向量经 lru_cache 复用:同一次未命中请求里 get 与紧随的 put
|
||||
不再重复做两遍分词+计数
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
|
||||
@@ -52,6 +58,17 @@ def _dot(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float:
|
||||
return sum(v * vec_b.get(k, 0.0) for k, v in vec_a.items())
|
||||
|
||||
|
||||
@lru_cache(maxsize=1024)
|
||||
def _embed(text: str) -> Tuple[Dict[str, float], float]:
|
||||
"""n-gram 向量 + 范数(带缓存)。
|
||||
|
||||
纯函数且同一次未命中请求会在 get/put 各用一次,缓存可省一遍分词。
|
||||
契约:返回的 dict 仅供只读(内部无任何修改点),调用方不得改写。
|
||||
"""
|
||||
vec = _tf_vector(_ngrams(text))
|
||||
return vec, _norm(vec)
|
||||
|
||||
|
||||
class RouterCache:
|
||||
"""L1 精确缓存 + L2 语义缓存。"""
|
||||
|
||||
@@ -63,7 +80,8 @@ class RouterCache:
|
||||
self.max_exact = max_exact
|
||||
self.max_semantic = max_semantic
|
||||
self._exact: Dict[str, CacheEntry] = {}
|
||||
self._semantic: List[Tuple[str, CacheEntry]] = [] # (query, entry)
|
||||
# (query -> entry),插入序即淘汰序;len 兼容原 list 断言
|
||||
self._semantic: "OrderedDict[str, CacheEntry]" = OrderedDict()
|
||||
self._sem_vecs: Dict[str, Dict[str, float]] = {}
|
||||
self._sem_norms: Dict[str, float] = {} # 预计算范数,避免查询期重算
|
||||
self.hits = {"exact": 0, "semantic": 0}
|
||||
@@ -78,57 +96,56 @@ class RouterCache:
|
||||
return ("exact", entry.result)
|
||||
|
||||
if self.semantic_enabled:
|
||||
q_vec = _tf_vector(_ngrams(query))
|
||||
q_norm = _norm(q_vec)
|
||||
q_vec, q_norm = _embed(query)
|
||||
best_sim = 0.0
|
||||
best_idx = -1
|
||||
best_q = ""
|
||||
if q_norm > 0.0:
|
||||
# 单遍扫描:同时跟踪最优相似度与条目位置
|
||||
for i, (q, _e) in enumerate(self._semantic):
|
||||
# 单遍扫描(插入序):同时跟踪最优条目,sim=1.0 提前终止
|
||||
for q, _e in self._semantic.items():
|
||||
n_q = self._sem_norms.get(q, 0.0)
|
||||
if n_q <= 0.0:
|
||||
continue
|
||||
sim = _dot(q_vec, self._sem_vecs.get(q, {})) / (q_norm * n_q)
|
||||
if sim > best_sim:
|
||||
best_sim = sim
|
||||
best_idx = i
|
||||
best_q = q
|
||||
if sim >= 1.0:
|
||||
break # 余弦相似度上界:完全相同查询,提前终止
|
||||
if best_idx >= 0 and best_sim >= self.similarity_threshold:
|
||||
best_q, best_entry = self._semantic[best_idx]
|
||||
if best_q and best_sim >= self.similarity_threshold:
|
||||
best_entry = self._semantic[best_q]
|
||||
# 完全相同查询(相似度=1.0)计为 exact 命中
|
||||
is_exact = best_sim >= 0.999
|
||||
level = "exact" if is_exact else "semantic"
|
||||
self.hits[level] += 1
|
||||
self._bump_semantic(best_idx, best_q, best_entry)
|
||||
self._bump_semantic(best_q, best_entry)
|
||||
return (level, best_entry.result)
|
||||
|
||||
self.misses += 1
|
||||
return None
|
||||
|
||||
def _bump_semantic(self, idx: int, query: str, entry: CacheEntry):
|
||||
"""语义命中:累计命中次数,达到阈值提升为精确缓存(O(1),无需二次查找)。"""
|
||||
def _bump_semantic(self, query: str, entry: CacheEntry):
|
||||
"""语义命中:累计命中次数,达到阈值提升为精确缓存(OrderedDict 删除 O(1))。"""
|
||||
entry.hits += 1
|
||||
if entry.hits >= self.promote_frequency:
|
||||
self._exact[query] = entry
|
||||
self._semantic.pop(idx)
|
||||
del self._semantic[query]
|
||||
self._sem_vecs.pop(query, None)
|
||||
self._sem_norms.pop(query, None)
|
||||
|
||||
# ---- 写入 ----
|
||||
def put(self, query: str, result: Dict[str, Any]):
|
||||
if query in self._exact:
|
||||
return
|
||||
if query in self._exact or query in self._semantic:
|
||||
return # 已缓存(语义区按 query 天然去重,避免重复条目参与扫描)
|
||||
entry = CacheEntry(result=result)
|
||||
if self.semantic_enabled:
|
||||
if len(self._semantic) >= self.max_semantic:
|
||||
old_q, _ = self._semantic.pop(0)
|
||||
old_q, _old = self._semantic.popitem(last=False) # O(1) 淘汰最旧
|
||||
self._sem_vecs.pop(old_q, None)
|
||||
self._sem_norms.pop(old_q, None)
|
||||
self._semantic.append((query, entry))
|
||||
vec = _tf_vector(_ngrams(query))
|
||||
self._semantic[query] = entry
|
||||
vec, norm = _embed(query) # 与 get miss 共享同一份缓存向量
|
||||
self._sem_vecs[query] = vec
|
||||
self._sem_norms[query] = _norm(vec)
|
||||
self._sem_norms[query] = norm
|
||||
else:
|
||||
self._exact[query] = entry
|
||||
if len(self._exact) > self.max_exact:
|
||||
|
||||
+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),
|
||||
("import", 0.7), ("loop", 0.7), ("recursion", 0.8), ("refactor", 0.8),
|
||||
("deploy", 0.8), ("docker", 0.8), ("kubernetes", 0.8),
|
||||
("async", 0.7), ("flask", 0.7), ("django", 0.7), ("api", 0.7),
|
||||
("async", 0.7), ("flask", 0.7), ("django", 0.7),
|
||||
("索引", 0.8), ("优化", 0.7), ("查询", 0.6),
|
||||
],
|
||||
"math": [
|
||||
@@ -98,24 +98,29 @@ class RuleClassifier(BaseClassifier):
|
||||
self.confidence_floor = confidence_floor
|
||||
self.rules = DOMAIN_RULES
|
||||
|
||||
def _score(self, query: str) -> Tuple[Dict[str, float], Dict[str, List[str]]]:
|
||||
q = query.lower()
|
||||
def _score(self, q: str) -> Dict[str, float]:
|
||||
"""对已 lowercase 的查询按领域规则打分,返回有命中的领域分数。
|
||||
|
||||
命中词列表只对最终胜出领域有意义(classify 的唯一消费点),
|
||||
故不在各领域上白建 list——胜出后由 _matched_rules 单独收集。
|
||||
"""
|
||||
scores: Dict[str, float] = {}
|
||||
matched: Dict[str, List[str]] = {}
|
||||
for domain, rules in self.rules.items():
|
||||
s = 0.0
|
||||
hits = []
|
||||
for kw, w in rules:
|
||||
if kw in q:
|
||||
s += w
|
||||
hits.append(kw)
|
||||
if s > 0:
|
||||
scores[domain] = s
|
||||
matched[domain] = hits
|
||||
return scores, matched
|
||||
return scores
|
||||
|
||||
def _matched_rules(self, q: str, domain: str) -> List[str]:
|
||||
"""收集指定领域命中的关键词(保持登记顺序)。"""
|
||||
return [kw for kw, _w in self.rules.get(domain, ()) if kw in q]
|
||||
|
||||
def classify(self, query: str) -> Classification:
|
||||
raw, matched = self._score(query)
|
||||
q = query.lower()
|
||||
raw = self._score(q)
|
||||
if not raw:
|
||||
# 完全无命中 -> general,低置信度
|
||||
diff, ds = estimate_difficulty(query)
|
||||
@@ -150,7 +155,7 @@ class RuleClassifier(BaseClassifier):
|
||||
difficulty=diff,
|
||||
difficulty_score=ds,
|
||||
raw_scores={k: round(v, 3) for k, v in raw.items()},
|
||||
matched_rules=matched.get(best_domain, []),
|
||||
matched_rules=self._matched_rules(q, best_domain),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -25,6 +25,10 @@ _MEDIUM_MARKERS = [
|
||||
"translate", "fix", "explain", "describe", "list",
|
||||
]
|
||||
|
||||
# 预编译正则(模块级一次,避免每次调用走 re 内部缓存查找)
|
||||
_RE_CODE_EXPR = re.compile(r"\b(def|class|function|import)\b")
|
||||
_RE_ARITH_EXPR = re.compile(r"[0-9]+\s*[+\-*/^=]\s*[0-9xya-z]")
|
||||
|
||||
|
||||
def estimate_difficulty(query: str) -> Tuple[str, float]:
|
||||
"""返回 (difficulty, score),score 属于 [0,1]。"""
|
||||
@@ -39,9 +43,9 @@ def estimate_difficulty(query: str) -> Tuple[str, float]:
|
||||
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
|
||||
if re.search(r"[0-9]+\s*[+\-*/^=]\s*[0-9xya-z]", q):
|
||||
if _RE_ARITH_EXPR.search(q):
|
||||
score += 0.15
|
||||
|
||||
score = max(0.0, min(1.0, score))
|
||||
|
||||
+25
-34
@@ -10,8 +10,11 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
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
|
||||
|
||||
# 按模型规模粗估的相对成本($ / 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)
|
||||
|
||||
|
||||
def extract_content_terms(query: str) -> List[str]:
|
||||
"""抽取查询中的"内容词"(中文词/英文单词),用于 Judge 覆盖度与 Mock 回显。"""
|
||||
@lru_cache(maxsize=2048)
|
||||
def _extract_content_terms_cached(query: str) -> Tuple[str, ...]:
|
||||
q = query.lower()
|
||||
terms: List[str] = []
|
||||
# 英文单词(>=2 字符)
|
||||
@@ -50,7 +53,16 @@ def extract_content_terms(query: str) -> List[str]:
|
||||
cn = re.findall(r"[\u4e00-\u9fff]{2,8}", q)
|
||||
for c in cn:
|
||||
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 = {
|
||||
@@ -202,33 +214,17 @@ class APIExpert(Expert):
|
||||
self.name = name
|
||||
self.domain = domain
|
||||
self.model = model
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self._client = None
|
||||
|
||||
def _get_client(self):
|
||||
if self._client is None:
|
||||
import httpx
|
||||
self._client = httpx.AsyncClient(timeout=60.0)
|
||||
return self._client
|
||||
self._client = OpenAICompatClient(base_url, api_key, timeout=60.0)
|
||||
|
||||
async def generate(self, query: str, difficulty: str) -> ExpertResponse:
|
||||
client = self._get_client()
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {self.api_key}"},
|
||||
json={
|
||||
"model": self.model,
|
||||
"messages": [{"role": "user", "content": query}],
|
||||
"temperature": 0.2,
|
||||
"max_tokens": 1024,
|
||||
},
|
||||
data = await self._client.chat(
|
||||
self.model,
|
||||
[{"role": "user", "content": query}],
|
||||
temperature=0.2,
|
||||
max_tokens=1024,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
body = data["choices"][0]["message"]["content"]
|
||||
usage = data.get("usage", {})
|
||||
tokens = usage.get("completion_tokens", int(len(body) / 2.2))
|
||||
body = self._client.completion_text(data)
|
||||
tokens = self._client.completion_tokens(data, body)
|
||||
return ExpertResponse(
|
||||
text=body,
|
||||
model_used=self.model,
|
||||
@@ -249,18 +245,13 @@ def build_expert(domain: str, cfg: Dict) -> Expert:
|
||||
return HFExpert(name, domain, model)
|
||||
if etype == "api":
|
||||
base_url = cfg.get("base_url", "https://api.deepseek.com/v1")
|
||||
api_key = cfg.get("api_key") or _env(cfg.get("api_key_env", ""))
|
||||
api_key = cfg.get("api_key") or get_api_key(cfg)
|
||||
if not api_key:
|
||||
raise RuntimeError(f"APIExpert({domain}) 缺少 API Key(env: {cfg.get('api_key_env')})")
|
||||
return APIExpert(name, domain, model, base_url, api_key)
|
||||
raise ValueError(f"未知专家后端类型: {etype}(支持 mock | hf | api)")
|
||||
|
||||
|
||||
def _env(name: str) -> Optional[str]:
|
||||
import os
|
||||
return os.environ.get(name) if name else None
|
||||
|
||||
|
||||
def build_expert_pool(experts_cfg: Dict[str, Dict], domains: List[str]) -> Dict[str, Expert]:
|
||||
"""构建完整专家池。"""
|
||||
pool: Dict[str, Expert] = {}
|
||||
|
||||
+12
-27
@@ -2,8 +2,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -45,34 +47,18 @@ class APIFallback(FallbackProvider):
|
||||
|
||||
def __init__(self, model: str, base_url: str, api_key: str):
|
||||
self.model = model
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.name = f"fallback-{model}"
|
||||
self._client = None
|
||||
|
||||
def _get_client(self):
|
||||
if self._client is None:
|
||||
import httpx
|
||||
self._client = httpx.AsyncClient(timeout=90.0)
|
||||
return self._client
|
||||
self._client = OpenAICompatClient(base_url, api_key, timeout=90.0)
|
||||
|
||||
async def generate(self, query: str) -> ExpertResponse:
|
||||
client = self._get_client()
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {self.api_key}"},
|
||||
json={
|
||||
"model": self.model,
|
||||
"messages": [{"role": "user", "content": query}],
|
||||
"temperature": 0.3,
|
||||
"max_tokens": 2048,
|
||||
},
|
||||
data = await self._client.chat(
|
||||
self.model,
|
||||
[{"role": "user", "content": query}],
|
||||
temperature=0.3,
|
||||
max_tokens=2048,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
body = data["choices"][0]["message"]["content"]
|
||||
usage = data.get("usage", {})
|
||||
tokens = usage.get("completion_tokens", int(len(body) / 2.2))
|
||||
body = self._client.completion_text(data)
|
||||
tokens = self._client.completion_tokens(data, body)
|
||||
return ExpertResponse(
|
||||
text=body,
|
||||
model_used=self.model,
|
||||
@@ -89,8 +75,7 @@ def build_fallback(cfg: Dict) -> FallbackProvider:
|
||||
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"))
|
||||
api_key = cfg.get("api_key") or get_api_key(cfg)
|
||||
if not api_key:
|
||||
raise RuntimeError(
|
||||
f"APIFallback 缺少 API Key:请设置环境变量 {cfg.get('api_key_env')} 或配置 api_key"
|
||||
|
||||
+13
-19
@@ -66,7 +66,8 @@ class RuleJudge(BaseJudge):
|
||||
# 1) 内容覆盖度:查询中的内容词有多少出现在响应里
|
||||
terms = extract_content_terms(query)
|
||||
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)
|
||||
scores["coverage"] = coverage
|
||||
if coverage < 0.4:
|
||||
@@ -120,33 +121,26 @@ class LLMJudge(BaseJudge):
|
||||
"""可选:LLM-as-Judge(API 后端)。"""
|
||||
|
||||
def __init__(self, model: str, base_url: str, api_key: str, fallback_threshold: float = 0.70):
|
||||
from .llm_client import OpenAICompatClient
|
||||
self.model = model
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.fallback_threshold = fallback_threshold
|
||||
self._client = None
|
||||
|
||||
def _get_client(self):
|
||||
if self._client is None:
|
||||
import httpx
|
||||
self._client = httpx.AsyncClient(timeout=60.0)
|
||||
return self._client
|
||||
# 请求体与旧实现一致:不下发 temperature / max_tokens
|
||||
self._client = OpenAICompatClient(base_url, api_key, timeout=60.0)
|
||||
|
||||
async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation:
|
||||
client = self._get_client()
|
||||
prompt = (
|
||||
f"你是质量评审员。评估以下回答对查询的满足程度,输出 0-1 分(相关性/正确性/完整性)。\n"
|
||||
f"查询: {query}\n领域: {domain}\n回答: {response[:2000]}\n"
|
||||
f"只输出一个 0 到 1 之间的数字。"
|
||||
)
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {self.api_key}"},
|
||||
json={"model": self.model, "messages": [{"role": "user", "content": prompt}]},
|
||||
data = await self._client.chat(
|
||||
self.model,
|
||||
[{"role": "user", "content": prompt}],
|
||||
temperature=None,
|
||||
max_tokens=None,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
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))
|
||||
except Exception:
|
||||
score = 0.5
|
||||
@@ -163,8 +157,8 @@ def build_judge(cfg: Dict, fallback_threshold: float = 0.70) -> BaseJudge:
|
||||
if jtype == "rule":
|
||||
return RuleJudge(fallback_threshold=fallback_threshold)
|
||||
if jtype == "llm":
|
||||
import os
|
||||
api_key = cfg.get("api_key") or os.environ.get(cfg.get("api_key_env", "API_KEY"))
|
||||
from .config import get_api_key
|
||||
api_key = cfg.get("api_key") or get_api_key(cfg)
|
||||
return LLMJudge(
|
||||
cfg.get("model", "deepseek-chat"),
|
||||
cfg.get("base_url", "https://api.deepseek.com/v1"),
|
||||
|
||||
@@ -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()
|
||||
def router():
|
||||
"""????????? mock ?????????????"""
|
||||
"""构建全 mock 链路的 Router 实例(离线可跑,供各测试复用)。"""
|
||||
return build_router()
|
||||
|
||||
+4
-4
@@ -49,9 +49,9 @@ def test_exact_hit():
|
||||
|
||||
def test_semantic_hit():
|
||||
c = RouterCache(semantic_enabled=True, similarity_threshold=0.5)
|
||||
c.put("?python?????", {"response": "code", "domain": "code"})
|
||||
# ?????????? L2
|
||||
hit = c.get("?python????????")
|
||||
c.put("用 python 实现快速排序", {"response": "code", "domain": "code"})
|
||||
# 相似但不完全相同的查询应命中 L2 语义缓存
|
||||
hit = c.get("用 python 实现快速排序的算法")
|
||||
assert hit is not None
|
||||
assert hit[0] == "semantic"
|
||||
|
||||
@@ -60,7 +60,7 @@ def test_promote_to_exact():
|
||||
c = RouterCache(promote_frequency=3)
|
||||
result = {"response": "x", "domain": "general"}
|
||||
c.put("query", result)
|
||||
# ?????? 3 ? ? ???????
|
||||
# 语义命中累计 3 次后提升为精确缓存
|
||||
for _ in range(3):
|
||||
hit = c.get("query")
|
||||
assert hit is not None
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""???????? fastapi + httpx??"""
|
||||
"""网关集成测试(需 fastapi + httpx,链路全 mock)。"""
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("fastapi")
|
||||
@@ -23,7 +23,7 @@ def test_health(client):
|
||||
|
||||
|
||||
def test_chat(client):
|
||||
resp = client.post("/chat", json={"query": "? Python ???????"})
|
||||
resp = client.post("/chat", json={"query": "用 Python 写一个快速排序函数"})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["response"]
|
||||
|
||||
Reference in New Issue
Block a user