Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/middleware/prometheus_auth_middleware.py: 42%

46 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2Prometheus Auth Middleware - Pure ASGI implementation 

3""" 

4 

5import json 

6from collections.abc import MutableMapping 

7from typing import Any, Final 

8 

9from fastapi import Request 

10from starlette.routing import get_route_path 

11from starlette.types import ASGIApp, Receive, Scope, Send 

12 

13import litellm 

14from litellm.proxy._types import SpecialHeaders 

15from litellm.proxy.auth.user_api_key_auth import user_api_key_auth 

16 

17# Cache the header name at module level to avoid repeated enum attribute access 

18_AUTHORIZATION_HEADER: Final = SpecialHeaders.openai_authorization.value # "Authorization" 

19_METRICS_MOUNT: Final = "/metrics" 

20 

21 

22def _is_metrics_route(scope: Scope) -> bool: 

23 route_path: Final = get_route_path(scope) 

24 return route_path == _METRICS_MOUNT or route_path.startswith(_METRICS_MOUNT + "/") 

25 

26 

27class PrometheusAuthMiddleware: 

28 """ 

29 Middleware to authenticate requests to the metrics endpoint. 

30 

31 By default, auth is run on the metrics endpoint. 

32 

33 To allow unauthenticated metrics in proxy_config.yaml: 

34 

35 ```yaml 

36 litellm_settings: 

37 require_auth_for_metrics_endpoint: false 

38 ``` 

39 """ 

40 

41 def __init__(self, app: ASGIApp) -> None: 

42 self.app = app 

43 

44 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: 

45 # Fast path: only inspect HTTP requests; pass through websocket/lifespan immediately 

46 if scope["type"] != "http" or not _is_metrics_route(scope): 46 ↛ 51line 46 didn't jump to line 51 because the condition on line 46 was always true

47 await self.app(scope, receive, send) 

48 return 

49 

50 # Run auth by default; allow legacy public metrics only when explicitly disabled. 

51 if litellm.require_auth_for_metrics_endpoint is not False: 

52 # user_api_key_auth reads the request body, which consumes ASGI `receive`. 

53 # Buffer those messages and replay them for the inner app; otherwise a 

54 # successful auth would forward an exhausted receive and /metrics hangs. 

55 buffered_messages: Final[list[MutableMapping[str, Any]]] = [] 

56 

57 async def receive_for_auth() -> MutableMapping[str, Any]: 

58 message: Final = await receive() 

59 buffered_messages.append(message) 

60 return message 

61 

62 request: Final = Request(scope, receive_for_auth) 

63 

64 try: 

65 await user_api_key_auth( 

66 request=request, 

67 api_key=request.headers.get(_AUTHORIZATION_HEADER) or "", 

68 azure_api_key_header=request.headers.get(SpecialHeaders.azure_authorization.value) or "", 

69 anthropic_api_key_header=request.headers.get(SpecialHeaders.anthropic_authorization.value), 

70 google_ai_studio_api_key_header=request.headers.get( 

71 SpecialHeaders.google_ai_studio_authorization.value 

72 ), 

73 azure_apim_header=request.headers.get(SpecialHeaders.azure_apim_authorization.value) or "", 

74 custom_litellm_key_header=request.headers.get(SpecialHeaders.custom_litellm_api_key.value), 

75 ) 

76 except Exception as e: 

77 # Send 401 response directly via ASGI protocol 

78 error_message: Final = getattr(e, "message", str(e)) 

79 body: Final = json.dumps( 

80 f"Unauthorized access to metrics endpoint: {error_message} " 

81 f"To allow unauthenticated access, set " 

82 f"`litellm_settings.require_auth_for_metrics_endpoint: false` " 

83 f"in your proxy_config.yaml." 

84 ).encode("utf-8") 

85 await send( 

86 { 

87 "type": "http.response.start", 

88 "status": 401, 

89 "headers": [ 

90 [b"content-type", b"application/json"], 

91 [b"content-length", str(len(body)).encode("ascii")], 

92 ], 

93 } 

94 ) 

95 await send( 

96 { 

97 "type": "http.response.body", 

98 "body": body, 

99 } 

100 ) 

101 return 

102 

103 replay_idx = 0 

104 

105 async def receive_replay() -> MutableMapping[str, Any]: 

106 nonlocal replay_idx 

107 if replay_idx < len(buffered_messages): 

108 msg: Final = buffered_messages[replay_idx] 

109 replay_idx += 1 

110 return msg 

111 return await receive() 

112 

113 await self.app(scope, receive_replay, send) 

114 return 

115 

116 # Pass through to the inner application 

117 await self.app(scope, receive, send)