Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/api/middleware.py: 0%
26 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
1import hmac
2from typing import Awaitable, Callable
4from fastapi import status
5from starlette.middleware.base import BaseHTTPMiddleware
6from starlette.requests import Request
7from starlette.responses import JSONResponse, Response
9from prefect import settings
10from prefect.server import models
11from prefect.server.database import provide_database_interface
13NextMiddlewareFunction = Callable[[Request], Awaitable[Response]]
16class CsrfMiddleware(BaseHTTPMiddleware):
17 """
18 Middleware for CSRF protection. This middleware will check for a CSRF token
19 in the headers of any POST, PUT, PATCH, or DELETE request. If the token is
20 not present or does not match the token stored in the database for the
21 client, the request will be rejected with a 403 status code.
22 """
24 async def dispatch(
25 self, request: Request, call_next: NextMiddlewareFunction
26 ) -> Response:
27 """
28 Dispatch method for the middleware. This method will check for the
29 presence of a CSRF token in the headers of the request and compare it
30 to the token stored in the database for the client. If the token is not
31 present or does not match, the request will be rejected with a 403
32 status code.
33 """
35 request_needs_csrf_protection = request.method in {
36 "POST",
37 "PUT",
38 "PATCH",
39 "DELETE",
40 }
42 if (
43 settings.PREFECT_SERVER_CSRF_PROTECTION_ENABLED.value()
44 and request_needs_csrf_protection
45 ):
46 incoming_token = request.headers.get("Prefect-Csrf-Token")
47 incoming_client = request.headers.get("Prefect-Csrf-Client")
49 if incoming_token is None:
50 return JSONResponse(
51 {"detail": "Missing CSRF token."},
52 status_code=status.HTTP_403_FORBIDDEN,
53 )
55 if incoming_client is None:
56 return JSONResponse(
57 {"detail": "Missing client identifier."},
58 status_code=status.HTTP_403_FORBIDDEN,
59 )
61 db = provide_database_interface()
62 async with db.session_context() as session:
63 token = await models.csrf_token.read_token_for_client(
64 session=session, client=incoming_client
65 )
67 if token is None or not hmac.compare_digest(
68 token.token, incoming_token
69 ):
70 return JSONResponse(
71 {"detail": "Invalid CSRF token or client identifier."},
72 status_code=status.HTTP_403_FORBIDDEN,
73 headers={"Access-Control-Allow-Origin": "*"},
74 )
76 return await call_next(request)