Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/auth/oauth2_check.py: 23%

65 statements  

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

1import base64 

2import os 

3from typing import Final, cast 

4 

5import httpx 

6 

7from litellm._logging import verbose_proxy_logger 

8from litellm.llms.custom_httpx.http_handler import ( 

9 get_async_httpx_client, 

10 httpxSpecialProvider, 

11) 

12from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth 

13 

14 

15class Oauth2Handler: 

16 """ 

17 Handles OAuth2 token validation. 

18 """ 

19 

20 @staticmethod 

21 def _is_introspection_endpoint( 

22 token_info_endpoint: str, 

23 oauth_client_id: str | None, 

24 oauth_client_secret: str | None, 

25 ) -> bool: 

26 """ 

27 Determine if this is an introspection endpoint (requires POST) or token info endpoint (uses GET). 

28 

29 Args: 

30 token_info_endpoint: The OAuth2 endpoint URL 

31 oauth_client_id: OAuth2 client ID 

32 oauth_client_secret: OAuth2 client secret 

33 

34 Returns: 

35 bool: True if this is an introspection endpoint 

36 """ 

37 return ( 

38 "introspect" in token_info_endpoint.lower() 

39 and oauth_client_id is not None 

40 and oauth_client_secret is not None 

41 ) 

42 

43 @staticmethod 

44 def _prepare_introspection_request( 

45 token: str, 

46 oauth_client_id: str | None, 

47 oauth_client_secret: str | None, 

48 ) -> tuple[dict[str, str], dict[str, str]]: 

49 """ 

50 Prepare headers and data for OAuth2 introspection endpoint (RFC 7662). 

51 

52 Args: 

53 token: The OAuth2 token to validate 

54 oauth_client_id: OAuth2 client ID 

55 oauth_client_secret: OAuth2 client secret 

56 

57 Returns: 

58 Tuple of (headers, data) for the introspection request 

59 """ 

60 headers: Final = {"Content-Type": "application/x-www-form-urlencoded"} 

61 data: Final = {"token": token} 

62 

63 # Add client authentication if credentials are provided 

64 if oauth_client_id and oauth_client_secret: 

65 # Use HTTP Basic authentication for client credentials 

66 credentials: Final = base64.b64encode(f"{oauth_client_id}:{oauth_client_secret}".encode()).decode() 

67 headers["Authorization"] = f"Basic {credentials}" 

68 elif oauth_client_id: 

69 # For public clients, include client_id in the request body 

70 data["client_id"] = oauth_client_id 

71 

72 return headers, data 

73 

74 @staticmethod 

75 def _prepare_token_info_request(token: str) -> dict[str, str]: 

76 """ 

77 Prepare headers for generic token info endpoint. 

78 

79 Args: 

80 token: The OAuth2 token to validate 

81 

82 Returns: 

83 Dict of headers for the token info request 

84 """ 

85 return {"Authorization": f"Bearer {token}", "Content-Type": "application/json"} 

86 

87 @staticmethod 

88 def _extract_user_info( 

89 response_data: dict, 

90 user_id_field_name: str, 

91 user_role_field_name: str, 

92 user_team_id_field_name: str, 

93 ) -> tuple[str | None, str | None, str | None]: 

94 """ 

95 Extract user information from OAuth2 response. 

96 

97 Args: 

98 response_data: The response data from OAuth2 endpoint 

99 user_id_field_name: Field name for user ID 

100 user_role_field_name: Field name for user role 

101 user_team_id_field_name: Field name for team ID 

102 

103 Returns: 

104 Tuple of (user_id, user_role, user_team_id) 

105 """ 

106 user_id: Final = response_data.get(user_id_field_name) 

107 user_team_id: Final = response_data.get(user_team_id_field_name) 

108 user_role: Final = response_data.get(user_role_field_name) 

109 

110 return user_id, user_role, user_team_id 

111 

112 @staticmethod 

113 async def check_oauth2_token(token: str) -> UserAPIKeyAuth: 

