"""
Self-RAG: Self-Reflective Retrieval Augmented Generation.

Implements the Self-RAG framework that uses reflection tokens to:
1. Decide when to retrieve (Retrieve token)
2. Assess retrieval relevance (ISREL token)
3. Evaluate response support (ISSUP token)
4. Judge response usefulness (ISUSE token)

Reference: Self-RAG: Learning to Retrieve, Generate, and Critique through Self-Reflection
"""

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 enum import Enum

from .base_retriever import DenseRetriever, RetrievalResult


class ReflectionToken(Enum):
    """Reflection tokens for Self-RAG."""
    # Retrieval decision
    RETRIEVE_YES = "[Retrieve=Yes]"
    RETRIEVE_NO = "[Retrieve=No]"

    # Relevance assessment
    ISREL_RELEVANT = "[ISREL=Relevant]"
    ISREL_IRRELEVANT = "[ISREL=Irrelevant]"

    # Support assessment
    ISSUP_FULLY = "[ISSUP=Fully]"
    ISSUP_PARTIALLY = "[ISSUP=Partially]"
    ISSUP_NOT = "[ISSUP=Not]"

    # Usefulness assessment
    ISUSE_5 = "[ISUSE=5]"  # Very useful
    ISUSE_4 = "[ISUSE=4]"
    ISUSE_3 = "[ISUSE=3]"
    ISUSE_2 = "[ISUSE=2]"
    ISUSE_1 = "[ISUSE=1]"  # Not useful


@dataclass
class SelfRAGOutput:
    """Output from Self-RAG generation."""
    response: str
    retrieve_decision: bool
    retrieved_docs: Optional[List[RetrievalResult]]
    relevance_scores: Optional[List[float]]
    support_scores: Optional[List[float]]
    usefulness_score: Optional[float]
    reflection_tokens: List[str]


