Topic 3: Caching
8 min read·21 Sept 2026
module12/caching.py
python
# module12/caching.py
"""Four caches with different risks: embeddings, retrieval, exact answers, and semantic answers."""
from __future__ import annotations
import hashlib
import json
import re
import sqlite3
import time
from dataclasses import dataclass, field
import numpy as np
def normalize_query(text: str) -> str:
"""Collapse the harmless differences: case, spacing, trailing punctuation, filler words."""
text = re.sub(r"\s+", " ", text.strip().lower()).rstrip("?!. ")
text = re.sub(r"^(hi|hello|hey)[, ]+", "", text)
return re.sub(r"^(please|can you|could you|i want to know)\s+", "", text)
@dataclass
class CacheStats:
hits: int = 0
misses: int = 0
stale_evictions: int = 0
@property
def hit_rate(self) -> float:
total = self.hits + self.misses
return self.hits / total if total else 0.0
class KeyValueCache:
"""Exact-match cache in SQLite, keyed by a namespace, a version, and the normalized input."""
def __init__(self, path: str = "data/cache.sqlite", namespace: str = "generic",
ttl_seconds: int = 3600, version: str = "v1"):
self.db = sqlite3.connect(path, check_same_thread=False)
self.db.execute("CREATE TABLE IF NOT EXISTS cache (key TEXT PRIMARY KEY, namespace TEXT, "
"value TEXT, created REAL, corpus_version TEXT)")
self.db.commit()
self.namespace, self.ttl, self.version = namespace, ttl_seconds, version
self.stats = CacheStats()
def key(self, text: str) -> str:
return hashlib.sha256(f"{self.namespace}|{self.version}|{normalize_query(text)}".encode()).hexdigest()
def get(self, text: str, corpus_version: str = ""):
row = self.db.execute("SELECT value, created, corpus_version FROM cache "
"WHERE key = ? AND namespace = ?", (self.key(text), self.namespace)).fetchone()
if not row:
self.stats.misses += 1
return None
value, created, stored_version = row
if time.time() - created > self.ttl or (corpus_version and stored_version != corpus_version):
self.stats.stale_evictions += 1
self.stats.misses += 1
self.db.execute("DELETE FROM cache WHERE key = ?", (self.key(text),))
self.db.commit()
return None
self.stats.hits += 1
return json.loads(value)
def put(self, text: str, value, corpus_version: str = "") -> None:
self.db.execute("INSERT OR REPLACE INTO cache VALUES (?, ?, ?, ?, ?)",
(self.key(text), self.namespace, json.dumps(value), time.time(), corpus_version))
self.db.commit()
def invalidate_all(self) -> int:
removed = self.db.execute("SELECT COUNT(*) FROM cache WHERE namespace = ?",
(self.namespace,)).fetchone()[0]
self.db.execute("DELETE FROM cache WHERE namespace = ?", (self.namespace,))
self.db.commit()
return removed
def invalidate_documents(self, doc_ids: set[str]) -> int:
"""Drop entries whose stored answer used any of these documents."""
removed = 0
rows = self.db.execute("SELECT key, value FROM cache WHERE namespace = ?",
(self.namespace,)).fetchall()
for key, value in rows:
payload = json.loads(value)
used = set(payload.get("doc_ids", [])) if isinstance(payload, dict) else set()
if used & doc_ids:
self.db.execute("DELETE FROM cache WHERE key = ?", (key,))
removed += 1
self.db.commit()
return removed
@dataclass
class SemanticCacheEntry:
query: str
vector: np.ndarray
payload: dict
created: float
corpus_version: str
class SemanticCache:
"""Answers reused for *similar* questions. Higher hit rate, real risk of the wrong answer."""
def __init__(self, embedder, threshold: float = 0.95, ttl_seconds: int = 1800,
max_entries: int = 5000):
self.embedder, self.threshold = embedder, threshold
self.ttl, self.max_entries = ttl_seconds, max_entries
self.entries: list[SemanticCacheEntry] = []
self.stats = CacheStats()
def _fresh(self, entry: SemanticCacheEntry, corpus_version: str) -> bool:
return (time.time() - entry.created <= self.ttl
and (not corpus_version or entry.corpus_version == corpus_version))
def get(self, query: str, corpus_version: str = "", guard=None):
if not self.entries:
self.stats.misses += 1
return None
vector = self.embedder.embed_query(query)
scores = np.array([float(vector @ e.vector) for e in self.entries])
best = int(np.argmax(scores))
entry = self.entries[best]
if scores[best] < self.threshold or not self._fresh(entry, corpus_version):
self.stats.misses += 1
return None
if guard and not guard(query, entry.query, entry.payload): # last line of defence
self.stats.misses += 1
return None
self.stats.hits += 1
return {**entry.payload, "cache_similarity": float(scores[best]), "cached_from": entry.query}
def put(self, query: str, payload: dict, corpus_version: str = "") -> None:
self.entries.append(SemanticCacheEntry(query, self.embedder.embed_query(query), payload,
time.time(), corpus_version))
if len(self.entries) > self.max_entries:
self.entries = self.entries[-self.max_entries:]
def invalidate_documents(self, doc_ids: set[str]) -> int:
before = len(self.entries)
self.entries = [e for e in self.entries if not (set(e.payload.get("doc_ids", [])) & doc_ids)]
return before - len(self.entries)
NUMBERS = re.compile(r"\d+(?:[.,]\d+)?")
NEGATION = re.compile(r"\b(not|without|except|unless|non|un)\w*\b", re.I)
def risky_semantic_hit(new_query: str, cached_query: str, payload: dict) -> bool:
"""Reject a semantic hit when the two questions differ in ways similarity does not see."""
if set(NUMBERS.findall(new_query)) != set(NUMBERS.findall(cached_query)):
return True # "7 days" vs "30 days"
if bool(NEGATION.search(new_query)) != bool(NEGATION.search(cached_query)):
return True # "covered" vs "not covered"
product = re.compile(r"\b[A-Z]{1,3}\d{2,4}\b")
if set(product.findall(new_query.upper())) != set(product.findall(cached_query.upper())):
return True # X200 vs X300
return False
def safe_semantic_guard(new_query: str, cached_query: str, payload: dict) -> bool:
return not risky_semantic_hit(new_query, cached_query, payload)
@dataclass
class CacheLayers:
"""The four caches, with the rule of thumb for each."""
embeddings: KeyValueCache # safest: text plus model determines the vector
retrieval: KeyValueCache # medium: invalidate when the corpus changes
answers: KeyValueCache # exact-match answers: safe if invalidated with the corpus
semantic: SemanticCache | None = None # riskiest: similar is not the same
stats: dict = field(default_factory=dict)
def invalidate_documents(self, doc_ids: set[str]) -> dict:
removed = {"retrieval": self.retrieval.invalidate_all(),
"answers": self.answers.invalidate_documents(doc_ids)}
if self.semantic:
removed["semantic"] = self.semantic.invalidate_documents(doc_ids)
return removed # embeddings survive: text is unchanged