"""
Audit logging system for security monitoring and compliance
"""

import asyncio
import hashlib
import json
import logging
import os
import queue
import threading
import time
from collections.abc import Callable
from dataclasses import asdict, dataclass
from datetime import UTC, datetime, timezone
from queue import Queue
from typing import Any, Dict, List, Optional

from .logging import get_structured_logger


@dataclass
class AuditEvent:
    """Audit event data structure."""

    event_id: str
    timestamp: str
    event_type: str
    category: str
    severity: str
    user_id: Optional[str]
    session_id: Optional[str]
    ip_address: Optional[str]
    user_agent: Optional[str]
    resource: str
    action: str
    outcome: str
    details: dict[str, Any]
    metadata: dict[str, Any] = None

    def __post_init__(self):
        if self.metadata is None:
            self.metadata = {}

    def to_dict(self) -> dict[str, Any]:
        """Convert to dictionary for serialization."""
        return asdict(self)

    def to_json(self) -> str:
        """Convert to JSON string."""
        return json.dumps(self.to_dict(), ensure_ascii=False, default=str)


class AuditLogger:
    """Asynchronous audit logger with multiple output targets."""

    def __init__(self, app_name: str = "web-parser-mcp"):
        self.app_name = app_name
        self.logger = get_structured_logger()
        self.event_queue = Queue(maxsize=10000)
        self.targets: list[Callable[[AuditEvent], None]] = []
        self.running = False
        self.worker_thread = None
        self._lock = threading.Lock()

        # Setup default targets
        self.add_target(self._log_to_structured_log)
        self.add_target(self._log_to_file)

    def add_target(self, target_func: Callable[[AuditEvent], None]) -> None:
        """Add an audit log target."""
        with self._lock:
            self.targets.append(target_func)

    def remove_target(self, target_func: Callable[[AuditEvent], None]) -> bool:
        """Remove an audit log target."""
        with self._lock:
            if target_func in self.targets:
                self.targets.remove(target_func)
                return True
            return False

    def start(self) -> None:
        """Start the audit logging worker."""
        if self.running:
            return

        self.running = True
        self.worker_thread = threading.Thread(target=self._process_queue, daemon=True)
        self.worker_thread.start()
        self.logger.log_performance_metric("audit_logger_started", 1)

    def stop(self) -> None:
        """Stop the audit logging worker."""
        self.running = False
        if self.worker_thread:
            self.worker_thread.join(timeout=5)
        self.logger.log_performance_metric("audit_logger_stopped", 1)

    def log_event(self, event: AuditEvent) -> None:
        """Log an audit event asynchronously."""
        try:
            # Add to queue (non-blocking)
            self.event_queue.put_nowait(event)
        except queue.Full:
            # If queue is full, log immediately to prevent loss
            self._log_immediately(event)
        except Exception as e:
            # Other queue errors
            print(f"Audit queue error: {e}")
            self._log_immediately(event)

    def log(
        self,
        event_type: str,
        category: str,
        severity: str,
        resource: str,
        action: str,
        outcome: str,
        user_id: Optional[str] = None,
        session_id: Optional[str] = None,
        ip_address: Optional[str] = None,
        user_agent: Optional[str] = None,
        details: Optional[dict[str, Any]] = None,
        metadata: Optional[dict[str, Any]] = None,
    ) -> str:
        """
        Log an audit event with automatic context gathering.

        Returns:
            Event ID for tracking
        """
        # Generate event ID
        event_id = self._generate_event_id()

        # Create event
        event = AuditEvent(
            event_id=event_id,
            timestamp=datetime.now(UTC).isoformat(),
            event_type=event_type,
            category=category,
            severity=severity,
            user_id=user_id,
            session_id=session_id,
            ip_address=ip_address,
            user_agent=user_agent,
            resource=resource,
            action=action,
            outcome=outcome,
            details=details or {},
            metadata=metadata or {},
        )

        # Log asynchronously
        self.log_event(event)
        return event_id

    def _generate_event_id(self) -> str:
        """Generate unique event ID."""
        timestamp = str(int(time.time() * 1000000))
        random_part = hashlib.md5(str(time.time()).encode()).hexdigest()[:8]
        return f"evt_{timestamp}_{random_part}"

    def _process_queue(self) -> None:
        """Process audit events from queue."""
        while self.running:
            try:
                # Get event from queue (blocking)
                event = self.event_queue.get(timeout=1)

                # Process event with all targets
                with self._lock:
                    for target in self.targets:
                        try:
                            target(event)
                        except Exception as e:
                            # Log target failure but continue
                            print(f"Audit target failed: {e}")

                self.event_queue.task_done()

            except queue.Empty:
                # Timeout, continue loop
                continue
            except Exception as e:
                # Log processing error but continue
                print(f"Audit queue processing error: {e}")
                continue

    def _log_immediately(self, event: AuditEvent) -> None:
        """Log event immediately (fallback when queue is full)."""
        try:
            for target in self.targets:
                target(event)
        except Exception as e:
            print(f"Audit logging failed: {e}")

    def _log_to_structured_log(self, event: AuditEvent) -> None:
        """Log to structured logger."""
        self.logger.log(
            logging.INFO
            if event.severity in ["info", "low"]
            else logging.WARNING
            if event.severity == "medium"
            else logging.ERROR,
            f"Audit event: {event.event_type}",
            extra={
                "event_id": event.event_id,
                "event_type": event.event_type,
                "category": event.category,
                "severity": event.severity,
                "user_id": event.user_id,
                "session_id": event.session_id,
                "resource": event.resource,
                "action": event.action,
                "outcome": event.outcome,
                "details": event.details,
                "metadata": event.metadata,
            },
        )

    def _log_to_file(self, event: AuditEvent) -> None:
        """Log to audit file."""
        try:
            # Ensure audit directory exists
            audit_dir = os.path.join(os.getcwd(), "logs", "audit")
            os.makedirs(audit_dir, exist_ok=True)

            # Create filename based on date
            date_str = datetime.now().strftime("%Y%m%d")
            filename = f"audit_{date_str}.jsonl"
            filepath = os.path.join(audit_dir, filename)

            # Append to file
            with open(filepath, "a", encoding="utf-8") as f:
                f.write(event.to_json() + "\n")

        except Exception as e:
            print(f"File audit logging failed: {e}")


class SecurityAuditLogger(AuditLogger):
    """Specialized audit logger for security events."""

    def log_authentication(
        self,
        user_id: str,
        ip_address: str,
        user_agent: str,
        success: bool,
        auth_method: str = "unknown",
    ) -> str:
        """Log authentication event."""
        return self.log(
            event_type="authentication",
            category="auth",
            severity="high" if not success else "info",
            resource="system",
            action="login",
            outcome="success" if success else "failure",
            user_id=user_id,
            ip_address=ip_address,
            user_agent=user_agent,
            details={"auth_method": auth_method, "success": success},
        )

    def log_authorization(self, user_id: str, resource: str, action: str, allowed: bool) -> str:
        """Log authorization event."""
        return self.log(
            event_type="authorization",
            category="access",
            severity="high" if not allowed else "low",
            resource=resource,
            action=action,
            outcome="granted" if allowed else "denied",
            user_id=user_id,
            details={"allowed": allowed, "required_permission": action},
        )

    def log_tool_execution(
        self,
        user_id: str,
        session_id: str,
        tool_name: str,
        arguments: dict[str, Any],
        success: bool,
        execution_time: float,
        error: Optional[str] = None,
    ) -> str:
        """Log tool execution event."""
        # Sanitize arguments for logging (remove sensitive data)
        safe_args = self._sanitize_arguments(arguments)

        return self.log(
            event_type="tool_execution",
            category="tool",
            severity="medium" if not success else "low",
            resource=tool_name,
            action="execute",
            outcome="success" if success else "failure",
            user_id=user_id,
            session_id=session_id,
            details={
                "tool_name": tool_name,
                "arguments": safe_args,
                "execution_time_ms": execution_time,
                "success": success,
                "error": error,
            },
        )

    def log_rate_limit(
        self, identifier: str, rule_name: str, violations: int, blocked: bool
    ) -> str:
        """Log rate limiting event."""
        return self.log(
            event_type="rate_limit",
            category="security",
            severity="medium" if blocked else "low",
            resource="rate_limiter",
            action="check",
            outcome="blocked" if blocked else "allowed",
            details={
                "identifier": identifier,
                "rule_name": rule_name,
                "violations": violations,
                "blocked": blocked,
            },
        )

    def log_suspicious_activity(
        self,
        ip_address: str,
        user_agent: str,
        activity_type: str,
        details: dict[str, Any],
    ) -> str:
        """Log suspicious activity."""
        return self.log(
            event_type="suspicious_activity",
            category="security",
            severity="high",
            resource="system",
            action=activity_type,
            outcome="detected",
            ip_address=ip_address,
            user_agent=user_agent,
            details=details,
        )

    def log_data_access(
        self,
        user_id: str,
        data_type: str,
        access_type: str,
        record_count: int,
        filters: Optional[dict[str, Any]] = None,
    ) -> str:
        """Log data access event."""
        return self.log(
            event_type="data_access",
            category="data",
            severity="info",
            resource=data_type,
            action=access_type,
            outcome="completed",
            user_id=user_id,
            details={
                "data_type": data_type,
                "access_type": access_type,
                "record_count": record_count,
                "filters": filters or {},
            },
        )

    def log_config_change(
        self, user_id: str, config_key: str, old_value: Any, new_value: Any
    ) -> str:
        """Log configuration change."""
        return self.log(
            event_type="config_change",
            category="config",
            severity="medium",
            resource="configuration",
            action="modify",
            outcome="completed",
            user_id=user_id,
            details={
                "config_key": config_key,
                "old_value": str(old_value),
                "new_value": str(new_value),
            },
        )

    def _sanitize_arguments(self, arguments: dict[str, Any]) -> dict[str, Any]:
        """Sanitize arguments to remove sensitive data."""
        sensitive_keys = {
            "password",
            "token",
            "api_key",
            "secret",
            "key",
            "auth_token",
            "session_token",
            "credentials",
        }

        sanitized = {}
        for key, value in arguments.items():
            key_lower = key.lower()
            if any(sensitive in key_lower for sensitive in sensitive_keys):
                sanitized[key] = "[REDACTED]"
            elif isinstance(value, dict):
                sanitized[key] = self._sanitize_arguments(value)
            elif isinstance(value, list):
                sanitized[key] = [
                    self._sanitize_arguments({"item": item})["item"]
                    if isinstance(item, dict)
                    else "[REDACTED]"
                    if any(sensitive in str(item).lower() for sensitive in sensitive_keys)
                    else item
                    for item in value
                ]
            else:
                sanitized[key] = value

        return sanitized


# Global audit logger instance
_audit_logger = SecurityAuditLogger()


def get_audit_logger() -> SecurityAuditLogger:
    """Get global audit logger instance."""
    return _audit_logger


def init_audit_logging():
    """Initialize audit logging system."""
    logger = get_audit_logger()
    logger.start()


def shutdown_audit_logging():
    """Shutdown audit logging system."""
    logger = get_audit_logger()
    logger.stop()


# Context manager for audit logging
class AuditContext:
    """Context manager for audit logging."""

    def __init__(
        self,
        user_id: Optional[str] = None,
        session_id: Optional[str] = None,
        ip_address: Optional[str] = None,
        user_agent: Optional[str] = None,
    ):
        self.user_id = user_id
        self.session_id = session_id
        self.ip_address = ip_address
        self.user_agent = user_agent
        self.logger = get_audit_logger()
        self.events: list[str] = []

    def log_event(
        self,
        event_type: str,
        category: str,
        severity: str,
        resource: str,
        action: str,
        outcome: str,
        details: Optional[dict[str, Any]] = None,
    ) -> str:
        """Log event within this context."""
        event_id = self.logger.log(
            event_type=event_type,
            category=category,
            severity=severity,
            resource=resource,
            action=action,
            outcome=outcome,
            user_id=self.user_id,
            session_id=self.session_id,
            ip_address=self.ip_address,
            user_agent=self.user_agent,
            details=details,
        )
        self.events.append(event_id)
        return event_id

    def get_event_count(self) -> int:
        """Get number of events logged in this context."""
        return len(self.events)


# Initialize audit logging
init_audit_logging()
