"""
Tool execution engine for MK3.

Handles:
- Parsing tool calls from model output
- Executing tools safely
- Managing tool context and state
- Error handling and retries
- Toolformer-style learning
"""

import re
import json
import torch
import torch.nn as nn
from typing import List, Dict, Optional, Tuple, Any, Callable
from dataclasses import dataclass, field
from enum import Enum

from .tool_schema import Tool, ToolRegistry, ToolSchema


class ExecutionStatus(Enum):
    """Status of tool execution."""
    SUCCESS = "success"
    FAILURE = "failure"
    TIMEOUT = "timeout"
    INVALID_CALL = "invalid_call"


@dataclass
class ToolCall:
    """Represents a tool call."""
    tool_name: str
    arguments: Dict[str, Any]
    call_id: Optional[str] = None
    position: Optional[int] = None  # Position in text where call appears


@dataclass
class ToolResult:
    """Result from tool execution."""
    call: ToolCall
    status: ExecutionStatus
    result: Any = None
    error: Optional[str] = None
    execution_time: float = 0.0


@dataclass
class ExecutionTrace:
    """Trace of tool execution."""
    calls: List[ToolCall] = field(default_factory=list)
    results: List[ToolResult] = field(default_factory=list)
    total_time: float = 0.0


class ToolCallParser:
    """
    Parses tool calls from model output.

    Supports multiple formats:
    - Function call: tool_name(arg1=val1, arg2=val2)
    - JSON: {"tool": "tool_name", "args": {...}}
    - API format: [tool_name] arg1: val1, arg2: val2 [/tool_name]
    """

    def __init__(self, registry: ToolRegistry):
        self.registry = registry

        # Regex patterns for different formats
        self.function_pattern = re.compile(
            r'(\w+)\((.*?)\)'
        )

        self.api_pattern = re.compile(
            r'\[(\w+)\](.*?)\[/\1\]'
        )

    def parse_function_call(self, text: str) -> Optional[ToolCall]:
        """Parse function-style call: tool_name(arg1=val1, arg2=val2)"""
        match = self.function_pattern.search(text)
        if not match:
            return None

        tool_name = match.group(1)
        args_str = match.group(2)

        # Check if tool exists
        if self.registry.get(tool_name) is None:
            return None

        # Parse arguments
        try:
            arguments = {}
            if args_str.strip():
                # Split by commas (but not within quotes)
                arg_parts = re.split(r',(?=(?:[^"]*"[^"]*")*[^"]*$)', args_str)

                for part in arg_parts:
                    part = part.strip()
                    if '=' in part:
                        key, value = part.split('=', 1)
                        key = key.strip()
                        value = value.strip()

                        # Parse value
                        if value.startswith('"') and value.endswith('"'):
                            value = value[1:-1]  # String
                        elif value.lower() == 'true':
                            value = True
                        elif value.lower() == 'false':
                            value = False
                        elif value.isdigit():
                            value = int(value)
                        else:
                            try:
                                value = float(value)
                            except ValueError:
                                pass  # Keep as string

                        arguments[key] = value

            return ToolCall(tool_name=tool_name, arguments=arguments)

        except Exception:
            return None

    def parse_json_call(self, text: str) -> Optional[ToolCall]:
        """Parse JSON-style call: {"tool": "name", "args": {...}}"""
        try:
            # Find JSON objects in text
            json_match = re.search(r'\{[^}]+\}', text)
            if not json_match:
                return None

            data = json.loads(json_match.group())

            if "tool" in data and "args" in data:
                tool_name = data["tool"]
                arguments = data["args"]

                # Check if tool exists
                if self.registry.get(tool_name) is not None:
                    return ToolCall(tool_name=tool_name, arguments=arguments)

        except json.JSONDecodeError:
            return None

        return None

    def parse_api_call(self, text: str) -> Optional[ToolCall]:
        """Parse API-style call: [tool_name] arg1: val1, arg2: val2 [/tool_name]"""
        match = self.api_pattern.search(text)
        if not match:
            return None

        tool_name = match.group(1)
        args_str = match.group(2).strip()

        # Check if tool exists
        if self.registry.get(tool_name) is None:
            return None

        # Parse arguments
        try:
            arguments = {}
            if args_str:
                # Split by commas
                arg_parts = args_str.split(',')
                for part in arg_parts:
                    if ':' in part:
                        key, value = part.split(':', 1)
                        arguments[key.strip()] = value.strip()

            return ToolCall(tool_name=tool_name, arguments=arguments)

        except Exception:
            return None

    def parse(self, text: str) -> List[ToolCall]:
        """Parse all tool calls from text."""
        calls = []

        # Try different formats
        parsers = [
            self.parse_function_call,
            self.parse_json_call,
            self.parse_api_call
        ]

        for parser in parsers:
            call = parser(text)
            if call is not None:
                calls.append(call)

        return calls


