diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 249cfdae4eb..bb5d5b119ab 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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, ) diff --git a/tests/test_litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py index 1a11deb6403..fb743002afb 100644 --- a/tests/test_litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -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", [