address greptile review feedback (greploop iteration 29)

This commit is contained in:
Ishaan Jaffer 2026-03-06 09:50:21 -08:00
parent e55e7546d0
commit 3e704e72f9
2 changed files with 95 additions and 46 deletions

View file

@ -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

View file

@ -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: