"""
PAL: Program-Aided Language Models.

Uses language models to generate programs that solve tasks,
then executes the programs to get answers.

Reference: PAL: Program-aided Language Models
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
import re
import ast
from typing import List, Dict, Optional, Tuple, Any, Callable, Union
from dataclasses import dataclass
from enum import Enum


class ProgrammingLanguage(Enum):
    """Supported programming languages."""
    PYTHON = "python"
    SQL = "sql"
    JAVASCRIPT = "javascript"
    MATH = "math"  # Mathematical expressions


@dataclass
class ProgramResult:
    """Result from program execution."""
    program: str
    language: ProgrammingLanguage
    output: Any
    success: bool
    error: Optional[str] = None
    execution_time: float = 0.0


class ProgramGenerator(nn.Module):
    """
    Generates programs from natural language descriptions.

    Uses a decoder-only transformer to generate code.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        hidden_dim: int = 1024,
        num_layers: int = 8,
        num_heads: int = 12,
        vocab_size: int = 50000,
        max_program_length: int = 512,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.hidden_dim = hidden_dim
        self.max_program_length = max_program_length
        self.vocab_size = vocab_size

        # Input projection
        self.input_proj = nn.Linear(embedding_dim, hidden_dim)

        # Program decoder
        self.program_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
        )

        # Output projection to vocabulary
        self.output_proj = nn.Linear(hidden_dim, vocab_size)

        # Language type predictor
        self.language_classifier = nn.Sequential(
            nn.Linear(embedding_dim, 512),
            nn.LayerNorm(512),
            nn.GELU(),
            nn.Linear(512, 4)  # [python, sql, javascript, math]
        )

        # Positional encoding
        self.register_buffer(
            'pos_encoding',
            self._create_positional_encoding(max_program_length, hidden_dim)
        )

    def _create_positional_encoding(self, max_len: int, d_model: int) -> torch.Tensor:
        """Create positional encoding."""
        position = torch.arange(max_len).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, d_model, 2) * (-torch.log(torch.tensor(10000.0)) / d_model))

        pe = torch.zeros(max_len, d_model)
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)

        return pe.unsqueeze(0)

    def predict_language(self, task_embeddings: torch.Tensor) -> ProgrammingLanguage:
        """
        Predict which programming language to use.

        Args:
            task_embeddings: [batch, seq_len, embed_dim]

        Returns:
            language: Predicted programming language
        """
        task_vec = task_embeddings.mean(dim=1)
        logits = self.language_classifier(task_vec)
        lang_idx = logits.argmax(dim=-1).item()

        languages = [
            ProgrammingLanguage.PYTHON,
            ProgrammingLanguage.SQL,
            ProgrammingLanguage.JAVASCRIPT,
            ProgrammingLanguage.MATH
        ]

        return languages[lang_idx]

    def forward(
        self,
        task_embeddings: torch.Tensor,
        max_length: Optional[int] = None
    ) -> torch.Tensor:
        """
        Generate program embeddings.

        Args:
            task_embeddings: [batch, seq_len, embed_dim] task description
            max_length: Maximum program length (uses default if None)

        Returns:
            program_logits: [batch, prog_len, vocab_size]
        """
        if max_length is None:
            max_length = self.max_program_length

        # Encode task as memory
        memory = self.input_proj(task_embeddings)

        # Initialize with task mean
        start = task_embeddings.mean(dim=1, keepdim=True)
        current = self.input_proj(start)

        # Add positional encoding
        current = current + self.pos_encoding[:, 0:1, :]

        # Generate program autoregressively
        all_logits = []

        for pos in range(max_length):
            # Add positional encoding
            current_with_pos = current + self.pos_encoding[:, :pos+1, :]

            # Decode
            output = self.program_decoder(current_with_pos, memory)

            # Get logits for next token
            next_logits = self.output_proj(output[:, -1:, :])
            all_logits.append(next_logits)

            # Sample next token (during inference)
            # For training, would use teacher forcing
            next_token_id = next_logits.argmax(dim=-1)

            # Continue generation (simplified - would use token embeddings in practice)
            next_embed = output[:, -1:, :]
            current = torch.cat([current, next_embed], dim=1)

        # Stack all logits
        program_logits = torch.cat(all_logits, dim=1)

        return program_logits


