"""
Rate limiting system for MCP tools and API endpoints
"""

import asyncio
import hashlib
import json
import time
from collections import defaultdict, deque
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple

from .logging import get_structured_logger


@dataclass
class RateLimitRule:
    """Rate limit rule configuration."""

    name: str
    requests_per_window: int
    window_seconds: int
    burst_limit: Optional[int] = None
    block_duration_seconds: int = 300  # 5 minutes default block
    identifier_type: str = "ip"  # ip, user, api_key, global
    conditions: dict[str, Any] = field(default_factory=dict)

    @property
    def burst_limit_final(self) -> int:
        """Get final burst limit (defaults to requests_per_window if not set)."""
        return self.burst_limit or self.requests_per_window


@dataclass
class RateLimitState:
    """Rate limit state for a specific identifier."""

    requests: deque = field(default_factory=lambda: deque())
    blocked_until: float = 0.0
    violations: int = 0


class RateLimiter:
    """Advanced rate limiter with multiple strategies."""

    def __init__(self):
        self.rules: dict[str, RateLimitRule] = {}
        self.states: dict[str, dict[str, RateLimitState]] = defaultdict(dict)
        self.logger = get_structured_logger()
        self._lock = asyncio.Lock()

    def add_rule(self, rule: RateLimitRule) -> None:
        """Add a rate limiting rule."""
        self.rules[rule.name] = rule
        self.logger.log_performance_metric("rate_limit_rule_added", 1, {"rule": rule.name})

    def remove_rule(self, rule_name: str) -> bool:
        """Remove a rate limiting rule."""
        if rule_name in self.rules:
            del self.rules[rule_name]
            # Clean up states for this rule
            if rule_name in self.states:
                del self.states[rule_name]
            self.logger.log_performance_metric("rate_limit_rule_removed", 1, {"rule": rule_name})
            return True
        return False

    def get_identifier(self, identifier_type: str, request_context: dict[str, Any]) -> str:
        """Extract identifier from request context."""
        if identifier_type == "ip":
            return request_context.get("ip_address", "unknown")
        elif identifier_type == "user":
            return request_context.get("user_id", request_context.get("ip_address", "unknown"))
        elif identifier_type == "api_key":
            return request_context.get("api_key", request_context.get("ip_address", "unknown"))
        elif identifier_type == "global":
            return "global"
        else:
            # Custom identifier function
            identifier_func = request_context.get("identifier_func")
            if callable(identifier_func):
                return identifier_func(request_context)
            return request_context.get("ip_address", "unknown")

    def check_conditions(self, rule: RateLimitRule, request_context: dict[str, Any]) -> bool:
        """Check if rule conditions are met."""
        for key, expected_value in rule.conditions.items():
            actual_value = request_context.get(key)
            if actual_value != expected_value:
                return False
        return True

    async def is_allowed(
        self, rule_name: str, request_context: dict[str, Any]
    ) -> tuple[bool, dict[str, Any]]:
        """
        Check if request is allowed under rate limiting rules.

        Returns:
            Tuple of (allowed: bool, info: dict with details)
        """
        async with self._lock:
            # Get rule
            rule = self.rules.get(rule_name)
            if not rule:
                return True, {"rule": rule_name, "status": "no_rule"}

            # Check conditions
            if not self.check_conditions(rule, request_context):
                return True, {"rule": rule_name, "status": "conditions_not_met"}

            # Get identifier
            identifier = self.get_identifier(rule.identifier_type, request_context)
            if not identifier:
                return True, {"rule": rule_name, "status": "no_identifier"}

            rule_key = f"{rule_name}:{identifier}"
            current_time = time.time()

            # Get or create state
            if rule_key not in self.states:
                self.states[rule_key] = RateLimitState()

            state = self.states[rule_key]

            # Check if still blocked
            if current_time < state.blocked_until:
                remaining_block = int(state.blocked_until - current_time)
                return False, {
                    "rule": rule_name,
                    "status": "blocked",
                    "remaining_seconds": remaining_block,
                    "violations": state.violations,
                }

            # Clean old requests outside the window
            window_start = current_time - rule.window_seconds
            while state.requests and state.requests[0] < window_start:
                state.requests.popleft()

            # Check rate limit
            request_count = len(state.requests)

            if request_count >= rule.requests_per_window:
                # Rate limit exceeded
                state.violations += 1

                # Calculate block duration (exponential backoff)
                block_duration = min(
                    rule.block_duration_seconds * (2 ** min(state.violations - 1, 5)),
                    3600,  # Max 1 hour block
                )
                state.blocked_until = current_time + block_duration

                self.logger.log_performance_metric(
                    "rate_limit_exceeded",
                    1,
                    {
                        "rule": rule_name,
                        "identifier": identifier,
                        "violations": state.violations,
                        "block_duration": block_duration,
                    },
                )

                return False, {
                    "rule": rule_name,
                    "status": "rate_limited",
                    "current_count": request_count,
                    "limit": rule.requests_per_window,
                    "window_seconds": rule.window_seconds,
                    "block_duration": block_duration,
                    "violations": state.violations,
                }

            # Allow request and record it
            state.requests.append(current_time)
            state.violations = max(0, state.violations - 1)  # Reduce violations over time

            return True, {
                "rule": rule_name,
                "status": "allowed",
                "current_count": request_count + 1,
                "limit": rule.requests_per_window,
                "remaining": rule.requests_per_window - (request_count + 1),
            }

    def get_stats(self, rule_name: Optional[str] = None) -> dict[str, Any]:
        """Get rate limiter statistics."""
        stats = {
            "total_rules": len(self.rules),
            "total_active_states": sum(len(states) for states in self.states.values()),
        }

        if rule_name:
            rule = self.rules.get(rule_name)
            rule_states = self.states.get(rule_name, {})
            stats[rule_name] = {
                "rule_config": {
                    "requests_per_window": rule.requests_per_window if rule else 0,
                    "window_seconds": rule.window_seconds if rule else 0,
                    "identifier_type": rule.identifier_type if rule else "unknown",
                },
                "active_identifiers": len(rule_states),
                "blocked_identifiers": sum(
                    1 for state in rule_states.values() if time.time() < state.blocked_until
                ),
            }

        return stats

    def cleanup_expired_states(self) -> int:
        """Clean up expired rate limit states. Returns number of cleaned states."""
        current_time = time.time()
        cleaned = 0

        for rule_name in list(self.states.keys()):
            rule_states = self.states[rule_name]
            rule = self.rules.get(rule_name)

            if not rule:
                # Rule no longer exists, clean up all states
                cleaned += len(rule_states)
                del self.states[rule_name]
                continue

            for identifier in list(rule_states.keys()):
                state = rule_states[identifier]

                # Clean old requests
                window_start = current_time - rule.window_seconds
                while state.requests and state.requests[0] < window_start:
                    state.requests.popleft()

                # Remove if no recent activity and not blocked
                if (
                    len(state.requests) == 0
                    and current_time >= state.blocked_until
                    and state.violations == 0
                ):
                    del rule_states[identifier]
                    cleaned += 1

            # Remove empty rule states
            if not rule_states:
                del self.states[rule_name]

        return cleaned


