"""
Tool schema and definitions for MK3 tool-use capabilities.

Defines:
- Tool types and schemas
- Parameter specifications
- Tool registry
- Built-in tool implementations
"""

import json
import re
from typing import Any, Dict, List, Optional, Callable, Union
from dataclasses import dataclass, field
from enum import Enum


class ToolType(Enum):
    """Types of tools."""
    CALCULATOR = "calculator"
    SEARCH = "search"
    CODE_EXECUTOR = "code_executor"
    QA = "qa"
    RETRIEVAL = "retrieval"
    API = "api"
    CUSTOM = "custom"


class ParameterType(Enum):
    """Parameter types for tool inputs."""
    STRING = "string"
    INTEGER = "integer"
    FLOAT = "float"
    BOOLEAN = "boolean"
    ARRAY = "array"
    OBJECT = "object"


@dataclass
class ParameterSpec:
    """Specification for a tool parameter."""
    name: str
    type: ParameterType
    description: str
    required: bool = True
    default: Optional[Any] = None
    enum_values: Optional[List[Any]] = None
    min_value: Optional[Union[int, float]] = None
    max_value: Optional[Union[int, float]] = None

    def validate(self, value: Any) -> bool:
        """Validate parameter value."""
        if value is None:
            return not self.required

        # Type validation
        if self.type == ParameterType.STRING:
            if not isinstance(value, str):
                return False
        elif self.type == ParameterType.INTEGER:
            if not isinstance(value, int):
                return False
            if self.min_value is not None and value < self.min_value:
                return False
            if self.max_value is not None and value > self.max_value:
                return False
        elif self.type == ParameterType.FLOAT:
            if not isinstance(value, (int, float)):
                return False
            if self.min_value is not None and value < self.min_value:
                return False
            if self.max_value is not None and value > self.max_value:
                return False
        elif self.type == ParameterType.BOOLEAN:
            if not isinstance(value, bool):
                return False
        elif self.type == ParameterType.ARRAY:
            if not isinstance(value, list):
                return False
        elif self.type == ParameterType.OBJECT:
            if not isinstance(value, dict):
                return False

        # Enum validation
        if self.enum_values is not None and value not in self.enum_values:
            return False

        return True

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary representation."""
        spec = {
            "type": self.type.value,
            "description": self.description
        }

        if not self.required:
            spec["required"] = False

        if self.default is not None:
            spec["default"] = self.default

        if self.enum_values is not None:
            spec["enum"] = self.enum_values

        if self.min_value is not None:
            spec["minimum"] = self.min_value

        if self.max_value is not None:
            spec["maximum"] = self.max_value

        return spec


@dataclass
class ToolSchema:
    """Schema for a tool."""
    name: str
    description: str
    tool_type: ToolType
    parameters: List[ParameterSpec] = field(default_factory=list)
    returns: Optional[str] = None
    examples: List[Dict[str, Any]] = field(default_factory=list)

    def validate_inputs(self, inputs: Dict[str, Any]) -> Tuple[bool, List[str]]:
        """
        Validate tool inputs.

        Returns:
            valid: Whether inputs are valid
            errors: List of validation errors
        """
        errors = []

        # Check required parameters
        for param in self.parameters:
            if param.required and param.name not in inputs:
                errors.append(f"Missing required parameter: {param.name}")
                continue

            # Validate value if present
            if param.name in inputs:
                value = inputs[param.name]
                if not param.validate(value):
                    errors.append(f"Invalid value for parameter {param.name}: {value}")

        # Check for unexpected parameters
        expected_params = {p.name for p in self.parameters}
        for key in inputs.keys():
            if key not in expected_params:
                errors.append(f"Unexpected parameter: {key}")

        return len(errors) == 0, errors

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary representation."""
        return {
            "name": self.name,
            "description": self.description,
            "type": self.tool_type.value,
            "parameters": {
                param.name: param.to_dict() for param in self.parameters
            },
            "returns": self.returns,
            "examples": self.examples
        }

    def to_json(self) -> str:
        """Convert to JSON string."""
        return json.dumps(self.to_dict(), indent=2)

    @classmethod
    def from_dict(cls, data: Dict[str, Any]) -> 'ToolSchema':
        """Create from dictionary."""
        parameters = []
        for name, spec in data.get("parameters", {}).items():
            param = ParameterSpec(
                name=name,
                type=ParameterType(spec["type"]),
                description=spec["description"],
                required=spec.get("required", True),
                default=spec.get("default"),
                enum_values=spec.get("enum"),
                min_value=spec.get("minimum"),
                max_value=spec.get("maximum")
            )
            parameters.append(param)

        return cls(
            name=data["name"],
            description=data["description"],
            tool_type=ToolType(data["type"]),
            parameters=parameters,
            returns=data.get("returns"),
            examples=data.get("examples", [])
        )


@dataclass
class Tool:
    """
    A tool that can be called by the model.
    """
    schema: ToolSchema
    function: Callable

    def __call__(self, **kwargs) -> Any:
        """Execute the tool."""
        # Validate inputs
        valid, errors = self.schema.validate_inputs(kwargs)
        if not valid:
            raise ValueError(f"Invalid inputs for tool {self.schema.name}: {errors}")

        # Execute
        return self.function(**kwargs)

    def get_signature(self) -> str:
        """Get tool signature as string."""
        params = []
        for param in self.schema.parameters:
            param_str = f"{param.name}: {param.type.value}"
            if not param.required:
                param_str += " = " + str(param.default)
            params.append(param_str)

        param_str = ", ".join(params)
        return_str = self.schema.returns or "Any"
        return f"{self.schema.name}({param_str}) -> {return_str}"


