diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index 9a1cf701142..c2c0597a637 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -67,6 +67,7 @@ class _TagRecord(Protocol): description: str | None models: Sequence[str] model_info: object + spend: float budget_id: str | None created_at: datetime updated_at: datetime @@ -342,6 +343,7 @@ async def new_tag( created_at=new_tag_record.created_at.isoformat(), updated_at=new_tag_record.updated_at.isoformat(), created_by=new_tag_record.created_by, + spend=new_tag_record.spend, ) return { @@ -411,6 +413,7 @@ async def update_tag( - description: Optional[str] - Updated description - models: List[str] - Updated list of allowed LLM models - budget_id: Optional[str] - The id for a budget to associate with the tag + - spend: Optional[float] - Set the tag's accumulated spend (e.g. 0 to reset it). Omit to leave unchanged. Does not change budget_reset_at; the current budget window is left as-is ### BUDGET UPDATE PARAMS ### - max_budget: Optional[float] - Max budget for tag @@ -421,11 +424,15 @@ async def update_tag( - model_max_budget: Optional[dict] - Max budget for a specific model - budget_duration: Optional[str] - Frequency of resetting tag budget """ + from litellm.proxy.management_endpoints.common_utils import validate_finite_spend from litellm.proxy.proxy_server import prisma_client if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") + # Reject NaN/±inf spend before it can reach the DB / spend counter. + validate_finite_spend(tag.spend) + try: # Check if tag exists existing_tag: Final = await _table(TagRepository(prisma_client)).find_unique(where={"tag_name": tag.name}) @@ -448,7 +455,7 @@ async def update_tag( model_info: Final = await _get_model_names(prisma_client, tag.models or []) # Prepare update data - update_data: Final = { + update_data: Final[dict[str, object]] = { "description": tag.description, "models": tag.models or [], "model_info": json.dumps(model_info), @@ -458,6 +465,9 @@ async def update_tag( if budget_id != existing_tag.budget_id: update_data["budget_id"] = budget_id + if tag.spend is not None: + update_data["spend"] = tag.spend + # Update tag in database updated_tag_record: Final = await _table(TagRepository(prisma_client)).update( where={"tag_name": tag.name}, @@ -466,6 +476,28 @@ async def update_tag( await _evict_tag_cache_keys((tag_cache_key(tag.name),)) + if tag.spend is not None: + # Refresh the live spend counter immediately, the same way /key/update does for + # spend:key: - otherwise the tag stays blocked on the stale cached value + # until the counter's TTL expires. + from litellm.proxy.proxy_server import spend_counter_cache + + counter_key: Final = f"spend:tag:{tag.name}" + spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=tag.spend, ttl=60) + if spend_counter_cache.redis_cache is not None: + try: + await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=tag.spend, ttl=60) + except Exception as redis_err: # noqa: BLE001 # best-effort refresh: a Redis failure must not fail the request + # tag.name is admin-supplied; strip CR/LF before it reaches the logs so it + # cannot forge additional log lines. + safe_counter_key: Final = counter_key.replace("\r", "").replace("\n", "") + verbose_proxy_logger.warning( + "Failed to update spend counter %s in Redis after tag spend update: %s. " + "Budget checks may use stale value until counter expires.", + safe_counter_key, + redis_err, + ) + # Build response tag_config: Final = TagConfig( name=updated_tag_record.tag_name, @@ -475,6 +507,7 @@ async def update_tag( created_at=updated_tag_record.created_at.isoformat(), updated_at=updated_tag_record.updated_at.isoformat(), created_by=updated_tag_record.created_by, + spend=updated_tag_record.spend, ) return { diff --git a/litellm/types/tag_management.py b/litellm/types/tag_management.py index f121f5bc562..db1069aafd3 100644 --- a/litellm/types/tag_management.py +++ b/litellm/types/tag_management.py @@ -12,6 +12,7 @@ class TagConfig(TagBase): created_at: str updated_at: str created_by: str | None = None + spend: float | None = None class TagNewRequest(TagBase): @@ -36,6 +37,7 @@ class TagUpdateRequest(TagBase): rpm_limit: int | None = None model_max_budget: dict | None = None budget_duration: str | None = None + spend: float | None = None class TagDeleteRequest(BaseModel): diff --git a/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py index d5dbd3df6a6..d927410e3ed 100644 --- a/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py @@ -241,6 +241,7 @@ async def test_new_tag_persists_a_budget(): tag_name="budget-tag", description=None, models=[], + spend=0.0, created_at=datetime(2024, 1, 1), updated_at=datetime(2024, 1, 1), created_by="admin", @@ -307,6 +308,7 @@ async def test_update_tag_clears_or_sets_only_provided_budget_fields( tag_name="budget-tag", description=None, models=[], + spend=0.0, created_at=datetime(2024, 1, 1), updated_at=datetime(2024, 1, 1), created_by="admin", @@ -374,6 +376,7 @@ async def test_update_tag_explicit_null_clears_budget_duration(): tag_name="budget-tag", description=None, models=[], + spend=0.0, created_at=datetime(2024, 1, 1), updated_at=datetime(2024, 1, 1), created_by="admin", @@ -609,6 +612,283 @@ async def test_update_tag_invalidates_only_the_tag_cache(): app.dependency_overrides.clear() +@pytest.mark.asyncio +async def test_update_tag_resets_spend(): + """POST /tag/update with spend writes it to the DB and returns it in the response.""" + from datetime import datetime + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + try: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, # test-quality-ok: the endpoint imports proxy_server.prisma_client itself; no parameter to inject + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id"), # test-quality-ok: the endpoint imports proxy_server.litellm_proxy_admin_name itself; no parameter to inject + patch("litellm.proxy.proxy_server.spend_counter_cache") as mock_spend_counter_cache, # test-quality-ok: the endpoint imports proxy_server.spend_counter_cache itself; no parameter to inject + ): + mock_spend_counter_cache.redis_cache = None + + mock_db = Mock() + mock_prisma.db = mock_db + + existing_tag = Mock() + existing_tag.tag_name = "batch-jobs" + existing_tag.description = "nightly batch jobs" + existing_tag.models = [] + existing_tag.budget_id = None + mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) + mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + updated_tag = Mock() + updated_tag.tag_name = "batch-jobs" + updated_tag.description = "nightly batch jobs" + updated_tag.models = [] + updated_tag.model_info = {} + updated_tag.spend = 0.0 + updated_tag.budget_id = None + updated_tag.created_at = datetime.now() + updated_tag.updated_at = datetime.now() + updated_tag.created_by = "test-user-123" + mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) + + response = client.post( + "/tag/update", + json={"name": "batch-jobs", "spend": 0}, + headers={"Authorization": "Bearer sk-1234"}, + ) + assert response.status_code == 200 + result = response.json() + assert result["tag"]["spend"] == 0.0 + + update_call_kwargs = mock_db.litellm_tagtable.update.call_args.kwargs + assert update_call_kwargs["data"]["spend"] == 0.0 + + mock_spend_counter_cache.in_memory_cache.set_cache.assert_called_once_with( + key="spend:tag:batch-jobs", value=0.0, ttl=60 + ) + finally: + app.dependency_overrides.clear() + + +@pytest.mark.asyncio +async def test_update_tag_resets_spend_in_redis(): + """POST /tag/update with spend also refreshes the counter in Redis when configured.""" + from datetime import datetime + from unittest.mock import AsyncMock, MagicMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + try: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, # test-quality-ok: the endpoint imports proxy_server.prisma_client itself; no parameter to inject + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id"), # test-quality-ok: the endpoint imports proxy_server.litellm_proxy_admin_name itself; no parameter to inject + patch("litellm.proxy.proxy_server.spend_counter_cache") as mock_spend_counter_cache, # test-quality-ok: the endpoint imports proxy_server.spend_counter_cache itself; no parameter to inject + ): + mock_spend_counter_cache.redis_cache = MagicMock() + mock_spend_counter_cache.redis_cache.async_set_cache = AsyncMock() + + mock_db = Mock() + mock_prisma.db = mock_db + + existing_tag = Mock() + existing_tag.tag_name = "batch-jobs" + existing_tag.description = "nightly batch jobs" + existing_tag.models = [] + existing_tag.budget_id = None + mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) + mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + updated_tag = Mock() + updated_tag.tag_name = "batch-jobs" + updated_tag.description = "nightly batch jobs" + updated_tag.models = [] + updated_tag.model_info = {} + updated_tag.spend = 0.0 + updated_tag.budget_id = None + updated_tag.created_at = datetime.now() + updated_tag.updated_at = datetime.now() + updated_tag.created_by = "test-user-123" + mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) + + response = client.post( + "/tag/update", + json={"name": "batch-jobs", "spend": 0}, + headers={"Authorization": "Bearer sk-1234"}, + ) + assert response.status_code == 200 + result = response.json() + assert result["tag"]["spend"] == 0.0 + + mock_spend_counter_cache.redis_cache.async_set_cache.assert_awaited_once_with( + key="spend:tag:batch-jobs", value=0.0, ttl=60 + ) + finally: + app.dependency_overrides.clear() + + +@pytest.mark.asyncio +async def test_update_tag_resets_spend_redis_failure_does_not_fail_request(): + """A Redis error while refreshing the counter must not fail the /tag/update request (best-effort).""" + from datetime import datetime + from unittest.mock import AsyncMock, MagicMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + try: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, # test-quality-ok: the endpoint imports proxy_server.prisma_client itself; no parameter to inject + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id"), # test-quality-ok: the endpoint imports proxy_server.litellm_proxy_admin_name itself; no parameter to inject + patch("litellm.proxy.proxy_server.spend_counter_cache") as mock_spend_counter_cache, # test-quality-ok: the endpoint imports proxy_server.spend_counter_cache itself; no parameter to inject + ): + mock_spend_counter_cache.redis_cache = MagicMock() + mock_spend_counter_cache.redis_cache.async_set_cache = AsyncMock( + side_effect=Exception("boom") + ) + + mock_db = Mock() + mock_prisma.db = mock_db + + existing_tag = Mock() + existing_tag.tag_name = "batch-jobs" + existing_tag.description = "nightly batch jobs" + existing_tag.models = [] + existing_tag.budget_id = None + mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) + mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + updated_tag = Mock() + updated_tag.tag_name = "batch-jobs" + updated_tag.description = "nightly batch jobs" + updated_tag.models = [] + updated_tag.model_info = {} + updated_tag.spend = 0.0 + updated_tag.budget_id = None + updated_tag.created_at = datetime.now() + updated_tag.updated_at = datetime.now() + updated_tag.created_by = "test-user-123" + mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) + + response = client.post( + "/tag/update", + json={"name": "batch-jobs", "spend": 0}, + headers={"Authorization": "Bearer sk-1234"}, + ) + assert response.status_code == 200 + result = response.json() + assert result["tag"]["spend"] == 0.0 + finally: + app.dependency_overrides.clear() + + +@pytest.mark.asyncio +async def test_update_tag_without_spend_does_not_touch_counter_cache(): + """A description-only update must not write to LiteLLM_TagTable.spend or the spend counter.""" + from datetime import datetime + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + try: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, # test-quality-ok: the endpoint imports proxy_server.prisma_client itself; no parameter to inject + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id"), # test-quality-ok: the endpoint imports proxy_server.litellm_proxy_admin_name itself; no parameter to inject + patch("litellm.proxy.proxy_server.spend_counter_cache") as mock_spend_counter_cache, # test-quality-ok: the endpoint imports proxy_server.spend_counter_cache itself; no parameter to inject + ): + mock_db = Mock() + mock_prisma.db = mock_db + + existing_tag = Mock() + existing_tag.tag_name = "batch-jobs" + existing_tag.description = "old description" + existing_tag.models = [] + existing_tag.budget_id = None + mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) + mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + updated_tag = Mock() + updated_tag.tag_name = "batch-jobs" + updated_tag.description = "new description" + updated_tag.models = [] + updated_tag.model_info = {} + updated_tag.spend = 5.0 + updated_tag.budget_id = None + updated_tag.created_at = datetime.now() + updated_tag.updated_at = datetime.now() + updated_tag.created_by = "test-user-123" + mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) + + response = client.post( + "/tag/update", + json={"name": "batch-jobs", "description": "new description"}, + headers={"Authorization": "Bearer sk-1234"}, + ) + assert response.status_code == 200 + result = response.json() + assert result["tag"]["spend"] == 5.0 + + update_call_kwargs = mock_db.litellm_tagtable.update.call_args.kwargs + assert "spend" not in update_call_kwargs["data"] + + mock_spend_counter_cache.in_memory_cache.set_cache.assert_not_called() + finally: + app.dependency_overrides.clear() + + +@pytest.mark.asyncio +async def test_update_tag_rejects_non_finite_spend(): + """NaN/inf must 400, not reach the DB or the spend counter.""" + from unittest.mock import AsyncMock, Mock + + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + try: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: the endpoint imports proxy_server.prisma_client itself; no parameter to inject + mock_db = Mock() + mock_prisma.db = mock_db + mock_db.litellm_tagtable.update = AsyncMock() + + # httpx's `json=` kwarg encodes with allow_nan=False and raises client-side + # before a request is even sent, so NaN must go over the wire as raw content + # (stdlib json.dumps allows NaN by default) to exercise the server's own check. + response = client.post( + "/tag/update", + content=json.dumps({"name": "batch-jobs", "spend": float("nan")}), + headers={ + "Authorization": "Bearer sk-1234", + "Content-Type": "application/json", + }, + ) + assert response.status_code == 400 + mock_db.litellm_tagtable.update.assert_not_called() + finally: + app.dependency_overrides.clear() + + @pytest.mark.asyncio async def test_delete_tag_invalidates_tag_and_registry_caches(): """Without this a deleted tag keeps its cached budget enforced until the TTL expires.""" diff --git a/tests/unit/types/test_tag_management.py b/tests/unit/types/test_tag_management.py new file mode 100644 index 00000000000..a496191346c --- /dev/null +++ b/tests/unit/types/test_tag_management.py @@ -0,0 +1,32 @@ +from litellm.types.tag_management import TagConfig, TagUpdateRequest + + +def test_tag_update_request_accepts_spend(): + request = TagUpdateRequest(name="batch-jobs", spend=0.0) + assert request.spend == 0.0 + assert "spend" in request.model_fields_set + + +def test_tag_update_request_spend_defaults_to_none(): + request = TagUpdateRequest(name="batch-jobs") + assert request.spend is None + assert "spend" not in request.model_fields_set + + +def test_tag_config_accepts_spend(): + config = TagConfig( + name="batch-jobs", + created_at="2026-01-01T00:00:00", + updated_at="2026-01-01T00:00:00", + spend=42.5, + ) + assert config.spend == 42.5 + + +def test_tag_config_spend_defaults_to_none(): + config = TagConfig( + name="batch-jobs", + created_at="2026-01-01T00:00:00", + updated_at="2026-01-01T00:00:00", + ) + assert config.spend is None diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 73336ca1823..fef5c935d0a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16338,6 +16338,7 @@ export interface paths { * - description: Optional[str] - Updated description * - models: List[str] - Updated list of allowed LLM models * - budget_id: Optional[str] - The id for a budget to associate with the tag + * - spend: Optional[float] - Set the tag's accumulated spend (e.g. 0 to reset it). Omit to leave unchanged. Does not change budget_reset_at; the current budget window is left as-is * * ### BUDGET UPDATE PARAMS ### * - max_budget: Optional[float] - Max budget for tag @@ -45443,6 +45444,8 @@ export interface components { rpm_limit?: number | null; /** Soft Budget */ soft_budget?: number | null; + /** Spend */ + spend?: number | null; /** Tpm Limit */ tpm_limit?: number | null; };