mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(budget): recompute budget_reset_at when budget_duration changes on /budget/update (#30555)
POST /budget/update did not recompute budget_reset_at when budget_duration changed and no explicit budget_reset_at was supplied, leaving shortened budgets pinned to the old (longer) schedule. The same defect reached POST /team/update via team_member_budget_duration, which delegates to update_budget. update_budget now recomputes budget_reset_at = get_budget_reset_time(duration) when the caller sets budget_duration without pinning budget_reset_at, mirroring /budget/new. Explicit budget_reset_at is preserved and updates that omit budget_duration leave the reset untouched. get_budget_reset_time now declares its datetime return type so the recomputed value stays concretely typed. Resolves LIT-3362
This commit is contained in:
parent
38a9a60484
commit
96b9437bdd
3 changed files with 110 additions and 2 deletions
|
|
@ -15,7 +15,7 @@ def get_budget_reset_timezone():
|
|||
return getattr(litellm, "timezone", None) or "UTC"
|
||||
|
||||
|
||||
def get_budget_reset_time(budget_duration: str):
|
||||
def get_budget_reset_time(budget_duration: str) -> datetime:
|
||||
"""
|
||||
Get the budget reset time based on the configured timezone.
|
||||
Falls back to UTC if not specified.
|
||||
|
|
|
|||
|
|
@ -183,10 +183,23 @@ async def update_budget(
|
|||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail={"error": str(e)})
|
||||
|
||||
# recompute budget_reset_at when the duration changes, unless the caller pinned a reset time explicitly
|
||||
recomputed_reset_at = (
|
||||
{
|
||||
"budget_reset_at": get_budget_reset_time(
|
||||
budget_duration=budget_obj.budget_duration
|
||||
)
|
||||
}
|
||||
if budget_obj.budget_duration is not None
|
||||
and "budget_reset_at" not in budget_obj.model_fields_set
|
||||
else {}
|
||||
)
|
||||
|
||||
response = await BudgetRepository(prisma_client).table.update(
|
||||
where={"budget_id": budget_obj.budget_id},
|
||||
data={
|
||||
**budget_obj.model_dump(exclude_unset=True), # type: ignore
|
||||
**recomputed_reset_at,
|
||||
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
}, # type: ignore
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
import os
|
||||
import sys
|
||||
import types
|
||||
from datetime import datetime, timedelta, timezone
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from fastapi.testclient import TestClient
|
||||
|
|
@ -11,7 +12,6 @@ import litellm.proxy.proxy_server as ps
|
|||
from litellm.proxy.proxy_server import app
|
||||
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles, CommonProxyErrors
|
||||
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../")
|
||||
) # Adds the parent directory to the system path
|
||||
|
|
@ -265,3 +265,98 @@ async def test_new_budget_invalid_model_max_budget(client_and_mocks, monkeypatch
|
|||
assert resp.status_code in (400, 422), resp.text
|
||||
detail = resp.json()["detail"]
|
||||
assert "model_max_budget" in str(detail) or "dictionary" in str(detail).lower()
|
||||
|
||||
|
||||
def _capture_update_data(mock_table):
|
||||
captured = {}
|
||||
|
||||
async def capture(*, where, data):
|
||||
captured.update(data)
|
||||
return {**where, **data}
|
||||
|
||||
mock_table.update = AsyncMock(side_effect=capture)
|
||||
return captured
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_budget_recomputes_reset_at_when_duration_changes(
|
||||
client_and_mocks,
|
||||
):
|
||||
"""
|
||||
Regression for LIT-3362: shortening budget_duration without an explicit
|
||||
budget_reset_at must bring the reset forward instead of leaving it pinned
|
||||
to the previous (longer) schedule.
|
||||
"""
|
||||
client, _, mock_table = client_and_mocks
|
||||
captured = _capture_update_data(mock_table)
|
||||
|
||||
before = datetime.now(timezone.utc)
|
||||
resp = client.post(
|
||||
"/budget/update",
|
||||
json={"budget_id": "budget_reset_recompute", "budget_duration": "1d"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
|
||||
assert (
|
||||
"budget_reset_at" in captured
|
||||
), "duration change must recompute budget_reset_at"
|
||||
reset_at = captured["budget_reset_at"]
|
||||
assert isinstance(reset_at, datetime)
|
||||
assert reset_at > before, "recomputed reset must be in the future"
|
||||
# "1d" resets at the next standardized day boundary, always within ~24h
|
||||
assert reset_at <= before + timedelta(days=1, hours=1), reset_at
|
||||
# and it must be far closer than a stale 30d schedule would have left it
|
||||
assert reset_at < before + timedelta(days=29)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_budget_preserves_explicit_reset_at(client_and_mocks):
|
||||
"""An explicit budget_reset_at from the caller always wins over recompute."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
captured = _capture_update_data(mock_table)
|
||||
|
||||
explicit = datetime(2027, 1, 1, tzinfo=timezone.utc)
|
||||
resp = client.post(
|
||||
"/budget/update",
|
||||
json={
|
||||
"budget_id": "budget_explicit_reset",
|
||||
"budget_duration": "1d",
|
||||
"budget_reset_at": explicit.isoformat(),
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
|
||||
assert captured["budget_reset_at"] == explicit
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_budget_without_duration_leaves_reset_at_untouched(
|
||||
client_and_mocks,
|
||||
):
|
||||
"""Updates that do not touch budget_duration must not introduce budget_reset_at."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
captured = _capture_update_data(mock_table)
|
||||
|
||||
resp = client.post(
|
||||
"/budget/update",
|
||||
json={"budget_id": "budget_other_field", "max_budget": 200.0},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
|
||||
assert "budget_reset_at" not in captured
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_budget_duration_none_does_not_recompute(client_and_mocks):
|
||||
"""Clearing budget_duration (explicit null) must not recompute against a None duration."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
captured = _capture_update_data(mock_table)
|
||||
|
||||
resp = client.post(
|
||||
"/budget/update",
|
||||
json={"budget_id": "budget_clear_duration", "budget_duration": None},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
|
||||
assert "budget_duration" in captured and captured["budget_duration"] is None
|
||||
assert "budget_reset_at" not in captured
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue