From 4859ac43da33b2bc4e41d5b2825b4bd45cb0e498 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 3 Sep 2026 13:43:25 -0700 Subject: [PATCH] fix(router): bound fusion search transcripts --- litellm/fusion_router.py | 57 ++++++++++++++++--- .../spend_tracking/budget_reservation.py | 52 ++++++++--------- .../proxy/test_budget_reservation.py | 4 +- tests/test_litellm/test_fusion_router.py | 53 +++++++++++++++++ 4 files changed, 127 insertions(+), 39 deletions(-) diff --git a/litellm/fusion_router.py b/litellm/fusion_router.py index c47563490c5..1139a915080 100644 --- a/litellm/fusion_router.py +++ b/litellm/fusion_router.py @@ -369,6 +369,37 @@ def _research_tool_calls(response: ModelResponse) -> tuple[ChatCompletionMessage ) +def _bounded_search_arguments(query: str | None, max_chars: int) -> str: + """Return valid search arguments whose serialized form fits the configured bound.""" + if query is None: + return "{}" + low = 0 + high = len(query) + while low < high: + midpoint = (low + high + 1) // 2 + serialized = json.dumps({"query": query[:midpoint]}, ensure_ascii=False, separators=(",", ":")) + if len(serialized) <= max_chars: + low = midpoint + else: + high = midpoint - 1 + return json.dumps({"query": query[:low]}, ensure_ascii=False, separators=(",", ":")) + + +def _bounded_research_tool_call( + tool_call: ChatCompletionMessageToolCall, + sequence: int, + max_chars: int, +) -> ChatCompletionMessageToolCall: + return ChatCompletionMessageToolCall( + id=f"fusion-search-{sequence}", + type="function", + function={ + "name": "litellm_fusion_search", + "arguments": _bounded_search_arguments(_fusion_query(tool_call), max_chars), + }, + ) + + def _response_text(response: ModelResponse) -> str | None: if not response.choices: return None @@ -679,17 +710,25 @@ class FusionRouter: selected_calls = search_calls[:remaining_searches] if not selected_calls: return response - current_messages.append(cast(AllMessageValues, response.choices[0].message.model_dump(exclude_none=True))) - current_messages.extend( - await asyncio.gather(*(self._execute_research_call(call, request_kwargs) for call in selected_calls)) + completed_searches = self.config.max_tool_calls - remaining_searches + bounded_calls = tuple( + _bounded_research_tool_call(call, completed_searches + index, self.config.max_candidate_chars) + for index, call in enumerate(selected_calls) + ) + # Keep the private transcript inside the reservation ceiling. The + # provider's prose and identifiers are not needed for continuation; + # only bounded, normalized search calls and their results are retained. + current_messages.append( + cast( + AllMessageValues, + { + "role": "assistant", + "tool_calls": [call.model_dump(exclude_none=True) for call in bounded_calls], + }, + ) ) current_messages.extend( - { - "role": "tool", - "tool_call_id": call.id, - "content": '{"status":"error","error":"search_call_limit_exceeded"}', - } - for call in search_calls[len(selected_calls) :] + await asyncio.gather(*(self._execute_research_call(call, request_kwargs) for call in bounded_calls)) ) remaining_searches -= len(selected_calls) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 853b9c96503..ac230bda8d9 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -1079,7 +1079,7 @@ def _estimate_request_model_max_cost( llm_router=llm_router, input_tokens=input_tokens, ) - internal_call_multiplier: Final = ( + internal_call_count: Final = ( fusion_router.config.max_tool_calls + 1 if fusion_router.config.search_tool_name is not None else 1 ) internal_request_body: Final = { @@ -1089,28 +1089,32 @@ def _estimate_request_model_max_cost( "max_completion_tokens": fusion_router.config.max_completion_tokens, } query_token_ceiling: Final = (4 * fusion_router.config.max_candidate_chars) + 1024 - search_context_token_ceiling: Final = ( - 4 * fusion_router.config.max_candidate_chars * fusion_router.config.max_tool_calls + # Each completed search can add one bounded assistant tool call and one + # bounded result. Price every progressively larger round independently; + # multiplying one flat context estimate misses cumulative transcript growth. + search_turn_token_ceiling: Final = ( + (8 * fusion_router.config.max_candidate_chars) + 1024 if fusion_router.config.search_tool_name is not None else 0 ) - panel_input_token_ceiling: Final = query_token_ceiling + search_context_token_ceiling - panel_estimates: Final = tuple( - ( - estimate * internal_call_multiplier - if ( - estimate := _estimate_request_max_cost_for_model( - request_body=internal_request_body, - route=route, - model=panel_model, - llm_router=llm_router, - input_tokens=panel_input_token_ceiling, - ) + + def estimate_internal_calls(model: str, base_input_tokens: int) -> float | None: + estimates = tuple( + _estimate_request_max_cost_for_model( + request_body=internal_request_body, + route=route, + model=model, + llm_router=llm_router, + input_tokens=base_input_tokens + (completed_searches * search_turn_token_ceiling), ) - is not None - else None + for completed_searches in range(internal_call_count) ) - for panel_model in fusion_router.config.panel_models + if any(estimate is None for estimate in estimates): + return None + return sum(cast("tuple[float, ...]", estimates)) + + panel_estimates: Final = tuple( + estimate_internal_calls(panel_model, query_token_ceiling) for panel_model in fusion_router.config.panel_models ) original_outer_tokens: Final = _count_input_tokens( request_body=request_body, @@ -1121,7 +1125,7 @@ def _estimate_request_model_max_cost( candidate_token_ceiling: Final = ( 4 * fusion_router.config.max_candidate_chars * len(fusion_router.config.panel_models) ) + 1024 - analyst_input_tokens: Final = candidate_token_ceiling + query_token_ceiling + search_context_token_ceiling + analyst_input_tokens: Final = candidate_token_ceiling + query_token_ceiling final_outer_input_tokens: Final = ( original_outer_tokens + candidate_token_ceiling @@ -1130,15 +1134,7 @@ def _estimate_request_model_max_cost( if original_outer_tokens is not None else None ) - analyst_estimate = _estimate_request_max_cost_for_model( - request_body=internal_request_body, - route=route, - model=fusion_router.config.resolved_analyst_model, - llm_router=llm_router, - input_tokens=analyst_input_tokens, - ) - if analyst_estimate is not None: - analyst_estimate *= internal_call_multiplier + analyst_estimate = estimate_internal_calls(fusion_router.config.resolved_analyst_model, analyst_input_tokens) final_outer_estimate: Final = _estimate_request_max_cost_for_model( request_body=request_body, route=route, diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 0ec8be1464d..a800d8be974 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -1248,8 +1248,8 @@ def test_fusion_reservation_expands_private_search_loops_and_context() -> None: ) assert estimated == pytest.approx(8.0) - assert ("panel", 13024) in observed - assert ("analyst", 18048) in observed + assert [tokens for model, tokens in observed if model == "panel"] == [5024, 14048, 23072] + assert [tokens for model, tokens in observed if model == "analyst"] == [10048, 19072, 28096] final_outer_tokens = [tokens for model, tokens in observed if model == "outer"][-1] assert final_outer_tokens is not None and final_outer_tokens >= 26048 diff --git a/tests/test_litellm/test_fusion_router.py b/tests/test_litellm/test_fusion_router.py index 38ee2710f4f..6b7ee7a8cbc 100644 --- a/tests/test_litellm/test_fusion_router.py +++ b/tests/test_litellm/test_fusion_router.py @@ -428,6 +428,59 @@ async def test_configured_search_tool_is_private_to_panel_and_analyst() -> None: assert completion.calls[-1].get("tools") is None +@pytest.mark.asyncio +async def test_search_continuation_drops_provider_prose_and_bounds_arguments() -> None: + search_queries: list[str] = [] + + async def search(*, query: str, **_: object) -> object: + search_queries.append(query) + return {"results": [{"snippet": "evidence"}]} + + oversized_query = '\\"' * 2000 + research_call = _response( + "unneeded provider prose" * 1000, + [ + { + "id": "provider-controlled-id" * 1000, + "type": "function", + "function": { + "name": "litellm_fusion_search", + "arguments": json.dumps({"query": oversized_query}), + }, + } + ], + ) + completion = RecordingCompletion( + { + "outer": [_fusion_call(), _response("Final")], + "panel-a": [research_call, _response("Evidence-backed answer")], + "panel-b": [_response("Independent answer")], + "analyst": [_response(_analysis())], + } + ) + + await _router( + completion, + search=search, + search_tool_name="web-search", + max_tool_calls=1, + max_candidate_chars=1000, + ).acompletion( + messages=[{"role": "user", "content": "Research this"}], + stream=False, + request_kwargs={}, + ) + + second_panel_call = [call for call in completion.calls if call["model"] == "panel-a"][1] + assistant_message = second_panel_call["messages"][-2] + bounded_tool_call = assistant_message["tool_calls"][0] + assert "content" not in assistant_message + assert bounded_tool_call["id"] == "fusion-search-0" + assert len(bounded_tool_call["function"]["arguments"]) <= 1000 + assert len(search_queries[0]) < len(oversized_query) + assert second_panel_call["messages"][-1]["tool_call_id"] == "fusion-search-0" + + @pytest.mark.asyncio async def test_reserved_tool_name_and_multiple_choices_are_rejected_before_calls() -> None: completion = RecordingCompletion({})