diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 57a43b43ee6..0fb4009b207 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, ) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index a511b715b75..4bdeb46fb1c 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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: diff --git a/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py b/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py index d1e65d781ce..6135a4c6996 100644 --- a/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py +++ b/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py @@ -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, diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 971edf12282..4f63f37c7a5 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -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"),