114 """ 

115 Makes a request to the token introspection endpoint to validate the OAuth2 token. 

116 

117 This function implements OAuth2 token introspection according to RFC 7662. 

118 It supports both generic token info endpoints (GET) and OAuth2 introspection endpoints (POST). 

119 

120 Args: 

121 token (str): The OAuth2 token to validate. 

122 

123 Returns: 

124 UserAPIKeyAuth: If the token is valid, containing user information. 

125 

126 Raises: 

127 ValueError: If the token is invalid, the request fails, or the token info endpoint is not set. 

128 """ 

129 from litellm.proxy.proxy_server import premium_user 

130 

131 if premium_user is not True: 

132 raise ValueError( 

133 "Oauth2 token validation is only available for premium users" + CommonProxyErrors.not_premium_user.value 

134 ) 

135 

136 verbose_proxy_logger.debug("Oauth2 token validation for token=[set=%s]", token is not None) 

137 

138 # Get the token info endpoint from environment variable 

139 token_info_endpoint: Final = os.getenv("OAUTH_TOKEN_INFO_ENDPOINT") 

140 user_id_field_name: Final = os.environ.get("OAUTH_USER_ID_FIELD_NAME", "sub") 

141 user_role_field_name: Final = os.environ.get("OAUTH_USER_ROLE_FIELD_NAME", "role") 

142 user_team_id_field_name: Final = os.environ.get("OAUTH_USER_TEAM_ID_FIELD_NAME", "team_id") 

143 

144 # OAuth2 client credentials for introspection endpoint authentication 

145 oauth_client_id: Final = os.environ.get("OAUTH_CLIENT_ID") 

146 oauth_client_secret: Final = os.environ.get("OAUTH_CLIENT_SECRET") 

147 

148 if not token_info_endpoint: 

149 raise ValueError("OAUTH_TOKEN_INFO_ENDPOINT environment variable is not set") 

150 

151 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) 

152 

153 # Determine if this is an introspection endpoint (requires POST) or token info endpoint (uses GET) 

154 is_introspection_endpoint: Final = Oauth2Handler._is_introspection_endpoint( 

155 token_info_endpoint=token_info_endpoint, 

156 oauth_client_id=oauth_client_id, 

157 oauth_client_secret=oauth_client_secret, 

158 ) 

159 

160 try: 

161 if is_introspection_endpoint: 

162 # OAuth2 Token Introspection (RFC 7662) - requires POST with form data 

163 verbose_proxy_logger.debug("Using OAuth2 introspection endpoint (POST)") 

164 

165 headers, data = Oauth2Handler._prepare_introspection_request( 

166 token=token, 

167 oauth_client_id=oauth_client_id, 

168 oauth_client_secret=oauth_client_secret, 

169 ) 

170 

171 response = await client.post(token_info_endpoint, headers=headers, data=data) 

172 else: 

173 # Generic token info endpoint - uses GET with Bearer token 

174 verbose_proxy_logger.debug("Using generic token info endpoint (GET)") 

175 headers = Oauth2Handler._prepare_token_info_request(token=token) 

176 response = await client.get(token_info_endpoint, headers=headers) 

177 

178 # if it's a bad token we expect it to raise an HTTPStatusError 

179 response.raise_for_status() 

180 

181 # If we get here, the request was successful 

182 data = response.json() 

183 

184 verbose_proxy_logger.debug( 

185 "Oauth2 token validation for token=%s, response from endpoint=%s", 

186 token, 

187 data, 

188 ) 

189 

190 # For introspection endpoints, check if token is active 

191 if is_introspection_endpoint and not data.get("active", True): 

192 raise ValueError("Token is not active") 

193 

194 # Extract user information from response 

195 user_id, user_role, user_team_id = Oauth2Handler._extract_user_info( 

196 response_data=data, 

197 user_id_field_name=user_id_field_name, 

198 user_role_field_name=user_role_field_name, 

199 user_team_id_field_name=user_team_id_field_name, 

200 ) 

201 

202 return UserAPIKeyAuth( 

203 api_key=token, 

204 team_id=user_team_id, 

205 user_id=user_id, 

206 user_role=cast(LitellmUserRoles, user_role), 

207 ) 

208 except httpx.HTTPStatusError as e: 

209 # This will catch any 4xx or 5xx errors 

210 raise ValueError(f"Oauth 2.0 Token validation failed: {e}") 

211 except Exception as e: 

212 # This will catch any other errors (like network issues) 

213 raise ValueError(f"An error occurred during token validation: {e}")