Documentation

Middleware Stack

ASGI middleware registration order, execution flow, and each layer — CORS, request context, JWT auth, logging, exception handlers.

Show:

Framework implementations

All middleware is registered in api/middleware/__init__.py via a single register_middleware(app) function called from api/main.py. Middleware is executed in reverse registration order (Starlette/ASGI convention).


Registration & Execution Order

# api/middleware/__init__.py
def register_middleware(app: FastAPI, settings: Settings) -> None:
    # Registered first = executes LAST (outermost)
    register_cors_middleware(app, settings)
 
    # Registered second = executes FOURTH
    app.add_middleware(JWTAuthMiddleware, jwt_service=get_jwt_service())
 
    # Registered third = executes THIRD
    app.add_middleware(RequestLoggingMiddleware, logging_service=get_logging_service())
 
    # Registered fourth = executes SECOND
    app.add_middleware(RequestContextMiddleware)
 
    # Exception handlers — closest to routes
    register_exception_handlers(app)

Actual execution order per request:

  1. RequestContextMiddleware — sets request ID, actor context shell, timing start
  2. RequestLoggingMiddleware — logs request start (has request ID available)
  3. JWTAuthMiddleware — validates token, populates actor context
  4. Route handler
  5. Exception handlers (on error)

On response (unwinding):

  1. Route handler returns
  2. JWTAuthMiddleware — no action on response
  3. RequestLoggingMiddleware — logs response with status code and duration
  4. RequestContextMiddleware — adds X-Request-ID response header
  5. CORS — adds Access-Control-Allow-* headers

Each Middleware

1. CORS (api/middleware/cors.py)

from starlette.middleware.cors import CORSMiddleware
 
def register_cors_middleware(app: FastAPI, settings: Settings) -> None:
    origins = [o.strip() for o in settings.cors_allowed_origins.split(",")]
    app.add_middleware(
        CORSMiddleware,
        allow_origins=origins,
        allow_credentials=True,
        allow_methods=["*"],
        allow_headers=["*"],
    )

Must be registered first so it's the outermost middleware and handles OPTIONS preflight before auth middleware runs.


2. Request Context (api/middleware/request_context.py)

Sets up request.state.ctx — a dict shared across the request lifecycle. Everything else reads from here.

import uuid, time
from starlette.middleware.base import BaseHTTPMiddleware
 
class RequestContextMiddleware(BaseHTTPMiddleware):
    async def dispatch(self, request: Request, call_next):
        request.state.ctx = {
            "request_id": str(uuid.uuid4()),
            "path": request.url.path,
            "method": request.method,
            "start_time": time.perf_counter(),
            "client_ip": request.client.host if request.client else None,
            "actor": {
                "actor_id": None,
                "actor_type": "anonymous",
                "roles": (),
                "account_ids": (),
            },
        }
        response = await call_next(request)
        response.headers["X-Request-ID"] = request.state.ctx["request_id"]
        return response

3. JWT Auth (api/middleware/jwt_auth.py)

Validates the Authorization: Bearer <token> header and updates request.state.ctx["actor"] with the authenticated identity.

class JWTAuthMiddleware(BaseHTTPMiddleware):
    PUBLIC_PATHS = {"/health", "/docs", "/openapi.json", "/auth/login", "/auth/signup"}
 
    def __init__(self, app, jwt_service: JWTService):
        super().__init__(app)
        self._jwt = jwt_service
 
    async def dispatch(self, request: Request, call_next):
        if request.url.path in self.PUBLIC_PATHS:
            return await call_next(request)
 
        auth_header = request.headers.get("Authorization", "")
        if not auth_header.startswith("Bearer "):
            return await call_next(request)  # Let route handle 401 if needed
 
        token = auth_header.removeprefix("Bearer ")
        if not self._jwt.validate_token(token):
            return JSONResponse({"detail": "Invalid or expired token"}, status_code=401)
 
        claims = self._jwt.extract_claims(token)
        request.state.ctx["actor"] = {
            "actor_id": claims["account_id"],
            "actor_type": claims.get("account_type", "user"),
            "roles": tuple(claims.get("roles", [])),
            "account_ids": (claims["account_id"],),
        }
        request.state.ctx["jwt"] = claims
        return await call_next(request)

4. Request Logging (api/middleware/logging.py)

Logs structured request start and end events. Uses request.state.ctx for request metadata.

class RequestLoggingMiddleware(BaseHTTPMiddleware):
    def __init__(self, app, logging_service: LoggingService):
        super().__init__(app)
        self._logger = logging_service
 
    async def dispatch(self, request: Request, call_next):
        ctx = getattr(request.state, "ctx", {})
        self._logger.info("request.start", extra={
            "request_id": ctx.get("request_id"),
            "path": ctx.get("path"),
            "method": ctx.get("method"),
        })
 
        response = await call_next(request)
 
        duration_ms = (time.perf_counter() - ctx.get("start_time", 0)) * 1000
        level = "error" if response.status_code >= 500 else \
                "warning" if response.status_code >= 400 else "info"
 
        getattr(self._logger, level)("request.end", extra={
            "request_id": ctx.get("request_id"),
            "status_code": response.status_code,
            "duration_ms": round(duration_ms, 2),
        })
        return response

5. Exception Handlers (api/middleware/error_handler.py)

Maps exceptions to structured JSON responses. Registered via register_exception_handlers(app).

def register_exception_handlers(app: FastAPI) -> None:
    @app.exception_handler(EntityNotFoundError)
    async def not_found_handler(request, exc):
        _log_warning("entity.not_found", exc, request)
        return JSONResponse({"detail": str(exc)}, status_code=404)
 
    @app.exception_handler(NotAuthorizedError)
    async def not_authorized_handler(request, exc):
        _log_warning("not.authorized", exc, request)
        return JSONResponse({"detail": str(exc)}, status_code=403)
 
    @app.exception_handler(RequestValidationError)
    async def validation_handler(request, exc):
        return JSONResponse({"detail": exc.errors()}, status_code=422)
 
    @app.exception_handler(Exception)
    async def catchall_handler(request, exc):
        _log_error("unhandled.exception", exc, request)
        ctx = getattr(request.state, "ctx", {})
        return JSONResponse(
            {"detail": "Internal server error", "request_id": ctx.get("request_id")},
            status_code=500,
        )

Webhook routes should return 200 on all errors to prevent provider retries:

@app.exception_handler(Exception)
async def webhook_error_handler(request, exc):
    if request.url.path.startswith("/webhooks/"):
        return JSONResponse({"status": "received"}, status_code=200)
    raise exc

Rules

  • Never handle auth, logging, or error formatting inside route handlers — that belongs in middleware.
  • Exception handlers are not technically ASGI middleware but are registered alongside them for the same reason: separation of concerns.
  • Adding a new middleware? Add it to register_middleware() in api/middleware/__init__.py and document its execution order position in this file.