class ToolRegistry:
    """Registry for managing tools."""

    def __init__(self):
        self.tools: Dict[str, Tool] = {}

    def register(self, tool: Tool):
        """Register a tool."""
        self.tools[tool.schema.name] = tool

    def register_function(
        self,
        name: str,
        description: str,
        function: Callable,
        parameters: List[ParameterSpec],
        tool_type: ToolType = ToolType.CUSTOM,
        returns: Optional[str] = None
    ):
        """Register a function as a tool."""
        schema = ToolSchema(
            name=name,
            description=description,
            tool_type=tool_type,
            parameters=parameters,
            returns=returns
        )
        tool = Tool(schema=schema, function=function)
        self.register(tool)

    def get(self, name: str) -> Optional[Tool]:
        """Get a tool by name."""
        return self.tools.get(name)

    def list_tools(self) -> List[str]:
        """List all registered tool names."""
        return list(self.tools.keys())

    def get_all_schemas(self) -> List[ToolSchema]:
        """Get schemas for all tools."""
        return [tool.schema for tool in self.tools.values()]

    def to_prompt(self) -> str:
        """Convert tools to prompt format."""
        lines = ["Available tools:"]
        for tool in self.tools.values():
            lines.append(f"\n{tool.get_signature()}")
            lines.append(f"  Description: {tool.schema.description}")

            if tool.schema.parameters:
                lines.append("  Parameters:")
                for param in tool.schema.parameters:
                    required = "required" if param.required else "optional"
                    lines.append(f"    - {param.name} ({param.type.value}, {required}): {param.description}")

        return "\n".join(lines)


# Built-in tool implementations

def calculator_tool(expression: str) -> Union[int, float, str]:
    """
    Evaluate a mathematical expression.

    Args:
        expression: Mathematical expression to evaluate

    Returns:
        result: Result of evaluation
    """
    try:
        # Sanitize expression (only allow safe operations)
        allowed_chars = set("0123456789+-*/(). ")
        if not all(c in allowed_chars for c in expression):
            return "Error: Invalid characters in expression"

        # Evaluate safely
        result = eval(expression, {"__builtins__": {}}, {})
        return result
    except Exception as e:
        return f"Error: {str(e)}"


def search_tool(query: str, num_results: int = 5) -> List[Dict[str, str]]:
    """
    Search for information (mock implementation).

    Args:
        query: Search query
        num_results: Number of results to return

    Returns:
        results: List of search results
    """
    # Mock implementation
    return [
        {
            "title": f"Result {i+1} for '{query}'",
            "snippet": f"This is a snippet about {query}...",
            "url": f"https://example.com/result_{i+1}"
        }
        for i in range(num_results)
    ]


def qa_tool(question: str, context: str) -> str:
    """
    Answer a question given context (mock implementation).

    Args:
        question: Question to answer
        context: Context for answering

    Returns:
        answer: Answer to the question
    """
    # Mock implementation
    return f"Based on the context, the answer to '{question}' is: [answer would go here]"


def code_executor_tool(code: str, language: str = "python") -> Dict[str, Any]:
    """
    Execute code (mock implementation for safety).

    Args:
        code: Code to execute
        language: Programming language

    Returns:
        result: Execution result
    """
    # Mock implementation (do NOT actually execute arbitrary code)
    return {
        "status": "success",
        "output": "[Code execution output would go here]",
        "language": language
    }


# Create default registry with built-in tools

def create_default_registry() -> ToolRegistry:
    """Create registry with built-in tools."""
    registry = ToolRegistry()

    # Calculator
    registry.register_function(
        name="calculator",
        description="Evaluate mathematical expressions",
        function=calculator_tool,
        parameters=[
            ParameterSpec(
                name="expression",
                type=ParameterType.STRING,
                description="Mathematical expression to evaluate"
            )
        ],
        tool_type=ToolType.CALCULATOR,
        returns="number or error message"
    )

    # Search
    registry.register_function(
        name="search",
        description="Search for information",
        function=search_tool,
        parameters=[
            ParameterSpec(
                name="query",
                type=ParameterType.STRING,
                description="Search query"
            ),
            ParameterSpec(
                name="num_results",
                type=ParameterType.INTEGER,
                description="Number of results to return",
                required=False,
                default=5,
                min_value=1,
                max_value=20
            )
        ],
        tool_type=ToolType.SEARCH,
        returns="List of search results"
    )

    # QA
    registry.register_function(
        name="qa",
        description="Answer questions given context",
        function=qa_tool,
        parameters=[
            ParameterSpec(
                name="question",
                type=ParameterType.STRING,
                description="Question to answer"
            ),
            ParameterSpec(
                name="context",
                type=ParameterType.STRING,
                description="Context for answering"
            )
        ],
        tool_type=ToolType.QA,
        returns="Answer string"
    )

    # Code executor
    registry.register_function(
        name="code_executor",
        description="Execute code in specified language",
        function=code_executor_tool,
        parameters=[
            ParameterSpec(
                name="code",
                type=ParameterType.STRING,
                description="Code to execute"
            ),
            ParameterSpec(
                name="language",
                type=ParameterType.STRING,
                description="Programming language",
                required=False,
                default="python",
                enum_values=["python", "javascript", "bash"]
            )
        ],
        tool_type=ToolType.CODE_EXECUTOR,
        returns="Execution result"
    )

    return registry
