"""Self-contained tests for the agent pipeline logic.

No model downloads, no FAISS index, no GPU: the embedder is faked, and the retriever + LLM are
replaced with stubs. Validates routing, caching, context compression, citation mapping, and
guardrails — the behavior we rely on for token efficiency and safety.

    python tests/test_agent.py      # or: pytest tests/test_agent.py
"""
from __future__ import annotations

import hashlib
import sys
import tempfile
from pathlib import Path

import numpy as np

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "src"))

import agent.embeddings as emb_mod  # noqa: E402


class _FakeSTModel:
    DIM = 64

    def _vec(self, text: str) -> np.ndarray:
        v = np.zeros(self.DIM, dtype="float32")
        for tok in text.lower().split():
            v[int(hashlib.md5(tok.encode()).hexdigest(), 16) % self.DIM] += 1.0
        n = np.linalg.norm(v)
        return v / n if n else v

    def encode(self, texts, normalize_embeddings=True, convert_to_numpy=True, show_progress_bar=False):
        if isinstance(texts, str):
            return self._vec(texts)
        return np.vstack([self._vec(t) for t in texts]).astype("float32")


# Patch the embedder so nothing is downloaded.
emb_mod._MODEL_CACHE["__fake__"] = _FakeSTModel()
emb_mod.Embedder._load = lambda self: emb_mod._MODEL_CACHE["__fake__"]

from agent.guardrails import Guardrails             # noqa: E402
from agent.llm import LLMResult                      # noqa: E402
from agent.pipeline import Agent, _citations_used    # noqa: E402
from agent.prompt import RetrievedChunk, compress_context  # noqa: E402
from agent.settings import load_config               # noqa: E402

CHUNKS = [
    RetrievedChunk(text="Trip cancellation from severe weather is covered.", source="p#cancel", score=0.7),
    RetrievedChunk(text="Emergency medical up to the certificate limit.", source="p#medical", score=0.6),
]


class _FakeRetriever:
    """Returns preset chunks for in-domain queries, nothing for 'unknown'."""

    def search(self, query):
        return [] if "unknown" in query.lower() else list(CHUNKS)


class _FakeLLM:
    def generate(self, messages):
        return LLMResult(text="Yes, covered [1].", prompt_tokens=100, completion_tokens=5)


def _agent():
    cfg = load_config(str(ROOT / "config" / "config.yaml"))
    a = Agent(cfg, lazy_llm=True)
    a.retriever = _FakeRetriever()
    a._llm = _FakeLLM()
    # Isolated, empty cache.
    a.cache.path = Path(tempfile.gettempdir()) / "test_cache_store.json"
    a.cache._vecs, a.cache._queries, a.cache._payloads = None, [], []
    if a.cache.path.exists():
        a.cache.path.unlink()
    return a


def test_router_zero_tokens():
    r = _agent().answer("hi")
    assert r.route == "router" and r.total_tokens == 0 and not r.grounded


def test_grounded_answer_maps_single_citation():
    r = _agent().answer("Is trip cancellation covered?")
    assert r.route == "llm" and r.grounded
    assert r.total_tokens == 105
    assert r.citations == ["p#cancel"]  # answer cited [1] -> first chunk only


def test_cache_hit_zero_tokens():
    a = _agent()
    q = "Is trip cancellation covered?"
    a.answer(q)
    r2 = a.answer(q)
    assert r2.cached and r2.route == "cache" and r2.total_tokens == 0


def test_no_context_escalates_without_llm():
    r = _agent().answer("some unknown thing")
    assert not r.grounded and r.guardrail_reason == "no_context" and r.total_tokens == 0


def test_compress_context_dedup_and_budget():
    dup = "same text here"
    chunks = [
        RetrievedChunk(text=dup, source="a", score=0.9),
        RetrievedChunk(text=dup, source="a", score=0.8),        # duplicate
        RetrievedChunk(text="x" * 4000, source="b", score=0.7),  # ~1000 tok, over budget
    ]
    kept = compress_context(chunks, budget_tokens=600, dedup=True)
    assert len(kept) == 1 and kept[0].source == "a"


def test_citations_used_fallback():
    chunks = [RetrievedChunk("t", "s1", 0.9), RetrievedChunk("t2", "s2", 0.8)]
    assert _citations_used("cited [2] here", chunks) == ["s2"]
    assert _citations_used("no markers", chunks) == ["s1", "s2"]  # fallback: all sources


def test_guardrails():
    cfg = load_config(str(ROOT / "config" / "config.yaml"))
    g = Guardrails(cfg)
    ctx = [RetrievedChunk("Plans info", "p#plans", 0.9)]
    assert g.check("You should buy Plus [1].", ctx).reason == "personalized_advice"
    assert g.check("Covered [1].", ctx).ok
    assert g.check("Covered.", ctx).reason == "missing_citation"
    assert g.check("Not here [All].", ctx).reason == "missing_citation"  # non-numeric not a citation


if __name__ == "__main__":
    tests = [v for k, v in sorted(globals().items()) if k.startswith("test_") and callable(v)]
    for t in tests:
        t()
        print(f"  PASS  {t.__name__}")
    print(f"\n{len(tests)} tests passed")
