This commit is contained in:
Arnold Gálovics 2026-10-05 20:10:19 +02:00 • committed by GitHub
commit 141c818cdb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 351 additions and 1 deletions

View file

@ -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:<hash> - 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 {

View file

@ -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):

View file

@ -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."""

View file

@ -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

View file

@ -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;
};