"""
Base retriever for RAG (Retrieval Augmented Generation).

Provides core dense retrieval functionality with vector search,
embedding generation, and document ranking.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import List, Dict, Optional, Tuple, Any
import numpy as np
from dataclasses import dataclass


@dataclass
class RetrievalResult:
    """Single retrieval result."""
    document: str
    score: float
    document_id: Optional[int] = None
    metadata: Optional[Dict[str, Any]] = None


@dataclass
class RetrievalBatch:
    """Batch of retrieval results."""
    queries: List[str]
    results: List[List[RetrievalResult]]
    retrieval_time: float

    def __len__(self) -> int:
        return len(self.queries)


class DenseRetriever(nn.Module):
    """
    Dense retriever using learned embeddings.

    Supports:
    - Dual encoder architecture (query encoder + document encoder)
    - Negative sampling for contrastive learning
    - Vector similarity search (cosine, dot product, L2)
    - Batch retrieval
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        hidden_dim: int = 512,
        num_layers: int = 3,
        dropout: float = 0.1,
        similarity_metric: str = "cosine",
        temperature: float = 0.05,
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.hidden_dim = hidden_dim
        self.similarity_metric = similarity_metric
        self.temperature = temperature

        # Query encoder
        self.query_encoder = self._build_encoder(
            embedding_dim, hidden_dim, num_layers, dropout
        )

        # Document encoder (separate from query encoder)
        self.document_encoder = self._build_encoder(
            embedding_dim, hidden_dim, num_layers, dropout
        )

        # Projection heads
        self.query_proj = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim),
            nn.LayerNorm(hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, embedding_dim)
        )

        self.doc_proj = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim),
            nn.LayerNorm(hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, embedding_dim)
        )

        # Document store (will be populated)
        self.document_embeddings: Optional[torch.Tensor] = None
        self.documents: List[str] = []

    def _build_encoder(
        self,
        embedding_dim: int,
        hidden_dim: int,
        num_layers: int,
        dropout: float
    ) -> nn.Module:
        """Build transformer encoder."""
        layers = []

        # Input projection
        layers.append(nn.Linear(embedding_dim, hidden_dim))
        layers.append(nn.LayerNorm(hidden_dim))
        layers.append(nn.Dropout(dropout))

        # Transformer layers
        for _ in range(num_layers):
            encoder_layer = nn.TransformerEncoderLayer(
                d_model=hidden_dim,
                nhead=8,
                dim_feedforward=hidden_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            )
            layers.append(encoder_layer)

        return nn.Sequential(*layers)

    def encode_query(self, query_embeddings: torch.Tensor) -> torch.Tensor:
        """
        Encode query into dense representation.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]

        Returns:
            query_vectors: [batch, embed_dim]
        """
        # Encode
        hidden = self.query_encoder(query_embeddings)

        # Mean pooling
        query_vec = hidden.mean(dim=1)  # [batch, hidden_dim]

        # Project
        query_vec = self.query_proj(query_vec)  # [batch, embed_dim]

        # L2 normalize
        query_vec = F.normalize(query_vec, p=2, dim=-1)

        return query_vec

    def encode_document(self, doc_embeddings: torch.Tensor) -> torch.Tensor:
        """
        Encode document into dense representation.

        Args:
            doc_embeddings: [batch, seq_len, embed_dim]

        Returns:
            doc_vectors: [batch, embed_dim]
        """
        # Encode
        hidden = self.document_encoder(doc_embeddings)

        # Mean pooling
        doc_vec = hidden.mean(dim=1)  # [batch, hidden_dim]

        # Project
        doc_vec = self.doc_proj(doc_vec)  # [batch, embed_dim]

        # L2 normalize
        doc_vec = F.normalize(doc_vec, p=2, dim=-1)

        return doc_vec

    def compute_similarity(
        self,
        query_vectors: torch.Tensor,
        doc_vectors: torch.Tensor
    ) -> torch.Tensor:
        """
        Compute similarity between queries and documents.

        Args:
            query_vectors: [batch_q, embed_dim]
            doc_vectors: [batch_d, embed_dim]

        Returns:
            scores: [batch_q, batch_d]
        """
        if self.similarity_metric == "cosine":
            # Cosine similarity (both already normalized)
            scores = torch.matmul(query_vectors, doc_vectors.t())
        elif self.similarity_metric == "dot":
            # Dot product
            scores = torch.matmul(query_vectors, doc_vectors.t())
        elif self.similarity_metric == "l2":
            # Negative L2 distance
            # [batch_q, 1, embed_dim] - [1, batch_d, embed_dim]
            diff = query_vectors.unsqueeze(1) - doc_vectors.unsqueeze(0)
            scores = -torch.norm(diff, p=2, dim=-1)
        else:
            raise ValueError(f"Unknown similarity metric: {self.similarity_metric}")

        # Scale by temperature
        scores = scores / self.temperature

        return scores

    def forward(
        self,
        query_embeddings: torch.Tensor,
        doc_embeddings: torch.Tensor,
        negative_doc_embeddings: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Forward pass for training with contrastive learning.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            doc_embeddings: [batch, seq_len, embed_dim] positive documents
            negative_doc_embeddings: [batch, num_neg, seq_len, embed_dim] negative docs

        Returns:
            loss: Contrastive loss
            scores: [batch, 1 + num_neg] similarity scores
        """
        # Encode queries and positive documents
        query_vecs = self.encode_query(query_embeddings)  # [batch, embed_dim]
        pos_doc_vecs = self.encode_document(doc_embeddings)  # [batch, embed_dim]

        # Positive scores
        pos_scores = (query_vecs * pos_doc_vecs).sum(dim=-1) / self.temperature

        # Negative scores
        if negative_doc_embeddings is not None:
            batch_size, num_neg, seq_len, embed_dim = negative_doc_embeddings.shape

            # Reshape and encode negatives
            neg_doc_flat = negative_doc_embeddings.view(
                batch_size * num_neg, seq_len, embed_dim
            )
            neg_doc_vecs = self.encode_document(neg_doc_flat)  # [batch * num_neg, embed_dim]
            neg_doc_vecs = neg_doc_vecs.view(batch_size, num_neg, embed_dim)

            # Compute negative scores [batch, num_neg]
            neg_scores = torch.matmul(
                query_vecs.unsqueeze(1),  # [batch, 1, embed_dim]
                neg_doc_vecs.transpose(1, 2)  # [batch, embed_dim, num_neg]
            ).squeeze(1) / self.temperature  # [batch, num_neg]

            # Combine scores
            all_scores = torch.cat([pos_scores.unsqueeze(1), neg_scores], dim=1)
        else:
            all_scores = pos_scores.unsqueeze(1)

        # Contrastive loss (positive doc should have highest score)
        labels = torch.zeros(query_vecs.size(0), dtype=torch.long, device=query_vecs.device)
        loss = F.cross_entropy(all_scores, labels)

        return loss, all_scores

    def index_documents(
        self,
        documents: List[str],
        doc_embeddings: torch.Tensor,
        batch_size: int = 32
    ):
        """
        Index documents for retrieval.

        Args:
            documents: List of document strings
            doc_embeddings: [num_docs, seq_len, embed_dim]
            batch_size: Batch size for encoding
        """
        self.documents = documents
        num_docs = len(documents)

        # Encode documents in batches
        all_doc_vecs = []

        with torch.no_grad():
            for i in range(0, num_docs, batch_size):
                batch_embeds = doc_embeddings[i:i+batch_size]
                batch_vecs = self.encode_document(batch_embeds)
                all_doc_vecs.append(batch_vecs)

        # Concatenate all vectors
        self.document_embeddings = torch.cat(all_doc_vecs, dim=0)

    def retrieve(
        self,
        query_embeddings: torch.Tensor,
        top_k: int = 5,
        return_scores: bool = True
    ) -> List[List[RetrievalResult]]:
        """
        Retrieve top-k documents for each query.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            top_k: Number of documents to retrieve
            return_scores: Whether to include similarity scores

        Returns:
            results: List of retrieval results for each query
        """
        if self.document_embeddings is None:
            raise ValueError("No documents indexed. Call index_documents() first.")

        with torch.no_grad():
            # Encode queries
            query_vecs = self.encode_query(query_embeddings)  # [batch, embed_dim]

            # Compute similarities
            scores = self.compute_similarity(
                query_vecs, self.document_embeddings
            )  # [batch, num_docs]

            # Get top-k
            top_scores, top_indices = scores.topk(k=top_k, dim=1)

            # Convert to list of results
            results = []
            for i in range(query_vecs.size(0)):
                query_results = []
                for j in range(top_k):
                    doc_idx = top_indices[i, j].item()
                    score = top_scores[i, j].item()

                    result = RetrievalResult(
                        document=self.documents[doc_idx],
                        score=score if return_scores else None,
                        document_id=doc_idx
                    )
                    query_results.append(result)

                results.append(query_results)

        return results


class SparseRetriever:
    """
    Sparse retriever using BM25 or TF-IDF.

    Lightweight retrieval based on term matching.
    """

    def __init__(
        self,
        method: str = "bm25",
        k1: float = 1.5,
        b: float = 0.75
    ):
        self.method = method
        self.k1 = k1  # BM25 parameter
        self.b = b  # BM25 parameter

        self.documents: List[str] = []
        self.term_freq: Dict[str, Dict[int, int]] = {}  # term -> {doc_id: freq}
        self.doc_lengths: List[int] = []
        self.avg_doc_length: float = 0.0
        self.num_docs: int = 0

    def _tokenize(self, text: str) -> List[str]:
        """Simple tokenization."""
        return text.lower().split()

    def index_documents(self, documents: List[str]):
        """Index documents for sparse retrieval."""
        self.documents = documents
        self.num_docs = len(documents)
        self.term_freq = {}
        self.doc_lengths = []

        # Build term frequency index
        for doc_id, doc in enumerate(documents):
            tokens = self._tokenize(doc)
            self.doc_lengths.append(len(tokens))

            # Count term frequencies
            for term in tokens:
                if term not in self.term_freq:
                    self.term_freq[term] = {}
                self.term_freq[term][doc_id] = self.term_freq[term].get(doc_id, 0) + 1

        # Compute average document length
        self.avg_doc_length = sum(self.doc_lengths) / max(len(self.doc_lengths), 1)

    def _bm25_score(self, query_terms: List[str], doc_id: int) -> float:
        """Compute BM25 score for a document."""
        score = 0.0
        doc_length = self.doc_lengths[doc_id]

        for term in query_terms:
            if term not in self.term_freq:
                continue

            # Document frequency
            df = len(self.term_freq[term])

            # IDF
            idf = np.log((self.num_docs - df + 0.5) / (df + 0.5) + 1.0)

            # Term frequency in document
            tf = self.term_freq[term].get(doc_id, 0)

            # Length normalization
            length_norm = 1 - self.b + self.b * (doc_length / self.avg_doc_length)

            # BM25 formula
            score += idf * (tf * (self.k1 + 1)) / (tf + self.k1 * length_norm)

        return score

    def retrieve(self, query: str, top_k: int = 5) -> List[RetrievalResult]:
        """Retrieve top-k documents using sparse retrieval."""
        query_terms = self._tokenize(query)

        # Score all documents
        scores = []
        for doc_id in range(self.num_docs):
            score = self._bm25_score(query_terms, doc_id)
            scores.append((doc_id, score))

        # Sort by score
        scores.sort(key=lambda x: x[1], reverse=True)

        # Get top-k
        results = []
        for doc_id, score in scores[:top_k]:
            results.append(
                RetrievalResult(
                    document=self.documents[doc_id],
                    score=score,
                    document_id=doc_id
                )
            )

        return results


class HybridRetriever:
    """
    Hybrid retriever combining dense and sparse retrieval.

    Uses weighted combination of dense and sparse scores.
    """

    def __init__(
        self,
        dense_retriever: DenseRetriever,
        sparse_retriever: SparseRetriever,
        dense_weight: float = 0.7,
        sparse_weight: float = 0.3
    ):
        self.dense_retriever = dense_retriever
        self.sparse_retriever = sparse_retriever
        self.dense_weight = dense_weight
        self.sparse_weight = sparse_weight

    def retrieve(
        self,
        query: str,
        query_embeddings: torch.Tensor,
        top_k: int = 5
    ) -> List[RetrievalResult]:
        """
        Retrieve using hybrid approach.

        Args:
            query: Query string
            query_embeddings: [1, seq_len, embed_dim]
            top_k: Number of documents to retrieve

        Returns:
            results: List of retrieval results
        """
        # Dense retrieval
        dense_results = self.dense_retriever.retrieve(
            query_embeddings, top_k=top_k * 2  # Get more candidates
        )[0]

        # Sparse retrieval
        sparse_results = self.sparse_retriever.retrieve(query, top_k=top_k * 2)

        # Combine scores
        doc_scores: Dict[int, float] = {}

        # Add dense scores
        for result in dense_results:
            doc_id = result.document_id
            doc_scores[doc_id] = self.dense_weight * result.score

        # Add sparse scores (normalize to [0, 1] range first)
        max_sparse_score = max([r.score for r in sparse_results]) if sparse_results else 1.0
        for result in sparse_results:
            doc_id = result.document_id
            normalized_score = result.score / max_sparse_score
            doc_scores[doc_id] = doc_scores.get(doc_id, 0) + self.sparse_weight * normalized_score

        # Sort by combined score
        sorted_docs = sorted(doc_scores.items(), key=lambda x: x[1], reverse=True)

        # Get top-k
        results = []
        for doc_id, score in sorted_docs[:top_k]:
            results.append(
                RetrievalResult(
                    document=self.dense_retriever.documents[doc_id],
                    score=score,
                    document_id=doc_id
                )
            )

        return results
