fix(router): make fusion budget updates idempotent

This commit is contained in:
moe-berri 2026-09-03 15:47:49 -07:00
parent 0a26e0cc51
commit 52e1c3abf4
4 changed files with 123 additions and 10 deletions

View file

@ -4380,9 +4380,22 @@ async def can_key_call_resolved_model(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if matched_model_access_groups:
# A logical request may already have reserved its worst-case cost against
# the group that also serves this resolved model (Fusion dependencies are
# one example). Re-reading that live counter would mistake this request's
# own reservation for exhausted capacity. The reservation already secured
# the request; only newly encountered groups still need a read-time check.
from litellm.proxy.spend_tracking.budget_reservation import get_reserved_counter_keys
reserved_counter_keys: Final = get_reserved_counter_keys(budget_reservation=valid_token.budget_reservation)
unreserved_model_access_groups: Final = tuple(
group
for group in matched_model_access_groups
if model_access_group_spend_counter_key(group) not in reserved_counter_keys
)
if unreserved_model_access_groups:
await _model_access_group_max_budget_check(
matched_model_access_groups=matched_model_access_groups,
matched_model_access_groups=unreserved_model_access_groups,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)

View file

@ -122,8 +122,8 @@ def _accumulate_fusion_cost(
budget_reservation: dict, # mutable-ok: SDK boundary
response_cost: float,
kwargs: dict, # mutable-ok: SDK boundary
) -> None:
"""Add one hidden call exactly once before its asynchronous DB write."""
) -> bool:
"""Add one hidden call exactly once and report whether this callback was new."""
call_id: Final = kwargs.get("litellm_call_id") or kwargs.get("id")
seen_call_ids: Final = budget_reservation.setdefault(
FUSION_BUDGET_ACCUMULATED_CALL_IDS_KEY,
@ -132,7 +132,7 @@ def _accumulate_fusion_cost(
if isinstance(seen_call_ids, list) and call_id is not None:
normalized_call_id: Final = str(call_id)
if normalized_call_id in seen_call_ids:
return
return False
seen_call_ids.append(normalized_call_id)
budget_reservation[ # rebind-ok: shared reservation ledger
FUSION_BUDGET_ACCUMULATED_COST_KEY
@ -142,6 +142,7 @@ def _accumulate_fusion_cost(
)
+ max(response_cost, 0.0)
)
return True
def _failure_should_leave_fusion_reservation_open(
@ -383,12 +384,15 @@ class _ProxyDBLogger(CustomLogger):
)
if response_cost is not None:
if defer_fusion_reconciliation and budget_reservation is not None:
fusion_call_should_charge_access_groups: Final = (
_accumulate_fusion_cost(
budget_reservation=budget_reservation,
response_cost=float(response_cost),
kwargs=kwargs,
)
if defer_fusion_reconciliation and budget_reservation is not None
else True
)
budget_counter_response_cost: Final = (
float(response_cost) + float(budget_reservation.get(FUSION_BUDGET_ACCUMULATED_COST_KEY) or 0.0)
if budget_reservation is not None
@ -431,7 +435,11 @@ class _ProxyDBLogger(CustomLogger):
budget_counter_response_cost=budget_counter_response_cost,
defer_budget_counter_update=defer_fusion_reconciliation,
request_tags=tags,
model_access_groups=model_access_groups,
# The accumulator's call-id ledger owns idempotency for
# the whole hidden call, including its deployment-group
# charge. A duplicate callback may still be persisted as
# before, but it must not debit the live budget twice.
model_access_groups=(model_access_groups if fusion_call_should_charge_access_groups else ()),
)
# update cache (fire-and-forget for backward compat:

View file

@ -410,6 +410,84 @@ async def test_resolved_model_authorization_enforces_its_access_group_budget(
assert seen == [MODEL_ACCESS_GROUP_COUNTER_KEY]
@pytest.mark.asyncio
async def test_resolved_model_authorization_does_not_recheck_its_own_reserved_group(
monkeypatch: pytest.MonkeyPatch,
):
"""A Fusion dependency can share the virtual model's already-reserved group."""
from litellm.proxy import proxy_server
cache = await _cache()
prisma = _RecordingPrismaClient(_MagBudgetRow("tier-a", spend=9.0, max_budget=10.0))
router = Router(model_list=MODEL_LIST)
valid_token = UserAPIKeyAuth(
api_key="hashed",
models=["tier-a"],
budget_reservation={
"entries": [{"counter_key": MODEL_ACCESS_GROUP_COUNTER_KEY, "reserved_cost": 1.0}],
},
)
read, seen = _spend_reader({MODEL_ACCESS_GROUP_COUNTER_KEY: 10.0})
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=cache))
monkeypatch.setattr(proxy_server, "get_current_spend", read)
assert await can_key_call_resolved_model(
model="gpt-4o",
llm_model_list=MODEL_LIST,
valid_token=valid_token,
llm_router=router,
) == ("tier-a",)
assert seen == []
@pytest.mark.asyncio
async def test_resolved_model_authorization_still_checks_a_new_dependency_group(
monkeypatch: pytest.MonkeyPatch,
):
"""A reservation for one shared group must not exempt another dependency group."""
from litellm.proxy import proxy_server
cache = await _cache()
prisma = _RecordingPrismaClient(
_MagBudgetRow("tier-a", spend=9.0, max_budget=10.0),
_MagBudgetRow("tier-b", spend=2.0, max_budget=2.0),
)
router = Router(model_list=MODEL_LIST)
valid_token = UserAPIKeyAuth(
api_key="hashed",
models=["tier-a", "tier-b"],
budget_reservation={
"entries": [{"counter_key": MODEL_ACCESS_GROUP_COUNTER_KEY, "reserved_cost": 1.0}],
},
)
tier_b_counter_key = model_access_group_spend_counter_key("tier-b")
read, seen = _spend_reader(
{
MODEL_ACCESS_GROUP_COUNTER_KEY: 10.0,
tier_b_counter_key: 2.0,
}
)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=cache))
monkeypatch.setattr(proxy_server, "get_current_spend", read)
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await can_key_call_resolved_model(
model="gpt-4o",
llm_model_list=MODEL_LIST,
valid_token=valid_token,
llm_router=router,
)
assert exc_info.value.entity_id == "tier-b"
assert seen == [tier_b_counter_key]
@pytest.mark.asyncio
async def test_group_just_under_its_max_budget_passes():
"""Asserting the counter was read is what keeps this honest: a group that got skipped entirely,

View file

@ -680,6 +680,7 @@ async def test_fusion_hidden_costs_accumulate_then_continuation_reconciles_once(
"user_api_key_user_id": "user-1",
"internal_call_origin": origin,
"user_api_key_budget_reservation": reservation,
MODEL_ACCESS_GROUP_METADATA_KEY: ["test-budget"],
}
},
"standard_logging_object": {
@ -690,9 +691,18 @@ async def test_fusion_hidden_costs_accumulate_then_continuation_reconciles_once(
}
with (
patch("litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock) as increment, # test-quality-ok: isolates proxy persistence while reservation state remains observable
patch("litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock), # test-quality-ok: isolates proxy persistence while reservation state remains observable
patch("litellm.proxy.proxy_server.proxy_logging_obj") as proxy_logging, # test-quality-ok: injects the callback persistence boundary
patch( # test-quality-ok: isolates proxy persistence while reservation state remains observable
"litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock
) as increment,
patch( # test-quality-ok: observes hidden-call group idempotency without touching global counters
"litellm.proxy.proxy_server.increment_fusion_model_access_group_spend_counters", new_callable=AsyncMock
) as increment_fusion_groups,
patch( # test-quality-ok: isolates proxy persistence while reservation state remains observable
"litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock
),
patch( # test-quality-ok: injects the callback persistence boundary
"litellm.proxy.proxy_server.proxy_logging_obj"
) as proxy_logging,
):
proxy_logging.db_spend_update_writer.update_database = AsyncMock()
proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
@ -721,6 +731,10 @@ async def test_fusion_hidden_costs_accumulate_then_continuation_reconciles_once(
assert reservation[FUSION_BUDGET_ACCUMULATED_COST_KEY] == pytest.approx(0.3)
increment.assert_not_awaited()
assert increment_fusion_groups.await_count == 2
assert [call.kwargs["response_cost"] for call in increment_fusion_groups.await_args_list] == pytest.approx(
[0.1, 0.2]
)
await logger._PROXY_track_cost_callback(
kwargs=kwargs_for("fusion_continuation", 0.4, "continuation-call"),