diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 3243e4ff084..aceb310646f 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -2685,9 +2685,10 @@ class SSOAuthenticationHandler: if client_secret: token_data["client_secret"] = client_secret - # Tighten the try/except to the POST call only, so httpx connection-pool - # teardown in __aexit__ (TLS close, etc.) does not get misclassified as a - # token-endpoint failure. + # Keep all response processing inside the async with block so that the + # response object (which httpx buffers) is always accessed while the client + # is still alive. Network errors on the POST are caught tightly here; TLS + # teardown errors from __aexit__ are NOT classified as token failures. async with httpx.AsyncClient() as http_client: try: response = await http_client.post(token_endpoint, **post_kwargs) @@ -2703,33 +2704,33 @@ class SSOAuthenticationHandler: code=status.HTTP_401_UNAUTHORIZED, ) from exc - 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 ProxyException( - message=f"Token endpoint returned invalid JSON: {json_err}", - 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 ProxyException( + message=f"Token endpoint returned invalid JSON: {json_err}", + type=ProxyErrorTypes.auth_error, + param="token_exchange", + code=status.HTTP_401_UNAUTHORIZED, + ) # Some providers return HTTP 200 with an error body (e.g. expired code, replay attack). # Also guard against JSON `null` for access_token — it passes key-existence checks diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 6ca9cf85891..171d94bd457 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -4826,3 +4826,34 @@ async def test_delete_pkce_verifier_swallows_deletion_errors(): await SSOAuthenticationHandler._delete_pkce_verifier("pkce_verifier:test_state") failing_cache.async_delete_cache.assert_called_once_with(key="pkce_verifier:test_state") + + +@pytest.mark.asyncio +async def test_pkce_cache_miss_unexpected_format_raises_proxy_exception(): + """When cached data exists but has an unrecognized format (not a dict with + code_verifier, not a plain string), prepare_token_exchange_parameters raises + ProxyException rather than silently falling through to a non-PKCE flow.""" + import os + from unittest.mock import AsyncMock, MagicMock, patch + + from starlette.requests import Request + + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Cache returns an integer — unexpected format + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=12345) + + mock_request = MagicMock(spec=Request) + mock_request.query_params = {"state": "bad_format_state"} + + with pytest.raises(ProxyException) as exc_info: + with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ), patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}): + await SSOAuthenticationHandler.prepare_token_exchange_parameters( + request=mock_request, generic_include_client_id=False + ) + + assert "cache" in exc_info.value.message.lower() or "verifier" in exc_info.value.message.lower()