"""Middleware for the Airflow upload plugin.

Provides request-id generation, correlation-id propagation, and HMAC
authentication as three separate BaseHTTPMiddleware classes.

Registration order on upload_app (last-added executes first in Starlette):
    upload_app.add_middleware(HMACVerificationMiddleware)
    upload_app.add_middleware(CorrelationIdMiddleware)
    upload_app.add_middleware(RequestIdMiddleware)

Execution order: RequestId → CorrelationId → HMAC → route handler.

NOTE: This module runs inside the Airflow api-server process.
Do NOT import from src.*.
"""

import logging
import time
import uuid

from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from starlette.responses import Response

from api.auth import VerificationError, verify_hmac_headers
from api.metrics import hmac_verify_duration_seconds
from api.responses import error_response

log = logging.getLogger(__name__)


class RequestIdMiddleware(BaseHTTPMiddleware):
    """Accept or generate X-Request-ID, store on request.state, echo in response."""

    async def dispatch(self, request: Request, call_next) -> Response:
        request_id = request.headers.get("X-Request-ID") or str(uuid.uuid4())
        request.state.request_id = request_id

        response = await call_next(request)
        response.headers["X-Request-ID"] = request_id
        return response


class CorrelationIdMiddleware(BaseHTTPMiddleware):
    """Forward or generate X-Correlation-ID, store on request.state, echo in response."""

    async def dispatch(self, request: Request, call_next) -> Response:
        correlation_id = request.headers.get("X-Correlation-ID")
        if not correlation_id:
            correlation_id = str(uuid.uuid4())
            log.warning(
                "Missing X-Correlation-ID on %s %s — generated fallback %s",
                request.method,
                request.url.path,
                correlation_id,
            )

        request.state.correlation_id = correlation_id

        response = await call_next(request)
        response.headers["X-Correlation-ID"] = correlation_id
        return response


class HMACVerificationMiddleware(BaseHTTPMiddleware):
    """Verify HMAC headers on every request. Short-circuits with 401/500 on failure."""

    async def dispatch(self, request: Request, call_next) -> Response:
        _t0 = time.perf_counter()
        try:
            ctx = verify_hmac_headers(request)
        except VerificationError as exc:
            hmac_verify_duration_seconds.labels(result="fail").observe(time.perf_counter() - _t0)
            log.warning(
                "HMAC verification failed: %s, request_id=%s, correlation_id=%s",
                exc,
                getattr(request.state, "request_id", None),
                getattr(request.state, "correlation_id", None),
            )
            return error_response("UNAUTHORIZED", str(exc), status_code=401)
        except Exception:
            hmac_verify_duration_seconds.labels(result="error").observe(time.perf_counter() - _t0)
            log.exception(
                "HMAC auth error, request_id=%s, correlation_id=%s",
                getattr(request.state, "request_id", None),
                getattr(request.state, "correlation_id", None),
            )
            return error_response(
                "INTERNAL_ERROR",
                "Authentication service unavailable",
                status_code=500,
            )

        hmac_verify_duration_seconds.labels(result="success").observe(time.perf_counter() - _t0)
        request.state.hmac_ctx = ctx
        request.state.auth_method = "hmac"
        log.info(
            "HMAC auth succeeded, method=hmac, request_id=%s, correlation_id=%s",
            getattr(request.state, "request_id", None),
            getattr(request.state, "correlation_id", None),
        )

        return await call_next(request)