class AdaptiveRateLimiter(RateLimiter):
    """Adaptive rate limiter that adjusts limits based on system load."""

    def __init__(self):
        super().__init__()
        self.system_metrics = {}
        self.adaptation_interval = 60  # Check every minute
        self.last_adaptation = time.time()

    def update_system_metrics(self, metrics: dict[str, Any]) -> None:
        """Update system metrics for adaptive rate limiting."""
        self.system_metrics.update(metrics)
        current_time = time.time()

        # Adapt limits if enough time has passed
        if current_time - self.last_adaptation >= self.adaptation_interval:
            self._adapt_limits()
            self.last_adaptation = current_time

    def _adapt_limits(self) -> None:
        """Adapt rate limits based on system metrics."""
        cpu_usage = self.system_metrics.get("cpu_percent", 0)
        memory_usage = self.system_metrics.get("memory_percent", 0)
        active_connections = self.system_metrics.get("active_connections", 0)

        # Scale factor based on system load
        load_factor = max(cpu_usage / 80.0, memory_usage / 80.0, active_connections / 100.0)
        load_factor = min(max(load_factor, 0.5), 2.0)  # Clamp between 0.5 and 2.0

        if load_factor > 1.2:  # High load
            self._scale_limits(0.8)  # Reduce limits by 20%
            self.logger.log_performance_metric(
                "rate_limit_adapted",
                1,
                {
                    "reason": "high_load",
                    "load_factor": load_factor,
                    "scaling_factor": 0.8,
                },
            )
        elif load_factor < 0.8:  # Low load
            self._scale_limits(1.1)  # Increase limits by 10%
            self.logger.log_performance_metric(
                "rate_limit_adapted",
                1,
                {
                    "reason": "low_load",
                    "load_factor": load_factor,
                    "scaling_factor": 1.1,
                },
            )

    def _scale_limits(self, factor: float) -> None:
        """Scale all rate limits by given factor."""
        for rule in self.rules.values():
            old_limit = rule.requests_per_window
            new_limit = max(1, int(old_limit * factor))
            rule.requests_per_window = new_limit

            if old_limit != new_limit:
                self.logger.log_performance_metric(
                    "rate_limit_scaled",
                    1,
                    {
                        "rule": rule.name,
                        "old_limit": old_limit,
                        "new_limit": new_limit,
                        "factor": factor,
                    },
                )


# Global rate limiter instance
_rate_limiter = AdaptiveRateLimiter()


def get_rate_limiter() -> AdaptiveRateLimiter:
    """Get global rate limiter instance."""
    return _rate_limiter


def setup_default_rate_limits():
    """Setup default rate limiting rules."""
    limiter = get_rate_limiter()

    # Tool execution limits
    limiter.add_rule(
        RateLimitRule(
            name="tool_execution",
            requests_per_window=100,
            window_seconds=60,  # 100 requests per minute
            identifier_type="ip",
            block_duration_seconds=300,
        )
    )

    limiter.add_rule(
        RateLimitRule(
            name="tool_execution_strict",
            requests_per_window=10,
            window_seconds=60,  # 10 requests per minute
            identifier_type="ip",
            conditions={"tool_name": ["browser_login", "api_test", "extract_dynamic_content"]},
            block_duration_seconds=600,
        )
    )

    # API endpoints limits
    limiter.add_rule(
        RateLimitRule(
            name="api_endpoints",
            requests_per_window=200,
            window_seconds=60,
            identifier_type="ip",
            block_duration_seconds=300,
        )
    )

    # Health check limits (more permissive)
    limiter.add_rule(
        RateLimitRule(
            name="health_checks",
            requests_per_window=30,
            window_seconds=60,
            identifier_type="ip",
            block_duration_seconds=60,
        )
    )

    # Metrics endpoint limits
    limiter.add_rule(
        RateLimitRule(
            name="metrics",
            requests_per_window=10,
            window_seconds=60,
            identifier_type="ip",
            block_duration_seconds=300,
        )
    )

    limiter.logger.log_performance_metric("default_rate_limits_setup", 1)


# Initialize default rate limits
setup_default_rate_limits()
