"""Evidence selection, deduplication, diversity, and token budgeting."""

from collections import Counter
from typing import List, Sequence

from .models import SearchResult
from .text import content_terms, jaccard, tokenize


class ContextPacker:
    """Greedy maximal-marginal-relevance context selection.

    Retrieval rank is not a license to fill the entire context window.  This
    component removes near duplicates, limits repeated chunks from one source,
    and balances relevance against diversity under an explicit token budget.
    """

    def __init__(
        self,
        max_tokens: int = 700,
        max_chunks: int = 6,
        relevance_weight: float = 0.75,
        duplicate_threshold: float = 0.88,
        max_chunks_per_document: int = 2,
    ) -> None:
        if max_tokens <= 0 or max_chunks <= 0:
            raise ValueError("context budgets must be positive")
        if not 0 <= relevance_weight <= 1:
            raise ValueError("relevance_weight must be between zero and one")
        self.max_tokens = max_tokens
        self.max_chunks = max_chunks
        self.relevance_weight = relevance_weight
        self.duplicate_threshold = duplicate_threshold
        self.max_chunks_per_document = max_chunks_per_document

    def pack(self, results: Sequence[SearchResult]) -> List[SearchResult]:
        if not results:
            return []
        low = min(result.score for result in results)
        high = max(result.score for result in results)

        def relevance(result: SearchResult) -> float:
            return (result.score - low) / (high - low) if high > low else 1.0

        remaining = list(results)
        selected: List[SearchResult] = []
        selected_terms: List[set] = []
        document_counts: Counter[str] = Counter()
        used_tokens = 0

        while remaining and len(selected) < self.max_chunks:
            candidates = []
            for result in remaining:
                terms = set(content_terms(result.chunk.text))
                redundancy = max((jaccard(terms, prior) for prior in selected_terms), default=0.0)
                mmr = self.relevance_weight * relevance(result) - (
                    1.0 - self.relevance_weight
                ) * redundancy
                candidates.append((mmr, redundancy, result, terms))
            candidates.sort(key=lambda item: (-item[0], item[2].rank, item[2].chunk.id))

            chosen = None
            for _, redundancy, result, terms in candidates:
                token_count = len(tokenize(result.chunk.text))
                if redundancy >= self.duplicate_threshold:
                    continue
                if document_counts[result.chunk.document_id] >= self.max_chunks_per_document:
                    continue
                if used_tokens + token_count > self.max_tokens:
                    continue
                chosen = (result, terms, token_count)
                break
            if chosen is None:
                break
            result, terms, token_count = chosen
            selected.append(result)
            selected_terms.append(terms)
            document_counts[result.chunk.document_id] += 1
            used_tokens += token_count
            remaining = [candidate for candidate in remaining if candidate.chunk.id != result.chunk.id]

        return [
            SearchResult(
                chunk=result.chunk,
                score=result.score,
                rank=rank,
                retriever=result.retriever,
                component_scores=result.component_scores,
            )
            for rank, result in enumerate(selected, start=1)
        ]

    @staticmethod
    def render(results: Sequence[SearchResult]) -> str:
        """Render an injection-resistant evidence envelope with source IDs."""

        blocks = []
        for result in results:
            blocks.append(
                "\n".join(
                    [
                        f"<evidence id=\"{result.chunk.id}\" source=\"{result.chunk.source}\">",
                        f"TITLE: {result.chunk.title}",
                        result.chunk.text,
                        "</evidence>",
                    ]
                )
            )
        return "\n\n".join(blocks)

