fix(sso): add direct PKCE token exchange and Redis cache wiring for multi-instance SSO

When PKCE is enabled, bypass fastapi-sso and perform direct token exchange so
code_verifier is correctly included. Store PKCE verifiers as dict in cache
for proper JSON serialization in Redis. Wire user_api_key_cache to Redis when
available so PKCE verifiers are shared across ECS tasks/pods.

Also adds clearer error messages when PKCE is required but not configured.
This commit is contained in:
Ishaan Jaffer 2026-03-05 11:31:48 -08:00
parent bf9c96b912
commit 08e249b6c7
2 changed files with 306 additions and 35 deletions

View file

@ -84,6 +84,7 @@ from litellm.proxy.utils import (
get_server_root_path,
)
from litellm.secret_managers.main import get_secret_bool, str_to_bool
from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403, F401
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
MicrosoftGraphAPIUserGroupDirectoryObject,
@ -92,7 +93,6 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
RoleMappings,
TeamMappings,
)
from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403, F401
from litellm.types.proxy.ui_sso import ParsedOpenIDResult
if TYPE_CHECKING:
@ -775,21 +775,228 @@ async def get_generic_sso_response(
additional_generic_sso_headers_dict[key] = value
try:
result = await generic_sso.verify_and_process(
request,
params=await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=request,
generic_include_client_id=generic_include_client_id,
),
headers=additional_generic_sso_headers_dict,
token_exchange_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=request,
generic_include_client_id=generic_include_client_id,
)
# Extract code_verifier before calling fastapi-sso
code_verifier = token_exchange_params.pop("code_verifier", None)
# Get authorization code from query params
authorization_code = request.query_params.get("code")
if not authorization_code:
raise ProxyException(
message="Missing authorization code in callback",
type=ProxyErrorTypes.auth_error,
param="code",
code=status.HTTP_400_BAD_REQUEST,
)
if code_verifier:
verbose_proxy_logger.info(
f"✔ PKCE: Performing direct token exchange with code_verifier. "
f"Length: {len(code_verifier)}"
)
# Direct token exchange to include code_verifier (fastapi-sso doesn't support PKCE)
import httpx
token_data = {
"grant_type": "authorization_code",
"code": authorization_code,
"redirect_uri": redirect_url,
"code_verifier": code_verifier,
}
# Add client credentials
if not generic_include_client_id:
# Use HTTP Basic Auth for client credentials
auth = httpx.BasicAuth(generic_client_id, generic_client_secret)
else:
# Include client credentials in body
token_data["client_id"] = generic_client_id
token_data["client_secret"] = generic_client_secret
auth = None
async with httpx.AsyncClient() as client:
post_kwargs: dict = {
"data": token_data,
"headers": {
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
**additional_generic_sso_headers_dict,
},
"timeout": 30.0,
}
if auth is not None:
post_kwargs["auth"] = auth
response = await client.post(generic_token_endpoint, **post_kwargs)
response_text = response.text
if response.status_code != 200:
verbose_proxy_logger.error(
f"Token exchange failed: Status={response.status_code}"
)
verbose_proxy_logger.error(
f"✘ Token exchange FAILED!\n"
f" Status: {response.status_code}\n"
f" Full Response:\n{response_text}"
)
raise ProxyException(
message=f"Token exchange failed: {response.status_code} - {response_text[:500]}",
type=ProxyErrorTypes.auth_error,
param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED,
)
# Parse token response
try:
token_response = response.json()
verbose_proxy_logger.info(
f"✔ Token exchange successful! "
f"access_token={bool(token_response.get('access_token'))}, "
f"id_token={bool(token_response.get('id_token'))}"
)
except Exception as json_err:
verbose_proxy_logger.error(
f"✘ Failed to parse token response as JSON!\n"
f" Error: {json_err}\n"
f" Response text: {response_text}"
)
raise
# Get user info - try userinfo endpoint first, fallback to id_token
access_token = token_response["access_token"]
id_token = token_response.get("id_token")
userinfo: dict = {}
try:
async with httpx.AsyncClient() as client:
userinfo_response = await client.get(
generic_userinfo_endpoint,
headers={
"Authorization": f"Bearer {access_token}",
**additional_generic_sso_headers_dict,
},
timeout=30.0,
)
if userinfo_response.status_code == 200:
userinfo = userinfo_response.json()
verbose_proxy_logger.debug(
"User info retrieved from userinfo endpoint"
)
else:
verbose_proxy_logger.warning(
f"⚠ Userinfo endpoint failed ({userinfo_response.status_code}), "
f"will extract from id_token instead"
)
except Exception as userinfo_err:
verbose_proxy_logger.warning(
f"⚠ Userinfo endpoint error: {userinfo_err}, will extract from id_token instead"
)
# If userinfo is empty, extract from id_token
if not userinfo and id_token:
import jwt
try:
# Decode id_token without verification (we already trust it from token exchange)
id_token_payload = jwt.decode(
id_token,
options={"verify_signature": False},
)
userinfo = id_token_payload
verbose_proxy_logger.debug("User info extracted from id_token")
except Exception as decode_err:
verbose_proxy_logger.error(
f"✘ Failed to decode id_token: {decode_err}"
)
raise ProxyException(
message="Failed to get user info from both userinfo endpoint and id_token",
type=ProxyErrorTypes.auth_error,
param="userinfo",
code=status.HTTP_401_UNAUTHORIZED,
)
# Build combined response for convertor
combined_response = {**token_response, **userinfo}
result = response_convertor(combined_response, generic_sso)
else:
verbose_proxy_logger.info("No PKCE code_verifier - using standard verify_and_process")
# No PKCE, use standard flow
result = await generic_sso.verify_and_process(
request,
params=token_exchange_params,
headers=additional_generic_sso_headers_dict,
)
access_token_str: Optional[str] = generic_sso.access_token
process_sso_jwt_access_token(
access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
)
except Exception as e:
error_message = str(e)
# Log detailed error information for debugging
verbose_proxy_logger.error(
f"Error verifying and processing generic SSO: {error_message}. "
f"Error type: {type(e).__name__}. "
f"Passed in headers: {additional_generic_sso_headers_dict}"
)
# Check if this is a JSON decode error (provider returned non-JSON response)
if "JSONDecodeError" in str(type(e).__name__) or "Expecting value" in error_message:
verbose_proxy_logger.error(
f"SSO provider returned invalid JSON response during token exchange. "
f"This typically means: (1) Token exchange request was rejected by the provider, "
f"(2) code_verifier format is incorrect, or (3) provider configuration issue."
)
# Check if this is a PKCE-related error
if "PKCE" in error_message or "code verifier" in error_message.lower():
# Detect if this is Okta by checking the endpoints
is_okta = (
generic_authorization_endpoint and "okta" in generic_authorization_endpoint.lower()
) or (
generic_token_endpoint and "okta" in generic_token_endpoint.lower()
)
provider_name = "Okta" if is_okta else "Your OAuth provider"
verbose_proxy_logger.error(
f"PKCE error detected: {error_message}. {provider_name} requires PKCE but it's not enabled. "
f"Set GENERIC_CLIENT_USE_PKCE=true in your environment variables. "
f"See https://docs.litellm.ai/docs/proxy/admin_ui_sso for more details."
)
detailed_message = (
f"SSO authentication failed: {provider_name} requires PKCE (Proof Key for Code Exchange) "
f"but it's not enabled in your LiteLLM configuration.\n\n"
f"SOLUTION: Add this environment variable and restart your proxy:\n"
f" GENERIC_CLIENT_USE_PKCE=true\n\n"
)
if is_okta:
detailed_message += (
f"For AWS ECS: Add the environment variable to your task definition.\n"
f"For Docker: Add -e GENERIC_CLIENT_USE_PKCE=true to your docker run command.\n"
f"For .env file: Add GENERIC_CLIENT_USE_PKCE=true to your .env file.\n\n"
)
detailed_message += f"Original error: {error_message}"
raise ProxyException(
message=detailed_message,
type=ProxyErrorTypes.auth_error,
param="GENERIC_CLIENT_USE_PKCE",
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
verbose_proxy_logger.exception(
f"Error verifying and processing generic SSO: {e}. Passed in headers: {additional_generic_sso_headers_dict}"
)
@ -1773,21 +1980,25 @@ class SSOAuthenticationHandler:
# 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). Use Redis when available
# so callbacks landing on another pod can retrieve it (multi-pod SSO).
# Store code_verifier in cache (10 min TTL). Wrap in dict for proper
# JSON serialization in Redis. Use Redis when available so callbacks
# landing on another pod can retrieve it (multi-pod SSO).
cache_key = f"pkce_verifier:{redirect_params['state']}"
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(
key=cache_key,
value=code_verifier,
value={"code_verifier": code_verifier},
ttl=600,
)
else:
await user_api_key_cache.async_set_cache(
key=cache_key,
value=code_verifier,
value={"code_verifier": code_verifier},
ttl=600,
)
verbose_proxy_logger.debug(
"PKCE code_verifier stored in cache (TTL: 600s)"
)
# Add PKCE parameters to the authorization URL
if pkce_params:
@ -1813,9 +2024,6 @@ class SSOAuthenticationHandler:
# Update the redirect response
redirect_response.headers["location"] = new_url
verbose_proxy_logger.debug(
"PKCE parameters added to authorization URL"
)
return redirect_response
@staticmethod
@ -1858,7 +2066,9 @@ class SSOAuthenticationHandler:
# 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"
pkce_env_value = os.getenv("GENERIC_CLIENT_USE_PKCE", "false")
use_pkce = pkce_env_value.lower() == "true"
if use_pkce:
(
code_verifier,
@ -1866,9 +2076,7 @@ class SSOAuthenticationHandler:
) = 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"
)
verbose_proxy_logger.info("PKCE enabled for authorization request")
return redirect_params, code_verifier
@ -2405,31 +2613,48 @@ class SSOAuthenticationHandler:
# Use same cache as store: Redis when available (multi-pod), else in-memory.
query_params = dict(request.query_params)
state = query_params.get("state")
if state:
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
cache_key = f"pkce_verifier:{state}"
if redis_usage_cache is not None:
code_verifier = await redis_usage_cache.async_get_cache(key=cache_key)
cached_data = await redis_usage_cache.async_get_cache(key=cache_key)
else:
code_verifier = await user_api_key_cache.async_get_cache(key=cache_key)
cached_data = await user_api_key_cache.async_get_cache(key=cache_key)
code_verifier = None
if cached_data:
# Extract code_verifier from dict (stored as dict for JSON serialization)
if isinstance(cached_data, dict) and "code_verifier" in cached_data:
code_verifier = cached_data["code_verifier"]
verbose_proxy_logger.debug("PKCE code_verifier retrieved from cache")
else:
# Handle legacy format (plain string) for backward compatibility
code_verifier = cached_data if isinstance(cached_data, str) else str(cached_data)
verbose_proxy_logger.warning(
"Retrieved code_verifier in legacy format (plain string). "
"Using it but future storage will use dict format."
)
if code_verifier:
# Add code_verifier to token exchange parameters (Redis returns decoded string)
token_params["code_verifier"] = (
code_verifier
if isinstance(code_verifier, str)
else str(code_verifier)
)
verbose_proxy_logger.debug(
"PKCE code_verifier retrieved and will be included in token exchange"
)
# Add code_verifier to token exchange parameters
token_params["code_verifier"] = code_verifier
# Clean up the cache entry (single-use verifier)
if redis_usage_cache is not None:
await redis_usage_cache.async_delete_cache(key=cache_key)
else:
await user_api_key_cache.async_delete_cache(key=cache_key)
else:
verbose_proxy_logger.error(
f"✙ CRITICAL: No PKCE code_verifier found in cache for state '{state}'. "
f"This indicates: (1) authorization request and callback handled by different instances without shared cache, "
f"(2) cache entry expired (TTL: 600s), or (3) Redis serialization error. "
f"SOLUTION: Ensure Redis is configured correctly. "
f"Cache type: {type(user_api_key_cache).__name__}. "
f"Cached data: {cached_data}"
)
return token_params
@staticmethod

