"""
Middleware system for tool processing
"""

import asyncio
import time
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Any, Dict, List, Optional

from .logging import log_tool_end, log_tool_start
from .security import SecurityMiddleware, get_security_middleware


@dataclass
class MiddlewareContext:
    """Context passed through middleware chain."""

    tool_name: str
    arguments: dict[str, Any]
    start_time: float
    metadata: dict[str, Any]
    user_context: Optional[dict[str, Any]] = None
    security_context: Optional[dict[str, Any]] = None


class MiddlewareResult:
    """Result from middleware processing."""

    def __init__(self, success: bool = True, data: Any = None, error: Optional[str] = None):
        self.success = success
        self.data = data
        self.error = error
        self.metadata = {}


class BaseMiddleware(ABC):
    """Base middleware class."""

    @abstractmethod
    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Process request before tool execution."""
        pass

    @abstractmethod
    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """Process response after tool execution."""
        pass


class ValidationMiddleware(BaseMiddleware):
    """Middleware for input validation."""

    def __init__(self, validators: Optional[dict[str, Callable]] = None):
        self.validators = validators or {}

    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Validate input arguments."""
        tool_validators = self.validators.get(context.tool_name, [])

        for validator in tool_validators:
            try:
                validation_result = validator(context.arguments)
                if not validation_result.get("valid", True):
                    return MiddlewareResult(
                        success=False,
                        error=f"Validation failed: {validation_result.get('message', 'Unknown error')}",
                    )
            except Exception as e:
                return MiddlewareResult(success=False, error=f"Validation error: {str(e)}")

        return MiddlewareResult(success=True)

    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """No processing needed for response."""
        return MiddlewareResult(success=True, data=result)


class LoggingMiddleware(BaseMiddleware):
    """Middleware for comprehensive logging."""

    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Log request start."""
        context.start_time = log_tool_start(context.tool_name, context.arguments)
        return MiddlewareResult(success=True)

    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """Log response completion."""
        success = (
            result and len(result) > 0 and not any("error" in str(r.text).lower() for r in result)
        )
        duration = time.time() - context.start_time

        # Extract error message if any
        error = None
        if not success and result:
            for r in result:
                if hasattr(r, "text") and "error" in str(r.text).lower():
                    error = str(r.text)
                    break

        log_tool_end(context.tool_name, context.start_time, result, success, error)
        return MiddlewareResult(success=True, data=result)


class CachingMiddleware(BaseMiddleware):
    """Middleware for response caching."""

    def __init__(self, cache_manager):
        self.cache_manager = cache_manager

    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Check cache before processing."""
        # For now, simple caching based on URL for fetch operations
        if context.tool_name in ["fetch_html", "extract_text"] and "url" in context.arguments:
            cache_key = f"{context.tool_name}:{context.arguments['url']}"
            cached_result = self.cache_manager.get(cache_key)

            if cached_result:
                context.metadata["cache_hit"] = True
                return MiddlewareResult(success=True, data=cached_result)

        return MiddlewareResult(success=True)

    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """Cache successful responses."""
        if (
            context.tool_name in ["fetch_html", "extract_text"]
            and "url" in context.arguments
            and result
            and len(result) > 0
        ):
            cache_key = f"{context.tool_name}:{context.arguments['url']}"
            # Cache for 5 minutes
            self.cache_manager.set(cache_key, result, ttl=300)

        return MiddlewareResult(success=True, data=result)


class RateLimitMiddleware(BaseMiddleware):
    """Middleware for rate limiting."""

    def __init__(self, requests_per_minute: int = 60):
        self.requests_per_minute = requests_per_minute
        self.requests = []
        self._lock = asyncio.Lock()

    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Check rate limit."""
        async with self._lock:
            current_time = time.time()

            # Remove old requests (older than 1 minute)
            self.requests = [req_time for req_time in self.requests if current_time - req_time < 60]

            if len(self.requests) >= self.requests_per_minute:
                return MiddlewareResult(
                    success=False,
                    error=f"Rate limit exceeded: {self.requests_per_minute} requests per minute",
                )

            self.requests.append(current_time)
            return MiddlewareResult(success=True)

    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """No processing needed for response."""
        return MiddlewareResult(success=True, data=result)


class TransformationMiddleware(BaseMiddleware):
    """Middleware for data transformation."""

    def __init__(self, transformers: Optional[dict[str, Callable]] = None):
        self.transformers = transformers or {}

    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Transform input arguments."""
        transformer = self.transformers.get(context.tool_name)
        if transformer:
            try:
                transformed_args = transformer(context.arguments)
                context.arguments = transformed_args
            except Exception as e:
                return MiddlewareResult(success=False, error=f"Transformation error: {str(e)}")

        return MiddlewareResult(success=True)

    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """Transform response data."""
        # Could add response transformers here if needed
        return MiddlewareResult(success=True, data=result)


