"""
Performance metrics and monitoring system
"""

import statistics
import threading
import time
from collections import defaultdict, deque
from dataclasses import dataclass, field
from datetime import datetime, timedelta
from typing import Any, Dict, List, Optional


@dataclass
class MetricPoint:
    """Single metric measurement point."""

    timestamp: float
    value: float
    labels: dict[str, str] = field(default_factory=dict)


class TimeSeriesMetrics:
    """Time series metrics with sliding window."""

    def __init__(self, window_size: int = 1000, max_age_seconds: int = 3600):
        self.window_size = window_size
        self.max_age_seconds = max_age_seconds
        self._data: dict[str, deque] = defaultdict(deque)
        self._lock = threading.RLock()

    def record(self, name: str, value: float, labels: Optional[dict[str, str]] = None):
        """Record a metric value."""
        with self._lock:
            point = MetricPoint(timestamp=time.time(), value=value, labels=labels or {})

            data_queue = self._data[name]
            data_queue.append(point)

            # Maintain window size
            while len(data_queue) > self.window_size:
                data_queue.popleft()

            # Remove old entries
            cutoff_time = time.time() - self.max_age_seconds
            while data_queue and data_queue[0].timestamp < cutoff_time:
                data_queue.popleft()

    def get_stats(self, name: str) -> dict[str, Any]:
        """Get statistics for a metric."""
        with self._lock:
            data_queue = self._data.get(name, deque())
            if not data_queue:
                return {"count": 0, "available": False}

            values = [point.value for point in data_queue]
            timestamps = [point.timestamp for point in data_queue]

            return {
                "count": len(values),
                "min": min(values),
                "max": max(values),
                "avg": statistics.mean(values),
                "median": statistics.median(values),
                "stddev": statistics.stdev(values) if len(values) > 1 else 0,
                "latest": values[-1],
                "oldest": values[0],
                "time_range_seconds": timestamps[-1] - timestamps[0] if len(timestamps) > 1 else 0,
                "rate_per_second": len(values) / (timestamps[-1] - timestamps[0])
                if len(timestamps) > 1
                else 0,
                "available": True,
            }

    def get_all_stats(self) -> dict[str, dict[str, Any]]:
        """Get statistics for all metrics."""
        with self._lock:
            return {name: self.get_stats(name) for name in self._data.keys()}

    def clear(self, name: Optional[str] = None):
        """Clear metrics."""
        with self._lock:
            if name:
                self._data.pop(name, None)
            else:
                self._data.clear()


class Counter:
    """Simple counter metric."""

    def __init__(self):
        self._value = 0
        self._lock = threading.RLock()

    def increment(self, amount: int = 1):
        """Increment counter."""
        with self._lock:
            self._value += amount

    def get(self) -> int:
        """Get current value."""
        with self._lock:
            return self._value

    def reset(self):
        """Reset counter."""
        with self._lock:
            self._value = 0


class Gauge:
    """Gauge metric for current values."""

    def __init__(self):
        self._value = 0.0
        self._lock = threading.RLock()

    def set(self, value: float):
        """Set gauge value."""
        with self._lock:
            self._value = value

    def increment(self, amount: float = 1.0):
        """Increment gauge."""
        with self._lock:
            self._value += amount

    def decrement(self, amount: float = 1.0):
        """Decrement gauge."""
        with self._lock:
            self._value -= amount

    def get(self) -> float:
        """Get current value."""
        with self._lock:
            return self._value


class MetricsCollector:
    """Central metrics collection system."""

    def __init__(self):
        self.time_series = TimeSeriesMetrics()
        self.counters: dict[str, Counter] = defaultdict(Counter)
        self.gauges: dict[str, Gauge] = defaultdict(Gauge)
        self.histograms: dict[str, list[float]] = defaultdict(list)

    def record_time(self, name: str, duration: float, labels: Optional[dict[str, str]] = None):
        """Record execution time."""
        self.time_series.record(name, duration, labels)

    def increment_counter(self, name: str, amount: int = 1):
        """Increment counter."""
        self.counters[name].increment(amount)

    def set_gauge(self, name: str, value: float):
        """Set gauge value."""
        self.gauges[name].set(value)

    def record_histogram(self, name: str, value: float):
        """Record value in histogram."""
        self.histograms[name].append(value)
        # Keep only last 1000 values
        if len(self.histograms[name]) > 1000:
            self.histograms[name] = self.histograms[name][-1000:]

    def measure_time(self, name: str):
        """Context manager for measuring execution time."""
        return Timer(self, name)

    def get_summary(self) -> dict[str, Any]:
        """Get comprehensive metrics summary."""
        time_series_stats = self.time_series.get_all_stats()

        # Calculate histogram statistics
        histogram_stats = {}
        for name, values in self.histograms.items():
            if values:
                histogram_stats[name] = {
                    "count": len(values),
                    "min": min(values),
                    "max": max(values),
                    "avg": statistics.mean(values),
                    "median": statistics.median(values),
                    "p95": sorted(values)[int(len(values) * 0.95)]
                    if len(values) > 1
                    else max(values),
                    "p99": sorted(values)[int(len(values) * 0.99)]
                    if len(values) > 1
                    else max(values),
                }

        return {
            "timestamp": datetime.utcnow().isoformat(),
            "time_series": time_series_stats,
            "counters": {name: counter.get() for name, counter in self.counters.items()},
            "gauges": {name: gauge.get() for name, gauge in self.gauges.items()},
            "histograms": histogram_stats,
            "uptime_seconds": time.time() - getattr(self, "_start_time", time.time()),
        }

    def reset(self):
        """Reset all metrics."""
        self.time_series.clear()
        self.counters.clear()
        self.gauges.clear()
        self.histograms.clear()


class Timer:
    """Context manager for timing operations."""

    def __init__(
        self,
        collector: MetricsCollector,
        name: str,
        labels: Optional[dict[str, str]] = None,
    ):
        self.collector = collector
        self.name = name
        self.labels = labels or {}
        self.start_time = None

    def __enter__(self):
        self.start_time = time.time()
        return self

    def __exit__(self, exc_type, exc_val, exc_tb):
        if self.start_time is not None:
            duration = time.time() - self.start_time
            self.collector.record_time(self.name, duration, self.labels)


# Global metrics collector
_metrics_collector = MetricsCollector()
_metrics_collector._start_time = time.time()


def get_metrics_collector() -> MetricsCollector:
    """Get global metrics collector."""
    return _metrics_collector


def record_http_request(
    method: str, url: str, status_code: int, duration: float, cache_hit: bool = False
):
    """Record HTTP request metrics."""
    collector = get_metrics_collector()

    # Record timing
    labels = {
        "method": method,
        "status_code": str(status_code),
        "cache_hit": str(cache_hit),
    }
    collector.record_time("http_request_duration", duration, labels)

    # Record counters
    collector.increment_counter("http_requests_total")
    collector.increment_counter(f"http_requests_{method.lower()}")
    collector.increment_counter(f"http_status_{status_code}")

    if cache_hit:
        collector.increment_counter("http_cache_hits")
    else:
        collector.increment_counter("http_cache_misses")


def record_tool_usage(tool_name: str, duration: float, success: bool):
    """Record tool usage metrics."""
    collector = get_metrics_collector()

    labels = {"tool": tool_name, "success": str(success)}
    collector.record_time("tool_execution_duration", duration, labels)

    collector.increment_counter(f"tool_{tool_name}_total")
    if success:
        collector.increment_counter(f"tool_{tool_name}_success")
    else:
        collector.increment_counter(f"tool_{tool_name}_errors")


def update_connection_stats(active_connections: int, pool_size: int):
    """Update connection pool statistics."""
    collector = get_metrics_collector()

    collector.set_gauge("connection_pool_active", active_connections)
    collector.set_gauge("connection_pool_size", pool_size)
    collector.set_gauge(
        "connection_pool_utilization",
        active_connections / pool_size if pool_size > 0 else 0,
    )


def get_performance_report() -> dict[str, Any]:
    """Get comprehensive performance report."""
    collector = get_metrics_collector()

    summary = collector.get_summary()

    # Calculate additional metrics
    time_series = summary.get("time_series", {})

    # HTTP metrics
    http_requests = time_series.get("http_request_duration", {})
    cache_hits = summary.get("counters", {}).get("http_cache_hits", 0)
    cache_misses = summary.get("counters", {}).get("http_cache_misses", 0)
    total_cache_requests = cache_hits + cache_misses

    cache_hit_rate = cache_hits / total_cache_requests if total_cache_requests > 0 else 0

    # Tool success rates
    counters = summary.get("counters", {})
    tool_success_rates = {}

    for key, value in counters.items():
        if key.endswith("_total"):
            tool_name = key[:-6]  # Remove "_total" suffix
            success_key = f"{tool_name}_success"
            success_count = counters.get(success_key, 0)

            if value > 0:
                tool_success_rates[tool_name] = success_count / value

    report = {
        **summary,
        "derived_metrics": {
            "http_cache_hit_rate": cache_hit_rate,
            "http_requests_per_second": http_requests.get("rate_per_second", 0),
            "tool_success_rates": tool_success_rates,
        },
    }

    return report


# Convenience functions for quick metrics
def time_function(name: str):
    """Decorator for timing functions."""

    def decorator(func):
        def wrapper(*args, **kwargs):
            start_time = time.time()
            try:
                result = func(*args, **kwargs)
                duration = time.time() - start_time
                record_tool_usage(name, duration, True)
                return result
            except Exception as e:
                duration = time.time() - start_time
                record_tool_usage(name, duration, False)
                raise e

        return wrapper

    return decorator
