diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 6d7f565fb85..4019d0d2f52 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -36,6 +36,7 @@ from litellm.proxy.management_endpoints.common_utils import ( _is_user_team_admin, _user_has_admin_view, require_caller_user_id_for_non_admin, + validate_finite_spend, ) from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, @@ -1336,6 +1337,9 @@ async def _update_single_user_helper( existing_metadata=existing_metadata or {}, ) + # Reject NaN/±inf spend before it can reach the DB / spend counter. + validate_finite_spend(non_default_values.get("spend")) + # Perform the update response: Optional[Dict[str, Any]] = None @@ -1384,6 +1388,18 @@ async def _update_single_user_helper( litellm_proxy_admin_name=litellm_proxy_admin_name, ) + # A direct `spend` change must also invalidate the cross-pod spend + # counter enforcement reads; the DB write alone leaves a warm counter + # at the stale value. `non_default_values["user_id"]` is populated in + # every branch above (incl. the email-new-user insert path, whose + # response is a bare model and not safely subscriptable). + if non_default_values.get("spend") is not None: + from litellm.proxy.proxy_server import _invalidate_spend_counter + + await _invalidate_spend_counter( + counter_key=f"spend:user:{non_default_values['user_id']}" + ) + if response is None: raise HTTPException( status_code=400, diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index b4602e0ad8b..d26b4226ccc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -3102,6 +3102,81 @@ async def test_ghsa_wvg4_proxy_admin_can_update_user_budget(mocker): assert result is not None +@pytest.mark.asyncio +async def test_admin_user_update_spend_invalidates_counter(mocker): + """A direct /user/update spend change must invalidate the cross-pod + spend counter so enforcement re-reads the new DB value.""" + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = mocker.MagicMock() + existing_user = mocker.MagicMock() + existing_user.model_dump.return_value = {"user_id": "target-user", "spend": 50.0} + existing_user.user_id = "target-user" + mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock( + return_value=existing_user + ) + mock_prisma_client.update_data = mocker.AsyncMock( + return_value={"user_id": "target-user", "spend": 0.0} + ) + mock_prisma_client.jsonify_object = mocker.MagicMock(side_effect=lambda x: x) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") + mock_invalidate = mocker.patch( + "litellm.proxy.proxy_server._invalidate_spend_counter", + new=mocker.AsyncMock(), + ) + + user_request = UpdateUserRequest(user_id="target-user", spend=0) + admin_caller = UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + await _update_single_user_helper( + user_request=user_request, user_api_key_dict=admin_caller + ) + mock_invalidate.assert_awaited_once_with(counter_key="spend:user:target-user") + + +@pytest.mark.asyncio +async def test_user_update_rejects_non_finite_spend(mocker): + """NaN/inf spend is rejected before any DB write or counter invalidation.""" + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = mocker.MagicMock() + existing_user = mocker.MagicMock() + existing_user.model_dump.return_value = {"user_id": "target-user", "spend": 50.0} + existing_user.user_id = "target-user" + mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock( + return_value=existing_user + ) + mock_prisma_client.update_data = mocker.AsyncMock() + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") + mock_invalidate = mocker.patch( + "litellm.proxy.proxy_server._invalidate_spend_counter", + new=mocker.AsyncMock(), + ) + + user_request = UpdateUserRequest(user_id="target-user", spend=float("nan")) + admin_caller = UserAPIKeyAuth( + user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + with pytest.raises(HTTPException) as exc: + await _update_single_user_helper( + user_request=user_request, user_api_key_dict=admin_caller + ) + assert exc.value.status_code == 400 + mock_prisma_client.update_data.assert_not_called() + mock_invalidate.assert_not_awaited() + + @pytest.mark.asyncio async def test_resolve_user_email_metadata_maps_page_user_ids_to_email(mocker): """Regression for LIT-3889.