mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
cb138ba92f
commit
91e9b1f06b
4 changed files with 244 additions and 17 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue