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>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 21:59:50 +00:00 • committed by GitHub
parent a3e15774ad
commit ee92183319
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 643 additions and 11 deletions

View file

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

View file

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

View file

@ -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 = "<title>Admin UI Disabled</title>"
LOGIN_FORM_TITLE: Final = "<title>LiteLLM Login</title>"
CLI_LOGIN_PAGE_TITLE: Final = "<title>LiteLLM CLI Login</title>"
CLI_SUCCESS_PAGE_TITLE: Final = "<title>CLI Authentication Successful - LiteLLM</title>"
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}"

View file

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