import os
import logging
import json
from logging.handlers import RotatingFileHandler, TimedRotatingFileHandler
from datetime import datetime, date, timezone, time
from flask import request, g
import traceback
from functools import wraps
from contextlib import contextmanager
from typing import Dict, Optional, Any, Union

def setup_logger(app):
    """Setup application logging with rotation and proper formatting"""

    # Ensure log directory exists
    log_dir = app.config.get('LOG_DIR')
    if not os.path.exists(log_dir):
        os.makedirs(log_dir, exist_ok=True)

    # Configure root logger
    root_logger = logging.getLogger()
    root_logger.setLevel(getattr(logging, app.config.get('LOG_LEVEL', 'INFO')))

    # Remove default handlers to avoid duplication
    for handler in root_logger.handlers[:]:
        root_logger.removeHandler(handler)

    # Setup daily rotating file handler
    log_file = os.path.join(log_dir, f"flask-api-{datetime.now().strftime('%Y-%m-%d')}.log")
    file_handler = TimedRotatingFileHandler(
        log_file,
        when='midnight',
        interval=1,
        backupCount=app.config.get('LOG_BACKUP_COUNT', 30),  # Keep 30 days of logs
        encoding='utf-8',
        utc=False
    )

    # Setup formatter
    formatter = logging.Formatter(
        fmt='[%(asctime)s] %(name)s.%(levelname)s: %(message)s',
        datefmt='%Y-%m-%d %H:%M:%S'
    )
    file_handler.setFormatter(formatter)

    # Add handler to root logger
    root_logger.addHandler(file_handler)

    # Add handler to pypandoc logger
    pypandoc_logger = logging.getLogger("pypandoc")
    pypandoc_logger.setLevel(logging.DEBUG)
    pypandoc_logger.addHandler(file_handler)

    # Setup console handler for development
    if app.config.get('DEBUG', False):
        console_handler = logging.StreamHandler()
        console_handler.setFormatter(formatter)
        root_logger.addHandler(console_handler)

    app.logger.info(f"Flask API logging initialized. Log file: {log_file}")

def get_logger(name=None):
    """Get a logger instance with optional name"""
    return logging.getLogger(name or 'flask-api')

def log_request_info():
    """Log incoming request information"""
    logger = get_logger('request')
    client_ip = request.environ.get('HTTP_X_FORWARDED_FOR', request.remote_addr)

    log_data = {
        'method': request.method,
        'path': request.path,
        'client_ip': client_ip,
        'user_agent': request.headers.get('User-Agent', 'Unknown'),
        'content_type': request.content_type
    }

    if request.is_json and request.get_json(silent=True):
        # Log request body for debugging (be careful with sensitive data)
        body = request.get_json()
        # Remove sensitive fields
        safe_body = {k: v for k, v in body.items() if k.lower() not in ['password', 'token', 'secret']}
        log_data['request_body'] = safe_body

    logger.info(f"Incoming request: {json.dumps(log_data)}")

def log_response_info(response):
    """Log response information"""
    logger = get_logger('response')

    log_data = {
        'status_code': response.status_code,
        'content_length': response.content_length,
        'content_type': response.content_type
    }

    logger.info(f"Response sent: {json.dumps(log_data)}")
    return response

def log_exception(exc_info=None):
    """Log exception with full traceback"""
    logger = get_logger('error')

    if exc_info is None:
        exc_info = True

    logger.error(
        f"Exception occurred: {traceback.format_exc()}",
        exc_info=exc_info,
        extra={
            'request_path': getattr(request, 'path', 'N/A'),
            'request_method': getattr(request, 'method', 'N/A'),
            'client_ip': getattr(request, 'remote_addr', 'N/A')
        }
    )

def performance_log(func):
    """Decorator to log function performance"""
    @wraps(func)
    def wrapper(*args, **kwargs):
        logger = get_logger('performance')
        start_time = datetime.now()

        try:
            result = func(*args, **kwargs)
            end_time = datetime.now()
            duration = (end_time - start_time).total_seconds()

            logger.info(f"Function {func.__name__} executed in {duration:.4f} seconds")
            return result

        except Exception as e:
            end_time = datetime.now()
            duration = (end_time - start_time).total_seconds()
            logger.error(f"Function {func.__name__} failed after {duration:.4f} seconds: {str(e)}")
            raise

    return wrapper

@contextmanager
def measure_performance(name: str):
    """Context manager to log performance of a code block"""
    logger = get_logger('performance')
    start_time = datetime.now()
    try:
        yield
    finally:
        end_time = datetime.now()
        duration = (end_time - start_time).total_seconds()
        logger.info(f"Block '{name}' executed in {duration:.4f} seconds")


def audit_log(action, user_id=None, details=None):
    """Log audit events for compliance"""
    logger = get_logger('audit')

    audit_data = {
        'timestamp': datetime.utcnow().isoformat(),
        'action': action,
        'user_id': user_id or getattr(g, 'user_id', 'anonymous'),
        'client_ip': getattr(request, 'remote_addr', 'N/A'),
        'user_agent': request.headers.get('User-Agent', 'Unknown') if request else 'N/A',
        'details': details or {}
    }

    logger.info(f"AUDIT: {json.dumps(audit_data)}")

def get_env_variable(var_name):
    import os
    """Get the environment variable or return exception."""
    try:
        return os.environ[var_name]
    except KeyError:
        raise Exception(f"Set the {var_name} environment variable")

def format_response(data, status_code=200):
    """Format the response to be returned by the API."""
    return {
        "status": "success" if status_code < 400 else "error",
        "data": data,
        "status_code": status_code
    }

def validate_request_data(request_data, required_fields):
    """Validate incoming request data against required fields."""
    missing_fields = [field for field in required_fields if field not in request_data]
    if missing_fields:
        raise ValueError(f"Missing fields: {', '.join(missing_fields)}")
