From bc633e15a7feb9040113b129c885a56a3a92f7f8 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 14 Jul 2026 19:54:35 +0000 Subject: [PATCH] fix(proxy): route CLI SSO login flow through coordination Redis for multi-worker deployments Store the pending CLI SSO login flow in redis_usage_cache when a coordination Redis is configured, falling back to the per-worker user_api_key_cache otherwise. This lets /sso/cli/start and /sso/key/generate land on different workers or replicas without losing the session, fixing the 'Invalid CLI login session' error, without changing the auth-cache semantics of user_api_key_cache. Fixes #33253 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/management_endpoints/ui_sso.py | 26 +++-- litellm/proxy/proxy_server.py | 26 ++--- .../proxy/management_endpoints/test_ui_sso.py | 106 ++++++++++++++++++ .../proxy/test_redis_auth_cache_flag.py | 73 ++---------- 4 files changed, 142 insertions(+), 89 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index d8015bb8031..f9384fe4eb2 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -240,6 +240,14 @@ def _check_cli_sso_start_rate_limit( ) +def _cli_sso_flow_cache() -> DualCache: + from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache + + if redis_usage_cache is not None: + return DualCache(redis_cache=redis_usage_cache) + return user_api_key_cache + + def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dict: if not _is_valid_cli_sso_login_id(login_id): raise HTTPException(status_code=400, detail="Invalid CLI login session") @@ -586,7 +594,7 @@ async def cli_sso_start(request: Request): "user_code_verified": False, "session_data": None, } - _set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow) + _set_cli_sso_flow(login_id=login_id, cache=_cli_sso_flow_cache(), flow=flow) verification_uri_complete: str | None = ( ( @@ -618,9 +626,8 @@ async def cli_sso_complete(request: Request, login_id: str): from litellm.proxy.common_utils.html_forms.cli_sso_success import ( render_cli_sso_success_page, ) - from litellm.proxy.proxy_server import user_api_key_cache - flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=user_api_key_cache) + flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=_cli_sso_flow_cache()) if not flow.get("sso_complete") or not flow.get("session_data"): raise HTTPException(status_code=400, detail="CLI login is not ready") @@ -644,7 +651,7 @@ async def cli_sso_complete(request: Request, login_id: str): raise HTTPException(status_code=400, detail="Invalid verification code") flow["user_code_verified"] = True - _set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow) + _set_cli_sso_flow(login_id=login_id, cache=_cli_sso_flow_cache(), flow=flow) html_content = render_cli_sso_success_page() return HTMLResponse(content=html_content, status_code=200) @@ -838,7 +845,6 @@ async def google_login( general_settings, premium_user, prisma_client, - user_api_key_cache, user_custom_ui_sso_sign_in_handler, ) @@ -886,7 +892,7 @@ async def google_login( ) if source == LITELLM_CLI_SOURCE_IDENTIFIER: - _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache) + _get_cli_sso_flow_or_raise(login_id=key, cache=_cli_sso_flow_cache()) # Store CLI login handle in state for OAuth flow cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state( @@ -1966,7 +1972,7 @@ async def _complete_cli_sso_callback_session( flow["sso_complete"] = True browser_complete_token = secrets.token_urlsafe(32) flow["browser_complete_token_hash"] = _hash_cli_sso_secret(browser_complete_token) - _set_cli_sso_flow(login_id=key, cache=user_api_key_cache, flow=flow) + _set_cli_sso_flow(login_id=key, cache=_cli_sso_flow_cache(), flow=flow) verbose_proxy_logger.info( f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}" @@ -2002,7 +2008,7 @@ async def cli_sso_callback( user_api_key_cache, ) - flow = _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache) + flow = _get_cli_sso_flow_or_raise(login_id=key, cache=_cli_sso_flow_cache()) if prisma_client is None: raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) @@ -2079,7 +2085,7 @@ async def cli_poll_key( from litellm.proxy.proxy_server import prisma_client, user_api_key_cache try: - flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache) + flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=_cli_sso_flow_cache()) if not _verify_cli_sso_poll_secret(flow=flow, poll_secret=x_litellm_cli_poll_secret): raise HTTPException(status_code=403, detail="Invalid CLI polling secret") @@ -2186,7 +2192,7 @@ async def cli_poll_key( ) # Delete cache entry (single-use) - user_api_key_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id)) + _cli_sso_flow_cache().delete_cache(key=_get_cli_sso_flow_cache_key(key_id)) verbose_proxy_logger.info(f"CLI JWT generated for user: {user_id}, team: {team_id}") poll_response = { diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 39c287d97eb..37f3d6e49e0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3634,7 +3634,7 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache: """ Wires an established coordination Redis into the proxy-level caches that consume it directly: the spend counter cache, the cluster-wide config - cache, and (unless explicitly opted out) the virtual-key auth cache. + cache, and (only when opted in) the virtual-key auth cache. """ spend_counter_cache.attach_redis_cache( redis_cache, @@ -3646,18 +3646,16 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache: default_redis_ttl=litellm.default_redis_ttl, ) verbose_proxy_logger.info( - "attached Redis to user_api_key_cache; virtual-key lookups and " - "short-lived cross-worker state (e.g. CLI SSO login sessions) are " - "now shared across all proxy workers. Set " - "litellm_settings.enable_redis_auth_cache: false to opt out and " - "keep the auth cache per-worker/DB-only." + "enable_redis_auth_cache=True: attached Redis to " + "user_api_key_cache — virtual-key lookups are now " + "shared across all proxy workers." ) else: verbose_proxy_logger.info( - "enable_redis_auth_cache is set to false: user_api_key_cache " - "remains in-memory only (per-worker). Cross-worker features that " - "rely on it (e.g. CLI SSO login) will not work on multi-worker " - "deployments." + "enable_redis_auth_cache is not set: user_api_key_cache " + "remains in-memory only (per-worker). Set " + "litellm_settings.enable_redis_auth_cache: true to share " + "the auth cache across workers and reduce DB load." ) litellm_config_cache.redis_cache = redis_cache @@ -3886,7 +3884,7 @@ class ProxyConfig: coordination_redis_cache = _build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) _attach_redis_usage_cache( coordination_redis_cache, - enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", True) is not False, + enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, ) verbose_proxy_logger.info( "coordination_redis: using a standalone Redis from general_settings " @@ -3897,7 +3895,7 @@ class ProxyConfig: def _init_cache( self, cache_params: dict, - enable_redis_auth_cache: bool = True, + enable_redis_auth_cache: bool = False, ) -> RedisCache | None: """ Initializes the response cache and resolves the coordination Redis. @@ -4271,7 +4269,7 @@ class ProxyConfig: _set_redis_usage_cache( self._init_cache( cache_params=cache_params, - enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", True) is not False, + enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, ) ) if litellm.cache is not None: @@ -7447,7 +7445,7 @@ class ProxyStartupEvent: coordination_redis_cache = _build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) _attach_redis_usage_cache( coordination_redis_cache, - enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", True) is not False, + enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, ) if llm_router is not None and llm_router.cache.redis_cache is None: llm_router._update_redis_cache(cache=coordination_redis_cache) 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 2a4e2ed6b25..28728c87010 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -7361,3 +7361,109 @@ async def test_auth_callback_without_oauth_error_proceeds_to_normal_flow(): assert exc_info.value.status_code == 500 assert "DB not connected" in str(exc_info.value.detail) + + +class _SharedFakeRedisCache: + """In-memory stand-in for a coordination Redis shared across proxy workers.""" + + def __init__(self, store): + self._store = store + + def set_cache(self, key, value, **kwargs): + self._store[key] = json.dumps(value) + return True + + def get_cache(self, key, parent_otel_span=None, **kwargs): + raw = self._store.get(key) + return None if raw is None else json.loads(raw) + + def delete_cache(self, key): + self._store.pop(key, None) + + +class TestCliSsoFlowCoordinationRedis: + """ + Regression for issue #33253: CLI SSO login fails on multi-worker / multi-replica + proxies because the pending login flow was stored in the per-worker in-memory + ``user_api_key_cache``. ``/sso/cli/start`` and ``/sso/key/generate`` can land on + different workers, so the second worker could not see the flow and returned + "Invalid CLI login session". The fix routes the flow through the coordination + Redis (``redis_usage_cache``) when it exists, falling back to ``user_api_key_cache``. + """ + + LOGIN_ID = "cli-abcdef012345" + FLOW = {"poll_secret_hash": "deadbeef", "sso_complete": False, "session_data": None} + + @pytest.mark.asyncio + async def test_cli_sso_start_persists_flow_to_coordination_redis(self): + from litellm.proxy.management_endpoints.ui_sso import ( + _get_cli_sso_flow_cache_key, + cli_sso_start, + ) + + store: dict = {} + fake_redis = _SharedFakeRedisCache(store) + per_worker_cache = MagicMock() + per_worker_cache.increment_cache.return_value = 1 + + mock_request = MagicMock(spec=Request) + mock_request.client = SimpleNamespace(host="127.0.0.1") + mock_request.headers = {} + + with ( + patch("litellm.proxy.proxy_server.redis_usage_cache", fake_redis), + patch("litellm.proxy.proxy_server.user_api_key_cache", per_worker_cache), + ): + result = await cli_sso_start(request=mock_request) + + assert _get_cli_sso_flow_cache_key(result["login_id"]) in store + per_worker_cache.set_cache.assert_not_called() + + def test_cli_sso_flow_shared_across_workers_via_redis(self): + from litellm.caching.caching import DualCache + from litellm.proxy.management_endpoints.ui_sso import ( + _get_cli_sso_flow_or_raise, + _set_cli_sso_flow, + ) + + store: dict = {} + worker_a = DualCache(redis_cache=_SharedFakeRedisCache(store)) + worker_b = DualCache(redis_cache=_SharedFakeRedisCache(store)) + + _set_cli_sso_flow(login_id=self.LOGIN_ID, cache=worker_a, flow=dict(self.FLOW)) + retrieved = _get_cli_sso_flow_or_raise(login_id=self.LOGIN_ID, cache=worker_b) + + assert retrieved["poll_secret_hash"] == "deadbeef" + + def test_cli_sso_flow_invisible_to_other_worker_without_redis(self): + from litellm.caching.caching import DualCache + from litellm.proxy.management_endpoints.ui_sso import ( + _get_cli_sso_flow_or_raise, + _set_cli_sso_flow, + ) + + worker_a = DualCache() + worker_b = DualCache() + + _set_cli_sso_flow(login_id=self.LOGIN_ID, cache=worker_a, flow=dict(self.FLOW)) + + with pytest.raises(HTTPException) as exc_info: + _get_cli_sso_flow_or_raise(login_id=self.LOGIN_ID, cache=worker_b) + assert exc_info.value.detail == "Invalid CLI login session" + + def test_cli_sso_flow_cache_prefers_coordination_redis_then_falls_back(self): + from litellm.caching.caching import DualCache + from litellm.proxy.management_endpoints.ui_sso import _cli_sso_flow_cache + + fake_redis = _SharedFakeRedisCache({}) + with patch("litellm.proxy.proxy_server.redis_usage_cache", fake_redis): + selected = _cli_sso_flow_cache() + assert isinstance(selected, DualCache) + assert selected.redis_cache is fake_redis + + per_worker_cache = DualCache() + with ( + patch("litellm.proxy.proxy_server.redis_usage_cache", None), + patch("litellm.proxy.proxy_server.user_api_key_cache", per_worker_cache), + ): + assert _cli_sso_flow_cache() is per_worker_cache diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py index 2b253f6b6ce..d0cb5ec5465 100644 --- a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -1,11 +1,8 @@ """ Tests for the enable_redis_auth_cache litellm_settings flag. -Verifies that _init_cache attaches Redis to user_api_key_cache by default -whenever a coordination Redis exists, and only leaves it in-memory-only when -the flag is explicitly set to False (opt-out). Also pins the cross-worker CLI -SSO login regression from issue #33253: with the default settings, a login -session written by one worker is readable by another. +Verifies that _init_cache attaches Redis to user_api_key_cache only when +the flag is explicitly set to True, and leaves it in-memory-only otherwise. """ from contextlib import contextmanager @@ -68,7 +65,7 @@ def _patched_init_cache(litellm_settings: dict, cache_params: dict): fresh_user_cache = DualCache() fresh_spend_cache = DualCache() - enable_redis_auth_cache = litellm_settings.get("enable_redis_auth_cache", True) is not False + enable_redis_auth_cache = litellm_settings.get("enable_redis_auth_cache", False) with ( patch.object(ps, "user_api_key_cache", fresh_user_cache), @@ -110,15 +107,15 @@ class TestRedisAuthCacheFlag: "enable_redis_auth_cache=False" ) - def test_flag_absent_attaches_redis_to_user_api_key_cache_by_default(self): - """When enable_redis_auth_cache is not set, Redis attaches by default (issue #33253).""" + def test_flag_absent_leaves_user_api_key_cache_in_memory_only(self): + """When enable_redis_auth_cache is not set at all, default is in-memory-only.""" with _patched_init_cache( litellm_settings={}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, ) as (user_cache, _): - assert user_cache.redis_cache is not None, ( - "user_api_key_cache must attach to the coordination Redis by " - "default when Redis is configured and the flag is absent" + assert user_cache.redis_cache is None, ( + "user_api_key_cache must remain in-memory-only when " + "enable_redis_auth_cache is absent from litellm_settings" ) def test_spend_counter_cache_always_gets_redis_regardless_of_flag(self): @@ -146,57 +143,3 @@ class TestRedisAuthCacheFlag: ) as (user_cache, spend_cache): assert spend_cache.redis_cache is not None assert user_cache.redis_cache is None - - -class TestCliSsoLoginCrossWorker: - """ - Regression for issue #33253: the CLI SSO login flow stores its pending - session in ``user_api_key_cache``. On a multi-worker deployment ``/sso/cli/start`` - and ``/sso/key/generate`` can land on different workers, so the session must - survive being read back from a different worker's cache instance. With the - default settings and a coordination Redis present, that now works because - ``user_api_key_cache`` shares the Redis backend across workers. - """ - - LOGIN_ID = "cli-abcdef012345" - FLOW = {"poll_secret_hash": "deadbeef", "sso_complete": False, "session_data": None} - - def _worker_cache(self, shared_redis, *, attach_redis: bool) -> DualCache: - cache = DualCache() - if attach_redis: - cache.attach_redis_cache(shared_redis) - return cache - - def test_default_flow_readable_from_other_worker(self): - from litellm.proxy.management_endpoints.ui_sso import ( - _get_cli_sso_flow_or_raise, - _set_cli_sso_flow, - ) - - shared_redis = _FakeRedisCache() - worker_a = self._worker_cache(shared_redis, attach_redis=True) - worker_b = self._worker_cache(shared_redis, attach_redis=True) - - _set_cli_sso_flow(login_id=self.LOGIN_ID, cache=worker_a, flow=dict(self.FLOW)) - - flow = _get_cli_sso_flow_or_raise(login_id=self.LOGIN_ID, cache=worker_b) - assert flow["poll_secret_hash"] == "deadbeef" - - def test_in_memory_only_flow_lost_across_workers(self): - """Without a shared Redis, the pre-fix behaviour reproduces (session lost).""" - from fastapi import HTTPException - - from litellm.proxy.management_endpoints.ui_sso import ( - _get_cli_sso_flow_or_raise, - _set_cli_sso_flow, - ) - - shared_redis = _FakeRedisCache() - worker_a = self._worker_cache(shared_redis, attach_redis=False) - worker_b = self._worker_cache(shared_redis, attach_redis=False) - - _set_cli_sso_flow(login_id=self.LOGIN_ID, cache=worker_a, flow=dict(self.FLOW)) - - with pytest.raises(HTTPException) as exc_info: - _get_cli_sso_flow_or_raise(login_id=self.LOGIN_ID, cache=worker_b) - assert exc_info.value.detail == "Invalid CLI login session"