fix: address greptile review feedback

- Fix access_token missing in PKCE path: read from combined_response directly
  instead of generic_sso.access_token (which is only set by verify_and_process)
- Fix PKCE error hint firing when PKCE is already enabled: only show
  'set GENERIC_CLIENT_USE_PKCE=true' advice when code_verifier was absent
- Fix unguarded KeyError on access_token: check for error field in HTTP 200
  responses before accessing token_response['access_token']
- Fix silent empty userinfo: raise ProxyException when both userinfo endpoint
  and id_token fallback produce no user data
- Fix backward-incompatible Redis wiring: only attach Redis to user_api_key_cache
  when GENERIC_CLIENT_USE_PKCE=true, preserving existing in-memory behaviour
This commit is contained in:
Ishaan Jaffer 2026-03-05 11:56:19 -08:00
parent da82ba4dc3
commit c94886d2b3
2 changed files with 65 additions and 21 deletions

View file

@ -776,6 +776,9 @@ async def get_generic_sso_response(
key, value = header.split("=")
additional_generic_sso_headers_dict[key] = value
# Initialized here so it's visible in the except block for error-hint logic
code_verifier: Optional[str] = None
try:
token_exchange_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=request,
@ -808,14 +811,18 @@ async def get_generic_sso_response(
additional_headers=additional_generic_sso_headers_dict,
)
result = response_convertor(combined_response, generic_sso)
# In the PKCE path verify_and_process is skipped, so generic_sso.access_token
# is never set. Read the token directly from the exchange response instead so
# process_sso_jwt_access_token can extract JWT-embedded roles/teams.
access_token_str: Optional[str] = combined_response.get("access_token")
else:
result = await generic_sso.verify_and_process(
request,
params=token_exchange_params,
headers=additional_generic_sso_headers_dict,
)
access_token_str = generic_sso.access_token
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
)
@ -823,8 +830,12 @@ async def get_generic_sso_response(
except Exception as e:
error_message = str(e)
# Detect PKCE misconfiguration and surface a helpful error
if "PKCE" in error_message or "code verifier" in error_message.lower():
# Only surface "enable PKCE" advice when PKCE was NOT already in use.
# If code_verifier is set, the token exchange itself failed — that's a
# provider-side error, not a configuration problem.
if code_verifier is None and (
"PKCE" in error_message or "code verifier" in error_message.lower()
):
is_okta = (
generic_authorization_endpoint
and "okta" in generic_authorization_endpoint.lower()
@ -2617,6 +2628,22 @@ class SSOAuthenticationHandler:
)
raise
# Some providers return HTTP 200 with an error body (e.g. expired code, replay attack).
if "access_token" not in token_response:
error = token_response.get("error", "unknown_error")
error_desc = token_response.get("error_description", "")
verbose_proxy_logger.error(
"Token response missing access_token. error=%s description=%s",
error,
error_desc,
)
raise ProxyException(
message=f"Token exchange error: {error} - {error_desc}",
type=ProxyErrorTypes.auth_error,
param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED,
)
verbose_proxy_logger.debug(
"PKCE token exchange successful. access_token=%s id_token=%s",
bool(token_response.get("access_token")),
@ -2678,6 +2705,17 @@ class SSOAuthenticationHandler:
code=status.HTTP_401_UNAUTHORIZED,
)
if not userinfo:
raise ProxyException(
message=(
"SSO user info unavailable: userinfo endpoint failed and no id_token "
"was present in the token response."
),
type=ProxyErrorTypes.auth_error,
param="userinfo",
code=status.HTTP_401_UNAUTHORIZED,
)
return userinfo

View file

@ -2473,15 +2473,16 @@ 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
## CONFIGURE USER API KEY CACHE TO USE REDIS FOR PKCE ##
# Only wire Redis when PKCE is explicitly enabled to avoid changing
# cache behaviour for deployments that don't use PKCE.
global user_api_key_cache
if user_api_key_cache.redis_cache is None:
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if use_pkce and 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."
"Configured user_api_key_cache to use Redis "
"(PKCE enabled — verifiers shared across instances)."
)
def switch_on_llm_response_caching(self):
@ -3032,17 +3033,18 @@ class ProxyConfig:
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:
### CONFIGURE USER API KEY CACHE TO USE REDIS FOR PKCE (if enabled) ###
# Only wire Redis to user_api_key_cache when PKCE is explicitly enabled.
# This avoids silently routing API key lookups through Redis for deployments
# that use Redis only for LLM response caching and not for session state.
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if use_pkce and 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(
@ -3053,18 +3055,22 @@ class ProxyConfig:
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."
"Configured user_api_key_cache to use Redis at %s:%s "
"(PKCE enabled — verifiers shared across instances).",
redis_host,
redis_port,
)
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."
"Failed to configure Redis for user_api_key_cache: %s. "
"Falling back to in-memory cache. Multi-instance PKCE will not work.",
e,
)
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."
verbose_proxy_logger.warning(
"GENERIC_CLIENT_USE_PKCE=true but REDIS_HOST is not set. "
"PKCE verifiers will not be shared across instances. "
"Configure Redis or enable sticky sessions for multi-instance deployments."
)
### STORE MODEL IN DB ### feature flag for `/model/new`
store_model_in_db = general_settings.get("store_model_in_db", False)