diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 1c3c9c12619..a1223474401 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -2543,7 +2543,9 @@ class SSOAuthenticationHandler: query_params = dict(request.query_params) state = query_params.get("state") - if os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true" and not state: + use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true" + + if use_pkce and not state: verbose_proxy_logger.warning( "PKCE is enabled (GENERIC_CLIENT_USE_PKCE=true) but no 'state' parameter " "was found in the callback. The PKCE verifier cannot be retrieved without " @@ -2552,7 +2554,7 @@ class SSOAuthenticationHandler: "in the callback redirect." ) - if state and os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true": + if state and use_pkce: from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache cache_key = f"pkce_verifier:{state}" @@ -2969,7 +2971,7 @@ class SSOAuthenticationHandler: except Exception as decode_err: verbose_proxy_logger.error("Failed to decode id_token: %s", decode_err) raise ProxyException( - message="Failed to get user info from both userinfo endpoint and id_token", + message=f"Failed to decode id_token JWT: {decode_err}", type=ProxyErrorTypes.auth_error, param="userinfo", code=status.HTTP_401_UNAUTHORIZED, 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 27b681f84c5..21ef3077e44 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -3383,18 +3383,19 @@ class TestPKCEFunctionality: mock_redis = MagicMock() mock_in_memory = MagicMock() - with patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis): - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory): - mock_request = MagicMock(spec=Request) - mock_request.query_params = {} - token_params = ( - await SSOAuthenticationHandler.prepare_token_exchange_parameters( - request=mock_request, generic_include_client_id=False - ) + with patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory + ), patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}, clear=False): + mock_request = MagicMock(spec=Request) + mock_request.query_params = {} + token_params = ( + await SSOAuthenticationHandler.prepare_token_exchange_parameters( + request=mock_request, generic_include_client_id=False ) - assert "code_verifier" not in token_params - mock_redis.async_get_cache.assert_not_called() - mock_in_memory.async_get_cache.assert_not_called() + ) + assert "code_verifier" not in token_params + mock_redis.async_get_cache.assert_not_called() + mock_in_memory.async_get_cache.assert_not_called() @pytest.mark.asyncio