"""API middleware for CORS, error handling, and request logging."""

import time
import uuid
from typing import Callable

import structlog
from fastapi import FastAPI, Request, Response
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from starlette.middleware.base import BaseHTTPMiddleware

from src.core.exceptions import DSOLError

logger = structlog.get_logger(__name__)


def add_cors_middleware(app: FastAPI) -> None:
    """Add CORS middleware with development settings.

    Args:
        app: FastAPI application instance.
    """
    app.add_middleware(
        CORSMiddleware,
        allow_origins=["*"],
        allow_credentials=True,
        allow_methods=["*"],
        allow_headers=["*"],
    )


class RequestIdMiddleware(BaseHTTPMiddleware):
    """Middleware to generate and propagate request IDs for tracing."""

    async def dispatch(
        self, request: Request, call_next: Callable
    ) -> Response:
        """Process request with request ID tracking.

        Args:
            request: The incoming request.
            call_next: The next middleware/handler in chain.

        Returns:
            Response with X-Request-ID header.
        """
        # Generate or use existing request ID
        request_id = request.headers.get("X-Request-ID", str(uuid.uuid4()))

        # Store in request state for handler access
        request.state.request_id = request_id

        # Bind request_id to structlog context for all logs in this request
        structlog.contextvars.clear_contextvars()
        structlog.contextvars.bind_contextvars(request_id=request_id)

        # Process request
        response = await call_next(request)

        # Add request ID to response headers
        response.headers["X-Request-ID"] = request_id

        return response


class RequestLoggingMiddleware(BaseHTTPMiddleware):
    """Middleware to log request/response information."""

    async def dispatch(
        self, request: Request, call_next: Callable
    ) -> Response:
        """Log request start and completion with timing.

        Args:
            request: The incoming request.
            call_next: The next middleware/handler in chain.

        Returns:
            Response from handler.
        """
        # Get request ID from state (set by RequestIdMiddleware)
        request_id = getattr(request.state, "request_id", "unknown")

        # Log request start
        logger.info(
            "request_started",
            method=request.method,
            path=request.url.path,
            request_id=request_id,
        )

        # Time the request
        start_time = time.perf_counter()

        try:
            response = await call_next(request)
        except Exception as exc:
            # Calculate duration even on error
            duration_ms = int((time.perf_counter() - start_time) * 1000)
            logger.error(
                "request_failed",
                method=request.method,
                path=request.url.path,
                request_id=request_id,
                duration_ms=duration_ms,
                error=str(exc),
            )
            raise

        # Calculate duration
        duration_ms = int((time.perf_counter() - start_time) * 1000)

        # Log request completion
        logger.info(
            "request_completed",
            method=request.method,
            path=request.url.path,
            status_code=response.status_code,
            duration_ms=duration_ms,
            request_id=request_id,
        )

        return response


def add_request_middleware(app: FastAPI) -> None:
    """Add request ID and logging middleware to the application.

    Note: Order matters - RequestIdMiddleware must be added last
    (executed first) so request_id is available for logging middleware.

    Args:
        app: FastAPI application instance.
    """
    # Add logging middleware first (executed second)
    app.add_middleware(RequestLoggingMiddleware)
    # Add request ID middleware last (executed first)
    app.add_middleware(RequestIdMiddleware)


async def dsol_error_handler(request: Request, exc: DSOLError) -> JSONResponse:
    """Handle DSOLError exceptions and return structured JSON response.

    Args:
        request: The incoming request.
        exc: The DSOLError exception.

    Returns:
        JSONResponse with error details.
    """
    request_id = getattr(request.state, "request_id", "unknown")

    logger.warning(
        "dsol_error",
        error_code=exc.code,
        error_message=exc.message,
        path=request.url.path,
        request_id=request_id,
    )
    return JSONResponse(
        status_code=400,
        content=exc.to_dict(),
    )


async def generic_exception_handler(request: Request, exc: Exception) -> JSONResponse:
    """Handle unhandled exceptions and return 500 error.

    Args:
        request: The incoming request.
        exc: The unhandled exception.

    Returns:
        JSONResponse with internal error response.
    """
    request_id = getattr(request.state, "request_id", "unknown")

    logger.exception(
        "unhandled_exception",
        path=request.url.path,
        method=request.method,
        request_id=request_id,
        error_type=type(exc).__name__,
    )
    return JSONResponse(
        status_code=500,
        content={
            "error": {
                "code": "INTERNAL_ERROR",
                "message": "An internal server error occurred",
                "details": {},
            }
        },
    )


def add_exception_handlers(app: FastAPI) -> None:
    """Register exception handlers on the FastAPI application.

    Args:
        app: FastAPI application instance.
    """
    app.add_exception_handler(DSOLError, dsol_error_handler)
    app.add_exception_handler(Exception, generic_exception_handler)
