mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
perf(scim): resolve group members with one user table read per member (#39228)
Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
04a25083a6
commit
47b9d838aa
2 changed files with 155 additions and 35 deletions
|
|
@ -585,6 +585,37 @@ async def _users_named_by_member_value(
|
|||
return tuple(dict.fromkeys(row.user_id for row in rows))
|
||||
|
||||
|
||||
async def _accounts_named_by_member_value(value: str, prisma_client: PrismaClient) -> tuple[str, ...]:
|
||||
"""Every user id this member value names, by user id, SSO identity or email.
|
||||
|
||||
Classification needs to know whether the value is one account's ``user_id`` and
|
||||
whether it names any other account, so all three fields are read in one pass. The
|
||||
id is compared exactly and unstripped, as a primary key lookup would; the
|
||||
identities compare as ``_users_named_by_member_value`` describes. Two rows are
|
||||
enough to tell one account from several, so the read stops there. Only a full
|
||||
read that lacks the row keyed by the value leaves that row's existence open, and
|
||||
only then is the id read on its own.
|
||||
"""
|
||||
subject: Final = value.strip()
|
||||
email: Final[_CaseInsensitiveMatch] = {"equals": subject, "mode": "insensitive"}
|
||||
users: Final = _table(UserRepository(prisma_client))
|
||||
rows: Final = await users.find_many(
|
||||
where={ # mutable-ok: Prisma filter
|
||||
"OR": [ # mutable-ok: Prisma filter
|
||||
{"user_id": value}, # mutable-ok: Prisma filter
|
||||
{"sso_user_id": subject}, # mutable-ok: Prisma filter
|
||||
{"user_email": email}, # mutable-ok: Prisma filter
|
||||
],
|
||||
},
|
||||
take=2,
|
||||
)
|
||||
named: Final = tuple(dict.fromkeys(row.user_id for row in rows))
|
||||
if len(named) < 2 or value in named:
|
||||
return named
|
||||
keyed: Final = await users.find_unique(where={"user_id": value})
|
||||
return named if keyed is None else (value, *named)
|
||||
|
||||
|
||||
async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient) -> _ClassifiedGroupMember:
|
||||
"""
|
||||
Decide what a single SCIM group member refers to.
|
||||
|
|
@ -627,11 +658,9 @@ async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient
|
|||
if member_type == "group":
|
||||
return _SkippedGroupMember(value=value, reason="nested_group")
|
||||
|
||||
user: Final = await _table(UserRepository(prisma_client)).find_unique(where={"user_id": value})
|
||||
if user is not None:
|
||||
shared_with: Final = tuple(
|
||||
other for other in await _users_named_by_member_value(value, prisma_client) if other != value
|
||||
)
|
||||
named: Final = await _accounts_named_by_member_value(value, prisma_client)
|
||||
if value in named:
|
||||
shared_with: Final = tuple(other for other in named if other != value)
|
||||
if shared_with:
|
||||
verbose_proxy_logger.warning(
|
||||
"SCIM: group member '%s' is one account's user id and is also account '%s' by SSO identity or email, "
|
||||
|
|
@ -651,7 +680,6 @@ async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient
|
|||
if team is not None and _team_metadata_has_scim_provenance(team.metadata):
|
||||
return _SkippedGroupMember(value=value, reason="existing_team")
|
||||
|
||||
named: Final = await _users_named_by_member_value(value, prisma_client)
|
||||
if len(named) == 1:
|
||||
verbose_proxy_logger.info(
|
||||
"SCIM: group member '%s' matched user_id '%s' by SSO identity or email",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import logging
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Callable, Mapping
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, call
|
||||
|
|
@ -1645,6 +1645,25 @@ async def test_update_group_e2e(mocker):
|
|||
ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team)
|
||||
|
||||
|
||||
def _rows_by_exact_id(
|
||||
user_row: Callable[[Mapping[str, str]], LiteLLM_UserTable | MagicMock | None],
|
||||
) -> Callable[..., tuple[LiteLLM_UserTable | MagicMock, ...]]:
|
||||
"""``find_many`` stand-in for the classifier's cross-field read on a table where a
|
||||
member value only ever matches as an exact ``user_id``."""
|
||||
|
||||
def rows(where: Mapping[str, object], take: int | None = None) -> tuple[LiteLLM_UserTable | MagicMock, ...]:
|
||||
clauses: Final = where["OR"]
|
||||
assert isinstance(clauses, list)
|
||||
found: Final = tuple(user_row(clause) for clause in clauses if "user_id" in clause)
|
||||
return tuple(row for row in found if row is not None)
|
||||
|
||||
return rows
|
||||
|
||||
|
||||
def _user_row_for(where: Mapping[str, str]) -> LiteLLM_UserTable:
|
||||
return LiteLLM_UserTable(user_id=where["user_id"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch):
|
||||
"""
|
||||
|
|
@ -1696,9 +1715,8 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch):
|
|||
return mock_user
|
||||
return None # new-user-1 and new-user-2 don't exist
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
|
||||
|
||||
# Mock dependencies
|
||||
mocker.patch(
|
||||
|
|
@ -1782,9 +1800,8 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch):
|
|||
return mock_user
|
||||
return None # new-user-3 and new-user-4 don't exist
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
|
||||
|
||||
# Mock dependencies
|
||||
mocker.patch(
|
||||
|
|
@ -1853,9 +1870,8 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true(mocker
|
|||
return mock_user
|
||||
return None # new-user-1 and new-user-2 don't exist
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
|
||||
|
||||
# Mock user creation
|
||||
created_user_1 = NewUserResponse(user_id="new-user-1", key="test-key-1")
|
||||
|
|
@ -1943,9 +1959,8 @@ async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, mon
|
|||
return mock_user
|
||||
return None # new-user-1 doesn't exist
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
|
||||
|
||||
# Mock user creation
|
||||
created_user = NewUserResponse(user_id="new-user-1", key="test-key-1")
|
||||
|
|
@ -2013,9 +2028,8 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa
|
|||
return mock_user
|
||||
return None # new-user-1 doesn't exist
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
|
||||
|
||||
# Mock dependencies
|
||||
mocker.patch(
|
||||
|
|
@ -3121,8 +3135,7 @@ async def test_process_group_patch_operations_add_retains_existing_members(mocke
|
|||
mock_prisma_client.db = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
|
||||
# new-user already exists in the DB
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock(user_id="new-user"))
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=(mocker.MagicMock(user_id="new-user"),))
|
||||
|
||||
_, final_members, _ = await _process_group_patch_operations(
|
||||
patch_ops=patch_ops,
|
||||
|
|
@ -3415,8 +3428,7 @@ async def test_patch_group_add_applies_delta_and_keeps_concurrent_add(mocker):
|
|||
)
|
||||
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=final_team)
|
||||
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock())
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(_user_row_for))
|
||||
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
|
||||
|
|
@ -3509,8 +3521,7 @@ async def test_patch_group_replace_stays_absolute_against_concurrent_roster(mock
|
|||
)
|
||||
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=final_team)
|
||||
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock())
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
|
||||
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(_user_row_for))
|
||||
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
|
||||
|
|
@ -3640,8 +3651,7 @@ async def test_process_group_patch_add_filtered_path_without_value(mocker):
|
|||
prisma_client = mocker.MagicMock()
|
||||
prisma_client.db = mocker.MagicMock()
|
||||
prisma_client.db.litellm_usertable = mocker.MagicMock()
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=LiteLLM_UserTable(user_id="user-3"))
|
||||
prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
|
||||
prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=(LiteLLM_UserTable(user_id="user-3"),))
|
||||
|
||||
_, final_members, _ = await _process_group_patch_operations(
|
||||
patch_ops=patch_ops,
|
||||
|
|
@ -3733,12 +3743,14 @@ def _member_resolution_prisma(
|
|||
starts folding it, fails here instead of passing.
|
||||
|
||||
A caller that must know which accounts match rather than merely how many
|
||||
passes take=None, so an unbounded read returns every match.
|
||||
passes take=None, so an unbounded read returns every match. The row keyed by
|
||||
the value comes last, the order a bounded read is least prepared for, since
|
||||
the database promises no order at all.
|
||||
"""
|
||||
clauses: Final = where["OR"]
|
||||
assert isinstance(clauses, list)
|
||||
fields: Final = tuple(next(iter(clause)) for clause in clauses)
|
||||
assert fields == ("sso_user_id", "user_email"), fields
|
||||
assert fields in (("user_id", "sso_user_id", "user_email"), ("sso_user_id", "user_email")), fields
|
||||
|
||||
def comparison(clause: Mapping[str, object]) -> tuple[str, bool]:
|
||||
"""The needle and whether production asked for a case-insensitive compare,
|
||||
|
|
@ -3749,8 +3761,9 @@ def _member_resolution_prisma(
|
|||
assert isinstance(criterion, dict), criterion
|
||||
return criterion["equals"], criterion.get("mode") == "insensitive"
|
||||
|
||||
sso_needle, sso_insensitive = comparison(clauses[0])
|
||||
email_needle, email_insensitive = comparison(clauses[1])
|
||||
by_field: Final = dict(zip(fields, (comparison(clause) for clause in clauses)))
|
||||
sso_needle, sso_insensitive = by_field["sso_user_id"]
|
||||
email_needle, email_insensitive = by_field["user_email"]
|
||||
|
||||
def same(stored: str, needle: str, insensitive: bool) -> bool:
|
||||
return stored.casefold() == needle.casefold() if insensitive else stored == needle
|
||||
|
|
@ -3768,6 +3781,11 @@ def _member_resolution_prisma(
|
|||
if same(email, email_needle, email_insensitive)
|
||||
for user_id in user_ids
|
||||
),
|
||||
(
|
||||
user_id
|
||||
for user_id in users
|
||||
if "user_id" in by_field and same(user_id, by_field["user_id"][0], by_field["user_id"][1])
|
||||
),
|
||||
)
|
||||
)
|
||||
found: Final = tuple(dict.fromkeys(matched))
|
||||
|
|
@ -4611,9 +4629,15 @@ async def test_resolve_group_member_ids_dedupes_repeated_member(mocker, scim_ups
|
|||
|
||||
|
||||
def _identity_lookup(value: str) -> object:
|
||||
"""The single cross-field lookup the classifier is expected to issue."""
|
||||
"""The single cross-field lookup the classifier is expected to issue per member."""
|
||||
return call(
|
||||
where={"OR": [{"sso_user_id": value}, {"user_email": {"equals": value, "mode": "insensitive"}}]},
|
||||
where={
|
||||
"OR": [
|
||||
{"user_id": value},
|
||||
{"sso_user_id": value},
|
||||
{"user_email": {"equals": value, "mode": "insensitive"}},
|
||||
]
|
||||
},
|
||||
take=2,
|
||||
)
|
||||
|
||||
|
|
@ -5152,6 +5176,77 @@ async def test_resolve_group_member_ids_refuses_a_user_id_that_names_another_acc
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_group_member_ids_reads_the_exact_id_when_two_other_accounts_fill_the_lookup(
|
||||
mocker, scim_upsert_user_enabled
|
||||
):
|
||||
"""A value that is one account's id and two other accounts' identities fills the
|
||||
bounded lookup with the other two. The account keyed by the value must still be
|
||||
found, or the id would lose its precedence and a non-canonical type would skip
|
||||
a member that names a real user."""
|
||||
prisma_client = _member_resolution_prisma(
|
||||
mocker,
|
||||
users={"shared"},
|
||||
teams=set(),
|
||||
sso_user_id_to_user_id={"shared": "by-sso"},
|
||||
email_to_user_id={"shared": "by-email"},
|
||||
)
|
||||
create_user_mock = mocker.patch( # test-quality-ok: user creation is module-level, not injectable into the resolver
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _resolve_group_member_ids(
|
||||
members=[SCIMMember(value="shared", type="direct")],
|
||||
created_via="scim_group_membership",
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "shared" in str(exc_info.value.detail)
|
||||
create_user_mock.assert_not_called()
|
||||
assert prisma_client.db.litellm_usertable.find_many.await_args_list == [_identity_lookup("shared")]
|
||||
prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with(where={"user_id": "shared"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_group_member_ids_reads_the_user_table_once_per_member(mocker, scim_upsert_user_enabled):
|
||||
"""Every member costs one read of the user table, however it resolves: by its exact
|
||||
id (which still outranks a non-canonical type), by identity, as a SCIM team, or not
|
||||
at all. Looking the exact id up on its own before the identity read doubled the
|
||||
reads of a push, and the identity read is a scan."""
|
||||
prisma_client = _member_resolution_prisma(
|
||||
mocker,
|
||||
users={"by-id"},
|
||||
teams={"by-team"},
|
||||
email_to_user_id={"by-email@example.com": "email-user"},
|
||||
)
|
||||
mocker.patch( # test-quality-ok: user creation is module-level, not injectable into the resolver
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
|
||||
AsyncMock(return_value=NewUserResponse(user_id="nobody", key="key")),
|
||||
)
|
||||
|
||||
result = await _resolve_group_member_ids(
|
||||
members=[
|
||||
SCIMMember(value="by-id", type="direct"),
|
||||
SCIMMember(value="by-email@example.com"),
|
||||
SCIMMember(value="by-team"),
|
||||
SCIMMember(value="nobody"),
|
||||
],
|
||||
created_via="scim_group_membership",
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
assert result.all_member_ids == ["by-id", "email-user", "nobody"]
|
||||
prisma_client.db.litellm_usertable.find_unique.assert_not_awaited()
|
||||
assert prisma_client.db.litellm_usertable.find_many.await_args_list == [
|
||||
_identity_lookup("by-id"),
|
||||
_identity_lookup("by-email@example.com"),
|
||||
_identity_lookup("by-team"),
|
||||
_identity_lookup("nobody"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_group_member_ids_warns_before_creating_unmatched_placeholder(
|
||||
|
|
@ -5536,10 +5631,7 @@ async def test_resolve_group_member_ids_admits_member_created_concurrently(mocke
|
|||
the member is still admitted: the id resolves to a real user row, so failing
|
||||
or dropping it would be wrong either way."""
|
||||
prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set())
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
side_effect=[None, LiteLLM_UserTable(user_id="raced-user")]
|
||||
)
|
||||
prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=LiteLLM_UserTable(user_id="raced-user"))
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
|
||||
AsyncMock(return_value=None),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue