diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index c23f444595c..b17251f4b55 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 481d88b8d04..51016126d35 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -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,