class ProgramExecutor:
    """
    Safely executes generated programs.

    Supports multiple languages with sandboxing.
    """

    def __init__(
        self,
        timeout: float = 5.0,
        max_memory: int = 100_000_000,  # 100 MB
        enable_sandbox: bool = True
    ):
        self.timeout = timeout
        self.max_memory = max_memory
        self.enable_sandbox = enable_sandbox

    def execute_python(self, code: str) -> Tuple[Any, Optional[str]]:
        """
        Execute Python code safely.

        Args:
            code: Python code to execute

        Returns:
            result: Execution result
            error: Error message if failed
        """
        try:
            # Create restricted namespace
            namespace = {
                '__builtins__': {
                    'abs': abs,
                    'max': max,
                    'min': min,
                    'sum': sum,
                    'len': len,
                    'range': range,
                    'sorted': sorted,
                    'int': int,
                    'float': float,
                    'str': str,
                    'list': list,
                    'dict': dict,
                    'set': set,
                    'True': True,
                    'False': False,
                    'None': None,
                }
            }

            # Execute code
            exec(code, namespace)

            # Get result (look for 'answer' or 'result' variable)
            if 'answer' in namespace:
                return namespace['answer'], None
            elif 'result' in namespace:
                return namespace['result'], None
            else:
                # Return last assigned variable
                vars_defined = {k: v for k, v in namespace.items()
                               if not k.startswith('_')}
                if vars_defined:
                    return list(vars_defined.values())[-1], None
                return None, "No result variable found"

        except Exception as e:
            return None, str(e)

    def execute_math(self, expression: str) -> Tuple[Any, Optional[str]]:
        """
        Execute mathematical expression.

        Args:
            expression: Math expression

        Returns:
            result: Computed result
            error: Error message if failed
        """
        try:
            # Sanitize expression
            allowed = set("0123456789+-*/(). ")
            if not all(c in allowed for c in expression):
                return None, "Invalid characters in expression"

            # Evaluate
            result = eval(expression, {"__builtins__": {}}, {})
            return result, None

        except Exception as e:
            return None, str(e)

    def execute_sql(self, query: str) -> Tuple[Any, Optional[str]]:
        """
        Execute SQL query (mock implementation).

        Args:
            query: SQL query

        Returns:
            result: Query result
            error: Error message if failed
        """
        # Mock implementation
        return {"mock": "SQL execution would go here"}, None

    def execute(
        self,
        code: str,
        language: ProgrammingLanguage
    ) -> ProgramResult:
        """
        Execute program in specified language.

        Args:
            code: Program code
            language: Programming language

        Returns:
            result: ProgramResult
        """
        import time

        start_time = time.time()

        # Execute based on language
        if language == ProgrammingLanguage.PYTHON:
            output, error = self.execute_python(code)
        elif language == ProgrammingLanguage.MATH:
            output, error = self.execute_math(code)
        elif language == ProgrammingLanguage.SQL:
            output, error = self.execute_sql(code)
        else:
            output, error = None, f"Unsupported language: {language}"

        execution_time = time.time() - start_time

        return ProgramResult(
            program=code,
            language=language,
            output=output,
            success=(error is None),
            error=error,
            execution_time=execution_time
        )


