From 3f40613c569b21d2daca14a956eee89f5a6778f6 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Sun, 2 Nov 2025 09:35:52 -0800 Subject: [PATCH] fix(ui_sso.py): support dot notation on ui sso (#16135) --- litellm/proxy/management_endpoints/ui_sso.py | 83 +++++++------ .../proxy/management_endpoints/test_ui_sso.py | 114 ++++++++++++++++-- 2 files changed, 148 insertions(+), 49 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index dec399c7f74..a4c1581a71b 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -24,6 +24,7 @@ from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.caching import DualCache from litellm.constants import MAX_SPENDLOG_ROWS_TO_QUERY +from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, get_async_httpx_client, @@ -273,12 +274,14 @@ def generic_response_convertor( all_teams.extend(team_ids) return CustomOpenID( - id=response.get(generic_user_id_attribute_name), - display_name=response.get(generic_user_display_name_attribute_name), - email=response.get(generic_user_email_attribute_name), - first_name=response.get(generic_user_first_name_attribute_name), - last_name=response.get(generic_user_last_name_attribute_name), - provider=response.get(generic_provider_attribute_name), + id=get_nested_value(response, generic_user_id_attribute_name), + display_name=get_nested_value( + response, generic_user_display_name_attribute_name + ), + email=get_nested_value(response, generic_user_email_attribute_name), + first_name=get_nested_value(response, generic_user_first_name_attribute_name), + last_name=get_nested_value(response, generic_user_last_name_attribute_name), + provider=get_nested_value(response, generic_provider_attribute_name), team_ids=all_teams, user_role=None, ) @@ -1081,7 +1084,7 @@ class SSOAuthenticationHandler: raise ValueError( "Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso" ) - + @staticmethod async def get_generic_sso_redirect_response( generic_sso: Any, @@ -1094,6 +1097,7 @@ class SSOAuthenticationHandler: from urllib.parse import parse_qs, urlencode, urlparse, urlunparse from litellm.proxy.proxy_server import user_api_key_cache + with generic_sso: # TODO: state should be a random string and added to the user session with cookie # or a cryptographicly signed state that we can verify stateless @@ -1133,22 +1137,24 @@ class SSOAuthenticationHandler: if pkce_params: parsed_url = urlparse(str(redirect_response.headers["location"])) query_params = parse_qs(parsed_url.query) - + # Add PKCE parameters for key, value in pkce_params.items(): query_params[key] = [value] - + # Reconstruct the URL with PKCE parameters new_query = urlencode(query_params, doseq=True) - new_url = urlunparse(( - parsed_url.scheme, - parsed_url.netloc, - parsed_url.path, - parsed_url.params, - new_query, - parsed_url.fragment - )) - + new_url = urlunparse( + ( + parsed_url.scheme, + parsed_url.netloc, + parsed_url.path, + parsed_url.params, + new_query, + parsed_url.fragment, + ) + ) + # Update the redirect response redirect_response.headers["location"] = new_url verbose_proxy_logger.debug( @@ -1175,7 +1181,7 @@ class SSOAuthenticationHandler: generic_authorization_endpoint: Authorization endpoint URL Returns: - Tuple[dict, Optional[str]]: + Tuple[dict, Optional[str]]: - Redirect parameters for SSO login (may include PKCE params) - code_verifier (if PKCE is enabled, None otherwise) """ @@ -1202,7 +1208,9 @@ class SSOAuthenticationHandler: # Set GENERIC_CLIENT_USE_PKCE=true to enable PKCE for enhanced OAuth security use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true" if use_pkce: - code_verifier, code_challenge = SSOAuthenticationHandler.generate_pkce_params() + code_verifier, code_challenge = ( + SSOAuthenticationHandler.generate_pkce_params() + ) redirect_params["code_challenge"] = code_challenge redirect_params["code_challenge_method"] = "S256" verbose_proxy_logger.debug( @@ -1693,11 +1701,10 @@ class SSOAuthenticationHandler: redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303) redirect_response.set_cookie(key="token", value=jwt_token) return redirect_response - @staticmethod def prepare_token_exchange_parameters( - request: Request, + request: Request, generic_include_client_id: bool, ) -> dict: """ @@ -1712,50 +1719,54 @@ class SSOAuthenticationHandler: """ # Prepare token exchange parameters token_params = {"include_client_id": generic_include_client_id} - + # Retrieve PKCE code_verifier if PKCE was used in authorization query_params = dict(request.query_params) state = query_params.get("state") if state: from litellm.proxy.proxy_server import user_api_key_cache - + cache_key = f"pkce_verifier:{state}" code_verifier = user_api_key_cache.get_cache(key=cache_key) - + if code_verifier: # Add code_verifier to token exchange parameters token_params["code_verifier"] = code_verifier verbose_proxy_logger.debug( "PKCE code_verifier retrieved and will be included in token exchange" ) - + # Clean up the cache entry (single-use verifier) user_api_key_cache.delete_cache(key=cache_key) return token_params - @staticmethod def generate_pkce_params() -> Tuple[str, str]: """ Generate PKCE (Proof Key for Code Exchange) parameters for OAuth 2.0. - + Returns: Tuple[str, str]: (code_verifier, code_challenge) - code_verifier: Random 43-128 character string (we use 43 for efficiency) - code_challenge: Base64-URL-encoded SHA256 hash of the code_verifier - + Reference: https://datatracker.ietf.org/doc/html/rfc7636 """ # Generate a cryptographically random code_verifier (43 characters) # Using 32 random bytes which becomes 43 characters when base64-url-encoded - code_verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).decode('utf-8').rstrip('=') - - # Generate code_challenge using S256 method (SHA256) - code_challenge_bytes = hashlib.sha256(code_verifier.encode('utf-8')).digest() - code_challenge = base64.urlsafe_b64encode(code_challenge_bytes).decode('utf-8').rstrip('=') - - return code_verifier, code_challenge + code_verifier = ( + base64.urlsafe_b64encode(secrets.token_bytes(32)) + .decode("utf-8") + .rstrip("=") + ) + # Generate code_challenge using S256 method (SHA256) + code_challenge_bytes = hashlib.sha256(code_verifier.encode("utf-8")).digest() + code_challenge = ( + base64.urlsafe_b64encode(code_challenge_bytes).decode("utf-8").rstrip("=") + ) + + return code_verifier, code_challenge class MicrosoftSSOHandler: 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 42cec94322c..cc23318505d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -1965,6 +1965,83 @@ class TestProcessSSOJWTAccessToken: assert result.team_ids == [] +class TestGenericResponseConvertorNestedAttributes: + """Test generic_response_convertor with nested attribute paths""" + + def test_generic_response_convertor_with_nested_attributes(self): + """ + Test that generic_response_convertor handles nested attributes with dotted notation + like "attributes.userId" + """ + from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor + + # Mock JWT handler + mock_jwt_handler = MagicMock(spec=JWTHandler) + mock_jwt_handler.get_team_ids_from_jwt.return_value = [] + + # Payload with nested attributes structure + nested_payload = { + "sub": "user-sub-123", + "service": "test-service", + "auth_time": 1234567890, + "attributes": { + "given_name": "John", + "oauthClientId": "client-123", + "family_name": "Doe", + "userId": "nested-user-456", + "email": "john.doe@example.com", + }, + "id": "top-level-id-789", + "client_id": "client-abc", + } + + # Test with nested user ID attribute + with patch.dict( + os.environ, + { + "GENERIC_USER_ID_ATTRIBUTE": "attributes.userId", + "GENERIC_USER_EMAIL_ATTRIBUTE": "attributes.email", + "GENERIC_USER_FIRST_NAME_ATTRIBUTE": "attributes.given_name", + "GENERIC_USER_LAST_NAME_ATTRIBUTE": "attributes.family_name", + "GENERIC_USER_DISPLAY_NAME_ATTRIBUTE": "sub", + }, + ): + # Act + result = generic_response_convertor( + response=nested_payload, + jwt_handler=mock_jwt_handler, + sso_jwt_handler=None, + ) + + # Assert + assert isinstance(result, CustomOpenID) + + # Note: The current implementation uses response.get() which doesn't support + # dotted notation for nested attributes. This test documents the current behavior. + # If nested attribute support is needed, the implementation would need to be updated + # to handle dotted paths like "attributes.userId" + + # Current behavior: returns None for nested paths + print(f"User ID result: {result.id}") + print(f"Email result: {result.email}") + print(f"First name result: {result.first_name}") + print(f"Last name result: {result.last_name}") + print(f"Display name result: {result.display_name}") + + # Expected behavior with current implementation (no nested path support): + assert result.id == "nested-user-456" + assert ( + result.email == "john.doe@example.com" + ) # Can't access "attributes.email" with .get() + assert ( + result.first_name == "John" + ) # Can't access "attributes.given_name" with .get() + assert ( + result.last_name == "Doe" + ) # Can't access "attributes.family_name" with .get() + assert result.display_name == "user-sub-123" # Top-level attribute works + + class TestPKCEFunctionality: """Test PKCE (Proof Key for Code Exchange) functionality""" @@ -1983,15 +2060,21 @@ class TestPKCEFunctionality: # Assert assert len(code_verifier) == 43 assert isinstance(code_verifier, str) - + # Verify code_challenge is correctly generated from code_verifier - expected_challenge_bytes = hashlib.sha256(code_verifier.encode('utf-8')).digest() - expected_challenge = base64.urlsafe_b64encode(expected_challenge_bytes).decode('utf-8').rstrip('=') + expected_challenge_bytes = hashlib.sha256( + code_verifier.encode("utf-8") + ).digest() + expected_challenge = ( + base64.urlsafe_b64encode(expected_challenge_bytes) + .decode("utf-8") + .rstrip("=") + ) assert code_challenge == expected_challenge - + # Verify both are base64url encoded (no padding) - assert '=' not in code_verifier - assert '=' not in code_challenge + assert "=" not in code_verifier + assert "=" not in code_challenge @pytest.mark.asyncio async def test_prepare_token_exchange_parameters_with_pkce(self): @@ -2013,17 +2096,20 @@ class TestPKCEFunctionality: with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): # Act token_params = SSOAuthenticationHandler.prepare_token_exchange_parameters( - request=mock_request, - generic_include_client_id=False + request=mock_request, generic_include_client_id=False ) # Assert assert token_params["include_client_id"] is False assert token_params["code_verifier"] == test_code_verifier - + # Verify cache was accessed and deleted - mock_cache.get_cache.assert_called_once_with(key=f"pkce_verifier:{test_state}") - mock_cache.delete_cache.assert_called_once_with(key=f"pkce_verifier:{test_state}") + mock_cache.get_cache.assert_called_once_with( + key=f"pkce_verifier:{test_state}" + ) + mock_cache.delete_cache.assert_called_once_with( + key=f"pkce_verifier:{test_state}" + ) @pytest.mark.asyncio async def test_get_generic_sso_redirect_response_with_pkce(self): @@ -2035,7 +2121,9 @@ class TestPKCEFunctionality: # Mock SSO provider mock_sso = MagicMock() mock_redirect_response = MagicMock() - original_location = "https://auth.example.com/authorize?state=test456&client_id=abc" + original_location = ( + "https://auth.example.com/authorize?state=test456&client_id=abc" + ) mock_redirect_response.headers = {"location": original_location} mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response) mock_sso.__enter__ = MagicMock(return_value=mock_sso) @@ -2050,7 +2138,7 @@ class TestPKCEFunctionality: result = await SSOAuthenticationHandler.get_generic_sso_redirect_response( generic_sso=mock_sso, state=test_state, - generic_authorization_endpoint="https://auth.example.com/authorize" + generic_authorization_endpoint="https://auth.example.com/authorize", ) # Assert