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>
This commit is contained in:
Krrish Dholakia 2026-07-14 19:54:35 +00:00
parent 22b4d55863
commit bc633e15a7
4 changed files with 142 additions and 89 deletions

View file

@ -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 = {

View file

@ -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)

View file

@ -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

View file

@ -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"