diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 15991e6c526..0deea219c8a 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6be57b3dcd7..397a64915d4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: