From 3e704e72f901769f90b8660ebdd50cafe8bec115 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Fri, 6 Mar 2026 09:50:21 -0800 Subject: [PATCH] address greptile review feedback (greploop iteration 29) --- litellm/proxy/management_endpoints/ui_sso.py | 105 ++++++++++-------- .../proxy/management_endpoints/test_ui_sso.py | 36 ++++++ 2 files changed, 95 insertions(+), 46 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index bb3dd0217e1..cac6607f302 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -2671,7 +2671,7 @@ class SSOAuthenticationHandler: token_endpoint: str, userinfo_endpoint: Optional[str], include_client_id: bool, - redirect_url: str, + redirect_url: Optional[str], additional_headers: Dict[str, str], ) -> dict: """ @@ -2691,9 +2691,12 @@ class SSOAuthenticationHandler: token_data: Dict[str, str] = { "grant_type": "authorization_code", "code": authorization_code, - "redirect_uri": redirect_url, "code_verifier": code_verifier, } + # Only include redirect_uri when set — omitting it avoids sending the + # literal string "None" to the provider if the env var is missing. + if redirect_url: + token_data["redirect_uri"] = redirect_url post_kwargs: Dict[str, Any] = { "data": token_data, @@ -2716,52 +2719,55 @@ class SSOAuthenticationHandler: if client_secret: token_data["client_secret"] = client_secret - # 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: + # Perform the POST inside async with; response is buffered by httpx so + # status_code, text, and json() are safe to access after __aexit__. + # Only the POST itself is wrapped in try/except — TLS teardown errors + # from __aexit__ propagate as-is and are NOT mis-labelled as "token + # endpoint request failed". + try: + async with httpx.AsyncClient() as http_client: response = await http_client.post(token_endpoint, **post_kwargs) - except Exception as exc: - # Catch all network-level errors (SSL, DNS, TCP, timeout, etc.) and - # wrap them as a clean ProxyException rather than leaking raw httpx/OS - # exceptions to callers. - 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 + except Exception as exc: + # Catch network-level errors (SSL, DNS, TCP, timeout, etc.) and + # wrap them as a clean ProxyException rather than leaking raw httpx + # or OS exceptions to callers. + 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, - ) + # All response processing happens outside the async with block — + # httpx buffers the full response so status_code / text / json() are safe. + 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 @@ -2803,10 +2809,17 @@ class SSOAuthenticationHandler: # But bearer credentials (access_token, id_token, refresh_token) must always come # from the token endpoint, not from userinfo (non-standard providers occasionally # include these fields in userinfo, which would otherwise shadow the real bearer token). + # Skip re-insertion when the token_response value is None (e.g. "id_token": null) — + # an absent key is a cleaner signal for "not present" than an explicit None and avoids + # diverging from the non-PKCE path where these fields are simply absent. merged = {**token_response, **userinfo} for field in _OAUTH_TOKEN_FIELDS: - if field in token_response: + if token_response.get(field) is not None: merged[field] = token_response[field] + elif field in merged: + # Remove the key entirely if token_response had it as null/None so that + # callers can use `field in response` as a reliable presence check. + del merged[field] return merged @staticmethod 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 f130dfe16bc..965cf0b90da 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -3825,6 +3825,42 @@ class TestPKCEFunctionality: assert "401" in exc_info.value.message or "token" in exc_info.value.message.lower() + @pytest.mark.asyncio + async def test_pkce_cache_miss_unexpected_format_non_strict_logs_warning(self): + """When cached data has an unexpected format (e.g. integer from corrupt Redis) + in non-strict mode, prepare_token_exchange_parameters logs a warning and + returns params without code_verifier rather than raising.""" + import os + from unittest.mock import AsyncMock, MagicMock, patch + + from starlette.requests import Request + + 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_non_strict"} + + # Non-strict mode: should log a warning and continue, not raise + 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"}, + clear=False, + ): + os.environ.pop("PKCE_STRICT_CACHE_MISS", None) + result = await SSOAuthenticationHandler.prepare_token_exchange_parameters( + request=mock_request, generic_include_client_id=False + ) + + # No raise in non-strict mode; verifier simply absent from params + assert "code_verifier" not in result + assert "_pkce_cache_key" not in result + # Tests for SSO user team assignment bug (Issue: SSO Users Not Added to Entra-Synced Teams on First Login) class TestAddMissingTeamMember: