import json
import logging
import time
from collections.abc import Awaitable, Callable
from uuid import uuid4

from fastapi import Request, Response
from starlette.middleware.base import BaseHTTPMiddleware

logger = logging.getLogger("nvr_gateway.requests")


class JsonRequestLoggingMiddleware(BaseHTTPMiddleware):
    async def dispatch(self, request: Request, call_next: Callable[[Request], Awaitable[Response]]) -> Response:
        request_id = request.headers.get("x-request-id", str(uuid4()))
        request.state.request_id = request_id
        start = time.perf_counter()
        status_code = 500
        response: Response | None = None
        try:
            response = await call_next(request)
            status_code = response.status_code
            return response
        finally:
            latency_ms = round((time.perf_counter() - start) * 1000, 2)
            log_record = {
                "request_id": request_id,
                "endpoint": request.url.path,
                "method": request.method,
                "latency_ms": latency_ms,
                "status_code": status_code,
                "user_id": getattr(request.state, "user_id", None),
            }
            logger.info(json.dumps(log_record))
            if response is not None:
                response.headers["x-request-id"] = request_id
