"""System prompt + RAG context assembly, tuned for minimal token usage.

Because QLoRA fine-tuning *internalizes* the behavioral rules (cite sources, use only context,
no personalized advice), the served system prompt can be short — every token here is paid on
EVERY request, so a compact prompt is a direct, permanent cost/energy saving. Training and serving
use the same compact prompt to avoid a train/serve mismatch.

Context assembly also minimizes input tokens: near-duplicate chunks are dropped and the context is
capped to a token budget (lowest-scoring chunks fall off first). This is the application-layer
analogue of the paper's sparse top-k context selection.
"""
from __future__ import annotations

from dataclasses import dataclass

# Compact behavioral contract (~45 tokens vs ~150 for the verbose form). The fine-tune teaches the
# detail; this just anchors the role and the non-negotiables.
SYSTEM_PROMPT = (
    "You are a travel-insurance assistant. Use ONLY the CONTEXT; if the answer isn't there, say so "
    "and suggest a licensed advisor. Cite the passage number(s) you use, e.g. [1]. "
    "Never say which plan to buy. Be brief."
)


@dataclass
class RetrievedChunk:
    text: str
    source: str
    score: float


def estimate_tokens(text: str) -> int:
    """Cheap, dependency-free token estimate (~4 chars/token). Good enough for budgeting."""
    return max(1, len(text) // 4)


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


def dedup_chunks(chunks: list[RetrievedChunk]) -> list[RetrievedChunk]:
    """Drop exact/subsumed duplicates (e.g. a table that also appears as flattened text)."""
    kept: list[RetrievedChunk] = []
    seen_norm: list[str] = []
    for c in chunks:
        norm = _normalize(c.text)
        if any(norm == s or norm in s or s in norm for s in seen_norm):
            continue
        kept.append(c)
        seen_norm.append(norm)
    return kept


def fit_token_budget(chunks: list[RetrievedChunk], budget_tokens: int) -> list[RetrievedChunk]:
    """Keep highest-scoring chunks until the token budget is hit (always keep at least one)."""
    ordered = sorted(chunks, key=lambda c: c.score, reverse=True)
    out, used = [], 0
    for c in ordered:
        cost = estimate_tokens(c.text)
        if out and used + cost > budget_tokens:
            continue
        out.append(c)
        used += cost
    return out


def compress_context(
    chunks: list[RetrievedChunk], budget_tokens: int, dedup: bool = True
) -> list[RetrievedChunk]:
    """Dedup + budget-trim retrieved chunks to minimize context tokens."""
    if dedup:
        chunks = dedup_chunks(chunks)
    return fit_token_budget(chunks, budget_tokens)


def build_context_block(chunks: list[RetrievedChunk]) -> str:
    """Render retrieved chunks as a numbered, citable CONTEXT block."""
    if not chunks:
        return "CONTEXT: (none found)"
    lines = ["CONTEXT:"]
    for i, c in enumerate(chunks, 1):
        lines.append(f"[{i}] ({c.source})\n{c.text.strip()}")
    return "\n\n".join(lines)


def build_messages(query: str, chunks: list[RetrievedChunk]) -> list[dict]:
    """Assemble the chat messages for the LLM (kept terse to save input tokens)."""
    context = build_context_block(chunks)
    user_content = f"{context}\n\nQ: {query.strip()}\nAnswer from CONTEXT only; cite the passage number(s) you use, e.g. [1]."
    return [
        {"role": "system", "content": SYSTEM_PROMPT},
        {"role": "user", "content": user_content},
    ]
