Merge pull request #40310 from BerriAI/litellm_lit7223_reconcile_before_db

fix(proxy): reconcile budget reservation before enqueuing spend to the DB
This commit is contained in:
Yassin Kortam 2026-09-15 13:57:55 -07:00 • committed by GitHub
commit 1debb438f5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 263 additions and 3 deletions

View file

@ -652,6 +652,10 @@ async def _update_database_and_spend_counters(
request_tags: list[str] | None = None,
model_access_groups: Sequence[str] | None = None,
) -> bool:
if budget_reservation is not None:
await _reconcile_budget_reservation_before_db_update(
budget_reservation=budget_reservation, response_cost=response_cost
)
try:
charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
@ -709,6 +713,30 @@ async def _update_database_and_spend_counters(
return True
async def _reconcile_budget_reservation_before_db_update(
budget_reservation: dict, # mutable-ok: reconcile_budget_reservation stamps applied_adjustment on the caller's shared reservation dict
response_cost: float,
) -> None:
from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation
try:
await reconcile_budget_reservation(
budget_reservation=budget_reservation, actual_cost=response_cost, finalize=False
)
except Exception: # noqa: BLE001 # a failed reconcile must not block the spend write; the counters are dropped instead
verbose_proxy_logger.warning(
"Failed to reconcile budget reservation before persisting spend; invalidating reserved counters"
)
try:
await _invalidate_budget_reservation_counters(budget_reservation=budget_reservation)
except Exception: # noqa: BLE001 # nothing left to try; the finalized stamp below keeps it from being reprocessed
verbose_proxy_logger.exception(
"Failed to invalidate budget reservation counters after pre-persist reconcile failed"
)
finally:
budget_reservation["finalized"] = True # rebind-ok: the counter update reads the stamp off the shared dict
async def _release_budget_reservation(budget_reservation: dict | None) -> None:
if budget_reservation is None:
return

View file

@ -3061,7 +3061,7 @@ async def _reconcile_budget_reservation_for_counter_update(
budget_reservation: dict | None,
response_cost: float | None,
) -> set[str]:
if budget_reservation is None:
if budget_reservation is None or budget_reservation.get("finalized") is True:
return set()
from litellm.proxy.spend_tracking.budget_reservation import (

View file

@ -959,8 +959,9 @@ async def _set_reserved_entries_actual_cost(
async def _reseed_reserved_entry(item: _EntryAdjustment, actual_cost: float) -> None:
"""Post-call reconcile / release of a counter that was flushed, expired or reseeded between reservation and
reconcile: the optimistic delta no longer applies, so reseed from the DB floor (which cannot include this
request's cost yet) and add the settled cost, since increment_spend_counters skips reserved keys."""
reconcile: the optimistic delta no longer applies, so reseed from the DB floor and add the settled cost, since
increment_spend_counters skips reserved keys. The reconcile runs before this request's spend is enqueued to the
DB, so the reseeded floor excludes it."""
from litellm.proxy.proxy_server import _increment_spend_counter_cache, reseed_spend_counter_from_db
reseeded: Final = await reseed_spend_counter_from_db(counter_key=item.counter_key)

View file

@ -678,6 +678,149 @@ async def test_update_database_and_spend_counters_preserves_counter_exception_wh
proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_reconciles_reservation_before_db_update():
call_order: list[str] = []
proxy_logging_obj = MagicMock()
async def _update_database(**kwargs):
call_order.append("update_database")
return True
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=_update_database)
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
async def _reconcile(**kwargs):
call_order.append("reconcile")
with patch( # test-quality-ok: the helper imports reconcile_budget_reservation in its body, no injection seam
"litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation",
new_callable=AsyncMock,
side_effect=_reconcile,
) as mock_reconcile_budget_reservation:
charged = await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
)
assert charged is True
assert call_order == ["reconcile", "update_database"]
mock_reconcile_budget_reservation.assert_awaited_once_with(
budget_reservation=budget_reservation,
actual_cost=0.2,
finalize=False,
)
increment_spend_counters.assert_awaited_once()
assert increment_spend_counters.await_args.kwargs["budget_reservation"] is budget_reservation
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails_after_early_reconcile():
proxy_logging_obj = MagicMock()
db_exception = RuntimeError("db unavailable")
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=db_exception)
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
with (
patch( # test-quality-ok: the helper imports reconcile_budget_reservation in its body, no injection seam
"litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation",
new_callable=AsyncMock,
) as mock_reconcile_budget_reservation,
patch( # test-quality-ok: _release_budget_reservation imports the release in its body, no injection seam
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation",
new_callable=AsyncMock,
) as mock_release_budget_reservation,
):
with pytest.raises(RuntimeError) as exc_info:
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
)
assert exc_info.value is db_exception
mock_reconcile_budget_reservation.assert_awaited_once_with(
budget_reservation=budget_reservation,
actual_cost=0.2,
finalize=False,
)
mock_release_budget_reservation.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
increment_spend_counters.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_invalidates_reservation_when_early_reconcile_fails():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(return_value=True)
increment_spend_counters = AsyncMock()
budget_reservation = {
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:test_api_key"}],
}
with (
patch( # test-quality-ok: the helper imports reconcile_budget_reservation in its body, no injection seam
"litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation",
new_callable=AsyncMock,
side_effect=RuntimeError("redis unavailable"),
) as mock_reconcile_budget_reservation,
patch( # test-quality-ok: _invalidate_budget_reservation_counters imports it in its body, no injection seam
"litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters",
new_callable=AsyncMock,
) as mock_invalidate_budget_reservation_counters,
):
charged = await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
)
assert charged is True
mock_reconcile_budget_reservation.assert_awaited_once()
mock_invalidate_budget_reservation_counters.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
assert budget_reservation["finalized"] is True
proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
increment_spend_counters.assert_awaited_once()
@pytest.mark.asyncio
async def test_track_cost_callback_skips_when_no_standard_logging_object():
"""

View file

@ -921,6 +921,30 @@ async def test_reconcile_budget_reservation_for_counter_update_failure_invalidat
assert fake_invalidate.called is True
@pytest.mark.asyncio
async def test_reconcile_budget_reservation_for_counter_update_finalized_reservation_falls_back_to_direct_increment(
monkeypatch,
):
"""A reservation already finalized before the counter update (the pre-persist
reconcile failed and dropped its counters) must not shield its keys from the
direct increment, or the settled cost is never added back after the drop."""
import litellm.proxy.spend_tracking.budget_reservation as br
fake_reconcile = AsyncMock()
monkeypatch.setattr(br, "reconcile_budget_reservation", fake_reconcile)
result = await ps._reconcile_budget_reservation_for_counter_update(
budget_reservation={
"finalized": True,
"entries": [{"counter_key": "spend:key:abc"}],
},
response_cost=1.0,
)
assert result == set()
fake_reconcile.assert_not_awaited()
# ---------------------------------------------------------------------------
# _prepare_end_user_and_tag_spend_increments
# ---------------------------------------------------------------------------

View file

@ -2230,6 +2230,17 @@ class _ExpiringRedisCache:
return None
class _TeamMembershipFloorDb:
"""Stands in for `prisma_client.db`: only the team-membership row exists and its spend is the DB floor."""
def __init__(self, spend: float) -> None:
self.spend = spend
def __getattr__(self, table_name: str) -> SimpleNamespace:
row = SimpleNamespace(spend=self.spend) if table_name == "litellm_teammembership" else None
return SimpleNamespace(find_unique=AsyncMock(return_value=row))
@pytest.mark.asyncio
async def test_reconcile_after_redis_counter_expiry_keeps_request_cost_enforced(
spend_counter_state,
@ -2275,6 +2286,59 @@ async def test_reconcile_after_redis_counter_expiry_keeps_request_cost_enforced(
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_reconcile_before_db_update_does_not_double_count_when_flush_lands_between_passes(
spend_counter_state,
):
"""The early reconcile (before the spend row is enqueued) reseeds from a DB
floor that cannot yet include this request. When the periodic flush commits
the row before increment_spend_counters runs its second reconcile, the
applied_adjustment early-return must keep the counter from adding the cost
a second time."""
import litellm.proxy.proxy_server as ps
from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation
counter_cache, _ = spend_counter_state
counter_key = "spend:team_member:user-flush:team-flush"
redis_cache = _ExpiringRedisCache()
counter_cache.redis_cache = redis_cache
counter_cache.in_memory_cache.set_cache(key=counter_key, value=0.6)
db_floor = _TeamMembershipFloorDb(spend=0.3)
ps.prisma_client = SimpleNamespace(db=db_floor)
reservation = {
"reserved_cost": 0.6,
"entries": [
{
"counter_key": counter_key,
"entity_type": "TeamMember",
"entity_id": "user-flush:team-flush",
"reserved_cost": 0.6,
"applied_adjustment": 0.0,
}
],
"finalized": False,
}
await reconcile_budget_reservation(budget_reservation=reservation, actual_cost=0.05, finalize=False)
assert redis_cache.store[counter_key] == pytest.approx(0.35)
assert reservation["entries"][0]["applied_adjustment"] == pytest.approx(-0.55)
assert reservation["finalized"] is False
db_floor.spend = 0.35
await ps.increment_spend_counters(
token="key-flush",
team_id="team-flush",
user_id="user-flush",
response_cost=0.05,
budget_reservation=reservation,
)
assert redis_cache.store[counter_key] == pytest.approx(0.35)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_should_invalidate_reserved_counters_after_persisted_spend_failure(
spend_counter_state,