class ToolExecutor:
    """
    Executes tools with safety and error handling.
    """

    def __init__(
        self,
        registry: ToolRegistry,
        timeout: float = 30.0,
        max_retries: int = 2,
        enable_logging: bool = True
    ):
        self.registry = registry
        self.timeout = timeout
        self.max_retries = max_retries
        self.enable_logging = enable_logging

        self.parser = ToolCallParser(registry)
        self.execution_traces: List[ExecutionTrace] = []

    def execute_call(
        self,
        call: ToolCall,
        retry_count: int = 0
    ) -> ToolResult:
        """
        Execute a single tool call.

        Args:
            call: ToolCall to execute
            retry_count: Current retry attempt

        Returns:
            result: ToolResult
        """
        import time

        start_time = time.time()

        # Get tool
        tool = self.registry.get(call.tool_name)
        if tool is None:
            return ToolResult(
                call=call,
                status=ExecutionStatus.INVALID_CALL,
                error=f"Tool '{call.tool_name}' not found"
            )

        # Execute with timeout
        try:
            result = tool(**call.arguments)
            execution_time = time.time() - start_time

            return ToolResult(
                call=call,
                status=ExecutionStatus.SUCCESS,
                result=result,
                execution_time=execution_time
            )

        except Exception as e:
            execution_time = time.time() - start_time

            # Retry if possible
            if retry_count < self.max_retries:
                return self.execute_call(call, retry_count + 1)

            return ToolResult(
                call=call,
                status=ExecutionStatus.FAILURE,
                error=str(e),
                execution_time=execution_time
            )

    def execute_from_text(self, text: str) -> ExecutionTrace:
        """
        Parse and execute all tool calls from text.

        Args:
            text: Text containing tool calls

        Returns:
            trace: ExecutionTrace with all calls and results
        """
        import time

        start_time = time.time()
        trace = ExecutionTrace()

        # Parse calls
        calls = self.parser.parse(text)

        # Execute each call
        for call in calls:
            result = self.execute_call(call)
            trace.calls.append(call)
            trace.results.append(result)

        trace.total_time = time.time() - start_time

        # Log trace
        if self.enable_logging:
            self.execution_traces.append(trace)

        return trace

    def format_result(self, result: ToolResult) -> str:
        """Format tool result as string."""
        if result.status == ExecutionStatus.SUCCESS:
            return f"[{result.call.tool_name}] Success: {result.result}"
        else:
            return f"[{result.call.tool_name}] {result.status.value}: {result.error}"


