"""
HyDE: Hypothetical Document Embeddings.

Generates hypothetical documents from queries, then uses them for retrieval.
This improves retrieval by bridging the query-document gap.

Reference: Precise Zero-Shot Dense Retrieval without Relevance Labels (HyDE)
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import List, Dict, Optional, Tuple, Any
from dataclasses import dataclass

from .base_retriever import DenseRetriever, RetrievalResult


@dataclass
class HyDEOutput:
    """Output from HyDE generation."""
    hypothetical_docs: List[str]
    retrieved_docs: List[RetrievalResult]
    hyde_scores: List[float]
    original_query: str


class HypotheticalDocGenerator(nn.Module):
    """
    Generator for creating hypothetical documents from queries.

    Uses a language model to generate documents that would answer the query.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        hidden_dim: int = 1024,
        num_layers: int = 6,
        num_heads: int = 8,
        dropout: float = 0.1,
        max_length: int = 256
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.hidden_dim = hidden_dim
        self.max_length = max_length

        # Query encoder
        self.query_encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(
                d_model=hidden_dim,
                nhead=num_heads,
                dim_feedforward=hidden_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            ),
            num_layers=num_layers // 2
        )

        # Document decoder
        self.doc_decoder = nn.TransformerDecoder(
            nn.TransformerDecoderLayer(
                d_model=hidden_dim,
                nhead=num_heads,
                dim_feedforward=hidden_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            ),
            num_layers=num_layers // 2
        )

        # Input projection
        self.input_proj = nn.Linear(embedding_dim, hidden_dim)

        # Output projection
        self.output_proj = nn.Linear(hidden_dim, embedding_dim)

        # Positional encoding
        self.pos_encoding = nn.Parameter(
            torch.randn(1, max_length, hidden_dim) * 0.02
        )

    def encode_query(self, query_embeddings: torch.Tensor) -> torch.Tensor:
        """
        Encode query into memory for document generation.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]

        Returns:
            memory: [batch, seq_len, hidden_dim]
        """
        # Project to hidden dimension
        hidden = self.input_proj(query_embeddings)

        # Add positional encoding
        seq_len = hidden.size(1)
        hidden = hidden + self.pos_encoding[:, :seq_len, :]

        # Encode
        memory = self.query_encoder(hidden)

        return memory

    def generate_step(
        self,
        memory: torch.Tensor,
        current_embeddings: torch.Tensor,
        position: int
    ) -> torch.Tensor:
        """
        Single generation step.

        Args:
            memory: [batch, mem_len, hidden_dim] encoded query
            current_embeddings: [batch, cur_len, embed_dim] current generated sequence
            position: Current position

        Returns:
            next_embedding: [batch, embed_dim]
        """
        # Project to hidden
        hidden = self.input_proj(current_embeddings)

        # Add positional encoding
        cur_len = hidden.size(1)
        hidden = hidden + self.pos_encoding[:, :cur_len, :]

        # Decode
        output = self.doc_decoder(hidden, memory)

        # Take last position
        last_hidden = output[:, -1, :]  # [batch, hidden_dim]

        # Project back to embedding dimension
        next_embedding = self.output_proj(last_hidden)

        return next_embedding

    def generate(
        self,
        query_embeddings: torch.Tensor,
        num_docs: int = 3,
        temperature: float = 0.8,
        top_p: float = 0.9
    ) -> torch.Tensor:
        """
        Generate hypothetical documents.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            num_docs: Number of hypothetical documents to generate
            temperature: Sampling temperature
            top_p: Nucleus sampling parameter

        Returns:
            doc_embeddings: [batch, num_docs, doc_len, embed_dim]
        """
        batch_size = query_embeddings.size(0)
        device = query_embeddings.device

        # Encode query
        memory = self.encode_query(query_embeddings)

        # Expand memory for multiple documents
        memory = memory.unsqueeze(1).expand(-1, num_docs, -1, -1)
        memory = memory.reshape(batch_size * num_docs, -1, self.hidden_dim)

        # Initialize with query mean as start token
        start_embed = query_embeddings.mean(dim=1, keepdim=True)  # [batch, 1, embed_dim]
        start_embed = start_embed.unsqueeze(1).expand(-1, num_docs, -1, -1)
        start_embed = start_embed.reshape(batch_size * num_docs, 1, self.embedding_dim)

        current_embeds = start_embed

        # Generate document embeddings autoregressively
        generated_embeds = [start_embed]

        for pos in range(self.max_length - 1):
            # Generate next embedding
            next_embed = self.generate_step(memory, current_embeds, pos)

            # Add noise for diversity (controlled by temperature)
            if temperature > 0:
                noise = torch.randn_like(next_embed) * temperature
                next_embed = next_embed + noise

            # Normalize
            next_embed = F.normalize(next_embed, p=2, dim=-1)

            # Add to sequence
            next_embed = next_embed.unsqueeze(1)  # [batch*num_docs, 1, embed_dim]
            generated_embeds.append(next_embed)
            current_embeds = torch.cat([current_embeds, next_embed], dim=1)

        # Concatenate all embeddings
        all_embeds = torch.cat(generated_embeds, dim=1)  # [batch*num_docs, doc_len, embed_dim]

        # Reshape to separate batch and num_docs
        all_embeds = all_embeds.view(batch_size, num_docs, -1, self.embedding_dim)

        return all_embeds


class HyDE(nn.Module):
    """
    HyDE: Hypothetical Document Embeddings for improved retrieval.

    Pipeline:
    1. Generate hypothetical documents from query
    2. Encode hypothetical documents
    3. Retrieve using hypothetical document embeddings
    4. Optional: Rerank using original query
    """

    def __init__(
        self,
        retriever: DenseRetriever,
        embedding_dim: int = 768,
        num_hypothetical_docs: int = 3,
        aggregation_method: str = "mean",
        use_reranking: bool = True,
        temperature: float = 0.8
    ):
        super().__init__()

        self.retriever = retriever
        self.embedding_dim = embedding_dim
        self.num_hypothetical_docs = num_hypothetical_docs
        self.aggregation_method = aggregation_method
        self.use_reranking = use_reranking
        self.temperature = temperature

        # Hypothetical document generator
        self.hyde_generator = HypotheticalDocGenerator(
            embedding_dim=embedding_dim
        )

        # Reranker (uses original query to rerank HyDE results)
        if use_reranking:
            self.reranker = nn.Sequential(
                nn.Linear(embedding_dim * 3, embedding_dim),  # query + hyde_doc + retrieved_doc
                nn.LayerNorm(embedding_dim),
                nn.GELU(),
                nn.Linear(embedding_dim, 1)
            )

    def aggregate_embeddings(
        self,
        embeddings: torch.Tensor,
        method: str = "mean"
    ) -> torch.Tensor:
        """
        Aggregate multiple document embeddings.

        Args:
            embeddings: [batch, num_docs, seq_len, embed_dim]
            method: Aggregation method ("mean", "max", "weighted")

        Returns:
            aggregated: [batch, seq_len, embed_dim]
        """
        if method == "mean":
            # Average across documents
            return embeddings.mean(dim=1)

        elif method == "max":
            # Max pooling across documents
            return embeddings.max(dim=1)[0]

        elif method == "weighted":
            # Weighted average based on self-attention
            batch, num_docs, seq_len, embed_dim = embeddings.shape

            # Compute attention weights
            flat_embeds = embeddings.view(batch, num_docs, -1)
            query = flat_embeds.mean(dim=2, keepdim=True)  # [batch, num_docs, 1]

            # Attention scores
            scores = torch.matmul(flat_embeds, query.transpose(1, 2))  # [batch, num_docs, num_docs]
            weights = F.softmax(scores.mean(dim=2), dim=1)  # [batch, num_docs]

            # Weighted sum
            weights = weights.view(batch, num_docs, 1, 1)
            aggregated = (embeddings * weights).sum(dim=1)

            return aggregated

        else:
            raise ValueError(f"Unknown aggregation method: {method}")

    def forward(
        self,
        query_embeddings: torch.Tensor,
        top_k: int = 5
    ) -> Tuple[List[List[RetrievalResult]], torch.Tensor]:
        """
        Forward pass for HyDE retrieval.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            top_k: Number of documents to retrieve

        Returns:
            results: List of retrieval results for each query
            hyde_embeddings: [batch, num_hyde_docs, doc_len, embed_dim]
        """
        batch_size = query_embeddings.size(0)

        # Step 1: Generate hypothetical documents
        hyde_embeddings = self.hyde_generator.generate(
            query_embeddings,
            num_docs=self.num_hypothetical_docs,
            temperature=self.temperature
        )  # [batch, num_hyde_docs, doc_len, embed_dim]

        # Step 2: Aggregate hypothetical documents
        aggregated_hyde = self.aggregate_embeddings(
            hyde_embeddings, method=self.aggregation_method
        )  # [batch, doc_len, embed_dim]

        # Step 3: Retrieve using aggregated HyDE embeddings
        results = self.retriever.retrieve(
            aggregated_hyde, top_k=top_k * 2 if self.use_reranking else top_k
        )

        # Step 4: Optional reranking with original query
        if self.use_reranking:
            reranked_results = self.rerank_results(
                query_embeddings, aggregated_hyde, results, top_k
            )
            results = reranked_results

        return results, hyde_embeddings

    def rerank_results(
        self,
        query_embeddings: torch.Tensor,
        hyde_embeddings: torch.Tensor,
        results: List[List[RetrievalResult]],
        top_k: int
    ) -> List[List[RetrievalResult]]:
        """
        Rerank retrieval results using original query.

        Args:
            query_embeddings: [batch, seq_len, embed_dim] original query
            hyde_embeddings: [batch, seq_len, embed_dim] aggregated HyDE
            results: Initial retrieval results
            top_k: Number of final results

        Returns:
            reranked_results: Reranked retrieval results
        """
        batch_size = query_embeddings.size(0)
        reranked_results = []

        with torch.no_grad():
            for i in range(batch_size):
                query_vec = query_embeddings[i].mean(dim=0)  # [embed_dim]
                hyde_vec = hyde_embeddings[i].mean(dim=0)  # [embed_dim]

                # Get document IDs from results
                doc_results = results[i]

                # Rerank each document
                rerank_scores = []
                for result in doc_results:
                    # Get document embedding from retriever
                    doc_id = result.document_id
                    doc_vec = self.retriever.document_embeddings[doc_id]  # [embed_dim]

                    # Concatenate features
                    features = torch.cat([query_vec, hyde_vec, doc_vec], dim=0).unsqueeze(0)

                    # Compute rerank score
                    score = self.reranker(features).item()
                    rerank_scores.append(score)

                # Sort by rerank score
                sorted_indices = sorted(
                    range(len(doc_results)),
                    key=lambda idx: rerank_scores[idx],
                    reverse=True
                )

                # Take top-k
                reranked = [doc_results[idx] for idx in sorted_indices[:top_k]]

                # Update scores
                for j, result in enumerate(reranked):
                    result.score = rerank_scores[sorted_indices[j]]

                reranked_results.append(reranked)

        return reranked_results

    def retrieve(
        self,
        query_embeddings: torch.Tensor,
        top_k: int = 5,
        return_hyde: bool = False
    ) -> HyDEOutput:
        """
        High-level retrieval interface.

        Args:
            query_embeddings: [1, seq_len, embed_dim]
            top_k: Number of documents to retrieve
            return_hyde: Whether to return HyDE embeddings

        Returns:
            output: HyDEOutput with retrieval results
        """
        # Generate and retrieve
        results, hyde_embeddings = self.forward(query_embeddings, top_k=top_k)

        # Extract hypothetical documents (would be decoded from embeddings in practice)
        hypothetical_docs = [
            f"Hypothetical document {i+1} for query"
            for i in range(self.num_hypothetical_docs)
        ]

        # Compute HyDE quality scores (how confident we are in each hypothetical doc)
        hyde_scores = []
        with torch.no_grad():
            for i in range(self.num_hypothetical_docs):
                # Score based on similarity to query
                hyde_doc = hyde_embeddings[0, i].mean(dim=0)
                query_vec = query_embeddings[0].mean(dim=0)
                score = F.cosine_similarity(
                    hyde_doc.unsqueeze(0), query_vec.unsqueeze(0), dim=1
                ).item()
                hyde_scores.append(score)

        return HyDEOutput(
            hypothetical_docs=hypothetical_docs,
            retrieved_docs=results[0],
            hyde_scores=hyde_scores,
            original_query="<query>"
        )


