mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
address greptile review feedback (greploop iteration 24)
This commit is contained in:
parent
b7e8eb4235
commit
0fe3cec376
2 changed files with 359 additions and 357 deletions
|
|
@ -2581,6 +2581,10 @@ class SSOAuthenticationHandler:
|
|||
if redis_usage_cache is None
|
||||
else ""
|
||||
)
|
||||
# Raise immediately — falling through to a non-PKCE flow would only
|
||||
# produce a confusing provider-side error (provider requires code_verifier).
|
||||
# Since PKCE support is new in this release, there is no prior behavior
|
||||
# to preserve: the verifier was never actually used before this PR.
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"PKCE verifier not found in cache for state '{state}'. "
|
||||
|
|
@ -2843,7 +2847,8 @@ 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.
|
||||
# An empty dict ({}) is treated as a failure (userinfo is set to None above)
|
||||
# so we also attempt the id_token fallback in that case.
|
||||
# Explicitly check for a non-empty string to avoid attempting JWT decode on
|
||||
# a blank or non-string id_token field from a misbehaving provider.
|
||||
if userinfo is None and isinstance(id_token, str) and id_token:
|
||||
|
|
|
|||
|
|
@ -3395,6 +3395,359 @@ class TestPKCEFunctionality:
|
|||
mock_in_memory.async_get_cache.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_basic_auth(self):
|
||||
"""When include_client_id=False, client credentials go via HTTP Basic Auth."""
|
||||
token_resp = {
|
||||
"access_token": "tok_abc",
|
||||
"id_token": None,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
userinfo_resp = {"sub": "user1", "email": "user@example.com"}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = token_resp
|
||||
|
||||
mock_userinfo_response = MagicMock()
|
||||
mock_userinfo_response.status_code = 200
|
||||
mock_userinfo_response.json.return_value = userinfo_resp
|
||||
|
||||
async def fake_post(*args, **kwargs):
|
||||
# Verify Basic Auth is set
|
||||
assert "auth" in kwargs
|
||||
assert isinstance(kwargs["auth"], httpx.BasicAuth)
|
||||
# Verify code_verifier is in the POST body (essential PKCE field)
|
||||
assert kwargs.get("data", {}).get("code_verifier") == "verifier_abc"
|
||||
return mock_response
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_client.post = AsyncMock(side_effect=fake_post)
|
||||
mock_client.get = AsyncMock(return_value=mock_userinfo_response)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="auth_code_123",
|
||||
code_verifier="verifier_abc",
|
||||
client_id="my_client",
|
||||
client_secret="my_secret",
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=False,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["access_token"] == "tok_abc"
|
||||
assert result["email"] == "user@example.com"
|
||||
# Verify userinfo GET used the correct Bearer token header
|
||||
get_call = mock_client.get.call_args
|
||||
assert get_call is not None
|
||||
assert get_call.kwargs["headers"]["Authorization"] == "Bearer tok_abc"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_credentials_in_body(self):
|
||||
"""When include_client_id=True, credentials go in the request body."""
|
||||
token_resp = {
|
||||
"access_token": "tok_body",
|
||||
"id_token": None,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
userinfo_resp = {"sub": "user2", "email": "user2@example.com"}
|
||||
|
||||
async def fake_post(*args, **kwargs):
|
||||
assert "auth" not in kwargs, "Should NOT use Basic Auth when include_client_id=True"
|
||||
data = kwargs.get("data", {})
|
||||
assert "client_id" in data
|
||||
assert "client_secret" in data
|
||||
assert data.get("code_verifier") == "verifier_xyz", "code_verifier must be in POST body"
|
||||
mock = MagicMock()
|
||||
mock.status_code = 200
|
||||
mock.json.return_value = token_resp
|
||||
return mock
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_client.post = AsyncMock(side_effect=fake_post)
|
||||
mock_userinfo = MagicMock()
|
||||
mock_userinfo.status_code = 200
|
||||
mock_userinfo.json.return_value = userinfo_resp
|
||||
mock_client.get = AsyncMock(return_value=mock_userinfo)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="auth_code_456",
|
||||
code_verifier="verifier_xyz",
|
||||
client_id="client_id_value",
|
||||
client_secret="client_secret_value",
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=True,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["access_token"] == "tok_body"
|
||||
assert result["sub"] == "user2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_http200_with_error_body(self):
|
||||
"""Provider returns HTTP 200 but with an error field instead of tokens."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
error_body = {"error": "invalid_grant", "error_description": "Code already used"}
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.json.return_value = error_body
|
||||
mock_client.post = AsyncMock(return_value=mock_resp)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="expired_code",
|
||||
code_verifier="verifier",
|
||||
client_id="cid",
|
||||
client_secret="csecret",
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=False,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert "invalid_grant" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_userinfo_falls_back_to_id_token(self):
|
||||
"""When the userinfo endpoint fails, decode the id_token as fallback."""
|
||||
import base64
|
||||
import json as _json
|
||||
|
||||
payload = {"sub": "user_from_jwt", "email": "jwt@example.com"}
|
||||
# Build a minimal JWT (header.payload.signature — signature not verified)
|
||||
encoded_payload = base64.urlsafe_b64encode(
|
||||
_json.dumps(payload).encode()
|
||||
).rstrip(b"=").decode()
|
||||
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_fail = MagicMock()
|
||||
mock_fail.status_code = 503
|
||||
mock_client.get = AsyncMock(return_value=mock_fail)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._get_pkce_userinfo(
|
||||
access_token="some_token",
|
||||
id_token=fake_id_token,
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["sub"] == "user_from_jwt"
|
||||
assert result["email"] == "jwt@example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_userinfo_uses_id_token_when_no_endpoint(self):
|
||||
"""When userinfo_endpoint is None, fall back to id_token directly without HTTP call."""
|
||||
import base64
|
||||
import json as _json
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
payload = {"sub": "id_token_user", "email": "id@example.com"}
|
||||
encoded_payload = (
|
||||
base64.urlsafe_b64encode(_json.dumps(payload).encode()).rstrip(b"=").decode()
|
||||
)
|
||||
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
|
||||
|
||||
# No httpx call should happen when userinfo_endpoint is None
|
||||
result = await SSOAuthenticationHandler._get_pkce_userinfo(
|
||||
access_token="some_token",
|
||||
id_token=fake_id_token,
|
||||
userinfo_endpoint=None,
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["sub"] == "id_token_user"
|
||||
assert result["email"] == "id@example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_userinfo_raises_when_both_sources_unavailable(self):
|
||||
"""When userinfo endpoint fails AND no id_token, raise ProxyException."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_fail = MagicMock()
|
||||
mock_fail.status_code = 503
|
||||
mock_client.get = AsyncMock(return_value=mock_fail)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await SSOAuthenticationHandler._get_pkce_userinfo(
|
||||
access_token="token",
|
||||
id_token=None, # no id_token available
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert "unavailable" in exc_info.value.message.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_cache_miss_raises_proxy_exception(self):
|
||||
"""prepare_token_exchange_parameters raises ProxyException when PKCE is enabled
|
||||
but no verifier is found in cache (cross-instance cache miss scenario)."""
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from starlette.requests import Request
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
|
||||
|
||||
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"}):
|
||||
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
|
||||
request=mock_request, generic_include_client_id=False
|
||||
)
|
||||
|
||||
assert "verifier not found" in exc_info.value.message.lower() or "cache" in exc_info.value.message.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_public_client_no_secret(self):
|
||||
"""Public PKCE client (include_client_id=False, no secret) sends client_id in
|
||||
POST body and does NOT include Basic Auth or client_secret."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
token_resp = {
|
||||
"access_token": "tok_public",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
userinfo_resp = {"sub": "pubuser", "email": "pub@example.com"}
|
||||
|
||||
async def fake_post(*args, **kwargs):
|
||||
assert "auth" not in kwargs, "Public client must not use Basic Auth"
|
||||
data = kwargs.get("data", {})
|
||||
assert data.get("client_id") == "public_client_id"
|
||||
assert "client_secret" not in data, "No secret should be sent for public client"
|
||||
assert data.get("code_verifier") == "public_verifier"
|
||||
mock = MagicMock()
|
||||
mock.status_code = 200
|
||||
mock.json.return_value = token_resp
|
||||
return mock
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_client.post = AsyncMock(side_effect=fake_post)
|
||||
mock_userinfo = MagicMock()
|
||||
mock_userinfo.status_code = 200
|
||||
mock_userinfo.json.return_value = userinfo_resp
|
||||
mock_client.get = AsyncMock(return_value=mock_userinfo)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="auth_pub",
|
||||
code_verifier="public_verifier",
|
||||
client_id="public_client_id",
|
||||
client_secret=None, # public client — no secret
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=False,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["access_token"] == "tok_public"
|
||||
assert result["sub"] == "pubuser"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_pkce_verifier_swallows_deletion_errors(self):
|
||||
"""_delete_pkce_verifier must not raise when the cache delete fails
|
||||
(best-effort cleanup — a leftover verifier must not abort a successful SSO login)."""
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
failing_cache = MagicMock()
|
||||
failing_cache.async_delete_cache = AsyncMock(side_effect=Exception("Redis down"))
|
||||
|
||||
# Should NOT raise even though the underlying cache delete fails
|
||||
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", failing_cache
|
||||
):
|
||||
await SSOAuthenticationHandler._delete_pkce_verifier("pkce_verifier:test_state")
|
||||
|
||||
failing_cache.async_delete_cache.assert_called_once_with(key="pkce_verifier:test_state")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_cache_miss_unexpected_format_raises_proxy_exception(self):
|
||||
"""When cached data exists but has an unrecognized format (not a dict with
|
||||
code_verifier, not a plain string), prepare_token_exchange_parameters raises
|
||||
ProxyException rather than silently falling through to a non-PKCE flow."""
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from starlette.requests import Request
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
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_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"}):
|
||||
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()
|
||||
|
||||
# Tests for SSO user team assignment bug (Issue: SSO Users Not Added to Entra-Synced Teams on First Login)
|
||||
class TestAddMissingTeamMember:
|
||||
"""Tests for the add_missing_team_member function"""
|
||||
|
|
@ -4501,359 +4854,3 @@ def test_generic_response_convertor_extra_attributes_missing_field(monkeypatch):
|
|||
assert result.extra_fields["missing_field"] is None
|
||||
assert result.extra_fields["another_missing"] is None
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Tests for SSOAuthenticationHandler PKCE token exchange methods
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_basic_auth():
|
||||
"""When include_client_id=False, client credentials go via HTTP Basic Auth."""
|
||||
token_resp = {
|
||||
"access_token": "tok_abc",
|
||||
"id_token": None,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
userinfo_resp = {"sub": "user1", "email": "user@example.com"}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = token_resp
|
||||
|
||||
mock_userinfo_response = MagicMock()
|
||||
mock_userinfo_response.status_code = 200
|
||||
mock_userinfo_response.json.return_value = userinfo_resp
|
||||
|
||||
async def fake_post(*args, **kwargs):
|
||||
# Verify Basic Auth is set
|
||||
assert "auth" in kwargs
|
||||
assert isinstance(kwargs["auth"], httpx.BasicAuth)
|
||||
# Verify code_verifier is in the POST body (essential PKCE field)
|
||||
assert kwargs.get("data", {}).get("code_verifier") == "verifier_abc"
|
||||
return mock_response
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_client.post = AsyncMock(side_effect=fake_post)
|
||||
mock_client.get = AsyncMock(return_value=mock_userinfo_response)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="auth_code_123",
|
||||
code_verifier="verifier_abc",
|
||||
client_id="my_client",
|
||||
client_secret="my_secret",
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=False,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["access_token"] == "tok_abc"
|
||||
assert result["email"] == "user@example.com"
|
||||
# Verify userinfo GET used the correct Bearer token header
|
||||
get_call = mock_client.get.call_args
|
||||
assert get_call is not None
|
||||
assert get_call.kwargs["headers"]["Authorization"] == "Bearer tok_abc"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_credentials_in_body():
|
||||
"""When include_client_id=True, credentials go in the request body."""
|
||||
token_resp = {
|
||||
"access_token": "tok_body",
|
||||
"id_token": None,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
userinfo_resp = {"sub": "user2", "email": "user2@example.com"}
|
||||
|
||||
async def fake_post(*args, **kwargs):
|
||||
assert "auth" not in kwargs, "Should NOT use Basic Auth when include_client_id=True"
|
||||
data = kwargs.get("data", {})
|
||||
assert "client_id" in data
|
||||
assert "client_secret" in data
|
||||
assert data.get("code_verifier") == "verifier_xyz", "code_verifier must be in POST body"
|
||||
mock = MagicMock()
|
||||
mock.status_code = 200
|
||||
mock.json.return_value = token_resp
|
||||
return mock
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_client.post = AsyncMock(side_effect=fake_post)
|
||||
mock_userinfo = MagicMock()
|
||||
mock_userinfo.status_code = 200
|
||||
mock_userinfo.json.return_value = userinfo_resp
|
||||
mock_client.get = AsyncMock(return_value=mock_userinfo)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="auth_code_456",
|
||||
code_verifier="verifier_xyz",
|
||||
client_id="client_id_value",
|
||||
client_secret="client_secret_value",
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=True,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["access_token"] == "tok_body"
|
||||
assert result["sub"] == "user2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_http200_with_error_body():
|
||||
"""Provider returns HTTP 200 but with an error field instead of tokens."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
error_body = {"error": "invalid_grant", "error_description": "Code already used"}
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.json.return_value = error_body
|
||||
mock_client.post = AsyncMock(return_value=mock_resp)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="expired_code",
|
||||
code_verifier="verifier",
|
||||
client_id="cid",
|
||||
client_secret="csecret",
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=False,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert "invalid_grant" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_userinfo_falls_back_to_id_token():
|
||||
"""When the userinfo endpoint fails, decode the id_token as fallback."""
|
||||
import base64
|
||||
import json as _json
|
||||
|
||||
payload = {"sub": "user_from_jwt", "email": "jwt@example.com"}
|
||||
# Build a minimal JWT (header.payload.signature — signature not verified)
|
||||
encoded_payload = base64.urlsafe_b64encode(
|
||||
_json.dumps(payload).encode()
|
||||
).rstrip(b"=").decode()
|
||||
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_fail = MagicMock()
|
||||
mock_fail.status_code = 503
|
||||
mock_client.get = AsyncMock(return_value=mock_fail)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._get_pkce_userinfo(
|
||||
access_token="some_token",
|
||||
id_token=fake_id_token,
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["sub"] == "user_from_jwt"
|
||||
assert result["email"] == "jwt@example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_userinfo_uses_id_token_when_no_endpoint():
|
||||
"""When userinfo_endpoint is None, fall back to id_token directly without HTTP call."""
|
||||
import base64
|
||||
import json as _json
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
payload = {"sub": "id_token_user", "email": "id@example.com"}
|
||||
encoded_payload = (
|
||||
base64.urlsafe_b64encode(_json.dumps(payload).encode()).rstrip(b"=").decode()
|
||||
)
|
||||
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
|
||||
|
||||
# No httpx call should happen when userinfo_endpoint is None
|
||||
result = await SSOAuthenticationHandler._get_pkce_userinfo(
|
||||
access_token="some_token",
|
||||
id_token=fake_id_token,
|
||||
userinfo_endpoint=None,
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["sub"] == "id_token_user"
|
||||
assert result["email"] == "id@example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_userinfo_raises_when_both_sources_unavailable():
|
||||
"""When userinfo endpoint fails AND no id_token, raise ProxyException."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_fail = MagicMock()
|
||||
mock_fail.status_code = 503
|
||||
mock_client.get = AsyncMock(return_value=mock_fail)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await SSOAuthenticationHandler._get_pkce_userinfo(
|
||||
access_token="token",
|
||||
id_token=None, # no id_token available
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert "unavailable" in exc_info.value.message.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_cache_miss_raises_proxy_exception():
|
||||
"""prepare_token_exchange_parameters raises ProxyException when PKCE is enabled
|
||||
but no verifier is found in cache (cross-instance cache miss scenario)."""
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from starlette.requests import Request
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
|
||||
|
||||
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"}):
|
||||
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
|
||||
request=mock_request, generic_include_client_id=False
|
||||
)
|
||||
|
||||
assert "verifier not found" in exc_info.value.message.lower() or "cache" in exc_info.value.message.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_token_exchange_public_client_no_secret():
|
||||
"""Public PKCE client (include_client_id=False, no secret) sends client_id in
|
||||
POST body and does NOT include Basic Auth or client_secret."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
token_resp = {
|
||||
"access_token": "tok_public",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
userinfo_resp = {"sub": "pubuser", "email": "pub@example.com"}
|
||||
|
||||
async def fake_post(*args, **kwargs):
|
||||
assert "auth" not in kwargs, "Public client must not use Basic Auth"
|
||||
data = kwargs.get("data", {})
|
||||
assert data.get("client_id") == "public_client_id"
|
||||
assert "client_secret" not in data, "No secret should be sent for public client"
|
||||
assert data.get("code_verifier") == "public_verifier"
|
||||
mock = MagicMock()
|
||||
mock.status_code = 200
|
||||
mock.json.return_value = token_resp
|
||||
return mock
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_client.post = AsyncMock(side_effect=fake_post)
|
||||
mock_userinfo = MagicMock()
|
||||
mock_userinfo.status_code = 200
|
||||
mock_userinfo.json.return_value = userinfo_resp
|
||||
mock_client.get = AsyncMock(return_value=mock_userinfo)
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
result = await SSOAuthenticationHandler._pkce_token_exchange(
|
||||
authorization_code="auth_pub",
|
||||
code_verifier="public_verifier",
|
||||
client_id="public_client_id",
|
||||
client_secret=None, # public client — no secret
|
||||
token_endpoint="https://example.com/token",
|
||||
userinfo_endpoint="https://example.com/userinfo",
|
||||
include_client_id=False,
|
||||
redirect_url="https://proxy.example.com/callback",
|
||||
additional_headers={},
|
||||
)
|
||||
|
||||
assert result["access_token"] == "tok_public"
|
||||
assert result["sub"] == "pubuser"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_pkce_verifier_swallows_deletion_errors():
|
||||
"""_delete_pkce_verifier must not raise when the cache delete fails
|
||||
(best-effort cleanup — a leftover verifier must not abort a successful SSO login)."""
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
failing_cache = MagicMock()
|
||||
failing_cache.async_delete_cache = AsyncMock(side_effect=Exception("Redis down"))
|
||||
|
||||
# Should NOT raise even though the underlying cache delete fails
|
||||
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", failing_cache
|
||||
):
|
||||
await SSOAuthenticationHandler._delete_pkce_verifier("pkce_verifier:test_state")
|
||||
|
||||
failing_cache.async_delete_cache.assert_called_once_with(key="pkce_verifier:test_state")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pkce_cache_miss_unexpected_format_raises_proxy_exception():
|
||||
"""When cached data exists but has an unrecognized format (not a dict with
|
||||
code_verifier, not a plain string), prepare_token_exchange_parameters raises
|
||||
ProxyException rather than silently falling through to a non-PKCE flow."""
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from starlette.requests import Request
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
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_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"}):
|
||||
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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue