"""
Security hardening utilities for input validation, sanitization, and protection
"""

import hashlib
import html
import ipaddress
import re
import secrets
import time
import urllib.parse
from collections.abc import Callable
from dataclasses import dataclass
from functools import wraps
from re import Pattern
from typing import Any, Dict, List, Optional, Union

from .audit import get_audit_logger
from .logging import get_structured_logger


@dataclass
class ValidationRule:
    """Input validation rule."""

    field_name: str
    validator: Callable[[Any], bool]
    error_message: str
    required: bool = True
    sanitizers: list[Callable[[Any], Any]] = None

    def __post_init__(self):
        if self.sanitizers is None:
            self.sanitizers = []


class SecurityValidator:
    """Comprehensive input validation and sanitization system."""

    def __init__(self):
        self.logger = get_structured_logger()
        self.audit_logger = get_audit_logger()
        self.validation_rules: dict[str, dict[str, ValidationRule]] = {}
        self._setup_default_validators()

    def _setup_default_validators(self):
        """Setup default validation rules for common inputs."""

        # URL validation
        self.add_rule(
            "url",
            ValidationRule(
                field_name="url",
                validator=self._validate_url,
                error_message="Invalid URL format or dangerous URL detected",
                required=True,
                sanitizers=[self._sanitize_url],
            ),
        )

        # HTML content validation
        self.add_rule(
            "html",
            ValidationRule(
                field_name="html",
                validator=self._validate_html,
                error_message="Invalid HTML content",
                required=True,
                sanitizers=[self._sanitize_html],
            ),
        )

        # CSS selector validation
        self.add_rule(
            "selector",
            ValidationRule(
                field_name="selector",
                validator=self._validate_css_selector,
                error_message="Invalid CSS selector",
                required=True,
                sanitizers=[self._sanitize_css_selector],
            ),
        )

        # Email validation
        self.add_rule(
            "email",
            ValidationRule(
                field_name="email",
                validator=self._validate_email,
                error_message="Invalid email format",
                required=False,
                sanitizers=[self._sanitize_email],
            ),
        )

        # IP address validation
        self.add_rule(
            "ip_address",
            ValidationRule(
                field_name="ip_address",
                validator=self._validate_ip_address,
                error_message="Invalid IP address",
                required=False,
                sanitizers=[self._sanitize_ip_address],
            ),
        )

        # File path validation
        self.add_rule(
            "file_path",
            ValidationRule(
                field_name="file_path",
                validator=self._validate_file_path,
                error_message="Invalid file path",
                required=False,
                sanitizers=[self._sanitize_file_path],
            ),
        )

        # SQL injection protection
        self.add_rule(
            "sql_safe",
            ValidationRule(
                field_name="sql_safe",
                validator=self._validate_sql_injection,
                error_message="Potential SQL injection detected",
                required=False,
                sanitizers=[self._sanitize_sql],
            ),
        )

        # XSS protection
        self.add_rule(
            "xss_safe",
            ValidationRule(
                field_name="xss_safe",
                validator=self._validate_xss,
                error_message="Potential XSS attack detected",
                required=False,
                sanitizers=[self._sanitize_xss],
            ),
        )

    def add_rule(self, rule_name: str, rule: ValidationRule) -> None:
        """Add a validation rule."""
        if rule_name not in self.validation_rules:
            self.validation_rules[rule_name] = {}
        self.validation_rules[rule_name][rule.field_name] = rule

    def validate_input(
        self,
        rule_name: str,
        data: dict[str, Any],
        context: Optional[dict[str, Any]] = None,
    ) -> dict[str, Any]:
        """
        Validate input data against rules.

        Returns:
            Dict with 'valid': bool, 'errors': list, 'sanitized_data': dict
        """
        if rule_name not in self.validation_rules:
            return {"valid": True, "errors": [], "sanitized_data": data}

        rules = self.validation_rules[rule_name]
        errors = []
        sanitized_data = {}
        valid = True

        for field_name, rule in rules.items():
            if field_name not in data and not rule.required:
                continue

            if field_name not in data and rule.required:
                errors.append(f"Required field '{field_name}' is missing")
                valid = False
                continue

            value = data[field_name]

            # Validate
            try:
                if not rule.validator(value):
                    errors.append(f"Field '{field_name}': {rule.error_message}")
                    valid = False
                    continue
            except Exception as e:
                errors.append(f"Field '{field_name}': Validation error - {str(e)}")
                valid = False
                continue

            # Sanitize
            sanitized_value = value
            for sanitizer in rule.sanitizers:
                try:
                    sanitized_value = sanitizer(sanitized_value)
                except Exception as e:
                    self.logger.log_performance_metric(
                        "sanitization_error", 1, {"field": field_name, "error": str(e)}
                    )

            sanitized_data[field_name] = sanitized_value

        # Audit logging for security events
        if not valid and context:
            self.audit_logger.log_suspicious_activity(
                ip_address=context.get("ip_address", "unknown"),
                user_agent=context.get("user_agent", "unknown"),
                activity_type="input_validation_failure",
                details={
                    "rule_name": rule_name,
                    "errors": errors,
                    "field_count": len(data),
                },
            )

        return {"valid": valid, "errors": errors, "sanitized_data": sanitized_data}

    # URL validation and sanitization
    def _validate_url(self, url: str) -> bool:
        """Validate URL for security."""
        if not isinstance(url, str) or len(url.strip()) == 0:
            return False

        url = url.strip()

        # Check length
        if len(url) > 2048:  # Common URL length limit
            return False

        # Parse URL
        try:
            parsed = urllib.parse.urlparse(url)
        except:
            return False

        # Must have scheme and netloc
        if not parsed.scheme or not parsed.netloc:
            return False

        # Only allow http/https
        if parsed.scheme.lower() not in ["http", "https"]:
            return False

        # Check for dangerous patterns
        dangerous_patterns = [
            r"\.\.",  # Directory traversal
            r"localhost",
            r"127\.0\.0\.1",
            r"0\.0\.0\.0",
            r"169\.254\.",  # Link-local
            r"10\.",  # Private network
            r"192\.168\.",  # Private network
            r"172\.(1[6-9]|2\d|3[01])\.",  # Private network
        ]

        url_lower = url.lower()
        for pattern in dangerous_patterns:
            if re.search(pattern, url_lower):
                return False

        return True

    def _sanitize_url(self, url: str) -> str:
        """Sanitize URL."""
        if not isinstance(url, str):
            return ""
        return urllib.parse.quote(url.strip(), safe=":/?=&")

    # HTML validation and sanitization
    def _validate_html(self, html_content: str) -> bool:
        """Validate HTML content."""
        if not isinstance(html_content, str):
            return False

        # Check length
        if len(html_content) > 10 * 1024 * 1024:  # 10MB limit
            return False

        # Check for dangerous patterns
        dangerous_patterns = [
            r"<script[^>]*>.*?</script>",  # Scripts
            r"javascript:",  # JavaScript URLs
            r"vbscript:",  # VBScript
            r"on\w+\s*=",  # Event handlers
            r"<iframe[^>]*>",  # Iframes
            r"<object[^>]*>",  # Objects
            r"<embed[^>]*>",  # Embeds
        ]

        html_lower = html_content.lower()
        for pattern in dangerous_patterns:
            if re.search(pattern, html_lower, re.IGNORECASE | re.DOTALL):
                return False

        return True

    def _sanitize_html(self, html_content: str) -> str:
        """Sanitize HTML content."""
        if not isinstance(html_content, str):
            return ""

        # Use html.escape for basic sanitization
        return html.escape(html_content)

    # CSS selector validation
    def _validate_css_selector(self, selector: str) -> bool:
        """Validate CSS selector."""
        if not isinstance(selector, str) or len(selector.strip()) == 0:
            return False

        selector = selector.strip()

        # Check length
        if len(selector) > 1000:
            return False

        # Basic pattern validation (simplified)
        valid_pattern = r"^[a-zA-Z0-9\s\[\]\'\"=\-\.\#\:\>\+\~\*\^\$\(\)]+$"
        return bool(re.match(valid_pattern, selector))

    def _sanitize_css_selector(self, selector: str) -> str:
        """Sanitize CSS selector."""
        if not isinstance(selector, str):
            return ""
        # Remove potentially dangerous characters
        return re.sub(r"[^\w\s\[\]\'\"=\-\.\#\:\>\+\~\*\^\$\(\)]", "", selector)

    # Email validation
    def _validate_email(self, email: str) -> bool:
        """Validate email address."""
        if not isinstance(email, str):
            return False

        pattern = r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$"
        return bool(re.match(pattern, email.strip()))

    def _sanitize_email(self, email: str) -> str:
        """Sanitize email address."""
        if not isinstance(email, str):
            return ""
        return email.strip().lower()

    # IP address validation
    def _validate_ip_address(self, ip: str) -> bool:
        """Validate IP address."""
        if not isinstance(ip, str):
            return False

        try:
            ipaddress.ip_address(ip.strip())
            return True
        except:
            return False

    def _sanitize_ip_address(self, ip: str) -> str:
        """Sanitize IP address."""
        if not isinstance(ip, str):
            return ""
        return ip.strip()

    # File path validation
    def _validate_file_path(self, path: str) -> bool:
        """Validate file path."""
        if not isinstance(path, str) or len(path.strip()) == 0:
            return False

        path = path.strip()

        # Check for dangerous patterns
        dangerous_patterns = [
            r"\.\./",  # Directory traversal
            r"\.\.",  # Parent directory
            r"~",  # Home directory
            r"/",  # Absolute path (depending on OS)
        ]

        for pattern in dangerous_patterns:
            if re.search(pattern, path):
                return False

        return True

    def _sanitize_file_path(self, path: str) -> str:
        """Sanitize file path."""
        if not isinstance(path, str):
            return ""
        # Remove dangerous characters
        return re.sub(r"[^\w\.\-\s/]", "", path)

    # SQL injection protection
    def _validate_sql_injection(self, sql_input: str) -> bool:
        """Validate SQL input for injection attacks."""
        if not isinstance(sql_input, str):
            return True  # Non-string inputs are safe

        # Common SQL injection patterns
        dangerous_patterns = [
            r";\s*(drop|delete|update|insert|alter)\s",
            r"--",  # Comments
            r"/\*.*\*/",  # Block comments
            r"\bunion\s+select\b",
            r"\bscript\s*>",
            r"\bexec\s*\(",
            r"\bxp_\w+",
        ]

        sql_lower = sql_input.lower()
        for pattern in dangerous_patterns:
            if re.search(pattern, sql_lower):
                return False

        return True

    def _sanitize_sql(self, sql_input: str) -> str:
        """Sanitize SQL input."""
        if not isinstance(sql_input, str):
            return str(sql_input) if sql_input is not None else ""

        # Basic sanitization - remove dangerous characters
        return re.sub(r'[;\'"\\]', "", sql_input)

    # XSS protection
    def _validate_xss(self, user_input: str) -> bool:
        """Validate input for XSS attacks."""
        if not isinstance(user_input, str):
            return True

        # XSS patterns
        xss_patterns = [
            r"<script[^>]*>.*?</script>",
            r"javascript:",
            r"vbscript:",
            r"on\w+\s*=",
            r"<iframe[^>]*>",
            r"<object[^>]*>",
            r"<embed[^>]*>",
            r"<form[^>]*>",
            r"<input[^>]*>",
            r"<img[^>]*onerror",
        ]

        input_lower = user_input.lower()
        for pattern in xss_patterns:
            if re.search(pattern, input_lower, re.IGNORECASE):
                return False

        return True

    def _sanitize_xss(self, user_input: str) -> str:
        """Sanitize input to prevent XSS."""
        if not isinstance(user_input, str):
            return str(user_input) if user_input is not None else ""

        # HTML escape
        return html.escape(user_input)


class SecurityMiddleware:
    """Security middleware for request processing."""

    def __init__(self):
        self.validator = SecurityValidator()
        self.logger = get_structured_logger()
        self.audit_logger = get_audit_logger()

    async def process_request(self, context):
        """Process request through security validation."""
        from .middleware import MiddlewareResult
        
        tool_name = context.tool_name
        arguments = context.arguments
        user_context = context.user_context

        # Get validation rule for tool
        validation_result = self.validator.validate_input(tool_name, arguments, user_context)

        if not validation_result["valid"]:
            # Log security violation
            self.audit_logger.log_suspicious_activity(
                ip_address=user_context.get("ip_address", "unknown") if user_context else "unknown",
                user_agent=user_context.get("user_agent", "unknown") if user_context else "unknown",
                activity_type="security_violation",
                details={
                    "tool_name": tool_name,
                    "validation_errors": validation_result["errors"],
                    "arguments_count": len(arguments),
                },
            )

            # Return error response
            return MiddlewareResult(
                success=False,
                error=f"Security validation failed: {validation_result['errors']}"
            )

        # Update context with sanitized data
        context.arguments = validation_result["sanitized_data"]
        return MiddlewareResult(success=True)

    async def process_response(self, context, result):
        """Process response through security filters."""
        from .middleware import MiddlewareResult
        
        # For now, just pass through the response
        # Could add response sanitization here
        return MiddlewareResult(success=True, data=result)


class ContentSecurityPolicy:
    """Content Security Policy utilities."""

    @staticmethod
    def generate_csp_header(directives: Optional[dict[str, list[str]]] = None) -> str:
        """Generate CSP header."""
        if directives is None:
            directives = {
                "default-src": ["'self'"],
                "script-src": ["'self'", "'unsafe-inline'"],
                "style-src": ["'self'", "'unsafe-inline'"],
                "img-src": ["'self'", "data:", "https:"],
                "font-src": ["'self'", "https:"],
                "connect-src": ["'self'"],
                "media-src": ["'self'"],
                "object-src": ["'none'"],
                "frame-src": ["'none'"],
                "base-uri": ["'self'"],
                "form-action": ["'self'"],
            }

        csp_parts = []
        for directive, sources in directives.items():
            csp_parts.append(f"{directive} {' '.join(sources)}")

        return "; ".join(csp_parts)


class SecurityHeaders:
    """Security headers utilities."""

    @staticmethod
    def get_security_headers() -> dict[str, str]:
        """Get comprehensive security headers."""
        return {
            "X-Content-Type-Options": "nosniff",
            "X-Frame-Options": "DENY",
            "X-XSS-Protection": "1; mode=block",
            "Strict-Transport-Security": "max-age=31536000; includeSubDomains",
            "Referrer-Policy": "strict-origin-when-cross-origin",
            "Permissions-Policy": "geolocation=(), microphone=(), camera=()",
            "Content-Security-Policy": ContentSecurityPolicy.generate_csp_header(),
            "Cross-Origin-Embedder-Policy": "require-corp",
            "Cross-Origin-Opener-Policy": "same-origin",
            "Cross-Origin-Resource-Policy": "cross-origin",
        }


class EncryptionUtils:
    """Encryption and hashing utilities."""

    @staticmethod
    def hash_sensitive_data(data: str, salt: Optional[str] = None) -> str:
        """Hash sensitive data."""
        if salt is None:
            salt = secrets.token_hex(16)

        return hashlib.sha256(f"{salt}:{data}".encode()).hexdigest()

    @staticmethod
    def generate_secure_token(length: int = 32) -> str:
        """Generate secure random token."""
        return secrets.token_urlsafe(length)

    @staticmethod
    def generate_api_key() -> str:
        """Generate API key."""
        return f"ak_{secrets.token_urlsafe(32)}"

    @staticmethod
    def generate_session_id() -> str:
        """Generate session ID."""
        return f"ses_{secrets.token_urlsafe(16)}_{int(time.time())}"


# Global instances
_security_validator = SecurityValidator()
_security_middleware = SecurityMiddleware()


def get_security_validator() -> SecurityValidator:
    """Get global security validator."""
    return _security_validator


def get_security_middleware() -> SecurityMiddleware:
    """Get global security middleware."""
    return _security_middleware


def secure_tool_wrapper(tool_func: Callable) -> Callable:
    """Decorator to add security validation to tools."""

    @wraps(tool_func)
    async def wrapper(arguments: dict[str, Any], context: Optional[dict[str, Any]] = None):
        middleware = get_security_middleware()

        # Validate and sanitize input
        sanitized_args = await middleware.process_request(tool_func.__name__, arguments, context)

        # Check if validation failed
        if "error" in sanitized_args:
            return [{"type": "text", "text": f"Security Error: {sanitized_args['error']}"}]

        # Execute tool with sanitized arguments
        return await tool_func(sanitized_args)

    return wrapper


# Common security patterns
SQL_INJECTION_PATTERNS = [
    r";\s*(drop|delete|update|insert|alter|exec|execute)\s",
    r"\bunion\s+select\b",
    r"\bscript\s*>",
    r"--",
    r"/\*.*\*/",
]

XSS_PATTERNS = [
    r"<script[^>]*>.*?</script>",
    r"javascript:",
    r"vbscript:",
    r"on\w+\s*=",
    r"<iframe[^>]*>",
    r"<object[^>]*>",
    r"<embed[^>]*>",
]

PATH_TRAVERSAL_PATTERNS = [
    r"\.\./",
    r"\.\.",
    r"~",
    r"\\",
]


def check_security_patterns(text: str, patterns: list[str]) -> bool:
    """Check text against security patterns."""
    if not isinstance(text, str):
        return True

    text_lower = text.lower()
    for pattern in patterns:
        if re.search(pattern, text_lower, re.IGNORECASE):
            return False
    return True
