diff --git a/litellm/proxy/auth/oauth2_check.py b/litellm/proxy/auth/oauth2_check.py index 245d6a9e84c..37973e52376 100644 --- a/litellm/proxy/auth/oauth2_check.py +++ b/litellm/proxy/auth/oauth2_check.py @@ -3,15 +3,11 @@ from typing import Literal import httpx -from litellm.llms.custom_httpx.http_handler import ( - AsyncHTTPHandler, - HTTPHandler, - _get_async_httpx_client, - _get_httpx_client, -) +from litellm.llms.custom_httpx.http_handler import _get_async_httpx_client +from litellm.proxy._types import UserAPIKeyAuth -async def check_oauth2_token(token: str) -> Literal[True]: +async def check_oauth2_token(token: str) -> UserAPIKeyAuth: """ Makes a request to the token info endpoint to validate the OAuth2 token. @@ -26,6 +22,9 @@ async def check_oauth2_token(token: str) -> Literal[True]: """ # Get the token info endpoint from environment variable token_info_endpoint = os.getenv("OAUTH_TOKEN_INFO_ENDPOINT") + user_id_field_name = os.environ.get("OAUTH_USER_ID_FIELD_NAME", "sub") + user_role_field_name = os.environ.get("OAUTH_USER_ROLE_FIELD_NAME", "role") + user_team_id_field_name = os.environ.get("OAUTH_USER_TEAM_ID_FIELD_NAME", "team_id") if not token_info_endpoint: raise ValueError("OAUTH_TOKEN_INFO_ENDPOINT environment variable is not set") @@ -45,11 +44,19 @@ async def check_oauth2_token(token: str) -> Literal[True]: # You might want to add additional checks here based on the response # For example, checking if the token is expired or has the correct scope + user_id = data.get(user_id_field_name) + user_team_id = data.get(user_team_id_field_name) + user_role = data.get(user_role_field_name) - return True + return UserAPIKeyAuth( + api_key=token, + team_id=user_team_id, + user_id=user_id, + user_role=user_role, + ) except httpx.HTTPStatusError as e: # This will catch any 4xx or 5xx errors - raise ValueError(f"Token validation failed: {e}") + raise ValueError(f"Oauth 2.0 Token validation failed: {e}") except Exception as e: # This will catch any other errors (like network issues) raise ValueError(f"An error occurred during token validation: {e}")