mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(team/member_add): atomic SQL for user.teams to prevent race condition duplicates
Replace non-atomic Prisma push on user.teams with execute_raw using PostgreSQL ARRAY dedup (SELECT DISTINCT unnest(...)). Row-level lock from UPDATE serialises concurrent writers for the same user. Change TeamMembership create to upsert so concurrent requests for the same (user, team) pair are idempotent. Fixes LIT-4168 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ef030235fd
commit
b6514c25fc
2 changed files with 149 additions and 25 deletions
|
|
@ -252,6 +252,28 @@ async def _resolve_member_budget_id(
|
|||
return response.budget_id
|
||||
|
||||
|
||||
async def _atomic_add_team_to_user(
|
||||
prisma_client: PrismaClient,
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
) -> None:
|
||||
"""Atomically add team_id to user.teams, skipping if already present.
|
||||
|
||||
Uses a single UPDATE with array dedup so concurrent calls for the
|
||||
same (user, team) pair never produce duplicate entries. The row-level
|
||||
lock acquired by UPDATE serialises concurrent writers.
|
||||
"""
|
||||
await prisma_client.db.execute_raw(
|
||||
'UPDATE "LiteLLM_UserTable" '
|
||||
"SET teams = ("
|
||||
" SELECT ARRAY(SELECT DISTINCT unnest(teams || ARRAY[$1]::text[]))"
|
||||
") "
|
||||
"WHERE user_id = $2",
|
||||
team_id,
|
||||
user_id,
|
||||
)
|
||||
|
||||
|
||||
async def add_new_member(
|
||||
new_member: Member,
|
||||
max_budget_in_team: Optional[float],
|
||||
|
|
@ -279,10 +301,11 @@ async def add_new_member(
|
|||
_returned_user = await UserRepository(prisma_client).table.upsert(
|
||||
where={"user_id": new_member.user_id},
|
||||
data={
|
||||
"update": {"teams": {"push": [team_id]}},
|
||||
"create": {"teams": [team_id], **new_user_defaults}, # type: ignore
|
||||
"update": {},
|
||||
"create": {**new_user_defaults}, # type: ignore
|
||||
},
|
||||
)
|
||||
await _atomic_add_team_to_user(prisma_client, new_member.user_id, team_id)
|
||||
if _returned_user is not None:
|
||||
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
|
||||
elif new_member.user_email is not None:
|
||||
|
|
@ -302,12 +325,8 @@ async def add_new_member(
|
|||
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
|
||||
elif len(existing_user_row) == 1:
|
||||
user_info = existing_user_row[0]
|
||||
_returned_user = await UserRepository(prisma_client).table.update(
|
||||
where={"user_id": user_info.user_id}, # type: ignore
|
||||
data={"teams": {"push": [team_id]}},
|
||||
)
|
||||
if _returned_user is not None:
|
||||
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
|
||||
await _atomic_add_team_to_user(prisma_client, user_info.user_id, team_id)
|
||||
returned_user = LiteLLM_UserTable(**user_info.model_dump())
|
||||
elif len(existing_user_row) > 1:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -325,11 +344,20 @@ async def add_new_member(
|
|||
)
|
||||
|
||||
if _budget_id and returned_user is not None and returned_user.user_id is not None:
|
||||
_returned_team_membership = await TeamMembershipRepository(prisma_client).table.create(
|
||||
_returned_team_membership = await TeamMembershipRepository(prisma_client).table.upsert(
|
||||
where={
|
||||
"user_id_team_id": {
|
||||
"user_id": returned_user.user_id,
|
||||
"team_id": team_id,
|
||||
}
|
||||
},
|
||||
data={
|
||||
"team_id": team_id,
|
||||
"user_id": returned_user.user_id,
|
||||
"budget_id": _budget_id,
|
||||
"create": {
|
||||
"team_id": team_id,
|
||||
"user_id": returned_user.user_id,
|
||||
"budget_id": _budget_id,
|
||||
},
|
||||
"update": {},
|
||||
},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -238,7 +238,7 @@ async def test_add_new_member_clones_default_team_budget_id():
|
|||
"budget_id": test_cloned_budget_id,
|
||||
"litellm_budget_table": None,
|
||||
}
|
||||
mock_prisma_client.db.litellm_teammembership.create = AsyncMock(
|
||||
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
|
||||
return_value=mock_team_membership_response
|
||||
)
|
||||
|
||||
|
|
@ -261,7 +261,7 @@ async def test_add_new_member_clones_default_team_budget_id():
|
|||
assert result_team_membership.budget_id != test_default_budget_id
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.upsert.assert_called_once()
|
||||
mock_prisma_client.db.litellm_teammembership.create.assert_called_once()
|
||||
mock_prisma_client.db.litellm_teammembership.upsert.assert_called_once()
|
||||
|
||||
# The clone must have happened: find_unique on the default, create for the clone.
|
||||
mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with(
|
||||
|
|
@ -278,9 +278,9 @@ async def test_add_new_member_clones_default_team_budget_id():
|
|||
assert cloned_create_data["created_by"] == user_api_key_dict.user_id
|
||||
|
||||
team_membership_call_args = (
|
||||
mock_prisma_client.db.litellm_teammembership.create.call_args
|
||||
mock_prisma_client.db.litellm_teammembership.upsert.call_args
|
||||
)
|
||||
create_data = team_membership_call_args.kwargs["data"]
|
||||
create_data = team_membership_call_args.kwargs["data"]["create"]
|
||||
assert create_data["budget_id"] == test_cloned_budget_id
|
||||
|
||||
|
||||
|
|
@ -335,7 +335,7 @@ async def test_add_new_member_budget_duration_only_clones_default_max_budget():
|
|||
"budget_id": "cloned-dc",
|
||||
"litellm_budget_table": None,
|
||||
}
|
||||
mock_prisma_client.db.litellm_teammembership.create = AsyncMock(
|
||||
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
|
||||
return_value=mock_team_membership_response
|
||||
)
|
||||
|
||||
|
|
@ -395,7 +395,7 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget():
|
|||
# Even though we mock these, they must NOT be called on the no-budget path.
|
||||
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock()
|
||||
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
|
||||
mock_prisma_client.db.litellm_teammembership.create = AsyncMock()
|
||||
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock()
|
||||
|
||||
result_user, result_team_membership = await add_new_member(
|
||||
new_member=new_member,
|
||||
|
|
@ -414,7 +414,7 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget():
|
|||
assert result_team_membership is None
|
||||
mock_prisma_client.db.litellm_budgettable.find_unique.assert_not_called()
|
||||
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
|
||||
mock_prisma_client.db.litellm_teammembership.create.assert_not_called()
|
||||
mock_prisma_client.db.litellm_teammembership.upsert.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -474,7 +474,7 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided():
|
|||
"budget_id": test_new_budget_id,
|
||||
"litellm_budget_table": None,
|
||||
}
|
||||
mock_prisma_client.db.litellm_teammembership.create = AsyncMock(
|
||||
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
|
||||
return_value=mock_team_membership_response
|
||||
)
|
||||
|
||||
|
|
@ -503,10 +503,10 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided():
|
|||
|
||||
# Verify the team membership was created with the correct budget_id
|
||||
team_membership_call_args = (
|
||||
mock_prisma_client.db.litellm_teammembership.create.call_args
|
||||
mock_prisma_client.db.litellm_teammembership.upsert.call_args
|
||||
)
|
||||
assert team_membership_call_args is not None
|
||||
create_data = team_membership_call_args.kwargs["data"]
|
||||
create_data = team_membership_call_args.kwargs["data"]["create"]
|
||||
assert create_data["budget_id"] == test_new_budget_id
|
||||
|
||||
|
||||
|
|
@ -546,7 +546,7 @@ async def test_add_new_member_persists_budget_duration():
|
|||
"budget_id": "budget-dur",
|
||||
"litellm_budget_table": None,
|
||||
}
|
||||
mock_prisma_client.db.litellm_teammembership.create = AsyncMock(
|
||||
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
|
||||
return_value=mock_team_membership_response
|
||||
)
|
||||
|
||||
|
|
@ -609,7 +609,7 @@ async def test_add_new_member_persists_budget_duration_without_max_budget():
|
|||
"budget_id": "budget-dur2",
|
||||
"litellm_budget_table": None,
|
||||
}
|
||||
mock_prisma_client.db.litellm_teammembership.create = AsyncMock(
|
||||
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
|
||||
return_value=mock_team_membership_response
|
||||
)
|
||||
|
||||
|
|
@ -699,7 +699,7 @@ async def test_add_new_member_with_user_email_clones_default_budget():
|
|||
"budget_id": test_cloned_budget_id,
|
||||
"litellm_budget_table": None,
|
||||
}
|
||||
mock_prisma_client.db.litellm_teammembership.create = AsyncMock(
|
||||
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
|
||||
return_value=mock_team_membership_response
|
||||
)
|
||||
|
||||
|
|
@ -739,6 +739,102 @@ async def test_add_new_member_with_user_email_clones_default_budget():
|
|||
mock_prisma_client.db.litellm_budgettable.create.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_new_member_uses_atomic_team_append():
|
||||
"""Regression for LIT-4168: concurrent team/member_add requests must not
|
||||
produce duplicate team_id entries in user.teams. add_new_member must call
|
||||
execute_raw with an atomic array-dedup SQL instead of Prisma's blind push.
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
new_member = Member(user_id="race-user", role="user")
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_user_response = MagicMock()
|
||||
mock_user_response.model_dump.return_value = {
|
||||
"user_id": "race-user",
|
||||
"user_email": None,
|
||||
"teams": [],
|
||||
"user_role": "internal_user",
|
||||
}
|
||||
mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(
|
||||
return_value=mock_user_response
|
||||
)
|
||||
|
||||
await add_new_member(
|
||||
new_member=new_member,
|
||||
max_budget_in_team=None,
|
||||
prisma_client=mock_prisma_client,
|
||||
team_id="team-race",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name="test_admin",
|
||||
)
|
||||
|
||||
# The upsert must NOT push team_id into user.teams (that was the race bug)
|
||||
upsert_args = mock_prisma_client.db.litellm_usertable.upsert.call_args
|
||||
update_clause = upsert_args.kwargs["data"]["update"]
|
||||
assert "teams" not in update_clause, (
|
||||
"upsert update clause must not touch teams; atomic SQL handles it"
|
||||
)
|
||||
|
||||
# execute_raw must be called with the atomic array-dedup query
|
||||
mock_prisma_client.db.execute_raw.assert_called_once()
|
||||
raw_call = mock_prisma_client.db.execute_raw.call_args
|
||||
sql = raw_call.args[0]
|
||||
assert "DISTINCT" in sql
|
||||
assert "unnest" in sql
|
||||
assert raw_call.args[1] == "team-race"
|
||||
assert raw_call.args[2] == "race-user"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_new_member_email_path_uses_atomic_team_append():
|
||||
"""Same as above but for the user_email lookup path where an existing user
|
||||
is found by email. The old code used Prisma update with push; the fix must
|
||||
use the same atomic execute_raw approach.
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
new_member = Member(user_email="race@example.com", role="user")
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
|
||||
existing_user = MagicMock()
|
||||
existing_user.user_id = "existing-uid"
|
||||
existing_user.user_email = "race@example.com"
|
||||
existing_user.model_dump.return_value = {
|
||||
"user_id": "existing-uid",
|
||||
"user_email": "race@example.com",
|
||||
"teams": ["other-team"],
|
||||
"user_role": "internal_user",
|
||||
}
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=[existing_user])
|
||||
|
||||
await add_new_member(
|
||||
new_member=new_member,
|
||||
max_budget_in_team=None,
|
||||
prisma_client=mock_prisma_client,
|
||||
team_id="team-race-email",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name="test_admin",
|
||||
)
|
||||
|
||||
# The old Prisma update with push must NOT be called
|
||||
mock_prisma_client.db.litellm_usertable.update.assert_not_called()
|
||||
|
||||
# execute_raw must be called with the atomic query
|
||||
mock_prisma_client.db.execute_raw.assert_called_once()
|
||||
raw_call = mock_prisma_client.db.execute_raw.call_args
|
||||
assert raw_call.args[1] == "team-race-email"
|
||||
assert raw_call.args[2] == "existing-uid"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attach_object_permission_to_dict_with_object_permission_id():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue