diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 687badfa09f..6a9c3a095e6 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -85,18 +85,21 @@ def _mapping_or_attribute(value: object, key: str) -> object: return getattr(value, key, None) -def _response_invoked_fusion(response: object) -> bool: +def _response_requires_fusion_continuation(response: object) -> bool: choices: Final = _mapping_or_attribute(response, "choices") if not isinstance(choices, Sequence) or isinstance(choices, (str, bytes)) or not choices: return False message: Final = _mapping_or_attribute(choices[0], "message") tool_calls: Final = _mapping_or_attribute(message, "tool_calls") - if not isinstance(tool_calls, Sequence) or isinstance(tool_calls, (str, bytes)): + if not isinstance(tool_calls, Sequence) or isinstance(tool_calls, (str, bytes)) or not tool_calls: return False - return any( - _mapping_or_attribute(_mapping_or_attribute(tool_call, "function"), "name") == _FUSION_TOOL_NAME - for tool_call in tool_calls + tool_names: Final = tuple( + _mapping_or_attribute(_mapping_or_attribute(tool_call, "function"), "name") for tool_call in tool_calls ) + # Fusion drops its private call and immediately returns any client tool call + # emitted alongside it. Only a Fusion-only response has a continuation whose + # cost callback can safely reconcile the shared reservation later. + return _FUSION_TOOL_NAME in tool_names and all(name == _FUSION_TOOL_NAME for name in tool_names) def _should_defer_fusion_budget_reconciliation( @@ -110,7 +113,9 @@ def _should_defer_fusion_budget_reconciliation( if origin != "fusion_initial": return False complete_stream: Final = kwargs.get("complete_streaming_response") - return _response_invoked_fusion(completion_response) or _response_invoked_fusion(complete_stream) + return _response_requires_fusion_continuation( + complete_stream if complete_stream is not None else completion_response + ) def _accumulate_fusion_cost( 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 71b8823be3d..c40fa9322fe 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 @@ -13,6 +13,7 @@ from litellm.proxy.hooks.proxy_track_cost_callback import ( _failure_should_leave_fusion_reservation_open, _get_budget_reservation_from_metadata, _ProxyDBLogger, + _should_defer_fusion_budget_reconciliation, _should_track_cost_callback, _update_database_and_spend_counters, ) @@ -732,6 +733,42 @@ async def test_fusion_hidden_costs_accumulate_then_continuation_reconciles_once( assert increment.await_args.kwargs["budget_reservation"] is reservation +def test_mixed_fusion_and_client_tool_calls_reconcile_on_the_initial_response(): + def response_with_tools(*tool_names: str) -> litellm.ModelResponse: + return litellm.ModelResponse( + choices=[ + { + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "tool_calls": [ + { + "id": f"call-{index}", + "type": "function", + "function": {"name": name, "arguments": "{}"}, + } + for index, name in enumerate(tool_names) + ], + }, + } + ] + ) + + pure_fusion_response = response_with_tools("litellm_fusion") + mixed_response = response_with_tools("litellm_fusion", "send_email") + metadata = {"internal_call_origin": "fusion_initial"} + + assert _should_defer_fusion_budget_reconciliation(metadata, pure_fusion_response, {}) + assert not _should_defer_fusion_budget_reconciliation(metadata, mixed_response, {}) + # Streaming callbacks may expose a partial chunk separately. The complete + # response is authoritative for whether Fusion will actually continue. + assert not _should_defer_fusion_budget_reconciliation( + metadata, + pure_fusion_response, + {"complete_streaming_response": mixed_response}, + ) + + @pytest.mark.asyncio async def test_cached_fusion_hidden_call_accumulates_zero_cost(): logger = _ProxyDBLogger()