"""End-to-end agent pipeline: router -> cache -> retrieve -> compress -> generate -> guardrail.

Every stage is designed to spend as few tokens as possible (the objective: minimize cost + energy):
  - router    : trivial/too-short inputs are answered with a canned reply — zero tokens.
  - cache     : semantically-repeated questions return instantly — zero tokens.
  - compress  : retrieved context is deduped and capped to a token budget — fewer input tokens.
  - generate  : compact prompt + capped max_tokens — fewer input and output tokens.
Token usage is recorded on every response so it can be measured and driven down.

Shared by app/server.py and scripts/evaluate.py so the exact same path is served and measured.
"""
from __future__ import annotations

import re
from dataclasses import asdict, dataclass, field

from .cache import SemanticCache
from .guardrails import Guardrails
from .llm import LLMClient
from .prompt import RetrievedChunk, build_messages, compress_context
from .retriever import Retriever
from .settings import Config

# Canned, zero-token replies for inputs the router short-circuits.
_GREETINGS = {"hi", "hello", "hey", "thanks", "thank you", "ok", "okay", "yo", "sup"}
_ROUTER_REPLY = "Hi! Ask me a travel-insurance question (coverage, claims, exclusions) and I'll answer from the policy documents."


@dataclass
class AgentResponse:
    answer: str
    citations: list[str] = field(default_factory=list)
    cached: bool = False
    grounded: bool = False
    guardrail_reason: str = ""
    similarity: float | None = None
    prompt_tokens: int = 0
    completion_tokens: int = 0
    total_tokens: int = 0
    route: str = "llm"  # llm | router | cache


class Agent:
    def __init__(self, cfg: Config, lazy_llm: bool = True):
        self.cfg = cfg
        self.cache = SemanticCache(cfg)
        self.retriever = Retriever(cfg)
        self.guardrails = Guardrails(cfg)
        self.router_enabled = bool(cfg.get("router.enabled", True))
        self.min_query_chars = int(cfg.get("router.min_query_chars", 3))
        self.context_budget = int(cfg.get("retrieval.context_token_budget", 600))
        self.dedup = bool(cfg.get("retrieval.dedup", True))
        self._llm: LLMClient | None = None
        if not lazy_llm:
            self._llm = LLMClient(cfg)

    @property
    def llm(self) -> LLMClient:
        if self._llm is None:
            self._llm = LLMClient(self.cfg)
        return self._llm

    def _route(self, query: str) -> AgentResponse | None:
        """Cheap pre-filter: handle trivial inputs without retrieval or the LLM."""
        if not self.router_enabled:
            return None
        q = query.strip()
        if len(q) < self.min_query_chars or _normalize(q) in _GREETINGS:
            return AgentResponse(answer=_ROUTER_REPLY, grounded=False, route="router")
        return None

    def answer(self, query: str) -> AgentResponse:
        # [0] Router — zero-token short-circuit for trivial inputs.
        routed = self._route(query)
        if routed is not None:
            return routed

        # [1] Semantic cache — zero-token hit for repeated questions.
        hit = self.cache.lookup(query)
        if hit is not None:
            return AgentResponse(
                answer=hit["answer"],
                citations=hit.get("citations", []),
                cached=True,
                grounded=hit.get("grounded", False),
                similarity=hit.get("_cache_similarity"),
                route="cache",
            )

        # [2] Retrieve, then [compress] to minimize context tokens.
        chunks: list[RetrievedChunk] = self.retriever.search(query)
        chunks = compress_context(chunks, self.context_budget, self.dedup)

        # [3] Generate only if we have grounding context.
        raw, p_tok, c_tok = "", 0, 0
        if chunks:
            result = self.llm.generate(build_messages(query, chunks))
            raw, p_tok, c_tok = result.text, result.prompt_tokens, result.completion_tokens

        # [4] Guardrail check.
        gr = self.guardrails.check(raw, chunks)
        citations = _citations_used(gr.answer, chunks) if gr.ok else []
        resp = AgentResponse(
            answer=gr.answer,
            citations=citations,
            cached=False,
            grounded=gr.ok,
            guardrail_reason=gr.reason,
            prompt_tokens=p_tok,
            completion_tokens=c_tok,
            total_tokens=p_tok + c_tok,
            route="llm",
        )

        # Only cache clean, grounded answers — never cache an escalation/refusal.
        if gr.ok:
            self.cache.store(query, {"answer": resp.answer, "citations": citations, "grounded": True})
        return resp

    def to_dict(self, resp: AgentResponse) -> dict:
        return asdict(resp)


def _normalize(text: str) -> str:
    return " ".join(text.lower().split()).strip("?.! ")


def _citations_used(answer: str, chunks: list[RetrievedChunk]) -> list[str]:
    """Map the [n] passage markers the answer actually cited back to their sources.

    Falls back to all retrieved sources if the answer cited no resolvable number.
    """
    nums = {int(n) for n in re.findall(r"\[(\d+)\]", answer)}
    used = [chunks[i - 1].source for i in sorted(nums) if 1 <= i <= len(chunks)]
    seen, ordered = set(), []
    for s in used:
        if s not in seen:
            seen.add(s)
            ordered.append(s)
    return ordered or sorted({c.source for c in chunks})