class MetricsMiddleware(BaseMiddleware):
    """Middleware for metrics collection."""

    def __init__(self, metrics_collector):
        self.metrics_collector = metrics_collector

    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Record request metrics."""
        self.metrics_collector.increment_counter(f"tool_{context.tool_name}_requests")
        return MiddlewareResult(success=True)

    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """Record response metrics."""
        duration = time.time() - context.start_time

        # Record timing
        self.metrics_collector.record_time(f"tool_{context.tool_name}_duration", duration)

        # Record success/failure
        if result and len(result) > 0:
            self.metrics_collector.increment_counter(f"tool_{context.tool_name}_success")
        else:
            self.metrics_collector.increment_counter(f"tool_{context.tool_name}_errors")

        return MiddlewareResult(success=True, data=result)


class MiddlewareChain:
    """Chain of middleware processors."""

    def __init__(self):
        self.middleware: list[BaseMiddleware] = []

    def add_middleware(self, middleware: BaseMiddleware):
        """Add middleware to chain."""
        self.middleware.append(middleware)

    def remove_middleware(self, middleware_type: type):
        """Remove middleware by type."""
        self.middleware = [m for m in self.middleware if not isinstance(m, middleware_type)]

    async def process_request(self, context: MiddlewareContext) -> MiddlewareResult:
        """Process request through middleware chain."""
        for middleware in self.middleware:
            result = await middleware.process_request(context)
            if not result.success:
                return result

        return MiddlewareResult(success=True)

    async def process_response(self, context: MiddlewareContext, result: Any) -> MiddlewareResult:
        """Process response through middleware chain."""
        current_result = result

        for middleware in reversed(self.middleware):
            response_result = await middleware.process_response(context, current_result)
            if not response_result.success:
                return response_result
            current_result = response_result.data

        return MiddlewareResult(success=True, data=current_result)


# Global middleware chain
_middleware_chain = MiddlewareChain()


def get_middleware_chain() -> MiddlewareChain:
    """Get global middleware chain."""
    return _middleware_chain


def setup_default_middleware():
    """Setup default middleware chain."""
    from .cache import get_response_cache
    from .logging import get_structured_logger
    from .metrics import get_metrics_collector
    from .rate_limiter import get_rate_limiter

    chain = get_middleware_chain()

    # Add default middleware in order
    chain.add_middleware(LoggingMiddleware())
    chain.add_middleware(SecurityMiddleware())
    chain.add_middleware(MetricsMiddleware(get_metrics_collector()))
    chain.add_middleware(CachingMiddleware(get_response_cache()))
    chain.add_middleware(RateLimitMiddleware(requests_per_minute=100))


async def process_with_middleware(
    tool_name: str,
    arguments: Dict[str, Any],
    tool_function: Callable,
    user_context: Optional[Dict[str, Any]] = None,
) -> List[Dict[str, Any]]:
    """
    Process tool request through middleware chain.

    Args:
        tool_name: Name of the tool
        arguments: Tool arguments
        tool_function: Actual tool function to execute
        user_context: Optional user context (IP, user agent, etc.)

    Returns:
        Tool results
    """
    chain = get_middleware_chain()

    # Create context
    context = MiddlewareContext(
        tool_name=tool_name,
        arguments=arguments,
        start_time=time.time(),
        metadata={},
        user_context=user_context,
    )

    # Process request through middleware
    request_result = await chain.process_request(context)
    if not request_result.success:
        return [{"type": "text", "text": f"Middleware error: {request_result.error}"}]

    # Check if middleware provided cached result
    if hasattr(request_result, "data") and request_result.data is not None:
        return request_result.data

    # Execute tool function with potentially modified arguments
    try:
        # Use arguments from context (may be modified by middleware)
        result = await tool_function(context.arguments)
    except Exception as e:
        # Process error through response middleware
        error_result = await chain.process_response(
            context, [{"type": "text", "text": f"Tool error: {str(e)}"}]
        )
        return (
            error_result.data
            if error_result.success
            else [{"type": "text", "text": f"Error: {str(e)}"}]
        )

    # Process response through middleware
    response_result = await chain.process_response(context, result)
    return response_result.data if response_result.success else result


# Example validators
def validate_url(args: dict[str, Any]) -> dict[str, Any]:
    """Validate URL argument."""
    if "url" not in args:
        return {"valid": False, "message": "URL parameter is required"}

    url = args["url"]
    if not isinstance(url, str) or not url.startswith(("http://", "https://")):
        return {"valid": False, "message": "Invalid URL format"}

    return {"valid": True}


def validate_html(args: dict[str, Any]) -> dict[str, Any]:
    """Validate HTML argument."""
    if "html" not in args:
        return {"valid": False, "message": "HTML parameter is required"}

    html = args["html"]
    if not isinstance(html, str) or len(html.strip()) == 0:
        return {"valid": False, "message": "HTML content cannot be empty"}

    return {"valid": True}


def validate_extraction_schema(args: dict[str, Any]) -> dict[str, Any]:
    """Validate extraction schema."""
    if "schema" not in args:
        return {"valid": False, "message": "Schema parameter is required"}

    schema = args["schema"]
    if not isinstance(schema, dict) or "fields" not in schema:
        return {"valid": False, "message": "Schema must contain 'fields' array"}

    fields = schema["fields"]
    if not isinstance(fields, list) or len(fields) == 0:
        return {"valid": False, "message": "Schema must contain at least one field"}

    for field in fields:
        if not isinstance(field, dict) or "name" not in field or "selector" not in field:
            return {
                "valid": False,
                "message": "Each field must have 'name' and 'selector' properties",
            }

    return {"valid": True}


# Setup default validators
_default_validators = {
    "fetch_html": [validate_url],
    "extract_text": [validate_html],
    "find_elements": [validate_html],
    "extract_links": [lambda args: validate_url(args) if "url" in args else validate_html(args)],
    "parse_advanced": [validate_extraction_schema],
    "extract_dynamic_content": [validate_url],
}


def create_validation_middleware() -> ValidationMiddleware:
    """Create validation middleware with default validators."""
    return ValidationMiddleware(_default_validators)


# Initialize default middleware
setup_default_middleware()
