From 0fe3cec376c1d5be0c2b958158b750a593e8970d Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Thu, 5 Mar 2026 19:48:54 -0800 Subject: [PATCH] address greptile review feedback (greploop iteration 24) --- litellm/proxy/management_endpoints/ui_sso.py | 7 +- .../proxy/management_endpoints/test_ui_sso.py | 709 +++++++++--------- 2 files changed, 359 insertions(+), 357 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 7a4ef8b8659..76faf647e80 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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: 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 171d94bd457..2582a532b12 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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()