From 04d3d552871e564dab2cc6a79ca92ed0b29133e0 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Fri, 6 Mar 2026 10:14:56 -0800 Subject: [PATCH] address greptile review feedback (greploop iteration 31) --- litellm/proxy/management_endpoints/ui_sso.py | 70 ++++++++++++------- .../proxy/management_endpoints/test_ui_sso.py | 30 ++++---- 2 files changed, 61 insertions(+), 39 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index e3df94c3257..daf9e2e44ab 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -852,6 +852,9 @@ async def get_generic_sso_response( # Strip bearer credentials from received_response after conversion. # received_response may appear in restricted-group error messages — # do not expose tokens to callers. + # Note: response_convertor always sets received_response via nonlocal, so it is + # non-None here at runtime. The guard satisfies Pyright's type narrowing since + # it cannot track nonlocal mutations inside closures. if received_response is not None: received_response = { k: v for k, v in received_response.items() if k not in _OAUTH_TOKEN_FIELDS @@ -2585,30 +2588,49 @@ class SSOAuthenticationHandler: os.getenv("PKCE_STRICT_CACHE_MISS", "false").lower() == "true" ) if strict_cache_miss: - verbose_proxy_logger.error( - "PKCE is enabled but no usable code_verifier found for state '%s'. " - "This usually means the authorization and callback were handled by different " - "instances without a shared cache, or the cached value had an unrecognized format. " - "Ensure Redis is configured. " - "Cache type: %s. Raw cache data present (may be unrecognized format): %s", - state, - type(active_cache).__name__, - cached_data is not None, - ) - redis_hint = ( - " Configure Redis and set REDIS_URL so all proxy instances share the PKCE verifier." - if redis_usage_cache is None - else "" - ) - raise ProxyException( - message=( - f"PKCE verifier not found in cache for state '{state}'. " - f"The login and callback requests were likely handled by different instances.{redis_hint}" - ), - type=ProxyErrorTypes.auth_error, - param="PKCE_CACHE_MISS", - code=status.HTTP_401_UNAUTHORIZED, - ) + # Distinguish corrupt-format entries from genuine cache misses + # so operators can investigate the correct root cause. + if cached_data is not None: + # Cache had data but in an unrecognised format (e.g. corrupt Redis value). + verbose_proxy_logger.error( + "PKCE verifier for state '%s' has an unrecognized format (type=%s); " + "treating as a cache miss. Investigate the cached value — it may be " + "a corrupt or stale entry.", + state, + type(cached_data).__name__, + ) + raise ProxyException( + message=( + f"PKCE verifier for state '{state}' has an unrecognized format " + f"(type={type(cached_data).__name__}). The cached entry may be corrupt." + ), + type=ProxyErrorTypes.auth_error, + param="PKCE_CACHE_MISS", + code=status.HTTP_401_UNAUTHORIZED, + ) + else: + # Genuine cache miss — verifier was never stored or already expired. + verbose_proxy_logger.error( + "PKCE is enabled but no verifier found in cache for state '%s'. " + "The authorization and callback were likely handled by different " + "instances without a shared cache. Cache type: %s.", + state, + type(active_cache).__name__, + ) + redis_hint = ( + " Configure Redis and set REDIS_URL so all proxy instances share the PKCE verifier." + if redis_usage_cache is None + else "" + ) + raise ProxyException( + message=( + f"PKCE verifier not found in cache for state '{state}'. " + f"The login and callback requests were likely handled by different instances.{redis_hint}" + ), + type=ProxyErrorTypes.auth_error, + param="PKCE_CACHE_MISS", + code=status.HTTP_401_UNAUTHORIZED, + ) else: verbose_proxy_logger.warning( "PKCE is enabled but verifier not found in cache for state '%s' " 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 965cf0b90da..5dfbf545860 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -3637,13 +3637,13 @@ class TestPKCEFunctionality: mock_request = MagicMock(spec=Request) mock_request.query_params = {"state": "missing_state_123"} - 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", "PKCE_STRICT_CACHE_MISS": "true"}, - ): + 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", "PKCE_STRICT_CACHE_MISS": "true"}, + ): + with pytest.raises(ProxyException) as exc_info: await SSOAuthenticationHandler.prepare_token_exchange_parameters( request=mock_request, generic_include_client_id=False ) @@ -3742,18 +3742,18 @@ class TestPKCEFunctionality: 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", "PKCE_STRICT_CACHE_MISS": "true"}, - ): + 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", "PKCE_STRICT_CACHE_MISS": "true"}, + ): + with pytest.raises(ProxyException) as exc_info: 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() + assert "cache" in exc_info.value.message.lower() or "verifier" in exc_info.value.message.lower() or "format" in exc_info.value.message.lower() @pytest.mark.asyncio async def test_pkce_cache_miss_non_strict_logs_warning_and_continues(self):