mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
feat(proxy): reset team member spend on /team/member_update
Persist LiteLLM_TeamMembership.spend directly (upsert, since a spend-only
request leaves no budget patch and the membership row may not exist), then
invalidate spend:team_member:{user_id}:{team_id} so enforcement re-reads
the new value. Reject non-finite spend before any write.
This commit is contained in:
parent
f6a05da3b1
commit
cc81bc9a77
2 changed files with 105 additions and 0 deletions
|
|
@ -83,6 +83,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
_update_metadata_fields,
|
||||
_upsert_budget_and_membership,
|
||||
_user_has_admin_view,
|
||||
validate_finite_spend,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.organization_endpoints import (
|
||||
add_member_to_organization,
|
||||
|
|
@ -2926,6 +2927,8 @@ async def team_member_update(
|
|||
)
|
||||
|
||||
_validate_budget_duration(data.budget_duration)
|
||||
# Reject NaN/±inf spend before it can reach the DB / spend counter.
|
||||
validate_finite_spend(data.spend)
|
||||
|
||||
_existing_team_row = await TeamRepository(prisma_client).table.find_unique(
|
||||
where={"team_id": data.team_id}
|
||||
|
|
@ -3013,6 +3016,36 @@ async def team_member_update(
|
|||
team_default_budget_id=team_default_budget_id,
|
||||
)
|
||||
|
||||
### reset member spend
|
||||
# spend lives on LiteLLM_TeamMembership, not the budget table that
|
||||
# _upsert_budget_and_membership writes, so persist it directly. Use an
|
||||
# upsert: a spend-only request produces an empty budget_patch, so the
|
||||
# membership row is not guaranteed to exist after the call above. Written
|
||||
# outside the budget tx and followed by a counter invalidation so
|
||||
# enforcement re-reads the new value (reseed-from-DB).
|
||||
if data.spend is not None:
|
||||
await prisma_client.db.litellm_teammembership.upsert(
|
||||
where={
|
||||
"user_id_team_id": {
|
||||
"user_id": received_user_id,
|
||||
"team_id": data.team_id,
|
||||
}
|
||||
},
|
||||
data={
|
||||
"create": {
|
||||
"user_id": received_user_id,
|
||||
"team_id": data.team_id,
|
||||
"spend": data.spend,
|
||||
},
|
||||
"update": {"spend": data.spend},
|
||||
},
|
||||
)
|
||||
from litellm.proxy.proxy_server import _invalidate_spend_counter
|
||||
|
||||
await _invalidate_spend_counter(
|
||||
counter_key=f"spend:team_member:{received_user_id}:{data.team_id}"
|
||||
)
|
||||
|
||||
### update team member role
|
||||
if data.role is not None:
|
||||
team_members: List[Member] = []
|
||||
|
|
@ -3045,6 +3078,7 @@ async def team_member_update(
|
|||
rpm_limit=data.rpm_limit,
|
||||
budget_duration=data.budget_duration,
|
||||
allowed_models=data.allowed_models,
|
||||
spend=data.spend,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -157,6 +157,77 @@ async def test_team_member_update_omits_unset_fields_from_patch(happy_path_upser
|
|||
assert happy_path_upsert.await_args.kwargs["budget_patch"] == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_update_spend_writes_membership_and_invalidates(
|
||||
happy_path_upsert, monkeypatch
|
||||
):
|
||||
"""A `spend` value resets LiteLLM_TeamMembership.spend (upsert, since the
|
||||
row may not exist) and invalidates the cross-pod team-member counter."""
|
||||
prisma_client = proxy_server.prisma_client
|
||||
membership_upsert = AsyncMock()
|
||||
prisma_client.db.litellm_teammembership.upsert = membership_upsert
|
||||
invalidate = AsyncMock()
|
||||
monkeypatch.setattr(proxy_server, "_invalidate_spend_counter", invalidate)
|
||||
|
||||
data, request, auth = _member_update_request(spend=0.0)
|
||||
|
||||
response = await team_member_update(data, request, auth)
|
||||
|
||||
membership_upsert.assert_awaited_once()
|
||||
kwargs = membership_upsert.await_args.kwargs
|
||||
assert kwargs["where"] == {
|
||||
"user_id_team_id": {"user_id": "user-1", "team_id": "team-1234"}
|
||||
}
|
||||
assert kwargs["data"]["update"] == {"spend": 0.0}
|
||||
assert kwargs["data"]["create"]["spend"] == 0.0
|
||||
invalidate.assert_awaited_once_with(
|
||||
counter_key="spend:team_member:user-1:team-1234"
|
||||
)
|
||||
assert response.spend == 0.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_update_rejects_non_finite_spend(
|
||||
happy_path_upsert, monkeypatch
|
||||
):
|
||||
"""NaN/inf spend is rejected with 400 before any membership write or
|
||||
counter invalidation."""
|
||||
prisma_client = proxy_server.prisma_client
|
||||
membership_upsert = AsyncMock()
|
||||
prisma_client.db.litellm_teammembership.upsert = membership_upsert
|
||||
invalidate = AsyncMock()
|
||||
monkeypatch.setattr(proxy_server, "_invalidate_spend_counter", invalidate)
|
||||
|
||||
data, request, auth = _member_update_request(spend=float("inf"))
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await team_member_update(data, request, auth)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
membership_upsert.assert_not_called()
|
||||
invalidate.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_update_without_spend_skips_membership_write(
|
||||
happy_path_upsert, monkeypatch
|
||||
):
|
||||
"""A budget-only update must not touch the membership spend row or the
|
||||
counter."""
|
||||
prisma_client = proxy_server.prisma_client
|
||||
membership_upsert = AsyncMock()
|
||||
prisma_client.db.litellm_teammembership.upsert = membership_upsert
|
||||
invalidate = AsyncMock()
|
||||
monkeypatch.setattr(proxy_server, "_invalidate_spend_counter", invalidate)
|
||||
|
||||
data, request, auth = _member_update_request(max_budget_in_team=10.0)
|
||||
|
||||
await team_member_update(data, request, auth)
|
||||
|
||||
membership_upsert.assert_not_called()
|
||||
invalidate.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_duration",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue