mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(router): make fusion budget updates idempotent
This commit is contained in:
parent
0a26e0cc51
commit
52e1c3abf4
4 changed files with 123 additions and 10 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue