"""
ReAct: Reasoning and Acting framework.

Combines reasoning traces and task-specific actions in an interleaved manner.

Reference: ReAct: Synergizing Reasoning and Acting in Language Models
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import List, Dict, Optional, Tuple, Any, Callable
from dataclasses import dataclass, field
from enum import Enum

from .tool_schema import ToolRegistry, Tool
from .tool_executor import ToolExecutor, ToolResult, ExecutionStatus


class StepType(Enum):
    """Type of step in ReAct trace."""
    THOUGHT = "Thought"
    ACTION = "Action"
    OBSERVATION = "Observation"


@dataclass
class ReActStep:
    """Single step in ReAct trace."""
    step_type: StepType
    content: str
    step_number: int
    metadata: Optional[Dict[str, Any]] = None


@dataclass
class ReActTrace:
    """Complete ReAct reasoning and acting trace."""
    question: str
    steps: List[ReActStep] = field(default_factory=list)
    answer: Optional[str] = None
    success: bool = False
    total_steps: int = 0

    def add_thought(self, thought: str, step_number: int):
        """Add a thought step."""
        self.steps.append(ReActStep(
            step_type=StepType.THOUGHT,
            content=thought,
            step_number=step_number
        ))
        self.total_steps += 1

    def add_action(self, action: str, step_number: int, metadata: Optional[Dict] = None):
        """Add an action step."""
        self.steps.append(ReActStep(
            step_type=StepType.ACTION,
            content=action,
            step_number=step_number,
            metadata=metadata
        ))
        self.total_steps += 1

    def add_observation(self, observation: str, step_number: int):
        """Add an observation step."""
        self.steps.append(ReActStep(
            step_type=StepType.OBSERVATION,
            content=observation,
            step_number=step_number
        ))
        self.total_steps += 1

    def format_trace(self) -> str:
        """Format trace as string."""
        lines = [f"Question: {self.question}\n"]

        for step in self.steps:
            lines.append(f"{step.step_type.value} {step.step_number}: {step.content}")

        if self.answer:
            lines.append(f"\nAnswer: {self.answer}")

        return "\n".join(lines)


class ThoughtGenerator(nn.Module):
    """
    Generates reasoning thoughts in ReAct framework.

    Given context, generates next thought about what to do.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        hidden_dim: int = 1024,
        num_layers: int = 4,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim

        # Thought generation network
        self.thought_generator = nn.TransformerDecoder(
            nn.TransformerDecoderLayer(
                d_model=hidden_dim,
                nhead=8,
                dim_feedforward=hidden_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            ),
            num_layers=num_layers
        )

        # Projections
        self.input_proj = nn.Linear(embedding_dim, hidden_dim)
        self.output_proj = nn.Linear(hidden_dim, embedding_dim)

        # Thought type classifier
        self.thought_type = nn.Sequential(
            nn.Linear(embedding_dim, 512),
            nn.LayerNorm(512),
            nn.GELU(),
            nn.Linear(512, 5)  # [analyze, plan, execute, verify, conclude]
        )

    def forward(
        self,
        context_embeddings: torch.Tensor,
        max_length: int = 64
    ) -> torch.Tensor:
        """
        Generate thought embeddings.

        Args:
            context_embeddings: [batch, seq_len, embed_dim] current context
            max_length: Maximum thought length

        Returns:
            thought_embeddings: [batch, thought_len, embed_dim]
        """
        # Project context to hidden
        memory = self.input_proj(context_embeddings)

        # Initialize with context mean
        start = context_embeddings.mean(dim=1, keepdim=True)
        current = self.input_proj(start)

        # Generate thought autoregressively
        generated = [current]

        for _ in range(max_length - 1):
            output = self.thought_generator(current, memory)
            next_hidden = output[:, -1:, :]
            generated.append(next_hidden)
            current = torch.cat([current, next_hidden], dim=1)

        # Concatenate and project back
        all_hidden = torch.cat(generated, dim=1)
        thought_embeddings = self.output_proj(all_hidden)

        return thought_embeddings

    def classify_thought_type(self, thought_embeddings: torch.Tensor) -> List[str]:
        """
        Classify the type of thought.

        Args:
            thought_embeddings: [batch, seq_len, embed_dim]

        Returns:
            thought_types: List of thought type names
        """
        # Pool thought
        thought_vec = thought_embeddings.mean(dim=1)

        # Classify
        logits = self.thought_type(thought_vec)
        type_indices = logits.argmax(dim=-1)

        # Map to names
        type_names = ["analyze", "plan", "execute", "verify", "conclude"]
        return [type_names[idx.item()] for idx in type_indices]


