mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
a3e15774ad
commit
ee92183319
4 changed files with 643 additions and 11 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue