CourseRAG · Module 12 : Production-Engineering · part 65 of 82
Part 65 · Module 12 : Production-Engineering

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

The rest of this course is yours to keep

This course is bought on its own, once, and stays readable afterwards, including the parts added to it later.