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