fix budget reservation window fallback races

This commit is contained in:
user 2026-04-30 16:26:43 -07:00
parent 1373ae1021
commit e034935b53
2 changed files with 98 additions and 1 deletions

View file

@ -85,6 +85,7 @@ async def reserve_budget_for_request(
route=route,
llm_router=llm_router,
)
using_remaining_budget_fallback = reservation_cost is None
if reservation_cost is None:
reservation_cost = await _get_smallest_remaining_budget(
counters=counters,
@ -111,6 +112,18 @@ async def reserve_budget_for_request(
cached_spend = await _get_current_counter_value(counter=counter)
current_spend = cached_spend + reservation_cost
if current_spend > counter.max_budget:
if using_remaining_budget_fallback:
remaining_before_reservation = counter.max_budget - (
current_spend - reservation_cost
)
if remaining_before_reservation > 0:
await _resize_applied_reservation(
entries=applied_entries,
current_reserved_cost=reservation_cost,
new_reserved_cost=remaining_before_reservation,
)
reservation_cost = remaining_before_reservation
continue
raise litellm.BudgetExceededError(
current_cost=current_spend,
max_budget=counter.max_budget,
@ -632,6 +645,19 @@ async def _set_reserved_entries_adjustment(
entry["applied_adjustment"] = target_adjustment
async def _resize_applied_reservation(
entries: List[dict],
current_reserved_cost: float,
new_reserved_cost: float,
) -> None:
await _set_reserved_entries_adjustment(
entries=entries,
target_adjustment=new_reserved_cost - current_reserved_cost,
)
for entry in entries:
entry["applied_adjustment"] = 0.0
def _counter_to_reservation_entry(counter: _BudgetCounter) -> Dict[str, Any]:
return {
"counter_key": counter.counter_key,
@ -653,7 +679,7 @@ def get_budget_window_start(window: Any) -> Optional[datetime]:
reset_at = _coerce_datetime(window_dict.get("reset_at"))
if reset_at is None:
reset_at = datetime.now(timezone.utc) + timedelta(seconds=duration_seconds)
return datetime.now(timezone.utc) - timedelta(seconds=duration_seconds)
if reset_at.tzinfo is None:
reset_at = reset_at.replace(tzinfo=timezone.utc)
return reset_at - timedelta(seconds=duration_seconds)

View file

@ -1,3 +1,4 @@
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -16,6 +17,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.spend_tracking.budget_reservation import (
estimate_request_max_cost,
get_budget_window_start,
invalidate_budget_reservation_counters,
release_budget_reservation,
reserve_budget_for_request,
@ -478,6 +480,75 @@ async def test_should_reserve_remaining_budget_when_output_cap_missing(
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_should_shrink_uncapped_reservation_when_counter_advances(
spend_counter_state,
monkeypatch,
):
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-budget-uncapped-race",
spend=0.2,
max_budget=1.0,
)
request_body = _request_body()
request_body.pop("max_tokens")
from litellm.proxy.spend_tracking import budget_reservation
async def stale_counter_read(counter):
await counter_cache.async_increment_cache(
key=counter.counter_key,
value=0.3,
)
return 0.2
monkeypatch.setattr(
budget_reservation,
"_get_current_counter_value",
stale_counter_read,
)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=None,
):
reservation = await reserve_budget_for_request(
request_body=request_body,
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert reservation is not None
assert reservation["reserved_cost"] == pytest.approx(0.7)
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-uncapped-race"
) == pytest.approx(1.0)
await release_budget_reservation(reservation)
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-uncapped-race"
) == pytest.approx(0.3)
def test_should_start_window_without_reset_at_at_duration_boundary():
before = datetime.now(timezone.utc) - timedelta(hours=1)
window_start = get_budget_window_start({"budget_duration": "1h"})
after = datetime.now(timezone.utc) - timedelta(hours=1)
assert window_start is not None
assert before <= window_start <= after
@pytest.mark.asyncio
async def test_should_seed_malformed_window_counter_from_parent_authoritative_spend(
spend_counter_state,