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 <mateo@berri.ai>
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>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 06:45:41 -07:00 • committed by GitHub
parent cb138ba92f
commit 91e9b1f06b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 244 additions and 17 deletions

View file

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

View file

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

View file

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

View file

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