From 91e9b1f06bf3f7021e671450c59977cbbf8b293e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 7 Oct 2026 06:45:41 -0700 Subject: [PATCH] fix(auth): clear the recent-miss user memo when /user/new creates the user (#45020) * fix(auth): clear the recent-miss user memo when /user/new creates the user Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(auth): pin the miss memo window and drop the class patch in the new_user regression test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): pin a first-time SSO user's first message and second sign-in on one worker * test(integration): cover /user/new clearing the user-miss memo on the creating worker The fake IdP now signs in the subject a login_hint names, so cells on the shared one-worker proxy can each use a fresh user. New cells: a plain-key miss followed by /user/new is budgeted at once (same on both legs, the auth prefetch loads the row) and the admin-created user's first SSO sign-in inside the window completes (500 at the callback before the fix). The first-sign-in cell moved onto the shared one-worker proxy fixture --------- Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 10 +- .../internal_user_endpoints.py | 4 + .../test_cli_sso_login_ui_disabled.py | 169 ++++++++++++++++-- .../test_internal_user_endpoints.py | 78 ++++++++ 4 files changed, 244 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3848f5965bf..bea62f40c76 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2583,6 +2583,14 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c return False +def _user_db_access_key(user_id: str) -> str: + return f"user_id:{user_id}" + + +def forget_missing_user(user_id: str) -> None: + last_db_access_time.pop(_user_db_access_key(user_id), None) + + def _update_last_db_access_time(key: str, value: object | None, last_db_access_time: LimitedSizeOrderedDict): last_db_access_time[key] = (value, time.time()) @@ -2751,7 +2759,7 @@ async def get_user_object( if prisma_client is None: raise Exception("No db connected") try: - db_access_time_key: Final = f"user_id:{user_id}" + db_access_time_key: Final = _user_db_access_key(user_id) should_check_db: Final = bool(check_db_only) or _should_check_db( key=db_access_time_key, last_db_access_time=last_db_access_time, diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 98df4ae46b1..e04d92a5398 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -33,6 +33,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import ( delete_cache_key_objects, + forget_missing_user, get_jwt_key_mapping_cache_keys_for_tokens, get_team_object, get_user_object, @@ -609,6 +610,9 @@ async def new_user( organization_ids: Final = cast(list[str] | None, data_json.pop("organizations", None)) response: Final = await generate_key_helper_fn(request_type="user", **data_json, llm_router=None) + created_user_id: Final = cast(str | None, response.get("user_id", None)) + if created_user_id is not None: + forget_missing_user(created_user_id) # Admin UI Logic # Add User to Team and Organization # if team_id passed add this user to the team diff --git a/tests/integration/authorization/test_cli_sso_login_ui_disabled.py b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py index ef960853a17..126fe27bef2 100644 --- a/tests/integration/authorization/test_cli_sso_login_ui_disabled.py +++ b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py @@ -10,12 +10,14 @@ import threading import uuid from collections.abc import Iterator, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path from typing import Final from urllib.parse import parse_qs, urlencode, urlparse import httpx +import jwt import psutil import pytest import yaml @@ -52,6 +54,7 @@ COMPLETE_FORM_ACTION: Final = re.compile(r'action="([^"]+/sso/cli/complete/[^"]+ class Idp: wire: Wire token_outage: threading.Event + subject: str @dataclass(frozen=True, slots=True) @@ -68,13 +71,22 @@ class DeviceGrant: verification_uri: str -def _idp_reply(request: Request, token_outage: threading.Event) -> Reply: +@dataclass(frozen=True, slots=True) +class OneWorkerProxy: + idp: Idp + owned: OwnedProxy + + +def _subject_inside(value: str, prefix: str) -> str: + return value.removeprefix(prefix).rsplit("-", 1)[0] + + +def _idp_reply(request: Request, token_outage: threading.Event, subject: str) -> 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]})}" - ) + issued: Final = f"code-{query.get('login_hint', [subject])[0]}-{uuid.uuid4().hex}" + location: Final = f"{query['redirect_uri'][0]}?{urlencode({'code': issued, 'state': query['state'][0]})}" return Reply(status=302, body=b"", headers={"location": location}) if target.path == "/token": if token_outage.is_set(): @@ -83,14 +95,23 @@ def _idp_reply(request: Request, token_outage: threading.Event) -> Reply: 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-"): + code: Final = form.get("code", [""])[0] + if form.get("grant_type") != ["authorization_code"] or not code.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()) + access_token: Final = f"idp-access-{_subject_inside(code, 'code-')}-{uuid.uuid4().hex}" + return Reply( + body=json.dumps({"access_token": access_token, "token_type": "Bearer", "expires_in": 3600}).encode() + ) if target.path == "/userinfo": - if not request.headers.get("authorization", "").startswith("Bearer idp-access-"): + bearer: Final = request.headers.get("authorization", "") + if not bearer.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()) + signed_in: Final = _subject_inside(bearer, "Bearer idp-access-") + return Reply( + body=json.dumps( + {"sub": signed_in, "preferred_username": signed_in, "email": f"{signed_in}@example.com"} + ).encode() + ) return Reply(status=404, body=b'{"error": "not_found"}') @@ -114,11 +135,17 @@ def _gateway_enabled_config(directory: Path) -> Path: return config +@contextmanager +def _fake_idp(subject: str) -> Iterator[Idp]: + outage: Final = threading.Event() + with wire_server(lambda request: _idp_reply(request, outage, subject)) as wire: + yield Idp(wire, outage, subject) + + @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) + with _fake_idp(SUBJECT) as fake: + yield fake @pytest.fixture(scope="module") @@ -138,6 +165,23 @@ def ui_disabled(idp: Idp, tmp_path_factory: pytest.TempPathFactory) -> Iterator[ yield owned +@pytest.fixture(scope="module") +def one_worker(tmp_path_factory: pytest.TempPathFactory) -> Iterator[OneWorkerProxy]: + directory: Final = tmp_path_factory.mktemp("cli-sso-one-worker") + with ( + _fake_idp(f"cli-sso-first-sign-in-{uuid.uuid4().hex[:12]}") as idp, + gateway_from_environment() as rig, + owned_proxy_process( + rig, + directory, + {"DISABLE_ADMIN_UI": "true", **_sso_environment(idp.wire.url)}, + remove_environment=("PROXY_BASE_URL",), + workers=1, + ) as owned, + ): + yield OneWorkerProxy(idp, owned) + + def _proxy_url(proxy: Gateway) -> str: return str(proxy.client.base_url).rstrip("/") @@ -194,18 +238,22 @@ def _poll(proxy: Gateway, session: CliSession, *, poll_secret: str | None = None ) -def _sign_in(proxy: Gateway, idp: Idp, browser: httpx.Client, session: CliSession) -> None: +def _sign_in( + proxy: Gateway, idp: Idp, browser: httpx.Client, session: CliSession, *, subject: str | None = None +) -> 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)) + authorize: Final = _assert_idp_redirect(proxy, idp, link, session.login_id) + hinted: Final = authorize if subject is None else f"{authorize}&{urlencode({'login_hint': subject})}" + callback: Final = _walk_idp(browser, hinted) 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: +def _ready_key(proxy: Gateway, session: CliSession, *, subject: str = SUBJECT) -> 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 + assert body["status"] == "ready" and body["user_id"] == subject, ready.text return string_value(body["key"]) @@ -489,6 +537,95 @@ def test_login_survives_a_worker_kill(idp: Idp, tmp_path: Path) -> None: assert victim not in respawned, respawned +def test_first_sign_in_user_serves_messages_and_signs_in_again_on_one_worker( + one_worker: OneWorkerProxy, provider: SharedProvider +) -> None: + proxy: Final = one_worker.owned.gateway + idp: Final = one_worker.idp + first: Final = _start_lite_login(proxy) + with _browser() as browser: + _sign_in(proxy, idp, browser, first) + _send_message(proxy, provider, _ready_key(proxy, first, subject=idp.subject)) + again: Final = _start_lite_login(proxy) + with _browser() as browser: + _sign_in(proxy, idp, browser, again) + _ready_key(proxy, again, subject=idp.subject) + + +def test_a_user_created_after_a_missed_lookup_is_budgeted_at_once_on_that_worker( + one_worker: OneWorkerProxy, provider: SharedProvider +) -> None: + proxy: Final = one_worker.owned.gateway + user_id: Final = f"cli-sso-late-user-{uuid.uuid4().hex[:12]}" + with proxy.scenario() as scenario: + key: Final = scenario.key(user_id=user_id, models=[MESSAGE_MODEL]) + _send_message(proxy, provider, key) + created: Final = proxy.request("POST", "/user/new", {"user_id": user_id, "max_budget": 0}) + assert created.status_code == 200, created.text + refused: Final = proxy.request( + "POST", + "/v1/messages", + {"model": MESSAGE_MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "over budget"}]}, + key=key, + ) + assert refused.status_code == 422 and f"ExceededBudget: User={user_id}" in refused.text, ( + f"{refused.status_code} {refused.text}" + ) + assert provider.received() == () + + +def test_a_user_an_admin_created_after_a_missed_lookup_signs_in_at_once_on_that_worker( + one_worker: OneWorkerProxy, provider: SharedProvider +) -> None: + proxy: Final = one_worker.owned.gateway + subject: Final = f"cli-sso-admin-created-{uuid.uuid4().hex[:12]}" + with proxy.scenario() as scenario: + _send_message(proxy, provider, scenario.key(user_id=subject, models=[MESSAGE_MODEL])) + created: Final = proxy.request( + "POST", + "/user/new", + {"user_id": subject, "user_email": f"{subject}@example.com", "user_role": "internal_user"}, + ) + assert created.status_code == 200, created.text + session: Final = _start_lite_login(proxy) + with _browser() as browser: + _sign_in(proxy, one_worker.idp, browser, session, subject=subject) + _send_message(proxy, provider, _ready_key(proxy, session, subject=subject)) + rows: Final = read_rows('SELECT user_email FROM "LiteLLM_UserTable" WHERE user_id = %s', (subject,)) + assert [row["user_email"] for row in rows] == [f"{subject}@example.com"], rows + + +def test_dashboard_sign_in_creates_the_user_and_hands_the_browser_a_session(tmp_path: Path) -> None: + subject: Final = f"ui-sso-first-sign-in-{uuid.uuid4().hex[:12]}" + with ( + _fake_idp(subject) as idp, + gateway_from_environment() as rig, + owned_proxy_process( + rig, + tmp_path, + _sso_environment(idp.wire.url), + remove_environment=("PROXY_BASE_URL", "DISABLE_ADMIN_UI"), + workers=1, + ) as owned, + ): + proxy: Final = owned.gateway + with _browser() as browser: + entry: Final = browser.get(f"{_proxy_url(proxy)}/sso/key/generate") + assert entry.is_redirect, f"{entry.status_code} {entry.text}" + signed_in: Final = _walk_idp(browser, entry.headers["location"]) + assert signed_in.status_code == 303, f"{signed_in.status_code} {signed_in.text}" + assert urlparse(signed_in.headers["location"]).query == "login=success", signed_in.headers["location"] + session: Final = JSON_OBJECT.validate_python( + jwt.decode(signed_in.cookies["token"], rig.key, algorithms=["HS256"]) + ) + assert session["user_id"] == subject and session["login_method"] == "sso", session + me: Final = proxy.request("GET", "/user/info", key=string_value(session["key"])) + assert me.status_code == 200, f"{me.status_code} {me.text}" + assert JSON_OBJECT.validate_json(me.content)["user_id"] == subject, me.text + rows: Final = read_rows('SELECT user_email FROM "LiteLLM_UserTable" WHERE user_id = %s', (subject,)) + assert [row["user_email"] for row in rows] == [f"{subject}@example.com"], rows + + @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 ( diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 912622da4bd..633c6738521 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -16,14 +16,18 @@ from pytest_mock import MockerFixture from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import ( + LiteLLM_UserTable, LiteLLM_UserTableFiltered, LitellmUserRoles, NewUserRequest, + NewUserResponse, ProxyErrorTypes, ProxyException, UpdateUserRequest, UserAPIKeyAuth, ) +from litellm.proxy.auth import auth_checks +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.management_endpoints.internal_user_endpoints import ( _authorize_user_list_request, _resolve_org_filter_for_user_search, @@ -34,6 +38,7 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import ( ui_view_users, ) from litellm.proxy.proxy_server import app +from litellm.types.proxy.auth.auth_checks import UserNotFoundError from litellm.types.proxy.management_endpoints.internal_user_endpoints import InsensitiveContains from tests.unit.proxy.management_endpoints.jwt_key_mapping_doubles import ( CascadingJWTMappingTable, @@ -1633,6 +1638,79 @@ async def test_new_user_default_teams_flow(mocker): litellm.default_internal_user_params = original_default_params +@pytest.mark.asyncio +async def test_new_user_clears_recent_missing_user_lookup(mocker: MockerFixture) -> None: + user_id: Final = "sso-created-user" + db_access_time_key: Final = f"user_id:{user_id}" + auth_checks.last_db_access_time.pop(db_access_time_key, None) + mocker.patch.object(auth_checks, "db_cache_expiry", 3600) + prisma_client: Final = mocker.MagicMock() + prisma_client.db.litellm_usertable.count = mocker.AsyncMock(return_value=1) + prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=None) + prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=None) + user_api_key_cache: Final = UserApiKeyCache() + mocker.patch("litellm.proxy.proxy_server.prisma_client", prisma_client) + license_check: Final = mocker.MagicMock() + license_check.is_over_limit.return_value = False + mocker.patch("litellm.proxy.proxy_server._license_check", license_check) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email", + new=mocker.AsyncMock(return_value=None), + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id", + new=mocker.AsyncMock(return_value=None), + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.check_if_default_team_set", + return_value=None, + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.generate_key_helper_fn", + new=mocker.AsyncMock( + return_value={"user_id": user_id, "token": "sk-sso-created-user", "expires": None} + ), + ) + + try: + with pytest.raises(UserNotFoundError): + await auth_checks.get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + sso_user_id="sso-user-id", + user_email="sso-created-user@example.com", + ) + + created_user: Final[NewUserResponse] = await new_user( + data=NewUserRequest( + user_id=user_id, + user_email="sso-created-user@example.com", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + assert created_user.user_id == user_id + + prisma_client.db.litellm_usertable.find_unique.return_value = LiteLLM_UserTable( + user_id=user_id, + user_email="sso-created-user@example.com", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + user: Final = await auth_checks.get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + ) + + assert user is not None + assert user.user_id == user_id + finally: + auth_checks.last_db_access_time.pop(db_access_time_key, None) + + def test_update_internal_new_user_params_proxy_admin_role(): """ Test that default_internal_user_params are NOT applied when user_role is PROXY_ADMIN