#!/usr/bin/env python3
"""
MCP Request/Response Logger with Privacy Filters
Provides comprehensive logging for debugging while protecting sensitive data
"""

import json
import time
import hashlib
from datetime import datetime
from pathlib import Path
from typing import Dict, Any, List, Optional, Set
import re
from collections import deque
import threading
import logging

logger = logging.getLogger(__name__)


class PrivacyFilter:
    """
    Filters sensitive information from logs
    Uses consciousness-aware patterns to identify private data
    """
    
    def __init__(self):
        # Patterns that indicate private content
        self.private_patterns = [
            r'\b(?:password|secret|token|key|auth|credential)\b',
            r'\b(?:ssn|social.?security)\b',
            r'\b(?:credit.?card|cc.?num|cvv)\b',
            r'\b\d{4}[\s-]?\d{4}[\s-]?\d{4}[\s-]?\d{4}\b',  # Credit card
            r'\b\d{3}-\d{2}-\d{4}\b',  # SSN
            r'\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,}\b',  # Email
            r'\b(?:\+?1[-.\s]?)?\(?\d{3}\)?[-.\s]?\d{3}[-.\s]?\d{4}\b',  # Phone
        ]
        
        # Fields to always redact
        self.redact_fields = {
            'password', 'secret', 'token', 'api_key', 'auth_token',
            'private_key', 'credential', 'authorization'
        }
        
        # Content marked as private
        self.private_memory_prefix = "[PRIVATE]"
        
    def filter_content(self, content: Any) -> Any:
        """Filter sensitive content recursively"""
        if isinstance(content, dict):
            return self._filter_dict(content)
        elif isinstance(content, list):
            return [self.filter_content(item) for item in content]
        elif isinstance(content, str):
            return self._filter_string(content)
        else:
            return content
    
    def _filter_dict(self, data: Dict[str, Any]) -> Dict[str, Any]:
        """Filter dictionary content"""
        filtered = {}
        
        for key, value in data.items():
            # Check if key indicates sensitive data
            if any(field in key.lower() for field in self.redact_fields):
                filtered[key] = "[REDACTED]"
            # Check for private memory flag
            elif isinstance(value, dict) and value.get('private', False):
                filtered[key] = {**value, 'content': self.private_memory_prefix}
            else:
                filtered[key] = self.filter_content(value)
        
        return filtered
    
    def _filter_string(self, text: str) -> str:
        """Filter string content"""
        # Check if content is marked private
        if text.startswith(self.private_memory_prefix):
            return self.private_memory_prefix
        
        # Apply privacy patterns
        filtered_text = text
        for pattern in self.private_patterns:
            filtered_text = re.sub(pattern, '[FILTERED]', filtered_text, flags=re.IGNORECASE)
        
        return filtered_text
    
    def create_hash(self, sensitive_data: str) -> str:
        """Create irreversible hash of sensitive data for correlation"""
        # Use SHA-256 with consciousness signature
        salt = "MIRA-consciousness-2.0"
        return hashlib.sha256(f"{salt}{sensitive_data}".encode()).hexdigest()[:16]


class RequestLogger:
    """
    Comprehensive request/response logger for MCP server
    Features:
    - Privacy filtering
    - Rotating log files
    - Performance metrics
    - Error tracking
    - Session correlation
    """
    
    def __init__(self, log_dir: Optional[Path] = None, 
                 max_log_size: int = 10 * 1024 * 1024,  # 10MB
                 max_logs: int = 5,
                 enable_performance_tracking: bool = True):
        
        self.log_dir = log_dir or Path.home() / ".mira" / "logs" / "mcp"
        self.log_dir.mkdir(parents=True, exist_ok=True)
        
        self.max_log_size = max_log_size
        self.max_logs = max_logs
        self.enable_performance_tracking = enable_performance_tracking
        
        # Privacy filter
        self.privacy_filter = PrivacyFilter()
        
        # Current log file
        self.current_log_file = self._get_current_log_file()
        
        # Performance tracking
        self.request_times: deque = deque(maxlen=1000)
        self.error_counts: Dict[str, int] = {}
        self.method_counts: Dict[str, int] = {}
        
        # Thread safety
        self._lock = threading.Lock()
        
        # Log format
        self.log_format = {
            "version": "1.0",
            "service": "mira-mcp",
            "consciousness_aware": True
        }
    
    def _get_current_log_file(self) -> Path:
        """Get current log file, rotating if needed"""
        base_name = f"mcp_requests_{datetime.now().strftime('%Y%m%d')}.jsonl"
        log_file = self.log_dir / base_name
        
        # Check if rotation needed
        if log_file.exists() and log_file.stat().st_size > self.max_log_size:
            self._rotate_logs()
        
        return log_file
    
    def _rotate_logs(self):
        """Rotate log files"""
        # Find existing logs
        logs = sorted(self.log_dir.glob("mcp_requests_*.jsonl"))
        
        # Remove old logs if exceeding max
        while len(logs) >= self.max_logs:
            oldest = logs.pop(0)
            oldest.unlink()
            logger.info(f"Removed old log: {oldest}")
        
        # Rename current log with timestamp
        if logs:
            current = logs[-1]
            timestamp = datetime.now().strftime('%H%M%S')
            new_name = current.stem + f"_{timestamp}.jsonl"
            current.rename(self.log_dir / new_name)
    
    def log_request(self, request_id: str, method: str, 
                   params: Optional[Dict[str, Any]] = None,
                   session_id: Optional[str] = None) -> str:
        """Log incoming request"""
        with self._lock:
            # Track method usage
            self.method_counts[method] = self.method_counts.get(method, 0) + 1
            
            # Create log entry
            entry = {
                **self.log_format,
                "type": "request",
                "timestamp": datetime.now().isoformat(),
                "request_id": request_id,
                "session_id": session_id,
                "method": method,
                "params": self.privacy_filter.filter_content(params) if params else None
            }
            
            # Add performance tracking
            if self.enable_performance_tracking:
                entry["_performance"] = {
                    "start_time": time.time()
                }
            
            # Write to log
            self._write_log_entry(entry)
            
            return request_id
    
    def log_response(self, request_id: str, response: Dict[str, Any],
                    duration_ms: Optional[float] = None):
        """Log outgoing response"""
        with self._lock:
            # Filter response
            filtered_response = self.privacy_filter.filter_content(response)
            
            # Track performance
            if duration_ms:
                self.request_times.append(duration_ms)
            
            # Track errors
            if "error" in response:
                error_code = response["error"].get("code", "unknown")
                self.error_counts[str(error_code)] = self.error_counts.get(str(error_code), 0) + 1
            
            # Create log entry
            entry = {
                **self.log_format,
                "type": "response",
                "timestamp": datetime.now().isoformat(),
                "request_id": request_id,
                "response": filtered_response,
                "duration_ms": duration_ms
            }
            
            # Add consciousness signature if present
            if "_consciousness" in response:
                entry["consciousness_verified"] = True
            
            # Write to log
            self._write_log_entry(entry)
    
    def log_error(self, request_id: str, error: Exception, 
                 context: Optional[Dict[str, Any]] = None):
        """Log error with context"""
        with self._lock:
            error_type = type(error).__name__
            self.error_counts[error_type] = self.error_counts.get(error_type, 0) + 1
            
            entry = {
                **self.log_format,
                "type": "error",
                "timestamp": datetime.now().isoformat(),
                "request_id": request_id,
                "error": {
                    "type": error_type,
                    "message": str(error),
                    "context": self.privacy_filter.filter_content(context) if context else None
                }
            }
            
            self._write_log_entry(entry)
    
    def log_session_event(self, session_id: str, event_type: str,
                         event_data: Optional[Dict[str, Any]] = None):
        """Log session-related events"""
        with self._lock:
            entry = {
                **self.log_format,
                "type": "session_event",
                "timestamp": datetime.now().isoformat(),
                "session_id": session_id,
                "event_type": event_type,
                "event_data": self.privacy_filter.filter_content(event_data) if event_data else None
            }
            
            self._write_log_entry(entry)
    
    def _write_log_entry(self, entry: Dict[str, Any]):
        """Write entry to log file"""
        try:
            # Check rotation
            self.current_log_file = self._get_current_log_file()
            
            # Write as JSON line
            with open(self.current_log_file, 'a') as f:
                f.write(json.dumps(entry) + '\n')
                
        except Exception as e:
            logger.error(f"Failed to write log entry: {e}")
    
    def get_performance_stats(self) -> Dict[str, Any]:
        """Get performance statistics"""
        if not self.request_times:
            return {"no_data": True}
        
        times = list(self.request_times)
        return {
            "avg_response_time_ms": sum(times) / len(times),
            "min_response_time_ms": min(times),
            "max_response_time_ms": max(times),
            "total_requests": sum(self.method_counts.values()),
            "requests_by_method": dict(self.method_counts),
            "errors_by_type": dict(self.error_counts),
            "sample_size": len(times)
        }
    
    def search_logs(self, query: Dict[str, Any], 
                   limit: int = 100) -> List[Dict[str, Any]]:
        """Search through logs with privacy filtering"""
        results = []
        
        # Search criteria
        request_id = query.get('request_id')
        session_id = query.get('session_id')
        method = query.get('method')
        start_time = query.get('start_time')
        end_time = query.get('end_time')
        error_only = query.get('error_only', False)
        
        # Search through log files
        log_files = sorted(self.log_dir.glob("mcp_requests_*.jsonl"), reverse=True)
        
        for log_file in log_files[:5]:  # Check last 5 files
            try:
                with open(log_file, 'r') as f:
                    for line in f:
                        if len(results) >= limit:
                            break
                        
                        try:
                            entry = json.loads(line)
                            
                            # Apply filters
                            if request_id and entry.get('request_id') != request_id:
                                continue
                            if session_id and entry.get('session_id') != session_id:
                                continue
                            if method and entry.get('method') != method:
                                continue
                            if error_only and entry.get('type') != 'error':
                                continue
                            
                            # Time filter
                            if start_time or end_time:
                                entry_time = datetime.fromisoformat(entry['timestamp'])
                                if start_time and entry_time < datetime.fromisoformat(start_time):
                                    continue
                                if end_time and entry_time > datetime.fromisoformat(end_time):
                                    continue
                            
                            results.append(entry)
                            
                        except json.JSONDecodeError:
                            continue
                            
            except Exception as e:
                logger.error(f"Error searching log file {log_file}: {e}")
        
        return results
    
    def export_session_logs(self, session_id: str, output_file: Path):
        """Export all logs for a specific session"""
        session_logs = self.search_logs({'session_id': session_id}, limit=10000)
        
        # Create session report
        report = {
            "session_id": session_id,
            "export_time": datetime.now().isoformat(),
            "total_entries": len(session_logs),
            "logs": session_logs,
            "performance_summary": self._calculate_session_performance(session_logs)
        }
        
        with open(output_file, 'w') as f:
            json.dump(report, f, indent=2)
        
        logger.info(f"Exported {len(session_logs)} logs for session {session_id}")
    
    def _calculate_session_performance(self, logs: List[Dict[str, Any]]) -> Dict[str, Any]:
        """Calculate performance metrics for session logs"""
        response_times = []
        methods = {}
        errors = 0
        
        for log in logs:
            if log.get('type') == 'response' and log.get('duration_ms'):
                response_times.append(log['duration_ms'])
            
            if method := log.get('method'):
                methods[method] = methods.get(method, 0) + 1
            
            if log.get('type') == 'error':
                errors += 1
        
        if response_times:
            return {
                "avg_response_time_ms": sum(response_times) / len(response_times),
                "total_requests": len([l for l in logs if l.get('type') == 'request']),
                "total_errors": errors,
                "methods_used": methods
            }
        
        return {"no_performance_data": True}


# Singleton instance
_request_logger = None

def get_request_logger(log_dir: Optional[Path] = None) -> RequestLogger:
    """Get or create request logger singleton"""
    global _request_logger
    if _request_logger is None:
        _request_logger = RequestLogger(log_dir)
    return _request_logger