class PAL(nn.Module):
    """
    PAL: Program-Aided Language Model.

    Generates programs to solve tasks, then executes them.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        vocab_size: int = 50000,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim

        # Program generator
        self.generator = ProgramGenerator(
            embedding_dim=embedding_dim,
            vocab_size=vocab_size,
            dropout=dropout
        )

        # Program executor
        self.executor = ProgramExecutor()

        # Program verifier (checks if program is safe)
        self.verifier = ProgramVerifier()

    def forward(
        self,
        task_embeddings: torch.Tensor,
        decode_fn: Callable,
        max_attempts: int = 3
    ) -> ProgramResult:
        """
        Generate and execute program to solve task.

        Args:
            task_embeddings: [1, seq_len, embed_dim]
            decode_fn: Function to decode embeddings to text
            max_attempts: Maximum generation attempts

        Returns:
            result: ProgramResult
        """
        # Predict language
        language = self.generator.predict_language(task_embeddings)

        # Generate program
        for attempt in range(max_attempts):
            # Generate program logits
            program_logits = self.generator(task_embeddings)

            # Decode to text (simplified - would use proper tokenizer)
            program_text = self._decode_program(program_logits, decode_fn)

            # Verify program safety
            is_safe, safety_error = self.verifier.verify(program_text, language)
            if not is_safe:
                continue  # Try again

            # Execute program
            result = self.executor.execute(program_text, language)

            if result.success:
                return result

        # All attempts failed
        return ProgramResult(
            program="# Failed to generate valid program",
            language=language,
            output=None,
            success=False,
            error="Exceeded maximum attempts"
        )

    def _decode_program(
        self,
        program_logits: torch.Tensor,
        decode_fn: Callable
    ) -> str:
        """
        Decode program logits to text.

        Args:
            program_logits: [batch, prog_len, vocab_size]
            decode_fn: Decode function

        Returns:
            program: Program text
        """
        # Get token IDs
        token_ids = program_logits.argmax(dim=-1)  # [batch, prog_len]

        # Decode (simplified - would use proper tokenizer)
        # For now, return placeholder
        return "# Generated program placeholder\nanswer = 42"


class ProgramVerifier:
    """
    Verifies that generated programs are safe to execute.
    """

    def __init__(self):
        # Blacklisted operations
        self.python_blacklist = [
            'import',
            'exec',
            'eval',
            'compile',
            '__import__',
            'open',
            'file',
            'input',
            'raw_input',
            'execfile',
            '__builtins__',
        ]

    def verify_python(self, code: str) -> Tuple[bool, Optional[str]]:
        """Verify Python code safety."""
        # Check for blacklisted operations
        for forbidden in self.python_blacklist:
            if forbidden in code.lower():
                return False, f"Forbidden operation: {forbidden}"

        # Try to parse as valid Python
        try:
            ast.parse(code)
        except SyntaxError as e:
            return False, f"Syntax error: {str(e)}"

        return True, None

    def verify_math(self, expression: str) -> Tuple[bool, Optional[str]]:
        """Verify mathematical expression safety."""
        # Check for allowed characters only
        allowed = set("0123456789+-*/(). ")
        if not all(c in allowed for c in expression):
            return False, "Invalid characters in expression"

        return True, None

    def verify(
        self,
        code: str,
        language: ProgrammingLanguage
    ) -> Tuple[bool, Optional[str]]:
        """
        Verify code safety.

        Args:
            code: Code to verify
            language: Programming language

        Returns:
            is_safe: Whether code is safe
            error: Error message if unsafe
        """
        if language == ProgrammingLanguage.PYTHON:
            return self.verify_python(code)
        elif language == ProgrammingLanguage.MATH:
            return self.verify_math(code)
        else:
            # Other languages - accept for now
            return True, None


class MathPAL(nn.Module):
    """
    Specialized PAL for mathematical reasoning.

    Generates Python code to solve math word problems.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim

        # Math-specific program generator
        self.math_generator = ProgramGenerator(
            embedding_dim=embedding_dim,
            dropout=dropout
        )

        # Executor
        self.executor = ProgramExecutor()

        # Problem parser (extracts numbers and operations)
        self.problem_parser = nn.Sequential(
            nn.Linear(embedding_dim, 512),
            nn.LayerNorm(512),
            nn.GELU(),
            nn.Linear(512, 256),
            nn.LayerNorm(256),
            nn.GELU(),
            nn.Linear(256, 64)  # Embedding of parsed problem
        )

    def parse_problem(self, problem_embeddings: torch.Tensor) -> torch.Tensor:
        """
        Parse math problem to extract structure.

        Args:
            problem_embeddings: [batch, seq_len, embed_dim]

        Returns:
            parsed: [batch, 64] parsed problem representation
        """
        problem_vec = problem_embeddings.mean(dim=1)
        parsed = self.problem_parser(problem_vec)
        return parsed

    def forward(
        self,
        problem_embeddings: torch.Tensor,
        decode_fn: Callable
    ) -> ProgramResult:
        """
        Solve math problem.

        Args:
            problem_embeddings: [1, seq_len, embed_dim]
            decode_fn: Decode function

        Returns:
            result: ProgramResult with answer
        """
        # Parse problem structure
        parsed = self.parse_problem(problem_embeddings)

        # Generate Python code to solve
        program_logits = self.math_generator(problem_embeddings)

        # For demo, use a template-based approach
        program = self._generate_math_program(problem_embeddings, decode_fn)

        # Execute
        result = self.executor.execute(program, ProgrammingLanguage.PYTHON)

        return result

    def _generate_math_program(
        self,
        problem_embeddings: torch.Tensor,
        decode_fn: Callable
    ) -> str:
        """
        Generate Python code for math problem.

        Args:
            problem_embeddings: Problem embeddings
            decode_fn: Decode function

        Returns:
            program: Python code
        """
        # Simplified - would use actual generation in practice
        # For demo, return a template
        return """
# Math problem solution
def solve():
    # Extract numbers and solve
    result = 42  # Placeholder
    return result

answer = solve()
"""


class SQLPAL(nn.Module):
    """
    Specialized PAL for SQL generation.

    Generates SQL queries from natural language.
    """

    def __init__(
        self,
        embedding_dim: int = 768,
        dropout: float = 0.1
    ):
        super().__init__()

        self.embedding_dim = embedding_dim

        # SQL-specific generator
        self.sql_generator = ProgramGenerator(
            embedding_dim=embedding_dim,
            dropout=dropout
        )

        # Schema encoder (encodes database schema)
        self.schema_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
        )

    def encode_schema(
        self,
        schema_embeddings: torch.Tensor
    ) -> torch.Tensor:
        """
        Encode database schema.

        Args:
            schema_embeddings: [batch, num_tables, embed_dim]

        Returns:
            encoded_schema: [batch, num_tables, embed_dim]
        """
        return self.schema_encoder(schema_embeddings)

    def forward(
        self,
        question_embeddings: torch.Tensor,
        schema_embeddings: torch.Tensor,
        decode_fn: Callable
    ) -> str:
        """
        Generate SQL query.

        Args:
            question_embeddings: [1, seq_len, embed_dim]
            schema_embeddings: [1, num_tables, embed_dim]
            decode_fn: Decode function

        Returns:
            sql_query: Generated SQL query
        """
        # Encode schema
        encoded_schema = self.encode_schema(schema_embeddings)

        # Combine question and schema
        combined = torch.cat([question_embeddings, encoded_schema], dim=1)

        # Generate SQL
        sql_logits = self.sql_generator(combined)

        # Decode (simplified)
        sql_query = "SELECT * FROM table WHERE condition;"  # Placeholder

        return sql_query