class MultiHopHyDE(nn.Module):
    """
    Multi-hop HyDE for complex queries requiring multiple retrieval steps.

    Iteratively generates hypothetical documents and retrieves,
    building up knowledge for complex questions.
    """

    def __init__(
        self,
        hyde: HyDE,
        max_hops: int = 3,
        aggregation_method: str = "concatenate"
    ):
        super().__init__()

        self.hyde = hyde
        self.max_hops = max_hops
        self.aggregation_method = aggregation_method

        # Context aggregator
        self.context_aggregator = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(
                d_model=hyde.embedding_dim,
                nhead=8,
                dim_feedforward=hyde.embedding_dim * 4,
                dropout=0.1,
                activation='gelu',
                batch_first=True
            ),
            num_layers=2
        )

    def forward(
        self,
        query_embeddings: torch.Tensor,
        top_k_per_hop: int = 3,
        final_top_k: int = 5
    ) -> Tuple[List[RetrievalResult], List[List[RetrievalResult]]]:
        """
        Multi-hop retrieval with HyDE.

        Args:
            query_embeddings: [1, seq_len, embed_dim]
            top_k_per_hop: Documents to retrieve per hop
            final_top_k: Final number of documents

        Returns:
            final_results: Final retrieval results
            hop_results: Results from each hop
        """
        current_query = query_embeddings
        all_retrieved_docs = []
        hop_results = []

        for hop in range(self.max_hops):
            # Retrieve using HyDE
            results, hyde_embeds = self.hyde(current_query, top_k=top_k_per_hop)
            hop_results.append(results[0])

            # Collect retrieved documents
            all_retrieved_docs.extend(results[0])

            # Update query based on retrieved documents
            if hop < self.max_hops - 1:
                # Aggregate retrieved document information
                # In practice, would encode actual document text
                doc_embeds = hyde_embeds[0]  # Use HyDE embeddings as proxy

                # Combine with original query
                combined = torch.cat([current_query, doc_embeds.mean(dim=0, keepdim=True).unsqueeze(0)], dim=1)

                # Process with context aggregator
                current_query = self.context_aggregator(combined)

        # Deduplicate and rank final results
        seen_docs = set()
        unique_results = []

        for result in all_retrieved_docs:
            doc_id = result.document_id
            if doc_id not in seen_docs:
                seen_docs.add(doc_id)
                unique_results.append(result)

        # Sort by score and take top-k
        unique_results.sort(key=lambda x: x.score, reverse=True)
        final_results = unique_results[:final_top_k]

        return final_results, hop_results
