fix: address fourth round of greptile review feedback

- Strip OAuth token credentials from response_convertor input to prevent
  access_token/id_token appearing in restricted-group error messages
- Reuse single httpx.AsyncClient for both token exchange and userinfo requests
  to avoid a second TCP/TLS handshake per SSO callback
- Revert Redis wiring to user_api_key_cache: PKCE code already uses
  redis_usage_cache directly; wiring would route all API-key lookups through
  Redis unnecessarily. Add startup warning instead when PKCE+Redis mismatch.
- Move _OAUTH_TOKEN_FIELDS to module level
This commit is contained in:
Ishaan Jaffer 2026-03-05 12:24:44 -08:00
parent c6f2446fa5
commit 1a16dcf598
2 changed files with 98 additions and 91 deletions

View file

@ -104,6 +104,12 @@ else:
router = APIRouter()
# OAuth token credential fields that must not appear in SSO debug responses
# (received_response is included in restricted-group error messages).
_OAUTH_TOKEN_FIELDS = frozenset(
{"access_token", "id_token", "refresh_token", "token_type", "expires_in", "scope"}
)
def normalize_email(email: Optional[str]) -> Optional[str]:
"""
@ -810,7 +816,13 @@ async def get_generic_sso_response(
redirect_url=redirect_url,
additional_headers=additional_generic_sso_headers_dict,
)
result = response_convertor(combined_response, generic_sso)
# Strip OAuth token credentials before passing to response_convertor.
# combined_response is stored as received_response and may appear in
# restricted-group error messages — do not expose tokens to callers.
result = response_convertor(
{k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS},
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.
@ -2607,60 +2619,63 @@ class SSOAuthenticationHandler:
token_data["client_id"] = client_id
token_data["client_secret"] = client_secret
async with httpx.AsyncClient() as client:
response = await client.post(token_endpoint, **post_kwargs)
async with httpx.AsyncClient() as http_client:
response = await http_client.post(token_endpoint, **post_kwargs)
if response.status_code != 200:
verbose_proxy_logger.error(
"PKCE token exchange failed. status=%s body=%s",
response.status_code,
response.text[:500],
)
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,
if response.status_code != 200:
verbose_proxy_logger.error(
"PKCE token exchange failed. status=%s body=%s",
response.status_code,
response.text[:500],
)
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,
)
try:
token_response: dict = response.json()
except Exception as json_err:
verbose_proxy_logger.error(
"Failed to parse token response as JSON: %s. Body: %s",
json_err,
response.text[:500],
)
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")),
bool(token_response.get("id_token")),
)
try:
token_response: dict = response.json()
except Exception as json_err:
verbose_proxy_logger.error(
"Failed to parse token response as JSON: %s. Body: %s",
json_err,
response.text[:500],
)
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,
# Reuse the same client for the userinfo request to avoid a second TCP/TLS handshake.
userinfo = await SSOAuthenticationHandler._get_pkce_userinfo(
access_token=token_response["access_token"],
id_token=token_response.get("id_token"),
userinfo_endpoint=userinfo_endpoint,
additional_headers=additional_headers,
http_client=http_client,
)
verbose_proxy_logger.debug(
"PKCE token exchange successful. access_token=%s id_token=%s",
bool(token_response.get("access_token")),
bool(token_response.get("id_token")),
)
userinfo = await SSOAuthenticationHandler._get_pkce_userinfo(
access_token=token_response["access_token"],
id_token=token_response.get("id_token"),
userinfo_endpoint=userinfo_endpoint,
additional_headers=additional_headers,
)
return {**token_response, **userinfo}
@staticmethod
@ -2669,16 +2684,22 @@ class SSOAuthenticationHandler:
id_token: Optional[str],
userinfo_endpoint: str,
additional_headers: Dict[str, str],
http_client: Optional[httpx.AsyncClient] = None,
) -> dict:
"""
Fetches user info from the userinfo endpoint.
Falls back to decoding the id_token if the endpoint is unavailable.
An existing ``http_client`` may be passed to reuse the connection pool
(e.g. from ``_pkce_token_exchange``).
"""
userinfo: dict = {}
try:
async with httpx.AsyncClient() as client:
resp = await client.get(
_own_client = http_client is None
_client = httpx.AsyncClient() if _own_client else http_client
try:
resp = await _client.get(
userinfo_endpoint,
headers={
"Authorization": f"Bearer {access_token}",
@ -2686,13 +2707,16 @@ class SSOAuthenticationHandler:
},
timeout=30.0,
)
if resp.status_code == 200:
userinfo = resp.json()
else:
verbose_proxy_logger.warning(
"Userinfo endpoint returned %s, falling back to id_token",
resp.status_code,
)
if resp.status_code == 200:
userinfo = resp.json()
else:
verbose_proxy_logger.warning(
"Userinfo endpoint returned %s, falling back to id_token",
resp.status_code,
)
finally:
if _own_client:
await _client.aclose()
except Exception as e:
verbose_proxy_logger.warning(
"Userinfo endpoint error: %s, falling back to id_token", e

View file

@ -2473,17 +2473,9 @@ class ProxyConfig:
## INIT PROXY REDIS USAGE CLIENT ##
redis_usage_cache = litellm.cache.cache
## 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
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(
"Configured user_api_key_cache to use Redis "
"(PKCE enabled — verifiers shared across instances)."
)
## INIT PROXY REDIS USAGE CLIENT ##
# Note: PKCE verifier storage uses redis_usage_cache directly (not
# user_api_key_cache) to avoid routing all API-key lookups through Redis.
def switch_on_llm_response_caching(self):
"""
@ -3033,28 +3025,19 @@ class ProxyConfig:
default_redis_ttl=None, # PKCE verifiers set explicit TTL on each store; Redis TTL not configured here
)
### 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.
### PKCE MULTI-INSTANCE PREREQUISITE CHECK ###
# PKCE verifiers are stored in redis_usage_cache when available so they can
# be read back by any instance (not just the one that started the auth flow).
# user_api_key_cache is intentionally left in-memory-only to avoid routing
# all API-key lookups through Redis.
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if use_pkce and user_api_key_cache.redis_cache is None:
if redis_usage_cache is not None:
# Reuse the existing Redis connection so all advanced connection
# options (SSL, db, timeouts) are inherited rather than re-created
# from a subset of env vars.
user_api_key_cache.redis_cache = redis_usage_cache
verbose_proxy_logger.info(
"Configured user_api_key_cache to use Redis "
"(PKCE enabled — verifiers shared across instances)."
)
else:
verbose_proxy_logger.warning(
"GENERIC_CLIENT_USE_PKCE=true but Redis is not configured for LiteLLM caching. "
"PKCE verifiers will not be shared across instances. "
"Configure Redis via the 'cache' section in your proxy config, "
"or enable sticky sessions for multi-instance deployments."
)
if use_pkce and redis_usage_cache is None:
verbose_proxy_logger.warning(
"GENERIC_CLIENT_USE_PKCE=true but Redis is not configured for LiteLLM caching. "
"PKCE verifiers will not be shared across instances. "
"Configure Redis via the 'cache' section in your proxy config, "
"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)
if store_model_in_db is None: