mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
address greptile review feedback (greploop iteration 14)
This commit is contained in:
parent
cbb6b6700f
commit
e72018db28
2 changed files with 95 additions and 74 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue