mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #15608 from BerriAI/litellm_sso_add_pkce
[Feat] UI SSO - Add PKCE for OKTA SSO
This commit is contained in:
commit
5009a43b0e
4 changed files with 283 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 `<token>`, Content-Type=application/json, etc.
|
||||
| GENERIC_INCLUDE_CLIENT_ID | Include client ID in requests for OAuth
|
||||
| GENERIC_SCOPE | Scope settings for generic OAuth providers
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue