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

1import hmac 

2from typing import Awaitable, Callable 

3 

4from fastapi import status 

5from starlette.middleware.base import BaseHTTPMiddleware 

6from starlette.requests import Request 

7from starlette.responses import JSONResponse, Response 

8 

9from prefect import settings 

10from prefect.server import models 

11from prefect.server.database import provide_database_interface 

12 

13NextMiddlewareFunction = Callable[[Request], Awaitable[Response]] 

14 

15 

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 """ 

23 

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 """ 

34 

35 request_needs_csrf_protection = request.method in { 

36 "POST", 

37 "PUT", 

38 "PATCH", 

39 "DELETE", 

40 } 

41 

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") 

48 

49 if incoming_token is None: 

50 return JSONResponse( 

51 {"detail": "Missing CSRF token."}, 

52 status_code=status.HTTP_403_FORBIDDEN, 

53 ) 

54 

55 if incoming_client is None: 

56 return JSONResponse( 

57 {"detail": "Missing client identifier."}, 

58 status_code=status.HTTP_403_FORBIDDEN, 

59 ) 

60 

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 ) 

66 

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 ) 

75 

76 return await call_next(request)