"""Composable baseline and advanced retrieval-augmented generation pipelines."""

from dataclasses import replace
from typing import List, Optional, Sequence

from .context import ContextPacker
from .generation import ExtractiveGenerator, Generator
from .models import Answer, Document, SearchResult, TraceEvent
from .rerankers import CrossFeatureReranker
from .retrievers import (
    AdaptiveRetriever,
    BM25Retriever,
    GraphExpandedRetriever,
    HashingSemanticRetriever,
    HybridRetriever,
    MultiQueryRetriever,
    Retriever,
    simple_query_variants,
)
from .text import chunk_documents


class RAGPipeline:
    """Retrieve, optionally rerank, pack evidence, and generate with a trace."""

    def __init__(
        self,
        retriever: Retriever,
        generator: Generator,
        context_packer: ContextPacker,
        reranker: Optional[CrossFeatureReranker] = None,
        retrieval_k: int = 12,
        rerank_k: int = 8,
    ) -> None:
        self.retriever = retriever
        self.generator = generator
        self.context_packer = context_packer
        self.reranker = reranker
        self.retrieval_k = retrieval_k
        self.rerank_k = rerank_k

    def search(self, query: str, k: Optional[int] = None) -> List[SearchResult]:
        """Run retrieval and reranking without context packing or generation."""

        requested = k or self.rerank_k
        candidates = self.retriever.search(query, max(self.retrieval_k, requested))
        if self.reranker is not None:
            return self.reranker.rerank(query, candidates, requested)
        return candidates[:requested]

    def ask(self, query: str) -> Answer:
        route = "fixed"
        route_method = getattr(self.retriever, "route_for", None)
        if callable(route_method):
            route = str(route_method(query))
        raw = self.retriever.search(query, self.retrieval_k)
        trace = [
            TraceEvent(
                stage="route",
                detail=f"selected {route} retrieval",
                values={"route": route},
            ),
            TraceEvent(
                stage="retrieve",
                detail=f"retrieved {len(raw)} candidates",
                values={"count": len(raw), "k": self.retrieval_k},
            ),
        ]
        ranked = raw
        if self.reranker is not None:
            ranked = self.reranker.rerank(query, raw, self.rerank_k)
            trace.append(
                TraceEvent(
                    stage="rerank",
                    detail=f"retained {len(ranked)} reranked candidates",
                    values={"count": len(ranked), "k": self.rerank_k},
                )
            )
        packed = self.context_packer.pack(ranked)
        trace.append(
            TraceEvent(
                stage="pack",
                detail=f"packed {len(packed)} non-redundant chunks",
                values={
                    "count": len(packed),
                    "max_tokens": self.context_packer.max_tokens,
                },
            )
        )
        answer = self.generator.generate(query, packed)
        trace.append(
            TraceEvent(
                stage="generate",
                detail="abstained" if answer.abstained else "returned grounded evidence",
                values={
                    "citations": len(answer.citations),
                    "confidence": round(answer.confidence, 4),
                    "abstained": answer.abstained,
                },
            )
        )
        return replace(answer, trace=tuple(trace))


def build_baseline_pipeline(documents: Sequence[Document]) -> RAGPipeline:
    """A 2020-style retrieve-then-generate baseline with sparse retrieval."""

    chunks = chunk_documents(documents, chunk_size=120, overlap=20)
    return RAGPipeline(
        retriever=BM25Retriever(chunks),
        generator=ExtractiveGenerator(max_sentences=2),
        context_packer=ContextPacker(max_tokens=450, max_chunks=5, relevance_weight=0.9),
        reranker=None,
        retrieval_k=5,
        rerank_k=5,
    )


def build_advanced_pipeline(documents: Sequence[Document]) -> RAGPipeline:
    """A modern modular stack: hybrid, multi-query, graph routing, reranking, MMR."""

    chunks = chunk_documents(documents, chunk_size=90, overlap=18)
    sparse = BM25Retriever(chunks)
    semantic = HashingSemanticRetriever(chunks)
    hybrid = HybridRetriever(
        (("sparse", sparse, 1.0), ("semantic", semantic, 1.0)),
        rrf_constant=30,
    )
    multi_query = MultiQueryRetriever(hybrid, simple_query_variants, rrf_constant=30)
    graph = GraphExpandedRetriever(multi_query, chunks, expansion_weight=0.30)
    adaptive = AdaptiveRetriever(sparse=sparse, hybrid=multi_query, graph=graph)
    return RAGPipeline(
        retriever=adaptive,
        generator=ExtractiveGenerator(max_sentences=2),
        context_packer=ContextPacker(
            max_tokens=560,
            max_chunks=6,
            relevance_weight=0.78,
            duplicate_threshold=0.86,
        ),
        reranker=CrossFeatureReranker(),
        retrieval_k=14,
        rerank_k=8,
    )
