mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(ui_sso.py): support dot notation on ui sso (#16135)
This commit is contained in:
parent
20b95e9a80
commit
3f40613c56
2 changed files with 148 additions and 49 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue