From e24819afefd385fcad4a8468dcbdcf09c628b648 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Fri, 27 Mar 2026 13:50:30 -0700 Subject: [PATCH] fix(sso): pass decoded JWT access token to role mapping during SSO login During SSO login, bearer tokens are stripped from the OAuth response before role mapping runs. Custom role claims encoded inside the JWT access token are lost, so map_jwt_role_to_litellm_role() returns None and the user falls back to internal_user_viewer. process_sso_jwt_access_token() now returns the decoded JWT payload, and a new _sync_user_role_from_jwt_role_map() receives it so jwt_litellm_role_map works correctly during SSO login. --- litellm/proxy/management_endpoints/ui_sso.py | 94 +++++++++- .../proxy/management_endpoints/test_ui_sso.py | 161 +++++++++++++++++- 2 files changed, 246 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 79a9e9bdf2..a5a62813b9 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -204,7 +204,7 @@ def process_sso_jwt_access_token( sso_jwt_handler: Optional[JWTHandler], result: Union[OpenID, dict, None], role_mappings: Optional["RoleMappings"] = None, -) -> None: +) -> Optional[dict]: """ Process SSO JWT access token and extract team IDs and user role if available. @@ -218,6 +218,12 @@ def process_sso_jwt_access_token( sso_jwt_handler: SSO-specific JWT handler for team ID extraction result: The SSO result object to update with team IDs and role role_mappings: Optional role mappings configuration for group-based role determination + + Returns: + The decoded access token payload dict, or None if decoding failed or + inputs were missing. Callers can pass this to _sync_user_role_from_jwt_role_map + so it has access to custom role claims (e.g. custom_roles) that are + encoded inside the JWT but stripped from received_response. """ if access_token_str and result: import jwt @@ -230,7 +236,7 @@ def process_sso_jwt_access_token( verbose_proxy_logger.debug( "Access token is not a valid JWT (possibly an opaque token), skipping JWT-based extraction" ) - return + return None # Extract team IDs from access token if sso_jwt_handler is available if sso_jwt_handler: @@ -306,6 +312,10 @@ def process_sso_jwt_access_token( f"Set user_role='{user_role}' from JWT access token" ) + return access_token_payload + + return None + @router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False) async def google_login( @@ -817,7 +827,7 @@ async def get_generic_sso_response( ], # sso specific jwt handler - used for restricted sso group access control generic_client_id: str, redirect_url: str, -) -> Tuple[Union[OpenID, dict], Optional[dict]]: # return received response +) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload) # make generic sso provider from fastapi_sso.sso.base import DiscoveryDocument from fastapi_sso.sso.generic import create_provider @@ -872,6 +882,7 @@ async def get_generic_sso_response( code_verifier: Optional[ str ] = None # assigned inside try; initialized for type tracking + access_token_payload: Optional[dict] = None # decoded JWT access token claims try: token_exchange_params = ( @@ -958,7 +969,7 @@ async def get_generic_sso_response( ) access_token_str = generic_sso.access_token - process_sso_jwt_access_token( + access_token_payload = process_sso_jwt_access_token( access_token_str, sso_jwt_handler, result, role_mappings=role_mappings ) # Delete the single-use PKCE verifier only after all downstream processing @@ -976,7 +987,7 @@ async def get_generic_sso_response( additional_generic_sso_headers_dict, ) verbose_proxy_logger.debug("generic result: %s", result) - return result or {}, received_response + return result or {}, received_response, access_token_payload async def create_team_member_add_task(team_id, user_info): @@ -1176,6 +1187,56 @@ def _build_sso_user_update_data( return update_data +async def _sync_user_role_from_jwt_role_map( + jwt_handler: Optional[JWTHandler], + received_response: Optional[dict], + user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]], + prisma_client: PrismaClient, + user_api_key_cache: DualCache, + user_defined_values: Optional[SSOUserDefinedValues], +) -> None: + """ + Apply jwt_litellm_role_map during SSO login. + + When jwt_litellm_role_map is configured with sync_user_role_and_teams=True, + this ensures SSO users get the same role mapping as API/JWT users. Without + this, the SSO path falls back to INTERNAL_USER_VIEW_ONLY for roles that + don't directly match LitellmUserRoles enum values. + """ + if jwt_handler is None or received_response is None: + return + if not jwt_handler.litellm_jwtauth.sync_user_role_and_teams: + return + if not jwt_handler.litellm_jwtauth.jwt_litellm_role_map: + return + + mapped_role = jwt_handler.map_jwt_role_to_litellm_role(received_response) + if mapped_role is None: + return + + verbose_proxy_logger.info( + f"SSO jwt_litellm_role_map matched role: {mapped_role.value}" + ) + + # Update user_defined_values so downstream code uses the mapped role + if user_defined_values is not None: + user_defined_values["user_role"] = mapped_role.value + + # Update existing DB record if role differs + if user_info is not None and user_info.user_role != mapped_role.value: + await prisma_client.db.litellm_usertable.update( + where={"user_id": user_info.user_id}, + data={"user_role": mapped_role.value}, + ) + user_info.user_role = mapped_role.value + await user_api_key_cache.async_set_cache( + key=user_info.user_id, + value=user_info.model_dump() + if hasattr(user_info, "model_dump") + else dict(user_info), + ) + + def apply_user_info_values_to_sso_user_defined_values( user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]], user_defined_values: Optional[SSOUserDefinedValues], @@ -1279,6 +1340,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: google_client_id = os.getenv("GOOGLE_CLIENT_ID", None) generic_client_id = os.getenv("GENERIC_CLIENT_ID", None) received_response: Optional[dict] = None + access_token_payload: Optional[dict] = None # get url from request if master_key is None: raise ProxyException( @@ -1307,7 +1369,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: ) elif generic_client_id is not None: - result, received_response = await get_generic_sso_response( + result, received_response, access_token_payload = await get_generic_sso_response( request=request, jwt_handler=jwt_handler, generic_client_id=generic_client_id, @@ -1345,6 +1407,8 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: received_response=received_response, generic_client_id=generic_client_id, ui_access_mode=ui_access_mode, + access_token_payload=access_token_payload, + jwt_handler=jwt_handler, return_to=cp_return_to, ) @@ -2417,6 +2481,8 @@ class SSOAuthenticationHandler: received_response: Optional[dict] = None, generic_client_id: Optional[str] = None, ui_access_mode: Optional[Dict] = None, + access_token_payload: Optional[dict] = None, + jwt_handler: Optional[JWTHandler] = None, return_to: Optional[str] = None, ) -> RedirectResponse: import jwt @@ -2498,6 +2564,20 @@ class SSOAuthenticationHandler: alternate_user_id=user_id, ) + # Sync user role from JWT claims via jwt_litellm_role_map (if configured). + # This ensures SSO users get the same role mapping as API/JWT users. + # Use the decoded access_token_payload (not received_response) because + # custom role claims (e.g. custom_roles) are encoded inside the JWT + # access token, which is stripped from received_response. + await _sync_user_role_from_jwt_role_map( + jwt_handler=jwt_handler, + received_response=access_token_payload or received_response, + user_info=user_info, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_defined_values=user_defined_values, + ) + user_defined_values = apply_user_info_values_to_sso_user_defined_values( user_info=user_info, user_defined_values=user_defined_values ) @@ -3703,7 +3783,7 @@ async def debug_sso_callback(request: Request): ) elif generic_client_id is not None: - result, _ = await get_generic_sso_response( + result, _, _ = await get_generic_sso_response( request=request, jwt_handler=jwt_handler, generic_client_id=generic_client_id, diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index ff636ca04a..f9c7cefcc4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -24,6 +24,7 @@ from litellm.proxy.management_endpoints.ui_sso import ( MicrosoftSSOHandler, SSOAuthenticationHandler, _setup_team_mappings, + _sync_user_role_from_jwt_role_map, determine_role_from_groups, normalize_email, process_sso_jwt_access_token, @@ -1321,7 +1322,7 @@ async def test_get_generic_sso_response_with_additional_headers(): "fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class ): # Act - result, received_response = await get_generic_sso_response( + result, received_response, _ = await get_generic_sso_response( request=mock_request, jwt_handler=mock_jwt_handler, generic_client_id=generic_client_id, @@ -1383,7 +1384,7 @@ async def test_get_generic_sso_response_with_empty_headers(): "fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class ): # Act - result, received_response = await get_generic_sso_response( + result, received_response, _ = await get_generic_sso_response( request=mock_request, jwt_handler=mock_jwt_handler, generic_client_id=generic_client_id, @@ -5254,3 +5255,159 @@ class TestValidateReturnTo: ) SSOAuthenticationHandler._validate_return_to("https://cp.example.com:3000/ui") + +class TestSyncUserRoleFromJwtRoleMap: + """Tests for _sync_user_role_from_jwt_role_map.""" + + @staticmethod + def _make_jwt_handler(): + from litellm.caching.caching import DualCache + from litellm.proxy._types import ( + JWTLiteLLMRoleMap, + LiteLLM_JWTAuth, + LitellmUserRoles, + ) + + handler = JWTHandler() + handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth( + roles_jwt_field="custom_roles", + user_id_upsert=True, + sync_user_role_and_teams=True, + jwt_litellm_role_map=[ + JWTLiteLLMRoleMap( + jwt_role="my-admin", + litellm_role=LitellmUserRoles.PROXY_ADMIN, + ), + JWTLiteLLMRoleMap( + jwt_role="my-viewer", + litellm_role=LitellmUserRoles.INTERNAL_USER, + ), + ], + ), + ) + return handler + + @staticmethod + def _make_sso_values(user_role=None): + from litellm.proxy._types import SSOUserDefinedValues + + user_id = "testuser@example.com" + return SSOUserDefinedValues( + models=[], + user_id=user_id, + user_email=user_id, + user_role=user_role, + max_budget=None, + budget_duration=None, + ) + + @pytest.mark.asyncio + async def test_stripped_response_has_no_roles(self): + """Bug repro: stripped received_response lacks role claims.""" + from litellm.caching.caching import DualCache + + handler = self._make_jwt_handler() + sso_values = self._make_sso_values() + + await _sync_user_role_from_jwt_role_map( + jwt_handler=handler, + received_response={"token_type": "Bearer", "expires_in": 3600}, + user_info=None, + prisma_client=AsyncMock(), + user_api_key_cache=DualCache(), + user_defined_values=sso_values, + ) + + assert sso_values["user_role"] is None + + @pytest.mark.asyncio + async def test_decoded_access_token_maps_role(self): + """Decoded JWT payload with role claims maps correctly.""" + from litellm.caching.caching import DualCache + from litellm.proxy._types import LitellmUserRoles + + handler = self._make_jwt_handler() + sso_values = self._make_sso_values() + + await _sync_user_role_from_jwt_role_map( + jwt_handler=handler, + received_response={"sub": "testuser@example.com", "custom_roles": ["my-admin"]}, + user_info=None, + prisma_client=AsyncMock(), + user_api_key_cache=DualCache(), + user_defined_values=sso_values, + ) + + assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value + + @pytest.mark.asyncio + async def test_existing_user_role_updated_in_db_and_cache(self): + """Existing user with stale role gets updated in DB and cache.""" + from litellm.caching.caching import DualCache + from litellm.proxy._types import LitellmUserRoles + + handler = self._make_jwt_handler() + cache = DualCache() + prisma = AsyncMock() + prisma.db.litellm_usertable.update = AsyncMock() + user_id = "testuser@example.com" + + existing_user = LiteLLM_UserTable( + user_id=user_id, + user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, + ) + await cache.async_set_cache(key=user_id, value=existing_user.model_dump(), ttl=60) + + sso_values = self._make_sso_values( + user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, + ) + + await _sync_user_role_from_jwt_role_map( + jwt_handler=handler, + received_response={"sub": user_id, "custom_roles": ["my-admin"]}, + user_info=existing_user, + prisma_client=prisma, + user_api_key_cache=cache, + user_defined_values=sso_values, + ) + + prisma.db.litellm_usertable.update.assert_called_once_with( + where={"user_id": user_id}, + data={"user_role": LitellmUserRoles.PROXY_ADMIN.value}, + ) + assert existing_user.user_role == LitellmUserRoles.PROXY_ADMIN.value + assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value + + @pytest.mark.asyncio + async def test_same_role_no_db_write(self): + """No DB update when the mapped role matches the existing role.""" + from litellm.caching.caching import DualCache + from litellm.proxy._types import LitellmUserRoles + + handler = self._make_jwt_handler() + prisma = AsyncMock() + prisma.db.litellm_usertable.update = AsyncMock() + + existing_user = LiteLLM_UserTable( + user_id="testuser@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN.value, + ) + + sso_values = self._make_sso_values( + user_role=LitellmUserRoles.PROXY_ADMIN.value, + ) + + await _sync_user_role_from_jwt_role_map( + jwt_handler=handler, + received_response={"sub": "testuser@example.com", "custom_roles": ["my-admin"]}, + user_info=existing_user, + prisma_client=prisma, + user_api_key_cache=DualCache(), + user_defined_values=sso_values, + ) + + prisma.db.litellm_usertable.update.assert_not_called() +