"""Semantic cache — the biggest energy lever.

Before we retrieve or call the GPU, we check whether a semantically-equivalent question was
already answered. If a stored query's embedding is within `similarity_threshold` cosine distance,
we return the cached answer and skip the model entirely.

The store is a small JSON file: parallel lists of embeddings, queries, and answer payloads.
For very large caches, swap this for a FAISS index — the interface stays the same.
"""
from __future__ import annotations

import json
from pathlib import Path
from typing import Any, Optional

import numpy as np

from .embeddings import Embedder
from .settings import Config


class SemanticCache:
    def __init__(self, cfg: Config):
        self.enabled = bool(cfg.get("cache.enabled", True))
        self.threshold = float(cfg.get("cache.similarity_threshold", 0.92))
        self.path = cfg.resolve("cache.path", "cache_store.json")
        self.embedder = Embedder(cfg)
        self._vecs: Optional[np.ndarray] = None  # shape (N, D), L2-normalized
        self._queries: list[str] = []
        self._payloads: list[dict] = []
        self._load()

    def _load(self):
        if not self.path.exists():
            return
        with open(self.path, "r", encoding="utf-8") as f:
            store = json.load(f)
        self._queries = store.get("queries", [])
        self._payloads = store.get("payloads", [])
        vecs = store.get("vectors", [])
        self._vecs = np.array(vecs, dtype="float32") if vecs else None

    def _persist(self):
        store = {
            "queries": self._queries,
            "payloads": self._payloads,
            "vectors": [] if self._vecs is None else self._vecs.tolist(),
        }
        self.path.parent.mkdir(parents=True, exist_ok=True)
        with open(self.path, "w", encoding="utf-8") as f:
            json.dump(store, f, ensure_ascii=False)

    def lookup(self, query: str) -> Optional[dict]:
        """Return the cached payload for a semantically-equivalent query, or None."""
        if not self.enabled or self._vecs is None or len(self._queries) == 0:
            return None
        qvec = self.embedder.embed_query(query)
        sims = self._vecs @ qvec  # cosine sim (all normalized)
        best = int(np.argmax(sims))
        if float(sims[best]) >= self.threshold:
            hit = dict(self._payloads[best])
            hit["_cache_similarity"] = float(sims[best])
            hit["_cache_matched_query"] = self._queries[best]
            return hit
        return None

    def store(self, query: str, payload: dict[str, Any]) -> None:
        if not self.enabled:
            return
        qvec = self.embedder.embed_query(query).reshape(1, -1)
        self._vecs = qvec if self._vecs is None else np.vstack([self._vecs, qvec])
        self._queries.append(query)
        self._payloads.append(payload)
        self._persist()
