"""RAG retriever: embed a query, search the FAISS index, return citable chunks.

The index and metadata are produced by scripts/build_index.py. Chunks below the configured
minimum similarity are dropped so the LLM never sees irrelevant context (which also keeps the
prompt short — an energy win).
"""
from __future__ import annotations

import json
from pathlib import Path

from .embeddings import Embedder
from .prompt import RetrievedChunk
from .settings import Config


class Retriever:
    def __init__(self, cfg: Config):
        self.cfg = cfg
        self.top_k = int(cfg.get("retrieval.top_k", 4))
        self.min_score = float(cfg.get("retrieval.min_score", 0.30))
        self.embedder = Embedder(cfg)

        kb_dir = cfg.resolve("data.kb_dir")
        self.index_path = kb_dir / "index.faiss"
        self.meta_path = kb_dir / "index.meta.json"
        self._index = None
        self._meta: list[dict] = []

    def _load(self):
        if self._index is not None:
            return
        import faiss

        if not self.index_path.exists() or not self.meta_path.exists():
            raise FileNotFoundError(
                f"RAG index not found at {self.index_path}. Run scripts/build_index.py first."
            )
        self._index = faiss.read_index(str(self.index_path))
        with open(self.meta_path, "r", encoding="utf-8") as f:
            self._meta = json.load(f)

    def search(self, query: str) -> list[RetrievedChunk]:
        self._load()
        qvec = self.embedder.embed_query(query).reshape(1, -1)
        # Vectors are L2-normalized, so inner product == cosine similarity.
        scores, idxs = self._index.search(qvec, self.top_k)
        results: list[RetrievedChunk] = []
        for score, idx in zip(scores[0], idxs[0]):
            if idx < 0:
                continue
            if float(score) < self.min_score:
                continue
            m = self._meta[int(idx)]
            results.append(
                RetrievedChunk(text=m["text"], source=m.get("source", "unknown"), score=float(score))
            )
        return results
