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,
|
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)
|
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
|
# Get authorization code from query params
|
||||||
authorization_code = request.query_params.get("code")
|
authorization_code = request.query_params.get("code")
|
||||||
|
|
@ -816,6 +817,10 @@ async def get_generic_sso_response(
|
||||||
redirect_url=redirect_url,
|
redirect_url=redirect_url,
|
||||||
additional_headers=additional_generic_sso_headers_dict,
|
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
|
# Pass the full response so custom response_convertor implementations
|
||||||
# can access all fields (including id_token for claim extraction).
|
# can access all fields (including id_token for claim extraction).
|
||||||
result = response_convertor(combined_response, generic_sso)
|
result = response_convertor(combined_response, generic_sso)
|
||||||
|
|
@ -2530,13 +2535,12 @@ class SSOAuthenticationHandler:
|
||||||
)
|
)
|
||||||
|
|
||||||
if code_verifier:
|
if code_verifier:
|
||||||
# Add code_verifier to token exchange parameters
|
# Add code_verifier to token exchange parameters.
|
||||||
token_params["code_verifier"] = code_verifier
|
token_params["code_verifier"] = code_verifier
|
||||||
# Clean up the cache entry (single-use verifier)
|
# Return the cache key so the caller can delete it *after* a
|
||||||
if redis_usage_cache is not None:
|
# successful token exchange (avoids losing the verifier on retry
|
||||||
await redis_usage_cache.async_delete_cache(key=cache_key)
|
# if the exchange fails partway through).
|
||||||
else:
|
token_params["_pkce_cache_key"] = cache_key
|
||||||
await user_api_key_cache.async_delete_cache(key=cache_key)
|
|
||||||
else:
|
else:
|
||||||
# PKCE is enabled (already checked above) but verifier is missing — likely a cross-instance cache miss.
|
# 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
|
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
|
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
|
@staticmethod
|
||||||
def generate_pkce_params() -> Tuple[str, str]:
|
def generate_pkce_params() -> Tuple[str, str]:
|
||||||
"""
|
"""
|
||||||
|
|
@ -2634,61 +2648,67 @@ class SSOAuthenticationHandler:
|
||||||
if client_secret:
|
if client_secret:
|
||||||
token_data["client_secret"] = client_secret
|
token_data["client_secret"] = client_secret
|
||||||
|
|
||||||
async with httpx.AsyncClient() as http_client:
|
try:
|
||||||
response = await http_client.post(token_endpoint, **post_kwargs)
|
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:
|
if response.status_code != 200:
|
||||||
verbose_proxy_logger.error(
|
verbose_proxy_logger.error(
|
||||||
"PKCE token exchange failed. status=%s body=%s",
|
"PKCE token exchange failed. status=%s body=%s",
|
||||||
response.status_code,
|
response.status_code,
|
||||||
response.text[:500],
|
response.text[:500],
|
||||||
)
|
)
|
||||||
raise ProxyException(
|
raise ProxyException(
|
||||||
message=f"Token exchange failed: {response.status_code} - {response.text[:500]}",
|
message=f"Token exchange failed: {response.status_code} - {response.text[:500]}",
|
||||||
type=ProxyErrorTypes.auth_error,
|
type=ProxyErrorTypes.auth_error,
|
||||||
param="token_exchange",
|
param="token_exchange",
|
||||||
code=status.HTTP_401_UNAUTHORIZED,
|
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")),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# token_response is set inside the async with block above and remains accessible here;
|
try:
|
||||||
# Python's scoping rules guarantee it is defined if no exception was raised.
|
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(
|
userinfo = await SSOAuthenticationHandler._get_pkce_userinfo(
|
||||||
access_token=token_response["access_token"],
|
access_token=token_response["access_token"],
|
||||||
id_token=token_response.get("id_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).
|
# 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.
|
# 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:
|
try:
|
||||||
userinfo = jwt.decode(id_token, options={"verify_signature": False})
|
userinfo = jwt.decode(id_token, options={"verify_signature": False})
|
||||||
except Exception as decode_err:
|
except Exception as decode_err:
|
||||||
|
|
|
||||||
|
|
@ -3161,14 +3161,15 @@ class TestPKCEFunctionality:
|
||||||
# Assert
|
# Assert
|
||||||
assert token_params["include_client_id"] is False
|
assert token_params["include_client_id"] is False
|
||||||
assert token_params["code_verifier"] == test_code_verifier
|
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(
|
mock_cache.async_get_cache.assert_called_once_with(
|
||||||
key=f"pkce_verifier:{test_state}"
|
key=f"pkce_verifier:{test_state}"
|
||||||
)
|
)
|
||||||
mock_cache.async_delete_cache.assert_called_once_with(
|
mock_cache.async_delete_cache.assert_not_called()
|
||||||
key=f"pkce_verifier:{test_state}"
|
|
||||||
)
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_get_generic_sso_redirect_response_with_pkce(self):
|
async def test_get_generic_sso_redirect_response_with_pkce(self):
|
||||||
|
|
@ -3292,11 +3293,11 @@ class TestPKCEFunctionality:
|
||||||
)
|
)
|
||||||
assert "code_verifier" in token_params
|
assert "code_verifier" in token_params
|
||||||
assert token_params["code_verifier"] == stored_dict["code_verifier"]
|
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()
|
mock_in_memory.async_get_cache.assert_not_called()
|
||||||
# delete_cache called; key removed (asserted below)
|
# Deletion is deferred — key still present until exchange succeeds
|
||||||
|
assert stored_key in mock_redis._store
|
||||||
# Verifier consumed (single-use); key removed from "Redis"
|
|
||||||
assert "pkce_verifier:multi_pod_state_xyz" not in mock_redis._store
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_pkce_fallback_in_memory_roundtrip_when_redis_none(self):
|
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 "code_verifier" in token_params
|
||||||
assert token_params["code_verifier"] == stored_value["code_verifier"]
|
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(
|
mock_in_memory.async_get_cache.assert_called_once_with(
|
||||||
key=stored_key
|
key=stored_key
|
||||||
)
|
)
|
||||||
mock_in_memory.async_delete_cache.assert_called_once_with(
|
# Deletion is deferred — not called by prepare_token_exchange_parameters
|
||||||
key=stored_key
|
mock_in_memory.async_delete_cache.assert_not_called()
|
||||||
)
|
|
||||||
|
|
||||||
# Verifier consumed; key removed from in-memory
|
|
||||||
assert "pkce_verifier:fallback_state_xyz" not in in_memory_store
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_pkce_prepare_token_exchange_returns_nothing_when_no_state(self):
|
async def test_pkce_prepare_token_exchange_returns_nothing_when_no_state(self):
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue