"""
Comprehensive logging system for training with file and console output.
"""

import logging
import os
import sys
from datetime import datetime
from typing import Optional


class TrainingLogger:
    """
    Logger for training with both file and console output.
    Provides structured logging with timestamps and context.
    """
    
    def __init__(self, log_dir: str = './logs', log_level: int = logging.INFO):
        """
        Initialize logger.
        
        Args:
            log_dir: Directory to save log files
            log_level: Logging level (logging.DEBUG, INFO, WARNING, ERROR)
        """
        self.log_dir = log_dir
        os.makedirs(log_dir, exist_ok=True)
        
        # Create logger
        self.logger = logging.getLogger('TrainingLogger')
        self.logger.setLevel(log_level)
        
        # Clear existing handlers
        self.logger.handlers = []
        
        # File handler
        timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
        log_file = os.path.join(log_dir, f'training_{timestamp}.log')
        file_handler = logging.FileHandler(log_file, encoding='utf-8')
        file_handler.setLevel(log_level)
        file_formatter = logging.Formatter(
            '%(asctime)s - %(levelname)s - [%(name)s] - %(message)s',
            datefmt='%Y-%m-%d %H:%M:%S'
        )
        file_handler.setFormatter(file_formatter)
        self.logger.addHandler(file_handler)
        
        # Console handler with different format
        console_handler = logging.StreamHandler(sys.stdout)
        console_handler.setLevel(log_level)
        console_formatter = logging.Formatter(
            '%(levelname)s - %(message)s'
        )
        console_handler.setFormatter(console_formatter)
        self.logger.addHandler(console_handler)
        
        self.log_file = log_file
        
        # Log initialization
        self.info(f"Training logger initialized. Log file: {log_file}")
    
    def debug(self, message: str):
        """Log debug message."""
        self.logger.debug(message)
    
    def info(self, message: str):
        """Log info message."""
        self.logger.info(message)
    
    def warning(self, message: str):
        """Log warning message."""
        self.logger.warning(message)
    
    def error(self, message: str, exc_info: bool = False):
        """Log error message."""
        self.logger.error(message, exc_info=exc_info)
    
    def critical(self, message: str, exc_info: bool = True):
        """Log critical error message."""
        self.logger.critical(message, exc_info=exc_info)
    
    def log_exception(self, exception: Exception, context: str = ""):
        """Log an exception with context."""
        context_msg = f"{context}: " if context else ""
        self.error(f"{context_msg}{type(exception).__name__}: {str(exception)}", exc_info=True)
    
    def log_training_step(self, step: int, loss: float, lr: float, grad_norm: Optional[float] = None):
        """Log training step information."""
        msg = f"Step {step}: Loss={loss:.6f}, LR={lr:.2e}"
        if grad_norm is not None:
            msg += f", GradNorm={grad_norm:.4f}"
        self.debug(msg)
    
    def log_epoch(self, epoch: int, train_loss: float, val_loss: Optional[float] = None):
        """Log epoch summary."""
        msg = f"Epoch {epoch}: TrainLoss={train_loss:.6f}"
        if val_loss is not None:
            msg += f", ValLoss={val_loss:.6f}"
        self.info(msg)
    
    def get_log_file(self) -> str:
        """Get the path to the log file."""
        return self.log_file


def setup_logging(log_dir: str = './logs', log_level: int = logging.INFO) -> TrainingLogger:
    """Set up and return a training logger."""
    return TrainingLogger(log_dir=log_dir, log_level=log_level)