class ActionSelector(nn.Module):
    """
    Selects actions based on thoughts and available tools.

    Maps thoughts to concrete tool calls.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        hidden_dim: int = 512,
        max_tools: int = 20,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.max_tools = max_tools

        # Action selection network
        self.action_selector = nn.Sequential(
            nn.Linear(embedding_dim * 2, hidden_dim),  # Thought + context
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, hidden_dim // 2),
            nn.LayerNorm(hidden_dim // 2),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim // 2, max_tools)
        )

        # Action confidence estimator
        self.confidence = nn.Sequential(
            nn.Linear(embedding_dim * 2, hidden_dim),
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Linear(hidden_dim, 1),
            nn.Sigmoid()
        )

    def forward(
        self,
        thought_embeddings: torch.Tensor,
        context_embeddings: torch.Tensor,
        available_tools: List[str]
    ) -> Tuple[List[str], torch.Tensor, torch.Tensor]:
        """
        Select action based on thought.

        Args:
            thought_embeddings: [batch, thought_len, embed_dim]
            context_embeddings: [batch, context_len, embed_dim]
            available_tools: List of available tool names

        Returns:
            selected_actions: List of selected tool names
            action_probs: [batch, num_tools] probabilities
            confidence_scores: [batch] confidence scores
        """
        # Pool embeddings
        thought_vec = thought_embeddings.mean(dim=1)
        context_vec = context_embeddings.mean(dim=1)

        # Combine
        combined = torch.cat([thought_vec, context_vec], dim=-1)

        # Select action
        logits = self.action_selector(combined)
        probs = F.softmax(logits, dim=-1)

        # Get confidence
        confidence_scores = self.confidence(combined).squeeze(-1)

        # Map to tool names
        top_indices = probs.argmax(dim=-1)
        selected_actions = []

        for idx in top_indices:
            if idx.item() < len(available_tools):
                selected_actions.append(available_tools[idx.item()])
            else:
                selected_actions.append(None)

        return selected_actions, probs, confidence_scores


class ReActAgent(nn.Module):
    """
    ReAct agent: Combines reasoning and acting.

    Interleaves thoughts, actions, and observations to solve tasks.
    """

    def __init__(
        self,
        tool_registry: ToolRegistry,
        embedding_dim: int = 768,
        max_steps: int = 10,
        confidence_threshold: float = 0.7,
        dropout: float = 0.1
    ):
        super().__init__()

        self.tool_registry = tool_registry
        self.embedding_dim = embedding_dim
        self.max_steps = max_steps
        self.confidence_threshold = confidence_threshold

        # Components
        self.thought_generator = ThoughtGenerator(embedding_dim, dropout=dropout)
        self.action_selector = ActionSelector(embedding_dim, dropout=dropout)

        # Tool executor
        self.executor = ToolExecutor(tool_registry)

        # Context encoder (maintains running context)
        self.context_encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(
                d_model=embedding_dim,
                nhead=8,
                dim_feedforward=embedding_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            ),
            num_layers=2
        )

        # Termination predictor
        self.termination = nn.Sequential(
            nn.Linear(embedding_dim, 512),
            nn.LayerNorm(512),
            nn.GELU(),
            nn.Linear(512, 2)  # [continue, terminate]
        )

    def should_terminate(self, context_embeddings: torch.Tensor) -> Tuple[bool, float]:
        """
        Decide whether to terminate reasoning.

        Args:
            context_embeddings: [batch, seq_len, embed_dim]

        Returns:
            should_stop: Whether to stop
            confidence: Confidence in termination
        """
        context_vec = context_embeddings.mean(dim=1)
        logits = self.termination(context_vec)
        probs = F.softmax(logits, dim=-1)

        should_stop = probs[:, 1] > self.confidence_threshold
        confidence = probs[:, 1]

        return should_stop[0].item(), confidence[0].item()

    def update_context(
        self,
        context_embeddings: torch.Tensor,
        new_embeddings: torch.Tensor
    ) -> torch.Tensor:
        """
        Update context with new information.

        Args:
            context_embeddings: [batch, context_len, embed_dim]
            new_embeddings: [batch, new_len, embed_dim]

        Returns:
            updated_context: [batch, total_len, embed_dim]
        """
        # Concatenate
        combined = torch.cat([context_embeddings, new_embeddings], dim=1)

        # Encode to update
        updated = self.context_encoder(combined)

        return updated

    def forward(
        self,
        question_embeddings: torch.Tensor,
        decode_fn: Callable,
        encode_fn: Callable
    ) -> ReActTrace:
        """
        Run ReAct loop to solve task.

        Args:
            question_embeddings: [1, seq_len, embed_dim] question
            decode_fn: Function to decode embeddings to text
            encode_fn: Function to encode text to embeddings

        Returns:
            trace: ReActTrace with full reasoning path
        """
        # Initialize trace
        question_text = decode_fn(question_embeddings)
        trace = ReActTrace(question=question_text)

        # Initialize context
        context = question_embeddings
        step_num = 1

        for _ in range(self.max_steps):
            # Step 1: Generate thought
            thought_embeddings = self.thought_generator(context)
            thought_text = decode_fn(thought_embeddings)
            trace.add_thought(thought_text, step_num)

            # Check termination
            should_stop, term_confidence = self.should_terminate(context)
            if should_stop and step_num > 1:  # Need at least one action
                break

            # Step 2: Select action
            available_tools = self.tool_registry.list_tools()
            selected_actions, _, confidence = self.action_selector(
                thought_embeddings, context, available_tools
            )

            if selected_actions[0] is None or confidence[0] < self.confidence_threshold:
                # No confident action, conclude
                break

            action_name = selected_actions[0]
            action_text = f"{action_name}()"
            trace.add_action(action_text, step_num)

            # Step 3: Execute action
            tool_result = self.executor.execute_call(
                type('ToolCall', (), {'tool_name': action_name, 'arguments': {}})()
            )

            # Step 4: Observe result
            observation_text = self.executor.format_result(tool_result)
            trace.add_observation(observation_text, step_num)

            # Update context with observation
            observation_embeddings = encode_fn(observation_text)
            context = self.update_context(context, observation_embeddings)

            step_num += 1

        # Generate final answer
        final_thought = self.thought_generator(context)
        answer_text = decode_fn(final_thought)
        trace.answer = answer_text
        trace.success = True

        return trace


class SelfAskAgent(nn.Module):
    """
    Self-Ask: Decomposes questions into sub-questions.

    Follows up on initial questions with follow-up questions before answering.
    """

    def __init__(
        self,
        tool_registry: ToolRegistry,
        embedding_dim: int = 768,
        max_depth: int = 3,
        dropout: float = 0.1
    ):
        super().__init__()

        self.tool_registry = tool_registry
        self.embedding_dim = embedding_dim
        self.max_depth = max_depth

        # Sub-question generator
        self.subquestion_generator = nn.TransformerDecoder(
            nn.TransformerDecoderLayer(
                d_model=embedding_dim,
                nhead=8,
                dim_feedforward=embedding_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            ),
            num_layers=4
        )

        # Question complexity estimator
        self.complexity = nn.Sequential(
            nn.Linear(embedding_dim, 512),
            nn.LayerNorm(512),
            nn.GELU(),
            nn.Linear(512, 1),
            nn.Sigmoid()
        )

        # ReAct agent for answering sub-questions
        self.react_agent = ReActAgent(tool_registry, embedding_dim)

    def estimate_complexity(self, question_embeddings: torch.Tensor) -> float:
        """Estimate question complexity (0=simple, 1=complex)."""
        question_vec = question_embeddings.mean(dim=1)
        complexity_score = self.complexity(question_vec)
        return complexity_score[0].item()

    def generate_subquestions(
        self,
        question_embeddings: torch.Tensor,
        num_subquestions: int = 2
    ) -> List[torch.Tensor]:
        """
        Generate sub-questions.

        Args:
            question_embeddings: [1, seq_len, embed_dim]
            num_subquestions: Number of sub-questions to generate

        Returns:
            subquestion_embeddings: List of sub-question embeddings
        """
        memory = question_embeddings
        subquestions = []

        for i in range(num_subquestions):
            # Generate sub-question
            start = question_embeddings.mean(dim=1, keepdim=True)
            current = start

            # Generate autoregressively
            generated = [current]
            for _ in range(31):  # 32 tokens total
                output = self.subquestion_generator(current, memory)
                next_token = output[:, -1:, :]
                generated.append(next_token)
                current = torch.cat([current, next_token], dim=1)

            subquestion = torch.cat(generated, dim=1)
            subquestions.append(subquestion)

        return subquestions

    def forward(
        self,
        question_embeddings: torch.Tensor,
        decode_fn: Callable,
        encode_fn: Callable,
        depth: int = 0
    ) -> Dict[str, Any]:
        """
        Self-Ask decomposition and answering.

        Args:
            question_embeddings: [1, seq_len, embed_dim]
            decode_fn: Decode embeddings to text
            encode_fn: Encode text to embeddings
            depth: Current recursion depth

        Returns:
            result: Dict with question, sub-questions, and answer
        """
        question_text = decode_fn(question_embeddings)

        # Base case: simple question or max depth
        complexity = self.estimate_complexity(question_embeddings)

        if complexity < 0.5 or depth >= self.max_depth:
            # Answer directly using ReAct
            trace = self.react_agent(question_embeddings, decode_fn, encode_fn)
            return {
                "question": question_text,
                "complexity": complexity,
                "sub_questions": [],
                "answer": trace.answer,
                "trace": trace
            }

        # Recursive case: decompose
        num_subs = min(2, self.max_depth - depth)
        subquestion_embeddings = self.generate_subquestions(
            question_embeddings, num_subquestions=num_subs
        )

        # Recursively answer sub-questions
        sub_results = []
        for sub_emb in subquestion_embeddings:
            sub_result = self.forward(sub_emb, decode_fn, encode_fn, depth + 1)
            sub_results.append(sub_result)

        # Combine sub-answers to answer original question
        # Concatenate context
        context_parts = [question_embeddings]
        for sub_result in sub_results:
            if sub_result["answer"]:
                answer_emb = encode_fn(sub_result["answer"])
                context_parts.append(answer_emb)

        combined_context = torch.cat(context_parts, dim=1)

        # Generate final answer
        final_trace = self.react_agent(combined_context, decode_fn, encode_fn)

        return {
            "question": question_text,
            "complexity": complexity,
            "sub_questions": sub_results,
            "answer": final_trace.answer,
            "trace": final_trace
        }


class ChainOfThought(nn.Module):
    """
    Chain-of-Thought reasoning.

    Generates explicit reasoning steps before answering.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        num_reasoning_steps: int = 5,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.num_reasoning_steps = num_reasoning_steps

        # Reasoning chain generator
        self.chain_generator = nn.TransformerDecoder(
            nn.TransformerDecoderLayer(
                d_model=embedding_dim,
                nhead=8,
                dim_feedforward=embedding_dim * 4,
                dropout=dropout,
                activation='gelu',
                batch_first=True
            ),
            num_layers=6
        )

    def forward(
        self,
        question_embeddings: torch.Tensor,
        max_length: int = 256
    ) -> torch.Tensor:
        """
        Generate chain of thought.

        Args:
            question_embeddings: [batch, seq_len, embed_dim]
            max_length: Maximum chain length

        Returns:
            chain_embeddings: [batch, chain_len, embed_dim]
        """
        memory = question_embeddings
        start = question_embeddings.mean(dim=1, keepdim=True)
        current = start

        # Generate chain
        generated = [current]

        for _ in range(max_length - 1):
            output = self.chain_generator(current, memory)
            next_token = output[:, -1:, :]
            generated.append(next_token)
            current = torch.cat([current, next_token], dim=1)

        chain_embeddings = torch.cat(generated, dim=1)
        return chain_embeddings
