From e72018db282d114feed73b4feabb9bf273cddc4e Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Thu, 5 Mar 2026 16:54:13 -0800 Subject: [PATCH] address greptile review feedback (greploop iteration 14) --- litellm/proxy/management_endpoints/ui_sso.py | 142 ++++++++++-------- .../proxy/management_endpoints/test_ui_sso.py | 27 ++-- 2 files changed, 95 insertions(+), 74 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index f07c2face9d..7d45e973562 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -791,8 +791,9 @@ async def get_generic_sso_response( generic_include_client_id=generic_include_client_id, ) - # Extract code_verifier before calling fastapi-sso + # Extract code_verifier (and the cache key for deferred deletion) before calling fastapi-sso code_verifier = token_exchange_params.pop("code_verifier", None) + pkce_cache_key = token_exchange_params.pop("_pkce_cache_key", None) # Get authorization code from query params authorization_code = request.query_params.get("code") @@ -816,6 +817,10 @@ async def get_generic_sso_response( redirect_url=redirect_url, additional_headers=additional_generic_sso_headers_dict, ) + # Token exchange succeeded — now it is safe to delete the single-use + # verifier. Deleting before the exchange would lose it on retries. + if pkce_cache_key: + await SSOAuthenticationHandler._delete_pkce_verifier(pkce_cache_key) # Pass the full response so custom response_convertor implementations # can access all fields (including id_token for claim extraction). result = response_convertor(combined_response, generic_sso) @@ -2530,13 +2535,12 @@ class SSOAuthenticationHandler: ) if code_verifier: - # Add code_verifier to token exchange parameters + # Add code_verifier to token exchange parameters. token_params["code_verifier"] = code_verifier - # Clean up the cache entry (single-use verifier) - if redis_usage_cache is not None: - await redis_usage_cache.async_delete_cache(key=cache_key) - else: - await user_api_key_cache.async_delete_cache(key=cache_key) + # Return the cache key so the caller can delete it *after* a + # successful token exchange (avoids losing the verifier on retry + # if the exchange fails partway through). + token_params["_pkce_cache_key"] = cache_key else: # PKCE is enabled (already checked above) but verifier is missing — likely a cross-instance cache miss. active_cache = redis_usage_cache if redis_usage_cache is not None else user_api_key_cache @@ -2552,6 +2556,16 @@ class SSOAuthenticationHandler: ) return token_params + @staticmethod + async def _delete_pkce_verifier(cache_key: str) -> None: + """Delete a single-use PKCE verifier from cache after a successful exchange.""" + from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache + + if redis_usage_cache is not None: + await redis_usage_cache.async_delete_cache(key=cache_key) + else: + await user_api_key_cache.async_delete_cache(key=cache_key) + @staticmethod def generate_pkce_params() -> Tuple[str, str]: """ @@ -2634,61 +2648,67 @@ class SSOAuthenticationHandler: if client_secret: token_data["client_secret"] = client_secret - async with httpx.AsyncClient() as http_client: - response = await http_client.post(token_endpoint, **post_kwargs) + try: + async with httpx.AsyncClient() as http_client: + response = await http_client.post(token_endpoint, **post_kwargs) + except httpx.HTTPError as exc: + verbose_proxy_logger.error("PKCE token endpoint unreachable: %s", exc) + raise ProxyException( + message=f"Token endpoint request failed: {exc}", + type=ProxyErrorTypes.auth_error, + param="token_exchange", + 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, - ) - - 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). - 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")), + 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, ) - # token_response is set inside the async with block above and remains accessible here; - # Python's scoping rules guarantee it is defined if no exception was raised. + 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). + 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")), + ) userinfo = await SSOAuthenticationHandler._get_pkce_userinfo( access_token=token_response["access_token"], id_token=token_response.get("id_token"), @@ -2744,7 +2764,9 @@ class SSOAuthenticationHandler: # Only fall back to id_token when the userinfo request failed (None). # An empty dict ({}) from the endpoint is a valid response and not retried. - if userinfo is None and id_token: + # Explicitly check for non-None and non-empty string to avoid attempting + # JWT decode on a blank id_token field. + if userinfo is None and id_token is not None and id_token != "": try: userinfo = jwt.decode(id_token, options={"verify_signature": False}) except Exception as decode_err: 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 4bac65d6d64..5512fd983c7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -3161,14 +3161,15 @@ class TestPKCEFunctionality: # Assert assert token_params["include_client_id"] is False assert token_params["code_verifier"] == test_code_verifier + # Cache key is returned for deferred deletion (after exchange succeeds) + assert token_params["_pkce_cache_key"] == f"pkce_verifier:{test_state}" - # Verify cache was accessed and deleted + # Verify cache was read but NOT deleted yet (deletion is deferred to after + # successful token exchange to preserve the verifier for retries) mock_cache.async_get_cache.assert_called_once_with( key=f"pkce_verifier:{test_state}" ) - mock_cache.async_delete_cache.assert_called_once_with( - key=f"pkce_verifier:{test_state}" - ) + mock_cache.async_delete_cache.assert_not_called() @pytest.mark.asyncio async def test_get_generic_sso_redirect_response_with_pkce(self): @@ -3292,11 +3293,11 @@ class TestPKCEFunctionality: ) assert "code_verifier" in token_params assert token_params["code_verifier"] == stored_dict["code_verifier"] + # Cache key returned for deferred deletion after successful exchange + assert token_params["_pkce_cache_key"] == stored_key mock_in_memory.async_get_cache.assert_not_called() - # delete_cache called; key removed (asserted below) - - # Verifier consumed (single-use); key removed from "Redis" - assert "pkce_verifier:multi_pod_state_xyz" not in mock_redis._store + # Deletion is deferred — key still present until exchange succeeds + assert stored_key in mock_redis._store @pytest.mark.asyncio async def test_pkce_fallback_in_memory_roundtrip_when_redis_none(self): @@ -3361,15 +3362,13 @@ class TestPKCEFunctionality: ) assert "code_verifier" in token_params assert token_params["code_verifier"] == stored_value["code_verifier"] + # Cache key returned for deferred deletion after successful exchange + assert token_params["_pkce_cache_key"] == stored_key mock_in_memory.async_get_cache.assert_called_once_with( key=stored_key ) - mock_in_memory.async_delete_cache.assert_called_once_with( - key=stored_key - ) - - # Verifier consumed; key removed from in-memory - assert "pkce_verifier:fallback_state_xyz" not in in_memory_store + # Deletion is deferred — not called by prepare_token_exchange_parameters + mock_in_memory.async_delete_cache.assert_not_called() @pytest.mark.asyncio async def test_pkce_prepare_token_exchange_returns_nothing_when_no_state(self):