From ee92183319a005ca9418b1fee651974ea06ba0f7 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Mon, 5 Oct 2026 21:59:50 +0000
Subject: [PATCH] fix(sso): let CLI and Claude Code gateway sign-in through on
DISABLE_ADMIN_UI nodes (#44620)
* fix(sso): let CLI and Claude Code gateway sign-in through on DISABLE_ADMIN_UI nodes
* test(proxy): cover CLI SSO sign-in on a UI-disabled node
---------
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
---
litellm/proxy/common_utils/admin_ui_utils.py | 7 +
litellm/proxy/management_endpoints/ui_sso.py | 23 +-
.../test_cli_sso_login_ui_disabled.py | 513 ++++++++++++++++++
.../proxy/management_endpoints/test_ui_sso.py | 111 ++++
4 files changed, 643 insertions(+), 11 deletions(-)
create mode 100644 tests/integration/authorization/test_cli_sso_login_ui_disabled.py
diff --git a/litellm/proxy/common_utils/admin_ui_utils.py b/litellm/proxy/common_utils/admin_ui_utils.py
index f279be36346..45b6ffae9c2 100644
--- a/litellm/proxy/common_utils/admin_ui_utils.py
+++ b/litellm/proxy/common_utils/admin_ui_utils.py
@@ -1,5 +1,12 @@
+import os
from typing import Final
+from litellm.secret_managers.main import str_to_bool
+
+
+def is_admin_ui_disabled() -> bool:
+ return bool(str_to_bool(value=os.getenv("DISABLE_ADMIN_UI")))
+
def show_missing_vars_in_env():
from fastapi.responses import HTMLResponse
diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py
index 01807fefd78..57f52ef4acf 100644
--- a/litellm/proxy/management_endpoints/ui_sso.py
+++ b/litellm/proxy/management_endpoints/ui_sso.py
@@ -100,6 +100,7 @@ from litellm.proxy.auth.team_grants import TeamModelAliasTable
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.admin_ui_utils import (
admin_ui_disabled,
+ is_admin_ui_disabled,
show_missing_vars_in_env,
)
from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint
@@ -137,7 +138,7 @@ from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import SSOConfigRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
-from litellm.secret_managers.main import get_secret_bool, get_secret_str, str_to_bool
+from litellm.secret_managers.main import get_secret_bool, get_secret_str
from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
@@ -1029,11 +1030,10 @@ async def google_login(
generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
####### Check if UI is disabled #######
- _disable_ui_flag: Final = os.getenv("DISABLE_ADMIN_UI")
- if _disable_ui_flag is not None:
- is_disabled: Final = str_to_bool(value=_disable_ui_flag)
- if is_disabled:
- return admin_ui_disabled()
+ admin_ui_is_disabled: Final = is_admin_ui_disabled()
+ is_cli_sso_login: Final = source == LITELLM_CLI_SOURCE_IDENTIFIER
+ if admin_ui_is_disabled and not is_cli_sso_login:
+ return admin_ui_disabled()
####### Check if user is a Enterprise / Premium User #######
if (
@@ -1055,7 +1055,7 @@ async def google_login(
sso_callback_route="sso/callback",
)
- if source == LITELLM_CLI_SOURCE_IDENTIFIER:
+ if is_cli_sso_login:
_get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache)
# Store CLI login handle in state for OAuth flow
@@ -1115,6 +1115,9 @@ async def google_login(
_persist_return_to_cookie(sso_redirect, return_to, request)
return sso_redirect
+ if admin_ui_is_disabled:
+ return admin_ui_disabled()
+
from fastapi.responses import HTMLResponse
hide_default_credentials_hint: Final = should_hide_default_credentials_hint(general_settings)
@@ -2158,8 +2161,7 @@ async def saml_login(request: Request, return_to: str | None = None):
"""SP-initiated SAML login. Redirects the user to the configured IdP."""
from litellm.proxy.proxy_server import user_api_key_cache
- _disable_ui_flag: Final = os.getenv("DISABLE_ADMIN_UI")
- if _disable_ui_flag is not None and str_to_bool(value=_disable_ui_flag):
+ if is_admin_ui_disabled():
return admin_ui_disabled()
return await SAMLAuthHandler.build_login_redirect(request=request, cache=user_api_key_cache, relay_state=return_to)
@@ -2186,8 +2188,7 @@ async def saml_callback(request: Request):
user_api_key_cache,
)
- _disable_ui_flag: Final = os.getenv("DISABLE_ADMIN_UI")
- if _disable_ui_flag is not None and str_to_bool(value=_disable_ui_flag):
+ if is_admin_ui_disabled():
return admin_ui_disabled()
if prisma_client is None:
diff --git a/tests/integration/authorization/test_cli_sso_login_ui_disabled.py b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py
new file mode 100644
index 00000000000..ef960853a17
--- /dev/null
+++ b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py
@@ -0,0 +1,513 @@
+from __future__ import annotations
+
+import base64
+import json
+import os
+import re
+import secrets
+import signal
+import threading
+import uuid
+from collections.abc import Iterator, Mapping, Sequence
+from concurrent.futures import ThreadPoolExecutor
+from dataclasses import dataclass
+from pathlib import Path
+from typing import Final
+from urllib.parse import parse_qs, urlencode, urlparse
+
+import httpx
+import psutil
+import pytest
+import yaml
+from pydantic import JsonValue
+
+from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, gateway_from_environment, string_value
+from tests.integration._support.database import read_rows
+from tests.integration._support.process import OwnedProxy, group_members, owned_proxy_process
+from tests.integration._support.provider import SharedProvider
+from tests.integration._support.wire import Reply, Request, Wire, wire_server
+
+pytestmark: Final = pytest.mark.timeout(300)
+
+CLI_SOURCE: Final = "litellm-cli"
+CLI_STATE_PREFIX: Final = "litellm-session-token"
+DEVICE_CODE_GRANT: Final = "urn:ietf:params:oauth:grant-type:device_code"
+POLL_SECRET_HEADER: Final = "x-litellm-cli-poll-secret"
+DISABLED_PAGE_TITLE: Final = "
Admin UI Disabled"
+LOGIN_FORM_TITLE: Final = "LiteLLM Login"
+CLI_LOGIN_PAGE_TITLE: Final = "LiteLLM CLI Login"
+CLI_SUCCESS_PAGE_TITLE: Final = "CLI Authentication Successful - LiteLLM"
+SESSION_GONE: Final = "CLI login session not found or expired"
+SUBJECT: Final = "cli-sso-audit-subject"
+SUBJECT_EMAIL: Final = "cli-sso-audit-subject@example.com"
+CLIENT_ID: Final = "integration-oidc-client"
+CLIENT_SECRET: Final = "integration-oidc-secret"
+MESSAGE_MODEL: Final = "anthropic/claude-haiku-4-5"
+BURST: Final = 8
+COMPLETE_TOKEN_FIELD: Final = re.compile(r'name="browser_complete_token" value="([^"]+)"')
+COMPLETE_FORM_ACTION: Final = re.compile(r'action="([^"]+/sso/cli/complete/[^"]+)"')
+
+
+@dataclass(frozen=True, slots=True)
+class Idp:
+ wire: Wire
+ token_outage: threading.Event
+
+
+@dataclass(frozen=True, slots=True)
+class CliSession:
+ login_id: str
+ poll_secret: str
+ user_code: str
+
+
+@dataclass(frozen=True, slots=True)
+class DeviceGrant:
+ device_code: str
+ user_code: str
+ verification_uri: str
+
+
+def _idp_reply(request: Request, token_outage: threading.Event) -> Reply:
+ target: Final = urlparse(request.target)
+ if target.path == "/authorize":
+ query: Final = parse_qs(target.query)
+ location: Final = (
+ f"{query['redirect_uri'][0]}?{urlencode({'code': f'code-{uuid.uuid4().hex}', 'state': query['state'][0]})}"
+ )
+ return Reply(status=302, body=b"", headers={"location": location})
+ if target.path == "/token":
+ if token_outage.is_set():
+ return Reply(status=503, body=b'{"error": "temporarily_unavailable"}')
+ expected: Final = base64.b64encode(f"{CLIENT_ID}:{CLIENT_SECRET}".encode()).decode()
+ if request.headers.get("authorization") != f"Basic {expected}":
+ return Reply(status=401, body=b'{"error": "invalid_client"}')
+ form: Final = parse_qs(request.body.decode())
+ if form.get("grant_type") != ["authorization_code"] or not form.get("code", [""])[0].startswith("code-"):
+ return Reply(status=400, body=b'{"error": "invalid_grant"}')
+ token: Final = {"access_token": f"idp-access-{uuid.uuid4().hex}", "token_type": "Bearer", "expires_in": 3600}
+ return Reply(body=json.dumps(token).encode())
+ if target.path == "/userinfo":
+ if not request.headers.get("authorization", "").startswith("Bearer idp-access-"):
+ return Reply(status=401, body=b'{"error": "invalid_token"}')
+ return Reply(body=json.dumps({"sub": SUBJECT, "preferred_username": SUBJECT, "email": SUBJECT_EMAIL}).encode())
+ return Reply(status=404, body=b'{"error": "not_found"}')
+
+
+def _sso_environment(idp_url: str) -> Mapping[str, str]:
+ return {
+ "GENERIC_CLIENT_ID": CLIENT_ID,
+ "GENERIC_CLIENT_SECRET": CLIENT_SECRET,
+ "GENERIC_AUTHORIZATION_ENDPOINT": f"{idp_url}/authorize",
+ "GENERIC_TOKEN_ENDPOINT": f"{idp_url}/token",
+ "GENERIC_USERINFO_ENDPOINT": f"{idp_url}/userinfo",
+ "OAUTHLIB_INSECURE_TRANSPORT": "1",
+ }
+
+
+def _gateway_enabled_config(directory: Path) -> Path:
+ stock: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text())
+ config: Final = directory / f"cli_sso_{uuid.uuid4().hex}.yaml"
+ config.write_text(
+ json.dumps({**stock, "general_settings": {**stock["general_settings"], "enable_claude_code_gateway": True}})
+ )
+ return config
+
+
+@pytest.fixture(scope="module")
+def idp() -> Iterator[Idp]:
+ outage: Final = threading.Event()
+ with wire_server(lambda request: _idp_reply(request, outage)) as wire:
+ yield Idp(wire, outage)
+
+
+@pytest.fixture(scope="module")
+def ui_disabled(idp: Idp, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedProxy]:
+ directory: Final = tmp_path_factory.mktemp("cli-sso-ui-disabled")
+ with (
+ gateway_from_environment() as rig,
+ owned_proxy_process(
+ rig,
+ directory,
+ {"DISABLE_ADMIN_UI": "true", **_sso_environment(idp.wire.url)},
+ config=_gateway_enabled_config(directory),
+ remove_environment=("PROXY_BASE_URL",),
+ workers=2,
+ ) as owned,
+ ):
+ yield owned
+
+
+def _proxy_url(proxy: Gateway) -> str:
+ return str(proxy.client.base_url).rstrip("/")
+
+
+def _browser() -> httpx.Client:
+ return httpx.Client(follow_redirects=False, trust_env=False, timeout=30)
+
+
+def _start_lite_login(proxy: Gateway) -> CliSession:
+ response: Final = proxy.client.post("/sso/cli/start")
+ assert response.status_code == 200, response.text
+ body: Final = JSON_OBJECT.validate_json(response.content)
+ return CliSession(
+ string_value(body["login_id"]), string_value(body["poll_secret"]), string_value(body["user_code"])
+ )
+
+
+def _cli_link(proxy: Gateway, login_id: str) -> str:
+ return f"{_proxy_url(proxy)}/sso/key/generate?{urlencode({'source': CLI_SOURCE, 'key': login_id})}"
+
+
+def _assert_idp_redirect(proxy: Gateway, idp: Idp, link: httpx.Response, login_id: str) -> str:
+ assert DISABLED_PAGE_TITLE not in link.text, link.text
+ assert link.is_redirect, f"{link.status_code} {link.text}"
+ location: Final = link.headers["location"]
+ parsed: Final = urlparse(location)
+ query: Final = parse_qs(parsed.query)
+ assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == f"{idp.wire.url}/authorize", location
+ assert query["state"] == [f"{CLI_STATE_PREFIX}:{login_id}"], location
+ assert query["redirect_uri"] == [f"{_proxy_url(proxy)}/sso/callback"], location
+ assert query["client_id"] == [CLIENT_ID], location
+ return location
+
+
+def _walk_idp(browser: httpx.Client, authorize_url: str) -> httpx.Response:
+ at_idp: Final = browser.get(authorize_url)
+ assert at_idp.status_code == 302, f"{at_idp.status_code} {at_idp.text}"
+ return browser.get(at_idp.headers["location"])
+
+
+def _complete_in_browser(browser: httpx.Client, callback: httpx.Response, user_code: str) -> httpx.Response:
+ assert callback.status_code == 200, f"{callback.status_code} {callback.text}"
+ assert CLI_LOGIN_PAGE_TITLE in callback.text, callback.text
+ token: Final = COMPLETE_TOKEN_FIELD.search(callback.text)
+ action: Final = COMPLETE_FORM_ACTION.search(callback.text)
+ assert token is not None and action is not None, callback.text
+ return browser.post(action.group(1), data={"user_code": user_code, "browser_complete_token": token.group(1)})
+
+
+def _poll(proxy: Gateway, session: CliSession, *, poll_secret: str | None = None) -> httpx.Response:
+ return proxy.client.get(
+ f"/sso/cli/poll/{session.login_id}",
+ headers={POLL_SECRET_HEADER: session.poll_secret if poll_secret is None else poll_secret},
+ )
+
+
+def _sign_in(proxy: Gateway, idp: Idp, browser: httpx.Client, session: CliSession) -> None:
+ link: Final = browser.get(_cli_link(proxy, session.login_id))
+ callback: Final = _walk_idp(browser, _assert_idp_redirect(proxy, idp, link, session.login_id))
+ done: Final = _complete_in_browser(browser, callback, session.user_code)
+ assert done.status_code == 200 and CLI_SUCCESS_PAGE_TITLE in done.text, f"{done.status_code} {done.text}"
+
+
+def _ready_key(proxy: Gateway, session: CliSession) -> str:
+ ready: Final = _poll(proxy, session)
+ assert ready.status_code == 200, f"{ready.status_code} {ready.text}"
+ body: Final = JSON_OBJECT.validate_json(ready.content)
+ assert body["status"] == "ready" and body["user_id"] == SUBJECT, ready.text
+ return string_value(body["key"])
+
+
+def _send_message(proxy: Gateway, provider: SharedProvider, key: str) -> str:
+ message_id: Final = f"msg_{uuid.uuid4().hex}"
+ marker: Final = f"audit-{uuid.uuid4().hex}"
+ provider.expect(
+ Reply(
+ body=json.dumps(
+ {
+ "id": message_id,
+ "type": "message",
+ "role": "assistant",
+ "model": "claude-haiku-4-5",
+ "content": [{"type": "text", "text": "scripted reply"}],
+ "stop_reason": "end_turn",
+ "usage": {"input_tokens": 9, "output_tokens": 5},
+ }
+ ).encode()
+ )
+ )
+ response: Final = proxy.request(
+ "POST",
+ "/v1/messages",
+ {"model": MESSAGE_MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": marker}]},
+ key=key,
+ )
+ assert response.status_code == 200, response.text
+ assert JSON_OBJECT.validate_json(response.content)["id"] == message_id, response.text
+ upstream: Final = provider.received()
+ assert len(upstream) == 1 and upstream[0].target == "/v1/messages", [item.target for item in upstream]
+ assert marker in upstream[0].body.decode(), upstream[0].body
+ return message_id
+
+
+def _user_rows() -> Sequence[Mapping[str, JsonValue]]:
+ return read_rows('SELECT user_id, user_email FROM "LiteLLM_UserTable" WHERE user_id = %s', (SUBJECT,))
+
+
+def _worker_pids(root_pid: int) -> frozenset[int]:
+ def is_worker(process: psutil.Process) -> bool:
+ try:
+ return "spawn_main" in " ".join(process.cmdline())
+ except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess):
+ return False
+
+ return frozenset(process.pid for process in group_members(root_pid) if is_worker(process))
+
+
+def test_cli_login_link_redirects_to_the_idp_when_the_ui_is_disabled(ui_disabled: OwnedProxy, idp: Idp) -> None:
+ proxy: Final = ui_disabled.gateway
+ session: Final = _start_lite_login(proxy)
+ with _browser() as browser:
+ link: Final = browser.get(_cli_link(proxy, session.login_id))
+ _assert_idp_redirect(proxy, idp, link, session.login_id)
+
+
+def test_lite_login_completes_and_the_key_serves_messages(
+ ui_disabled: OwnedProxy, idp: Idp, provider: SharedProvider
+) -> None:
+ proxy: Final = ui_disabled.gateway
+ session: Final = _start_lite_login(proxy)
+ with _browser() as browser:
+ early: Final = browser.post(
+ f"{_proxy_url(proxy)}/sso/cli/complete/{session.login_id}",
+ data={"user_code": session.user_code, "browser_complete_token": "x"},
+ )
+ assert early.status_code == 400 and "CLI login is not ready" in early.text, f"{early.status_code} {early.text}"
+ assert JSON_OBJECT.validate_json(_poll(proxy, session).content) == {"status": "pending"}
+ _sign_in(proxy, idp, browser, session)
+ forged: Final = _poll(proxy, session, poll_secret="not-the-poll-secret")
+ assert forged.status_code == 403, f"{forged.status_code} {forged.text}"
+ key: Final = _ready_key(proxy, session)
+ users: Final = _user_rows()
+ assert [(row["user_id"], row["user_email"]) for row in users] == [(SUBJECT, SUBJECT_EMAIL)], users
+ message_id: Final = _send_message(proxy, provider, key)
+ spend: Final = eventually(
+ lambda: read_rows('SELECT "user" FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (message_id,)),
+ lambda rows: len(rows) == 1,
+ seconds=70,
+ )
+ assert spend[0]["user"] == SUBJECT, spend
+
+
+def _device_authorization(proxy: Gateway) -> DeviceGrant:
+ response: Final = proxy.client.post("/claude_code_gateway/oauth/device_authorization")
+ assert response.status_code == 200, response.text
+ body: Final = JSON_OBJECT.validate_json(response.content)
+ return DeviceGrant(
+ string_value(body["device_code"]), string_value(body["user_code"]), string_value(body["verification_uri"])
+ )
+
+
+def _device_token(proxy: Gateway, grant: DeviceGrant) -> httpx.Response:
+ return proxy.client.post(
+ "/claude_code_gateway/oauth/token", data={"grant_type": DEVICE_CODE_GRANT, "device_code": grant.device_code}
+ )
+
+
+def test_claude_code_device_flow_completes_when_the_ui_is_disabled(
+ ui_disabled: OwnedProxy, idp: Idp, provider: SharedProvider
+) -> None:
+ proxy: Final = ui_disabled.gateway
+ grant: Final = _device_authorization(proxy)
+ login_id: Final = parse_qs(urlparse(grant.verification_uri).query)["key"][0]
+ pending: Final = _device_token(proxy, grant)
+ assert pending.status_code == 400, f"{pending.status_code} {pending.text}"
+ assert JSON_OBJECT.validate_json(pending.content)["error"] == "authorization_pending", pending.text
+ with _browser() as browser:
+ link: Final = browser.get(grant.verification_uri)
+ callback: Final = _walk_idp(browser, _assert_idp_redirect(proxy, idp, link, login_id))
+ done: Final = _complete_in_browser(browser, callback, grant.user_code)
+ assert done.status_code == 200 and CLI_SUCCESS_PAGE_TITLE in done.text, f"{done.status_code} {done.text}"
+ issued: Final = _device_token(proxy, grant)
+ assert issued.status_code == 200, f"{issued.status_code} {issued.text}"
+ access_token: Final = string_value(JSON_OBJECT.validate_json(issued.content)["access_token"])
+ replay: Final = _device_token(proxy, grant)
+ assert replay.status_code == 400, f"{replay.status_code} {replay.text}"
+ assert JSON_OBJECT.validate_json(replay.content)["error"] == "expired_token", replay.text
+ _send_message(proxy, provider, access_token)
+
+
+@pytest.mark.parametrize(
+ ("method", "path", "params"),
+ (
+ ("GET", "/sso/key/generate", ()),
+ ("GET", "/sso/key/generate", (("source", "1"),)),
+ ("GET", "/sso/key/generate", (("source", ""),)),
+ ("GET", "/sso/key/generate", (("source", "LITELLM-CLI"), ("key", f"cli-{'a' * 32}"))),
+ ("GET", "/sso/key/generate", (("return_to", "/mcp/"),)),
+ ("GET", "/sso/saml/login", ()),
+ ("POST", "/sso/saml/callback", ()),
+ ),
+ ids=("no-source", "source-int", "source-empty", "source-case", "return-to", "saml-login", "saml-callback"),
+)
+def test_non_cli_entries_stay_refused_when_the_ui_is_disabled(
+ ui_disabled: OwnedProxy, method: str, path: str, params: tuple[tuple[str, str], ...]
+) -> None:
+ response: Final = ui_disabled.gateway.client.request(method, path, params=params)
+ assert response.status_code == 200 and DISABLED_PAGE_TITLE in response.text, (
+ f"{response.status_code} {response.text}"
+ )
+
+
+@pytest.mark.parametrize(
+ ("params", "detail"),
+ (
+ ((("source", CLI_SOURCE),), "Invalid CLI login session id"),
+ ((("source", CLI_SOURCE), ("key", "1")), "Invalid CLI login session id"),
+ ((("source", CLI_SOURCE), ("key", "cli-" + "k" * 5000)), "Invalid CLI login session id"),
+ ((("source", CLI_SOURCE), ("key", f"cli-{'a' * 32}"), ("key", f"cli-{'b' * 32}")), SESSION_GONE),
+ ((("source", CLI_SOURCE), ("key", "sk-legacy-cli-key")), "Your litellm CLI is out of date"),
+ ((("source", CLI_SOURCE), ("key", f"cli-{secrets.token_urlsafe(24)}")), SESSION_GONE),
+ ((("source", CLI_SOURCE), ("source", CLI_SOURCE), ("key", f"cli-{secrets.token_urlsafe(24)}")), SESSION_GONE),
+ ),
+ ids=("missing", "int", "five-kb", "twice", "legacy-sk", "unknown", "source-twice"),
+)
+def test_malformed_cli_keys_answer_400_when_the_ui_is_disabled(
+ ui_disabled: OwnedProxy, params: tuple[tuple[str, str], ...], detail: str
+) -> None:
+ response: Final = ui_disabled.gateway.client.get("/sso/key/generate", params=params)
+ assert response.status_code == 400, f"{response.status_code} {response.text}"
+ assert DISABLED_PAGE_TITLE not in response.text, response.text
+ assert detail in string_value(JSON_OBJECT.validate_json(response.content)["detail"]), response.text
+
+
+def test_login_form_and_cli_validation_when_the_flag_is_unset(gateway: Gateway) -> None:
+ form: Final = gateway.client.get("/sso/key/generate")
+ assert form.status_code == 200 and LOGIN_FORM_TITLE in form.text, f"{form.status_code} {form.text}"
+ unknown: Final = gateway.client.get(
+ "/sso/key/generate", params={"source": CLI_SOURCE, "key": f"cli-{secrets.token_urlsafe(24)}"}
+ )
+ assert unknown.status_code == 400 and SESSION_GONE in unknown.text, f"{unknown.status_code} {unknown.text}"
+ for method, path in (("GET", "/sso/saml/login"), ("POST", "/sso/saml/callback")):
+ saml: Final = gateway.client.request(method, path)
+ assert DISABLED_PAGE_TITLE not in saml.text and saml.status_code != 200, (
+ f"{path}: {saml.status_code} {saml.text}"
+ )
+
+
+def test_cli_link_is_reentrant_until_the_poll_consumes_the_session(ui_disabled: OwnedProxy, idp: Idp) -> None:
+ proxy: Final = ui_disabled.gateway
+ session: Final = _start_lite_login(proxy)
+ with _browser() as browser:
+ first: Final = _assert_idp_redirect(
+ proxy, idp, browser.get(_cli_link(proxy, session.login_id)), session.login_id
+ )
+ second: Final = _assert_idp_redirect(
+ proxy, idp, browser.get(_cli_link(proxy, session.login_id)), session.login_id
+ )
+ assert parse_qs(urlparse(first).query)["state"] == parse_qs(urlparse(second).query)["state"]
+ done: Final = _complete_in_browser(browser, _walk_idp(browser, second), session.user_code)
+ assert done.status_code == 200, f"{done.status_code} {done.text}"
+ _ready_key(proxy, session)
+ gone: Final = browser.get(_cli_link(proxy, session.login_id))
+ assert gone.status_code == 400 and SESSION_GONE in gone.text, f"{gone.status_code} {gone.text}"
+ consumed: Final = _poll(proxy, session)
+ assert consumed.status_code == 400 and SESSION_GONE in consumed.text, f"{consumed.status_code} {consumed.text}"
+
+
+def test_burst_of_logins_recovers_from_an_idp_token_outage(ui_disabled: OwnedProxy, idp: Idp) -> None:
+ proxy: Final = ui_disabled.gateway
+ sessions: Final = tuple(_start_lite_login(proxy) for _ in range(BURST))
+
+ def failed_exchange(session: CliSession) -> httpx.Response:
+ with _browser() as browser:
+ link: Final = browser.get(_cli_link(proxy, session.login_id))
+ return _walk_idp(browser, _assert_idp_redirect(proxy, idp, link, session.login_id))
+
+ def liveliness() -> int:
+ return proxy.client.get("/health/liveliness").status_code
+
+ idp.wire.drain()
+ idp.token_outage.set()
+ try:
+ with ThreadPoolExecutor(max_workers=BURST + 1) as pool:
+ probe: Final = pool.submit(liveliness)
+ failures: Final = tuple(pool.map(failed_exchange, sessions))
+ assert probe.result() == 200
+ finally:
+ idp.token_outage.clear()
+ token_attempts: Final = tuple(request for request in idp.wire.drain() if request.target == "/token")
+ assert len(token_attempts) == BURST, len(token_attempts)
+ for failure in failures:
+ assert failure.status_code >= 400, f"{failure.status_code} {failure.text!r}"
+ assert CLI_LOGIN_PAGE_TITLE not in failure.text, failure.text
+ for session in sessions:
+ assert JSON_OBJECT.validate_json(_poll(proxy, session).content) == {"status": "pending"}, session.login_id
+
+ def recovered_login(session: CliSession) -> str:
+ with _browser() as browser:
+ _sign_in(proxy, idp, browser, session)
+ return _ready_key(proxy, session)
+
+ with ThreadPoolExecutor(max_workers=BURST) as pool:
+ keys: Final = tuple(pool.map(recovered_login, sessions))
+ assert len(set(keys)) == BURST, keys
+ for session in sessions:
+ again: Final = _poll(proxy, session)
+ assert again.status_code == 400 and SESSION_GONE in again.text, f"{again.status_code} {again.text}"
+ assert [row["user_id"] for row in _user_rows()] == [SUBJECT]
+
+
+def test_login_survives_a_worker_kill(idp: Idp, tmp_path: Path) -> None:
+ with (
+ gateway_from_environment() as rig,
+ owned_proxy_process(
+ rig,
+ tmp_path,
+ {"DISABLE_ADMIN_UI": "true", **_sso_environment(idp.wire.url)},
+ remove_environment=("PROXY_BASE_URL",),
+ workers=2,
+ ) as owned,
+ ):
+ proxy: Final = owned.gateway
+ workers: Final = eventually(lambda: _worker_pids(owned.process.pid), lambda pids: len(pids) == 2, seconds=30)
+ victim: Final = min(workers)
+ session: Final = _start_lite_login(proxy)
+ with _browser() as browser:
+ link: Final = browser.get(_cli_link(proxy, session.login_id))
+ at_idp: Final = browser.get(_assert_idp_redirect(proxy, idp, link, session.login_id))
+ assert at_idp.status_code == 302, f"{at_idp.status_code} {at_idp.text}"
+ os.kill(victim, signal.SIGKILL)
+
+ def callback_attempt() -> httpx.Response | None:
+ try:
+ return browser.get(at_idp.headers["location"])
+ except httpx.TransportError:
+ return None
+
+ callback: Final = eventually(
+ callback_attempt, lambda response: response is not None and response.status_code == 200, seconds=60
+ )
+ assert callback is not None
+ done: Final = _complete_in_browser(browser, callback, session.user_code)
+ assert done.status_code == 200 and CLI_SUCCESS_PAGE_TITLE in done.text, f"{done.status_code} {done.text}"
+ _ready_key(proxy, session)
+ respawned: Final = eventually(
+ lambda: _worker_pids(owned.process.pid), lambda pids: len(pids) == 2 and victim not in pids, seconds=60
+ )
+ assert victim not in respawned, respawned
+
+
+@pytest.mark.parametrize("flag", ("false", ""), ids=("false", "empty"))
+def test_flag_values_that_do_not_disable_keep_the_gates_open(idp: Idp, tmp_path: Path, flag: str) -> None:
+ with (
+ gateway_from_environment() as rig,
+ owned_proxy_process(
+ rig,
+ tmp_path,
+ {"DISABLE_ADMIN_UI": flag, **_sso_environment(idp.wire.url)},
+ remove_environment=("PROXY_BASE_URL",),
+ workers=2,
+ ) as owned,
+ ):
+ proxy: Final = owned.gateway
+ form_entry: Final = proxy.client.get("/sso/key/generate")
+ assert form_entry.is_redirect, f"{form_entry.status_code} {form_entry.text}"
+ assert urlparse(form_entry.headers["location"]).path == "/authorize", form_entry.headers["location"]
+ session: Final = _start_lite_login(proxy)
+ with _browser() as browser:
+ _assert_idp_redirect(proxy, idp, browser.get(_cli_link(proxy, session.login_id)), session.login_id)
+ for method, path in (("GET", "/sso/saml/login"), ("POST", "/sso/saml/callback")):
+ saml: Final = proxy.client.request(method, path)
+ assert DISABLED_PAGE_TITLE not in saml.text and saml.status_code != 200, f"{path}: {saml.status_code}"
diff --git a/tests/unit/proxy/management_endpoints/test_ui_sso.py b/tests/unit/proxy/management_endpoints/test_ui_sso.py
index 8ff0b24982f..65426b8d9cb 100644
--- a/tests/unit/proxy/management_endpoints/test_ui_sso.py
+++ b/tests/unit/proxy/management_endpoints/test_ui_sso.py
@@ -9659,3 +9659,114 @@ async def test_cli_sign_in_enrolls_only_verified_subjects_before_completing(
)
else:
table.upsert.assert_not_awaited()
+
+
+_GOOGLE_DISCOVERY_DOCUMENT = {
+ "authorization_endpoint": "https://accounts.google.com/o/oauth2/v2/auth",
+ "token_endpoint": "https://oauth2.googleapis.com/token",
+ "userinfo_endpoint": "https://openidconnect.googleapis.com/v1/userinfo",
+}
+
+
+async def _sso_key_generate_on_ui_disabled_node(*, source, key, google_sso_configured, known_login_ids):
+ """Drives GET /sso/key/generate on a node running with DISABLE_ADMIN_UI=true, the worker
+ shape of a control plane deployment, with a real Google redirect builder behind a mocked
+ discovery document."""
+ from litellm.proxy.management_endpoints.ui_sso import _get_cli_sso_flow_cache_key, google_login
+
+ env_without_sso_providers = {name: value for name, value in os.environ.items() if name not in _SSO_PROVIDER_ENV_VARS}
+ env = {
+ **env_without_sso_providers,
+ "DISABLE_ADMIN_UI": "true",
+ "PROXY_BASE_URL": "https://worker.example.com",
+ **(
+ {"GOOGLE_CLIENT_ID": "google-client-id", "GOOGLE_CLIENT_SECRET": "google-client-secret"}
+ if google_sso_configured
+ else {}
+ ),
+ }
+ flows = {_get_cli_sso_flow_cache_key(login_id): {"poll_secret_hash": "h"} for login_id in known_login_ids}
+ cli_cache = MagicMock(redis_cache=None)
+ cli_cache.get_cache.side_effect = lambda key: flows.get(key)
+ mock_request = MagicMock(spec=Request)
+ mock_request.base_url = "https://worker.example.com/"
+ mock_request.url.scheme = "https"
+ mock_request.cookies = {}
+
+ with (
+ patch.dict(os.environ, env, clear=True),
+ patch("litellm.proxy.proxy_server.premium_user", True),
+ patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
+ patch("litellm.proxy.proxy_server.master_key", "sk-1234"),
+ patch("litellm.proxy.proxy_server.general_settings", {}),
+ patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
+ patch("litellm.proxy.proxy_server.cli_sso_session_cache", cli_cache),
+ patch("litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", None),
+ patch("litellm.proxy.management_endpoints.ui_sso.show_missing_vars_in_env", return_value=None),
+ respx.mock(assert_all_called=False) as router,
+ ):
+ router.get("https://accounts.google.com/.well-known/openid-configuration").mock(
+ return_value=httpx.Response(200, json=_GOOGLE_DISCOVERY_DOCUMENT)
+ )
+ return await google_login(request=mock_request, source=source, key=key)
+
+
+@pytest.mark.asyncio
+async def test_cli_sso_login_reaches_the_idp_on_a_ui_disabled_node():
+ """Regression: a Claude Code gateway or `lite login` sign-in whose verification link lands on a
+ worker running DISABLE_ADMIN_UI=true used to get the "Admin UI is Disabled" page instead of the
+ IdP redirect, so sign-in never completed off the admin node."""
+ from urllib.parse import parse_qs, urlparse
+
+ from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
+
+ login_id = "cli-worker-login-session-0001"
+
+ response = await _sso_key_generate_on_ui_disabled_node(
+ source="litellm-cli", key=login_id, google_sso_configured=True, known_login_ids=(login_id,)
+ )
+
+ assert response.status_code == 303
+ location = urlparse(response.headers["location"])
+ assert location.hostname == "accounts.google.com"
+ query = parse_qs(location.query)
+ assert query["state"] == [f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{login_id}"]
+ assert query["redirect_uri"] == ["https://worker.example.com/sso/callback"]
+
+
+@pytest.mark.asyncio
+async def test_admin_ui_login_stays_refused_on_a_ui_disabled_node():
+ """The gate still covers the admin UI: the same SSO-configured worker refuses a plain UI login."""
+ response = await _sso_key_generate_on_ui_disabled_node(
+ source=None, key=None, google_sso_configured=True, known_login_ids=()
+ )
+
+ assert response.status_code == 200
+ assert "Admin UI is Disabled" in response.body.decode()
+
+
+@pytest.mark.asyncio
+async def test_cli_sso_login_with_an_unknown_session_is_rejected_on_a_ui_disabled_node():
+ """Only a login session the proxy issued passes the gate; a made-up key is refused before any redirect."""
+ with pytest.raises(HTTPException) as exc:
+ await _sso_key_generate_on_ui_disabled_node(
+ source="litellm-cli", key="cli-never-issued-session-00", google_sso_configured=True, known_login_ids=()
+ )
+
+ assert exc.value.status_code == 400
+
+
+@pytest.mark.asyncio
+async def test_cli_sso_login_never_serves_the_admin_login_form_on_a_ui_disabled_node():
+ """Without an SSO provider the endpoint falls back to the admin username/password form, which a
+ UI-disabled node must not serve even to a valid CLI login session."""
+ login_id = "cli-worker-login-session-0002"
+
+ response = await _sso_key_generate_on_ui_disabled_node(
+ source="litellm-cli", key=login_id, google_sso_configured=False, known_login_ids=(login_id,)
+ )
+
+ assert response.status_code == 200
+ body = response.body.decode()
+ assert "Admin UI is Disabled" in body
+ assert 'name="username"' not in body