mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
c6f2446fa5
commit
1a16dcf598
2 changed files with 98 additions and 91 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue