fix(ui_sso.py): support dot notation on ui sso (#16135)

This commit is contained in:
Krish Dholakia 2025-11-02 09:35:52 -08:00 • committed by GitHub
parent 20b95e9a80
commit 3f40613c56
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 148 additions and 49 deletions

View file

@ -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:

View file

@ -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