class RetrievalCritic(nn.Module):
    """
    Critic that decides when to retrieve and evaluates retrieval quality.

    Produces:
    - Retrieve decision: Should we retrieve documents?
    - Relevance scores: Are retrieved documents relevant?
    """

    def __init__(
        self,
        hidden_dim: int = 768,
        num_layers: int = 2,
        dropout: float = 0.1
    ):
        super().__init__()

        # Retrieval decision head
        self.retrieve_head = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim // 2),
            nn.LayerNorm(hidden_dim // 2),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim // 2, 2)  # [No, Yes]
        )

        # Relevance assessment head
        self.relevance_head = nn.Sequential(
            nn.Linear(hidden_dim * 2, hidden_dim),  # Query + doc
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, 2)  # [Irrelevant, Relevant]
        )

    def forward(
        self,
        query_hidden: torch.Tensor,
        doc_hidden: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Forward pass.

        Args:
            query_hidden: [batch, hidden_dim] query representation
            doc_hidden: [batch, num_docs, hidden_dim] document representations

        Returns:
            retrieve_logits: [batch, 2] retrieval decision logits
            relevance_logits: [batch, num_docs, 2] relevance logits (if docs provided)
        """
        # Retrieval decision
        retrieve_logits = self.retrieve_head(query_hidden)  # [batch, 2]

        # Relevance assessment
        relevance_logits = None
        if doc_hidden is not None:
            batch_size, num_docs, hidden_dim = doc_hidden.shape

            # Expand query for each document
            query_expanded = query_hidden.unsqueeze(1).expand(-1, num_docs, -1)

            # Concatenate query and document
            combined = torch.cat([query_expanded, doc_hidden], dim=-1)

            # Compute relevance
            relevance_logits = self.relevance_head(combined)  # [batch, num_docs, 2]

        return retrieve_logits, relevance_logits


class ResponseCritic(nn.Module):
    """
    Critic that evaluates generated responses.

    Produces:
    - Support scores: Is response supported by retrieved documents?
    - Usefulness scores: Is response useful for the query?
    """

    def __init__(
        self,
        hidden_dim: int = 768,
        dropout: float = 0.1
    ):
        super().__init__()

        # Support assessment head
        self.support_head = nn.Sequential(
            nn.Linear(hidden_dim * 3, hidden_dim),  # Query + doc + response
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, 3)  # [Not, Partially, Fully]
        )

        # Usefulness assessment head
        self.usefulness_head = nn.Sequential(
            nn.Linear(hidden_dim * 2, hidden_dim),  # Query + response
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, 5)  # [1, 2, 3, 4, 5]
        )

    def forward(
        self,
        query_hidden: torch.Tensor,
        response_hidden: torch.Tensor,
        doc_hidden: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Forward pass.

        Args:
            query_hidden: [batch, hidden_dim]
            response_hidden: [batch, hidden_dim]
            doc_hidden: [batch, num_docs, hidden_dim] (optional)

        Returns:
            usefulness_logits: [batch, 5]
            support_logits: [batch, num_docs, 3] (if docs provided)
        """
        # Usefulness assessment
        query_response = torch.cat([query_hidden, response_hidden], dim=-1)
        usefulness_logits = self.usefulness_head(query_response)  # [batch, 5]

        # Support assessment
        support_logits = None
        if doc_hidden is not None:
            batch_size, num_docs, hidden_dim = doc_hidden.shape

            # Expand query and response
            query_expanded = query_hidden.unsqueeze(1).expand(-1, num_docs, -1)
            response_expanded = response_hidden.unsqueeze(1).expand(-1, num_docs, -1)

            # Concatenate
            combined = torch.cat([query_expanded, doc_hidden, response_expanded], dim=-1)

            # Compute support
            support_logits = self.support_head(combined)  # [batch, num_docs, 3]

        return usefulness_logits, support_logits


class SelfRAG(nn.Module):
    """
    Self-RAG: Self-Reflective Retrieval Augmented Generation.

    Integrates retrieval and generation with self-reflection mechanisms.
    """

    def __init__(
        self,
        retriever: DenseRetriever,
        embedding_dim: int = 768,
        hidden_dim: int = 768,
        num_retrieve_docs: int = 5,
        retrieval_threshold: float = 0.5,
        relevance_threshold: float = 0.5,
        dropout: float = 0.1
    ):
        super().__init__()

        self.retriever = retriever
        self.embedding_dim = embedding_dim
        self.num_retrieve_docs = num_retrieve_docs
        self.retrieval_threshold = retrieval_threshold
        self.relevance_threshold = relevance_threshold

        # Critics
        self.retrieval_critic = RetrievalCritic(hidden_dim, dropout=dropout)
        self.response_critic = ResponseCritic(hidden_dim, dropout=dropout)

        # Document encoder for critic
        self.doc_encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(
                d_model=hidden_dim,
                nhead=8,
                dim_feedforward=hidden_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            ),
            num_layers=2
        )

    def encode_hidden(self, embeddings: torch.Tensor) -> torch.Tensor:
        """
        Encode embeddings to hidden representation.

        Args:
            embeddings: [batch, seq_len, embed_dim]

        Returns:
            hidden: [batch, hidden_dim]
        """
        # Mean pooling
        hidden = embeddings.mean(dim=1)
        return hidden

    def decide_retrieve(
        self,
        query_embeddings: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Decide whether to retrieve documents.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]

        Returns:
            retrieve_decision: [batch] boolean tensor
            retrieve_probs: [batch, 2] probabilities
        """
        # Encode query
        query_hidden = self.encode_hidden(query_embeddings)

        # Get retrieval decision
        retrieve_logits, _ = self.retrieval_critic(query_hidden)
        retrieve_probs = F.softmax(retrieve_logits, dim=-1)

        # Decision: retrieve if P(yes) > threshold
        retrieve_decision = retrieve_probs[:, 1] > self.retrieval_threshold

        return retrieve_decision, retrieve_probs

    def assess_relevance(
        self,
        query_embeddings: torch.Tensor,
        doc_embeddings: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Assess relevance of retrieved documents.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            doc_embeddings: [batch, num_docs, seq_len, embed_dim]

        Returns:
            relevance_scores: [batch, num_docs]
            relevance_logits: [batch, num_docs, 2]
        """
        batch_size, num_docs, seq_len, embed_dim = doc_embeddings.shape

        # Encode query
        query_hidden = self.encode_hidden(query_embeddings)  # [batch, hidden_dim]

        # Encode documents
        doc_flat = doc_embeddings.view(batch_size * num_docs, seq_len, embed_dim)
        doc_encoded = self.doc_encoder(doc_flat)  # [batch * num_docs, seq_len, hidden_dim]
        doc_hidden = doc_encoded.mean(dim=1)  # [batch * num_docs, hidden_dim]
        doc_hidden = doc_hidden.view(batch_size, num_docs, -1)  # [batch, num_docs, hidden_dim]

        # Assess relevance
        _, relevance_logits = self.retrieval_critic(query_hidden, doc_hidden)
        relevance_probs = F.softmax(relevance_logits, dim=-1)

        # Relevance scores (probability of being relevant)
        relevance_scores = relevance_probs[:, :, 1]  # [batch, num_docs]

        return relevance_scores, relevance_logits

    def assess_support(
        self,
        query_embeddings: torch.Tensor,
        response_embeddings: torch.Tensor,
        doc_embeddings: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Assess how well documents support the response.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            response_embeddings: [batch, seq_len, embed_dim]
            doc_embeddings: [batch, num_docs, seq_len, embed_dim]

        Returns:
            support_scores: [batch, num_docs]
            support_logits: [batch, num_docs, 3]
        """
        batch_size, num_docs, seq_len, embed_dim = doc_embeddings.shape

        # Encode
        query_hidden = self.encode_hidden(query_embeddings)
        response_hidden = self.encode_hidden(response_embeddings)

        # Encode documents
        doc_flat = doc_embeddings.view(batch_size * num_docs, seq_len, embed_dim)
        doc_encoded = self.doc_encoder(doc_flat)
        doc_hidden = doc_encoded.mean(dim=1).view(batch_size, num_docs, -1)

        # Assess support
        _, support_logits = self.response_critic(query_hidden, response_hidden, doc_hidden)
        support_probs = F.softmax(support_logits, dim=-1)

        # Support scores (weighted average: fully=1.0, partially=0.5, not=0.0)
        weights = torch.tensor([0.0, 0.5, 1.0], device=support_probs.device)
        support_scores = (support_probs * weights.view(1, 1, 3)).sum(dim=-1)

        return support_scores, support_logits

    def assess_usefulness(
        self,
        query_embeddings: torch.Tensor,
        response_embeddings: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Assess usefulness of response.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            response_embeddings: [batch, seq_len, embed_dim]

        Returns:
            usefulness_scores: [batch] averaged score in [1, 5]
            usefulness_logits: [batch, 5]
        """
        # Encode
        query_hidden = self.encode_hidden(query_embeddings)
        response_hidden = self.encode_hidden(response_embeddings)

        # Assess usefulness
        usefulness_logits, _ = self.response_critic(query_hidden, response_hidden)
        usefulness_probs = F.softmax(usefulness_logits, dim=-1)

        # Usefulness scores (expected value: sum of prob * rating)
        ratings = torch.arange(1, 6, dtype=torch.float, device=usefulness_probs.device)
        usefulness_scores = (usefulness_probs * ratings).sum(dim=-1)

        return usefulness_scores, usefulness_logits

    def forward(
        self,
        query_embeddings: torch.Tensor,
        response_embeddings: Optional[torch.Tensor] = None,
        retrieved_doc_embeddings: Optional[torch.Tensor] = None,
        mode: str = "train"
    ) -> Dict[str, torch.Tensor]:
        """
        Forward pass for Self-RAG.

        Args:
            query_embeddings: [batch, seq_len, embed_dim]
            response_embeddings: [batch, seq_len, embed_dim] (optional, for training)
            retrieved_doc_embeddings: [batch, num_docs, seq_len, embed_dim] (optional)
            mode: "train" or "inference"

        Returns:
            outputs: Dictionary of outputs and losses
        """
        outputs = {}

        # Retrieval decision
        retrieve_decision, retrieve_probs = self.decide_retrieve(query_embeddings)
        outputs["retrieve_decision"] = retrieve_decision
        outputs["retrieve_probs"] = retrieve_probs

        # Relevance assessment (if documents provided)
        if retrieved_doc_embeddings is not None:
            relevance_scores, relevance_logits = self.assess_relevance(
                query_embeddings, retrieved_doc_embeddings
            )
            outputs["relevance_scores"] = relevance_scores
            outputs["relevance_logits"] = relevance_logits

        # Response assessment (if response provided)
        if response_embeddings is not None:
            usefulness_scores, usefulness_logits = self.assess_usefulness(
                query_embeddings, response_embeddings
            )
            outputs["usefulness_scores"] = usefulness_scores
            outputs["usefulness_logits"] = usefulness_logits

            # Support assessment (if both response and documents provided)
            if retrieved_doc_embeddings is not None:
                support_scores, support_logits = self.assess_support(
                    query_embeddings, response_embeddings, retrieved_doc_embeddings
                )
                outputs["support_scores"] = support_scores
                outputs["support_logits"] = support_logits

        return outputs

    def generate_with_reflection(
        self,
        query_embeddings: torch.Tensor,
        generate_fn: callable,
        max_length: int = 512
    ) -> SelfRAGOutput:
        """
        Generate response with self-reflection.

        Args:
            query_embeddings: [1, seq_len, embed_dim]
            generate_fn: Function to generate response (takes embeddings, returns text + embeddings)
            max_length: Maximum generation length

        Returns:
            output: SelfRAGOutput with response and reflection data
        """
        reflection_tokens = []

        # Step 1: Decide whether to retrieve
        retrieve_decision, retrieve_probs = self.decide_retrieve(query_embeddings)
        should_retrieve = retrieve_decision[0].item()

        if should_retrieve:
            reflection_tokens.append(ReflectionToken.RETRIEVE_YES.value)

            # Step 2: Retrieve documents
            retrieved_docs = self.retriever.retrieve(
                query_embeddings, top_k=self.num_retrieve_docs
            )[0]

            # Get document embeddings (would come from actual documents in practice)
            # For now, use placeholder
            doc_embeddings = torch.randn(
                1, self.num_retrieve_docs, 128, self.embedding_dim,
                device=query_embeddings.device
            )

            # Step 3: Assess relevance
            relevance_scores, _ = self.assess_relevance(query_embeddings, doc_embeddings)
            relevance_list = relevance_scores[0].tolist()

            # Filter relevant documents
            relevant_mask = relevance_scores[0] > self.relevance_threshold
            relevant_docs = [
                doc for i, doc in enumerate(retrieved_docs) if relevant_mask[i]
            ]

            for score in relevance_list:
                if score > self.relevance_threshold:
                    reflection_tokens.append(ReflectionToken.ISREL_RELEVANT.value)
                else:
                    reflection_tokens.append(ReflectionToken.ISREL_IRRELEVANT.value)

        else:
            reflection_tokens.append(ReflectionToken.RETRIEVE_NO.value)
            retrieved_docs = None
            relevant_docs = None
            relevance_list = None
            doc_embeddings = None

        # Step 4: Generate response
        response_text, response_embeddings = generate_fn(
            query_embeddings, retrieved_docs, max_length
        )

        # Step 5: Assess support (if documents available)
        support_list = None
        if doc_embeddings is not None:
            support_scores, _ = self.assess_support(
                query_embeddings, response_embeddings, doc_embeddings
            )
            support_list = support_scores[0].tolist()

            # Add support tokens
            avg_support = support_scores[0].mean().item()
            if avg_support > 0.8:
                reflection_tokens.append(ReflectionToken.ISSUP_FULLY.value)
            elif avg_support > 0.4:
                reflection_tokens.append(ReflectionToken.ISSUP_PARTIALLY.value)
            else:
                reflection_tokens.append(ReflectionToken.ISSUP_NOT.value)

        # Step 6: Assess usefulness
        usefulness_score, _ = self.assess_usefulness(query_embeddings, response_embeddings)
        usefulness_val = usefulness_score[0].item()

        # Add usefulness token
        usefulness_rating = int(round(usefulness_val))
        usefulness_rating = max(1, min(5, usefulness_rating))
        reflection_tokens.append(f"[ISUSE={usefulness_rating}]")

        return SelfRAGOutput(
            response=response_text,
            retrieve_decision=should_retrieve,
            retrieved_docs=relevant_docs,
            relevance_scores=relevance_list,
            support_scores=support_list,
            usefulness_score=usefulness_val,
            reflection_tokens=reflection_tokens
        )
