Files

142 lines
5.3 KiB
Python

"""两阶段路由缓存(对齐实现方案):
- L1 精确缓存:完全相同的查询 -> 直接命中
- L2 语义缓存:字符 n-gram 余弦相似度(零依赖)-> 相似查询命中
- 命中 N 次(promote_frequency)后提升为精确缓存
说明:语义缓存中的"完全相同查询"(相似度=1.0)直接计为 exact 命中;
高频语义命中会提升为 O(1) 的精确缓存条目。
只缓存"未升级"的结果(升级路径每次都走大模型,不缓存,避免陈旧)。
"""
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 _cosine(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float:
if not vec_a or not vec_b:
return 0.0
common = set(vec_a) & set(vec_b)
dot = sum(vec_a[k] * vec_b[k] for k in common)
na = sum(v * v for v in vec_a.values()) ** 0.5
nb = sum(v * v for v in vec_b.values()) ** 0.5
if na == 0 or nb == 0:
return 0.0
return dot / (na * nb)
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
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.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))
best_sim = 0.0
best_query: Optional[str] = None
best_result: Optional[Dict[str, Any]] = None
for q, e in self._semantic:
sim = _cosine(q_vec, self._sem_vecs.get(q, {}))
if sim > best_sim:
best_sim = sim
best_query = q
best_result = e.result
if best_query is not None and best_sim >= self.similarity_threshold:
# 完全相同查询(相似度=1.0)计为 exact 命中
is_exact = best_sim >= 0.999
level = "exact" if is_exact else "semantic"
self.hits[level] += 1
self._semantic_hit(best_query)
return (level, best_result)
self.misses += 1
return None
def _semantic_hit(self, query: str):
"""语义命中:累计命中次数,达到阈值提升为精确缓存。"""
for i, (q, e) in enumerate(self._semantic):
if q == query:
e.hits += 1
if e.hits >= self.promote_frequency:
self._exact[query] = e
self._semantic.pop(i)
self._sem_vecs.pop(query, None)
break
# ---- 写入 ----
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._semantic.append((query, entry))
self._sem_vecs[query] = _tf_vector(_ngrams(query))
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.hits = {"exact": 0, "semantic": 0}
self.misses = 0