mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 19a854aeaa into a99bccacea
This commit is contained in:
commit
141c818cdb
5 changed files with 351 additions and 1 deletions
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
32
tests/unit/types/test_tag_management.py
Normal file
32
tests/unit/types/test_tag_management.py
Normal 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
|
||||
3
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
3
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue