"""两阶段路由缓存(对齐实现方案): - L1 精确缓存:完全相同的查询 -> 直接命中 - L2 语义缓存:字符 n-gram 余弦相似度(零依赖)-> 相似查询命中 - 命中 N 次(promote_frequency)后提升为精确缓存 说明:语义缓存中的"完全相同查询"(相似度=1.0)直接计为 exact 命中; 高频语义命中会提升为 O(1) 的精确缓存条目。 只缓存"未升级"的结果(升级路径每次都走大模型,不缓存,避免陈旧)。 性能设计(2026-09 优化): - 每条语义缓存条目在写入时预计算并缓存向量范数,查询时免重复计算(原来每对比较都重算) - 语义查找单遍完成:扫描即跟踪最优条目与命中计数,命中后不再二次线性查找 - 相似度达到 1.0(完全相同查询)时提前终止扫描(余弦相似度上界,不可能更优) """ from __future__ import annotations import re from dataclasses import dataclass from typing import Any, Dict, List, Optional, Tuple @dataclass class CacheEntry: result: Dict[str, Any] hits: int = 1 def _ngrams(text: str, n: int = 3) -> List[str]: """字符 n-gram(去空白、小写),用于轻量语义相似度。""" cleaned = re.sub(r"\s+", "", text.lower()) if len(cleaned) < n: return [cleaned] return [cleaned[i:i + n] for i in range(len(cleaned) - n + 1)] def _tf_vector(grams: List[str]) -> Dict[str, float]: vec: Dict[str, float] = {} for g in grams: vec[g] = vec.get(g, 0.0) + 1.0 return vec def _norm(vec: Dict[str, float]) -> float: return sum(v * v for v in vec.values()) ** 0.5 def _dot(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float: """点积:遍历较小的一方,另一侧用 get 兜底。""" if len(vec_a) > len(vec_b): vec_a, vec_b = vec_b, vec_a return sum(v * vec_b.get(k, 0.0) for k, v in vec_a.items()) class RouterCache: """L1 精确缓存 + L2 语义缓存。""" def __init__(self, semantic_enabled: bool = True, similarity_threshold: float = 0.88, promote_frequency: int = 5, max_exact: int = 10000, max_semantic: int = 5000): self.semantic_enabled = semantic_enabled self.similarity_threshold = similarity_threshold self.promote_frequency = promote_frequency self.max_exact = max_exact self.max_semantic = max_semantic self._exact: Dict[str, CacheEntry] = {} self._semantic: List[Tuple[str, CacheEntry]] = [] # (query, entry) self._sem_vecs: Dict[str, Dict[str, float]] = {} self._sem_norms: Dict[str, float] = {} # 预计算范数,避免查询期重算 self.hits = {"exact": 0, "semantic": 0} self.misses = 0 # ---- 查询 ---- def get(self, query: str) -> Optional[Tuple[Optional[str], Dict[str, Any]]]: """返回 (level, result);未命中返回 None。level: 'exact' | 'semantic'""" entry = self._exact.get(query) if entry is not None: self.hits["exact"] += 1 return ("exact", entry.result) if self.semantic_enabled: q_vec = _tf_vector(_ngrams(query)) q_norm = _norm(q_vec) best_sim = 0.0 best_idx = -1 if q_norm > 0.0: # 单遍扫描:同时跟踪最优相似度与条目位置 for i, (q, _e) in enumerate(self._semantic): 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 if sim >= 1.0: break # 余弦相似度上界:完全相同查询,提前终止 if best_idx >= 0 and best_sim >= self.similarity_threshold: best_q, best_entry = self._semantic[best_idx] # 完全相同查询(相似度=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) return (level, best_entry.result) self.misses += 1 return None def _bump_semantic(self, idx: int, query: str, entry: CacheEntry): """语义命中:累计命中次数,达到阈值提升为精确缓存(O(1),无需二次查找)。""" entry.hits += 1 if entry.hits >= self.promote_frequency: self._exact[query] = entry self._semantic.pop(idx) 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 entry = CacheEntry(result=result) if self.semantic_enabled: if len(self._semantic) >= self.max_semantic: old_q, _ = self._semantic.pop(0) 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._sem_vecs[query] = vec self._sem_norms[query] = _norm(vec) else: self._exact[query] = entry if len(self._exact) > self.max_exact: self._exact.pop(next(iter(self._exact))) # ---- 统计 ---- def stats(self) -> Dict[str, Any]: total = self.hits["exact"] + self.hits["semantic"] + self.misses return { "exact_hits": self.hits["exact"], "semantic_hits": self.hits["semantic"], "misses": self.misses, "hit_rate": round((self.hits["exact"] + self.hits["semantic"]) / total, 4) if total else 0.0, "exact_size": len(self._exact), "semantic_size": len(self._semantic), } def clear(self): self._exact.clear() self._semantic.clear() self._sem_vecs.clear() self._sem_norms.clear() self.hits = {"exact": 0, "semantic": 0} self.misses = 0