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:
devin-ai-integration[bot] 2026-09-01 17:49:40 -07:00 committed by GitHub
parent 04a25083a6
commit 47b9d838aa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 155 additions and 35 deletions

View file

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

View file

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