mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
22b4d55863
commit
bc633e15a7
4 changed files with 142 additions and 89 deletions
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue