From 287468ba0cb83243869491778f282bb14e2d29a1 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Mon, 8 Jun 2026 14:36:15 -0700 Subject: [PATCH] fix(identity): guard OAuth2 IdP role against unknown enum values --- litellm/identity/oauth2.py | 15 ++++++++--- .../identity/test_oauth2_builder.py | 27 ++++++++++++++++--- 2 files changed, 35 insertions(+), 7 deletions(-) diff --git a/litellm/identity/oauth2.py b/litellm/identity/oauth2.py index c6fd8b9aba3..4d6dd16f363 100644 --- a/litellm/identity/oauth2.py +++ b/litellm/identity/oauth2.py @@ -10,7 +10,7 @@ OAuth2 paths converge on the same construction surface. from __future__ import annotations -from typing import TYPE_CHECKING, Optional, cast +from typing import TYPE_CHECKING, Optional if TYPE_CHECKING: from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth @@ -38,12 +38,21 @@ def build_user_api_key_auth_from_oauth2_response( from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth user_id: Optional[str] = response_data.get(user_id_field_name) - user_role: Optional[str] = response_data.get(user_role_field_name) + raw_role = response_data.get(user_role_field_name) user_team_id: Optional[str] = response_data.get(user_team_id_field_name) + user_role: Optional[LitellmUserRoles] + if raw_role is None: + user_role = None + else: + try: + user_role = LitellmUserRoles(raw_role) + except ValueError: + user_role = LitellmUserRoles.INTERNAL_USER + return UserAPIKeyAuth( api_key=token, team_id=user_team_id, user_id=user_id, - user_role=cast("LitellmUserRoles", user_role), + user_role=user_role, ) diff --git a/tests/test_litellm/identity/test_oauth2_builder.py b/tests/test_litellm/identity/test_oauth2_builder.py index 84d5e24df42..af0c0a448b0 100644 --- a/tests/test_litellm/identity/test_oauth2_builder.py +++ b/tests/test_litellm/identity/test_oauth2_builder.py @@ -4,7 +4,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) from litellm.identity import build_user_api_key_auth_from_oauth2_response -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth def test_default_field_names_extract_from_introspection_response(): @@ -37,14 +37,33 @@ def test_custom_field_names_override_defaults(): def test_missing_fields_default_to_none(): - uak = build_user_api_key_auth_from_oauth2_response( - token="t", response_data={} - ) + uak = build_user_api_key_auth_from_oauth2_response(token="t", response_data={}) assert uak.user_id is None assert uak.user_role is None assert uak.team_id is None +def test_unknown_idp_role_defaults_to_internal_user(): + uak = build_user_api_key_auth_from_oauth2_response( + token="t", response_data={"sub": "u", "role": "definitely-not-a-role"} + ) + assert uak.user_role == LitellmUserRoles.INTERNAL_USER + + +def test_known_idp_role_passes_through(): + uak = build_user_api_key_auth_from_oauth2_response( + token="t", response_data={"sub": "u", "role": "proxy_admin"} + ) + assert uak.user_role == LitellmUserRoles.PROXY_ADMIN + + +def test_missing_role_field_stays_none(): + uak = build_user_api_key_auth_from_oauth2_response( + token="t", response_data={"sub": "u"} + ) + assert uak.user_role is None + + def test_token_is_hashed_into_token_field(): """The api_key is hashed by the UserAPIKeyAuth validator; the OAuth2 builder must not bypass that path."""