class ToolLearner(nn.Module):
    """
    Toolformer-style learning: learns when and how to use tools.

    Uses a scoring mechanism to decide whether to call a tool.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        hidden_dim: int = 512,
        num_tools: int = 10,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.num_tools = num_tools

        # Tool decision network
        self.tool_decision = nn.Sequential(
            nn.Linear(embedding_dim, hidden_dim),
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, 2)  # [no_tool, use_tool]
        )

        # Tool selection network
        self.tool_selector = nn.Sequential(
            nn.Linear(embedding_dim, hidden_dim),
            nn.LayerNorm(hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, num_tools)
        )

        # Argument generator
        self.arg_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=2
        )

        self.input_proj = nn.Linear(embedding_dim, hidden_dim)
        self.output_proj = nn.Linear(hidden_dim, embedding_dim)

    def decide_tool_use(
        self,
        context_embeddings: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Decide whether to use a tool.

        Args:
            context_embeddings: [batch, seq_len, embed_dim]

        Returns:
            use_tool: [batch] binary decision
            tool_probs: [batch, 2] probabilities
        """
        # Pool context
        context_vec = context_embeddings.mean(dim=1)  # [batch, embed_dim]

        # Decide
        logits = self.tool_decision(context_vec)
        probs = torch.softmax(logits, dim=-1)

        # Binary decision
        use_tool = probs[:, 1] > 0.5

        return use_tool, probs

    def select_tool(
        self,
        context_embeddings: torch.Tensor,
        tool_names: List[str]
    ) -> Tuple[List[str], torch.Tensor]:
        """
        Select which tool to use.

        Args:
            context_embeddings: [batch, seq_len, embed_dim]
            tool_names: List of available tool names

        Returns:
            selected_tools: List of selected tool names
            tool_probs: [batch, num_tools] probabilities
        """
        # Pool context
        context_vec = context_embeddings.mean(dim=1)

        # Select tool
        logits = self.tool_selector(context_vec)
        probs = torch.softmax(logits, dim=-1)

        # Get top tool for each batch
        top_indices = probs.argmax(dim=-1)

        selected_tools = []
        for idx in top_indices:
            if idx.item() < len(tool_names):
                selected_tools.append(tool_names[idx.item()])
            else:
                selected_tools.append(None)

        return selected_tools, probs

    def generate_arguments(
        self,
        context_embeddings: torch.Tensor,
        tool_schema: ToolSchema,
        max_length: int = 32
    ) -> torch.Tensor:
        """
        Generate arguments for tool call.

        Args:
            context_embeddings: [batch, seq_len, embed_dim]
            tool_schema: Schema of selected tool
            max_length: Maximum argument length

        Returns:
            arg_embeddings: [batch, arg_len, embed_dim]
        """
        batch_size = context_embeddings.size(0)
        device = context_embeddings.device

        # Project context to hidden
        memory = self.input_proj(context_embeddings)

        # Initialize with context mean
        current = context_embeddings.mean(dim=1, keepdim=True)  # [batch, 1, embed_dim]
        current = self.input_proj(current)

        # Generate autoregressively
        generated = [current]

        for _ in range(max_length - 1):
            # Decode
            output = self.arg_generator(current, memory)

            # Project back
            next_embed = self.output_proj(output[:, -1:, :])
            next_hidden = self.input_proj(next_embed)

            generated.append(next_hidden)
            current = torch.cat([current, next_hidden], dim=1)

        # Project all back to embedding space
        all_hidden = torch.cat(generated, dim=1)
        all_embeds = self.output_proj(all_hidden)

        return all_embeds

    def forward(
        self,
        context_embeddings: torch.Tensor,
        tool_registry: ToolRegistry
    ) -> Tuple[bool, Optional[str], Optional[torch.Tensor]]:
        """
        Full forward pass: decide, select, generate arguments.

        Args:
            context_embeddings: [1, seq_len, embed_dim]
            tool_registry: Available tools

        Returns:
            should_use_tool: Whether to use a tool
            selected_tool: Name of selected tool
            arg_embeddings: Argument embeddings
        """
        # Decide whether to use tool
        use_tool, _ = self.decide_tool_use(context_embeddings)

        if not use_tool[0]:
            return False, None, None

        # Select tool
        tool_names = tool_registry.list_tools()
        selected_tools, _ = self.select_tool(context_embeddings, tool_names)
        selected_tool = selected_tools[0]

        if selected_tool is None:
            return False, None, None

        # Generate arguments
        tool = tool_registry.get(selected_tool)
        arg_embeddings = self.generate_arguments(
            context_embeddings, tool.schema, max_length=32
        )

        return True, selected_tool, arg_embeddings


class AdaptiveToolExecutor:
    """
    Adaptive tool executor that learns from execution feedback.

    Tracks success rates and adjusts tool selection.
    """

    def __init__(
        self,
        executor: ToolExecutor,
        learner: Optional[ToolLearner] = None
    ):
        self.executor = executor
        self.learner = learner

        # Track statistics
        self.tool_stats: Dict[str, Dict[str, int]] = {}

    def execute_with_learning(
        self,
        text: str,
        context_embeddings: Optional[torch.Tensor] = None
    ) -> Tuple[ExecutionTrace, Optional[float]]:
        """
        Execute tools and update statistics.

        Args:
            text: Text with tool calls
            context_embeddings: Context for learning (optional)

        Returns:
            trace: Execution trace
            reward: Reward signal for learning
        """
        # Execute
        trace = self.executor.execute_from_text(text)

        # Update statistics
        for result in trace.results:
            tool_name = result.call.tool_name

            if tool_name not in self.tool_stats:
                self.tool_stats[tool_name] = {
                    "success": 0,
                    "failure": 0,
                    "total": 0
                }

            self.tool_stats[tool_name]["total"] += 1

            if result.status == ExecutionStatus.SUCCESS:
                self.tool_stats[tool_name]["success"] += 1
            else:
                self.tool_stats[tool_name]["failure"] += 1

        # Compute reward
        reward = None
        if trace.results:
            success_count = sum(1 for r in trace.results if r.status == ExecutionStatus.SUCCESS)
            reward = success_count / len(trace.results)

        return trace, reward

    def get_tool_success_rate(self, tool_name: str) -> float:
        """Get success rate for a tool."""
        if tool_name not in self.tool_stats:
            return 0.0

        stats = self.tool_stats[tool_name]
        if stats["total"] == 0:
            return 0.0

        return stats["success"] / stats["total"]

    def get_best_tools(self, top_k: int = 5) -> List[Tuple[str, float]]:
        """Get tools with highest success rates."""
        tool_rates = [
            (name, self.get_tool_success_rate(name))
            for name in self.tool_stats.keys()
        ]

        tool_rates.sort(key=lambda x: x[1], reverse=True)
        return tool_rates[:top_k]
