diff --git a/docs/my-website/docs/proxy/admin_ui_sso.md b/docs/my-website/docs/proxy/admin_ui_sso.md index bd18dd9c690..ae082848b6b 100644 --- a/docs/my-website/docs/proxy/admin_ui_sso.md +++ b/docs/my-website/docs/proxy/admin_ui_sso.md @@ -320,6 +320,16 @@ Okta requires the `GENERIC_CLIENT_STATE` parameter: GENERIC_CLIENT_STATE="random-string" # Required for Okta ``` +### Okta PKCE + +If your Okta application is configured to require PKCE (Proof Key for Code Exchange), enable it by setting: + +```bash +GENERIC_CLIENT_USE_PKCE="true" +``` + +This is required when your Okta app settings enforce PKCE for enhanced security. LiteLLM will automatically handle PKCE parameter generation and verification during the OAuth flow. + ### Common Configuration Issues #### Missing Protocol in Base URL diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 627b1bb2374..0dbecad9d4a 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -533,6 +533,7 @@ router_settings: | GENERIC_CLIENT_ID | Client ID for generic OAuth providers | GENERIC_CLIENT_SECRET | Client secret for generic OAuth providers | GENERIC_CLIENT_STATE | State parameter for generic client authentication +| GENERIC_CLIENT_USE_PKCE | Enable PKCE (Proof Key for Code Exchange) for generic OAuth providers. Set to "true" when your OAuth provider requires PKCE. **Default is false** | GENERIC_SSO_HEADERS | Comma-separated list of additional headers to add to the request - e.g. Authorization=Bearer ``, Content-Type=application/json, etc. | GENERIC_INCLUDE_CLIENT_ID | Include client ID in requests for OAuth | GENERIC_SCOPE | Scope settings for generic OAuth providers diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 541f7656d20..dec399c7f74 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -9,7 +9,10 @@ Has all /sso/* routes """ import asyncio +import base64 +import hashlib import os +import secrets from copy import deepcopy from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast @@ -381,7 +384,10 @@ async def get_generic_sso_response( try: result = await generic_sso.verify_and_process( request, - params={"include_client_id": generic_include_client_id}, + params=SSOAuthenticationHandler.prepare_token_exchange_parameters( + request=request, + generic_include_client_id=generic_include_client_id, + ), headers=additional_generic_sso_headers_dict, ) @@ -1067,30 +1073,97 @@ class SSOAuthenticationHandler: allow_insecure_http=True, scope=generic_scope, ) - 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 - # For simplification we are using a static state, this is not perfect but some - # SSO providers do not allow stateless verification - redirect_params = ( - SSOAuthenticationHandler._get_generic_sso_redirect_params( - state=state, - generic_authorization_endpoint=generic_authorization_endpoint, - ) - ) - - return await generic_sso.get_login_redirect(**redirect_params) # type: ignore + return await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=generic_sso, + state=state, + generic_authorization_endpoint=generic_authorization_endpoint, + ) 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, + state: Optional[str] = None, + generic_authorization_endpoint: Optional[str] = None, + ) -> Optional[RedirectResponse]: + """ + Get the redirect response for Generic SSO + """ + 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 + # For simplification we are using a static state, this is not perfect but some + # SSO providers do not allow stateless verification + redirect_params, code_verifier = ( + SSOAuthenticationHandler._get_generic_sso_redirect_params( + state=state, + generic_authorization_endpoint=generic_authorization_endpoint, + ) + ) + + # Separate PKCE params from state params (fastapi-sso doesn't accept code_challenge) + pkce_params = {} + state_only_params = {} + for key, value in redirect_params.items(): + if key in ("code_challenge", "code_challenge_method"): + pkce_params[key] = value + else: + state_only_params[key] = value + + # Get the redirect response from fastapi-sso with only state param + redirect_response = await generic_sso.get_login_redirect(**state_only_params) # type: ignore + + # If PKCE is enabled, add PKCE parameters to the redirect URL + if code_verifier and "state" in redirect_params: + + # Store code_verifier in cache (10 min TTL) + cache_key = f"pkce_verifier:{redirect_params['state']}" + user_api_key_cache.set_cache( + key=cache_key, + value=code_verifier, + ttl=600, + ) + + # Add PKCE parameters to the authorization URL + 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 + )) + + # Update the redirect response + redirect_response.headers["location"] = new_url + verbose_proxy_logger.debug( + "PKCE parameters added to authorization URL" + ) + return redirect_response @staticmethod def _get_generic_sso_redirect_params( state: Optional[str] = None, generic_authorization_endpoint: Optional[str] = None, - ) -> dict: + ) -> Tuple[dict, Optional[str]]: """ Get redirect parameters for Generic SSO with proper state priority handling. + Optionally generates PKCE parameters if GENERIC_CLIENT_USE_PKCE is enabled. Priority order: 1. CLI state (if provided) @@ -1102,9 +1175,12 @@ class SSOAuthenticationHandler: generic_authorization_endpoint: Authorization endpoint URL Returns: - dict: Redirect parameters for SSO login + Tuple[dict, Optional[str]]: + - Redirect parameters for SSO login (may include PKCE params) + - code_verifier (if PKCE is enabled, None otherwise) """ redirect_params = {} + code_verifier: Optional[str] = None if state: # CLI state takes priority @@ -1122,7 +1198,18 @@ class SSOAuthenticationHandler: uuid.uuid4().hex ) # set state param for okta - required - return redirect_params + # Handle PKCE (Proof Key for Code Exchange) if enabled + # 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() + redirect_params["code_challenge"] = code_challenge + redirect_params["code_challenge_method"] = "S256" + verbose_proxy_logger.debug( + "PKCE enabled - code_challenge added to authorization request" + ) + + return redirect_params, code_verifier @staticmethod def should_use_sso_handler( @@ -1606,6 +1693,69 @@ 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, + generic_include_client_id: bool, + ) -> dict: + """ + Prepare token exchange parameters for Generic SSO. + + Args: + request: Request object + generic_include_client_id: Generic OAuth Client ID + + Returns: + dict: Token exchange parameters + """ + # 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 + 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 bdc0852b151..0f403ae5d65 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -1968,3 +1968,108 @@ class TestProcessSSOJWTAccessToken: # Even empty team IDs should be set assert result.team_ids == [] + + +class TestPKCEFunctionality: + """Test PKCE (Proof Key for Code Exchange) functionality""" + + def test_generate_pkce_params(self): + """ + Test that generate_pkce_params generates valid PKCE parameters + """ + import base64 + import hashlib + + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Act + code_verifier, code_challenge = SSOAuthenticationHandler.generate_pkce_params() + + # 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('=') + assert code_challenge == expected_challenge + + # Verify both are base64url encoded (no padding) + assert '=' not in code_verifier + assert '=' not in code_challenge + + @pytest.mark.asyncio + async def test_prepare_token_exchange_parameters_with_pkce(self): + """ + Test prepare_token_exchange_parameters retrieves PKCE code_verifier from cache + """ + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Mock request with state parameter + mock_request = MagicMock(spec=Request) + test_state = "test_oauth_state_123" + mock_request.query_params = {"state": test_state} + + # Mock cache + mock_cache = MagicMock() + test_code_verifier = "test_code_verifier_abc123xyz" + mock_cache.get_cache.return_value = test_code_verifier + + 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 + ) + + # 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}") + + @pytest.mark.asyncio + async def test_get_generic_sso_redirect_response_with_pkce(self): + """ + Test get_generic_sso_redirect_response with PKCE enabled stores verifier and adds challenge to URL + """ + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Mock SSO provider + mock_sso = MagicMock() + mock_redirect_response = MagicMock() + 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) + mock_sso.__exit__ = MagicMock(return_value=False) + + test_state = "test456" + mock_cache = MagicMock() + + with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}): + with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + # Act + result = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=mock_sso, + state=test_state, + generic_authorization_endpoint="https://auth.example.com/authorize" + ) + + # Assert + # Verify cache was called to store code_verifier + mock_cache.set_cache.assert_called_once() + cache_call = mock_cache.set_cache.call_args + assert cache_call.kwargs["key"] == f"pkce_verifier:{test_state}" + assert cache_call.kwargs["ttl"] == 600 + assert len(cache_call.kwargs["value"]) == 43 + + # Verify PKCE parameters were added to the redirect URL + assert result is not None + updated_location = str(result.headers["location"]) + assert "code_challenge=" in updated_location + assert "code_challenge_method=S256" in updated_location + assert f"state={test_state}" in updated_location +