View file

@ -372,6 +372,9 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import (
from litellm.proxy.management_endpoints.internal_user_endpoints import (
user_update,
)
from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import (
router as jwt_key_mapping_router,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
delete_verification_tokens,
duration_in_seconds,
@ -380,9 +383,6 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
from litellm.proxy.management_endpoints.key_management_endpoints import (
router as key_management_router,
)
from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import (
router as jwt_key_mapping_router,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
router as mcp_management_router,
)
@ -2473,6 +2473,17 @@ class ProxyConfig:
## INIT PROXY REDIS USAGE CLIENT ##
redis_usage_cache = litellm.cache.cache
## CONFIGURE USER API KEY CACHE TO USE REDIS ##
# This is critical for multi-task deployments (e.g., multiple ECS tasks)
# to share cached data like PKCE code_verifiers across all tasks
global user_api_key_cache
if user_api_key_cache.redis_cache is None:
user_api_key_cache.redis_cache = redis_usage_cache
verbose_proxy_logger.info(
"\u2713 Configured user_api_key_cache to use Redis. "
"PKCE and other cached data will now be shared across all tasks/instances."
)
def switch_on_llm_response_caching(self):
"""
Enable caching on the router by setting cache_responses=True.
@ -3018,8 +3029,43 @@ class ProxyConfig:
if user_api_key_cache_ttl is not None:
user_api_key_cache.update_cache_ttl(
default_in_memory_ttl=float(user_api_key_cache_ttl),
default_redis_ttl=None, # user_api_key_cache is an in-memory cache
default_redis_ttl=None, # will be set below if Redis is available
)
### CONFIGURE USER API KEY CACHE TO USE REDIS (if available) ###
# This is critical for multi-task/multi-instance deployments (e.g., multiple ECS tasks)
# to share cached data like PKCE code_verifiers, API keys, etc. across all instances
if user_api_key_cache.redis_cache is None:
redis_host = get_secret("REDIS_HOST", None)
redis_port = get_secret("REDIS_PORT", None)
redis_password = get_secret("REDIS_PASSWORD", None)
if redis_host is not None:
try:
# Initialize Redis for user_api_key_cache
from litellm.caching.caching import RedisCache
user_redis_cache = RedisCache(
host=redis_host,
port=redis_port,
password=redis_password,
)
user_api_key_cache.redis_cache = user_redis_cache
verbose_proxy_logger.info(
f"\u2713 Configured user_api_key_cache to use Redis at {redis_host}:{redis_port}. "
f"PKCE verifiers and other session data will now be shared across all tasks/instances."
)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to configure Redis for user_api_key_cache: {e}. "
f"Falling back to in-memory cache only. Multi-task PKCE will not work."
)
else:
verbose_proxy_logger.debug(
"REDIS_HOST not configured. user_api_key_cache will use in-memory cache only. "
"For multi-task deployments with PKCE, configure Redis or enable sticky sessions."
)
### STORE MODEL IN DB ### feature flag for `/model/new`
store_model_in_db = general_settings.get("store_model_in_db", False)
if store_model_in_db is None: