From ba3536520f836b8916740983b7f0053a47f66d89 Mon Sep 17 00:00:00 2001 From: Moe Khalil Date: Tue, 22 Sep 2026 03:15:46 +0000 Subject: [PATCH] fix(fusion): retain unknown cancelled-call costs and pass current gates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/fusion_router.py | 7 ++- litellm/litellm_core_utils/fusion_budget.py | 12 ----- .../proxy/hooks/proxy_track_cost_callback.py | 46 +++++++++---------- .../model_management_endpoints.py | 3 +- tests/test_litellm/test_fusion_router.py | 4 +- 5 files changed, 29 insertions(+), 43 deletions(-) diff --git a/litellm/fusion_router.py b/litellm/fusion_router.py index 0dc43f38240..af7638b0032 100644 --- a/litellm/fusion_router.py +++ b/litellm/fusion_router.py @@ -15,7 +15,6 @@ from litellm.constants import ( INTERNAL_CALL_ORIGIN_METADATA_KEY, ) from litellm.litellm_core_utils.fusion_budget import ( - cancel_fusion_budget_call, complete_fusion_budget_call, register_fusion_budget_call, wait_for_fusion_budget_calls, @@ -910,7 +909,7 @@ class FusionRouter: _fusion_proxy_auth_required=isinstance(request_kwargs.get("proxy_server_request"), Mapping), ) except asyncio.CancelledError: - cancel_fusion_budget_call(metadata) + complete_fusion_budget_call(metadata, cost_known=False) raise except Exception: complete_fusion_budget_call(metadata, cost_known=True) @@ -968,7 +967,7 @@ class FusionRouter: try: response = await self._completion(model=model, messages=current_messages, stream=False, **call_kwargs) except asyncio.CancelledError: - cancel_fusion_budget_call(call_metadata) + complete_fusion_budget_call(call_metadata, cost_known=False) raise except Exception: complete_fusion_budget_call(call_metadata, cost_known=True) @@ -1050,7 +1049,7 @@ class FusionRouter: **kwargs, ) except asyncio.CancelledError: - cancel_fusion_budget_call(metadata) + complete_fusion_budget_call(metadata, cost_known=False) raise except Exception: complete_fusion_budget_call(metadata, cost_known=True) diff --git a/litellm/litellm_core_utils/fusion_budget.py b/litellm/litellm_core_utils/fusion_budget.py index 978509eca64..06550dc4b9d 100644 --- a/litellm/litellm_core_utils/fusion_budget.py +++ b/litellm/litellm_core_utils/fusion_budget.py @@ -71,18 +71,6 @@ def complete_fusion_budget_call( unpriced.append(token) -def cancel_fusion_budget_call(metadata: Mapping[str, object]) -> None: - """Finish a call deliberately cancelled by Fusion without marking the whole request unpriced. - - Panel and analyst timeouts actively cancel their in-flight child call. They - are different from a completed provider call whose cost callback went - missing: the latter must retain the conservative full-reservation fallback, - while the former must not turn one timed-out advisory member into a charge - for every possible Fusion call. - """ - complete_fusion_budget_call(metadata, cost_known=True) - - async def wait_for_fusion_budget_calls( metadata: Mapping[str, object], *, diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index f898f53b8e3..6308fc6f5b2 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -106,6 +106,21 @@ def _response_requires_fusion_continuation(response: object) -> bool: return bool(tool_calls) and all(tool_call["name"] == _FUSION_TOOL_NAME for tool_call in tool_calls) +async def _fusion_budget_counter_cost(metadata: Mapping[str, object], response_cost: float) -> float: + reservation: Final = budget_reservation_from_metadata(metadata) + if ( + reservation is None + or reservation.get(FUSION_BUDGET_ACTIVE_KEY) is not True + or metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != "fusion_continuation" + ): + return response_cost + await wait_for_fusion_budget_calls(metadata) + return fusion_budget_reconciliation_cost( + budget_reservation=reservation, + known_cost=response_cost + float(reservation.get(FUSION_BUDGET_ACCUMULATED_COST_KEY) or 0.0), + ) + + def _should_defer_fusion_budget_reconciliation( metadata: dict, # mutable-ok: SDK boundary completion_response: object, @@ -237,8 +252,13 @@ class _ProxyDBLogger(CustomLogger): traceback_str: str | None = None, ): try: - if not _failure_should_leave_fusion_reservation_open(request_data): - await _release_budget_reservation(budget_reservation=user_api_key_dict.budget_reservation) + await _release_budget_reservation( + budget_reservation=( + None + if _failure_should_leave_fusion_reservation_open(request_data) + else user_api_key_dict.budget_reservation + ) + ) except Exception: verbose_proxy_logger.exception("Failed to release budget reservation during failure handling") try: @@ -407,12 +427,6 @@ class _ProxyDBLogger(CustomLogger): ) _write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata) budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata) - if ( - budget_reservation is not None - and budget_reservation.get(FUSION_BUDGET_ACTIVE_KEY) is True - and metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == "fusion_continuation" - ): - await wait_for_fusion_budget_calls(metadata) if ( isinstance(completion_response, LiteLLMBatch) and kwargs.get("call_type") == CallTypes.aretrieve_batch.value @@ -468,21 +482,7 @@ class _ProxyDBLogger(CustomLogger): # sufficient for the final-call barrier; persistence and alerts # can continue without delaying the model orchestration. complete_fusion_budget_call(metadata, cost_known=True) - known_fusion_cost: Final = float(response_cost) + ( - float(budget_reservation.get(FUSION_BUDGET_ACCUMULATED_COST_KEY) or 0.0) - if budget_reservation is not None - else 0.0 - ) - budget_counter_response_cost: Final = ( - fusion_budget_reconciliation_cost( - budget_reservation=budget_reservation, - known_cost=known_fusion_cost, - ) - if budget_reservation is not None - and budget_reservation.get(FUSION_BUDGET_ACTIVE_KEY) is True - and metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == "fusion_continuation" - else float(response_cost) - ) + budget_counter_response_cost: Final = await _fusion_budget_counter_cost(metadata, float(response_cost)) user_api_key: Final = metadata.get("user_api_key", None) verbose_proxy_logger.debug( "user_api_key %s, user_id %s, team_id %s, end_user_id %s", diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 6d99bebc1e7..4ea969a1b4c 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -299,7 +299,6 @@ def _strategy_router_write_violation( return None incoming_model: Final = incoming_params.model existing_model: Final = existing_params.model if existing_params is not None else None - effective_model: Final = incoming_model or existing_model incoming_fusion_config: Final = incoming_params.fusion_router_config existing_fusion_config: Final = existing_params.fusion_router_config if existing_params is not None else None if ( @@ -308,7 +307,7 @@ def _strategy_router_write_violation( or (isinstance(existing_model, str) and is_fusion_router_model(existing_model)) ): fusion_violation: Final = validate_fusion_router_write( - model=effective_model, + model=incoming_model or existing_model, raw_config=(incoming_fusion_config if incoming_fusion_config is not None else existing_fusion_config), ) if fusion_violation is not None: diff --git a/tests/test_litellm/test_fusion_router.py b/tests/test_litellm/test_fusion_router.py index f4dbeab2d22..45e121d4077 100644 --- a/tests/test_litellm/test_fusion_router.py +++ b/tests/test_litellm/test_fusion_router.py @@ -657,8 +657,8 @@ async def test_analyst_timeout_degrades_to_raw_panel_responses() -> None: assert completion.analyst_started.is_set() assert completion.analyst_cancelled.is_set() assert response._hidden_params["fusion"]["analysis_available"] is False - assert reservation.get(FUSION_BUDGET_UNPRICED_CALL_IDS_KEY) in (None, []) - assert fusion_budget_reconciliation_cost(reservation, known_cost=0.4) == pytest.approx(0.4) + assert len(reservation[FUSION_BUDGET_UNPRICED_CALL_IDS_KEY]) == 1 + assert fusion_budget_reconciliation_cost(reservation, known_cost=0.4) == pytest.approx(reservation["reserved_cost"]) payload = json.loads(completion.calls[-1]["messages"][-1]["content"]) assert [item["content"] for item in payload["responses"]] == ["Panel A", "Panel B"]