diff --git a/litellm/fusion_router.py b/litellm/fusion_router.py index b4ba1e9fd7b..7fcfb075b6d 100644 --- a/litellm/fusion_router.py +++ b/litellm/fusion_router.py @@ -127,6 +127,44 @@ class FusionCandidate: } +def _serialized_prompt_value(value: Mapping[str, object]) -> str: + return json.dumps(value, ensure_ascii=False, separators=(",", ":"), default=str) + + +def _bounded_candidate_prompt_value(candidate: FusionCandidate, max_candidate_chars: int) -> Mapping[str, object]: + """Bound the complete candidate payload, including advisory tool arguments.""" + prompt_value: Final = candidate.as_prompt_value() + if len(_serialized_prompt_value(prompt_value)) <= max_candidate_chars: + return prompt_value + + advisory_json: Final = _serialized_prompt_value( + { + "content": candidate.content, + "tool_proposals": candidate.tool_proposals, + } + ) + marker: Final = "Truncated candidate advisory JSON: " + + def truncated_value(prefix_length: int) -> Mapping[str, object]: + return { + "candidate": candidate.label, + "content": f"{marker}{advisory_json[:prefix_length]}", + "tool_proposals": (), + "finish_reason": candidate.finish_reason, + "truncated": True, + } + + low = 0 + high = len(advisory_json) + while low < high: + midpoint = (low + high + 1) // 2 + if len(_serialized_prompt_value(truncated_value(midpoint))) <= max_candidate_chars: + low = midpoint + else: + high = midpoint - 1 + return truncated_value(low) + + @dataclass(frozen=True, slots=True) class FusionPanelSuccess: candidate: FusionCandidate @@ -231,18 +269,17 @@ def _candidate_from_response(label: str, response: ModelResponse, max_candidate_ proposals: Final = _tool_proposals(response) if not content and not proposals: return None - bounded_content: Final = content[:max_candidate_chars] if content is not None else None return FusionCandidate( label=label, - content=bounded_content, + content=content, tool_proposals=proposals, finish_reason=choice.finish_reason, ) -def _aggregator_instruction(candidates: tuple[FusionCandidate, ...]) -> str: +def _aggregator_instruction(candidates: tuple[FusionCandidate, ...], max_candidate_chars: int) -> str: candidate_json: Final = json.dumps( - tuple(candidate.as_prompt_value() for candidate in candidates), + tuple(_bounded_candidate_prompt_value(candidate, max_candidate_chars) for candidate in candidates), ensure_ascii=False, separators=(",", ":"), ) @@ -262,6 +299,7 @@ def _aggregator_instruction(candidates: tuple[FusionCandidate, ...]) -> str: def _aggregator_messages( messages: list[AllMessageValues], # mutable-ok: Router completion requires its public message-list shape candidates: tuple[FusionCandidate, ...], + max_candidate_chars: int, ) -> list[AllMessageValues]: # mutable-ok: Router completion requires its public message-list shape prefix_length: Final = next( (index for index, message in enumerate(messages) if message["role"] not in ("system", "developer")), @@ -269,7 +307,7 @@ def _aggregator_messages( ) instruction: Final[AllMessageValues] = { "role": "developer", - "content": _aggregator_instruction(candidates), + "content": _aggregator_instruction(candidates, max_candidate_chars), } return [ # mutable-ok: Router completion requires its public message-list shape *messages[:prefix_length], @@ -420,7 +458,9 @@ class FusionRouter: model=self.model_name, llm_provider="", ) - aggregator_messages: Final = _aggregator_messages(messages, candidates) if quorum_met else messages + aggregator_messages: Final = ( + _aggregator_messages(messages, candidates, self.config.max_candidate_chars) if quorum_met else messages + ) aggregator_kwargs: Final = { # mutable-ok: aggregator kwargs require a native mapping for keyword expansion key: value for key, value in request_kwargs.items() diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index c3dd919064b..cf12063de14 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -1028,7 +1028,7 @@ def estimate_request_max_cost( def _estimate_request_model_max_cost( - request_body: dict, + request_body: dict, # mutable-ok: mirrors the existing public reservation request shape route: str, model: str, llm_router: Router | None, @@ -1040,9 +1040,7 @@ def _estimate_request_model_max_cost( if llm_router is not None else model ) - fusion_router: Final = ( - llm_router.fusion_routers.get(registered_model_name) if llm_router is not None else None - ) + fusion_router: Final = llm_router.fusion_routers.get(registered_model_name) if llm_router is not None else None if fusion_router is None: return _estimate_request_max_cost_for_model( request_body=request_body, @@ -1099,7 +1097,7 @@ def estimate_request_input_cost( reconciled to this instead of being refunded to zero. """ estimates = [ - _estimate_request_input_cost_for_model( + _estimate_request_model_input_cost( request_body=request_body, route=route, model=model_name, @@ -1114,6 +1112,60 @@ def estimate_request_input_cost( return max(cast("list[float]", estimates)) +def _estimate_request_model_input_cost( + request_body: dict, # mutable-ok: mirrors the existing public reservation request shape + route: str, + model: str, + llm_router: Router | None, + input_tokens: int | None = None, +) -> float | None: + """Estimate one selectable model's billed input, expanding Fusion children.""" + registered_model_name: Final = ( + llm_router._get_model_from_alias(model=model) or model # pyright: ignore[reportPrivateUsage] # cancellation must price the routed group + if llm_router is not None + else model + ) + fusion_router: Final = llm_router.fusion_routers.get(registered_model_name) if llm_router is not None else None + if fusion_router is None: + return _estimate_request_input_cost_for_model( + request_body=request_body, + route=route, + model=model, + llm_router=llm_router, + input_tokens=input_tokens, + ) + + panel_estimates: Final = tuple( + _estimate_request_input_cost_for_model( + request_body=request_body, + route=route, + model=panel_model, + llm_router=llm_router, + ) + for panel_model in fusion_router.config.panel_models + ) + original_aggregator_tokens: Final = _count_input_tokens( + request_body=request_body, + model=fusion_router.config.aggregator_model, + ) + candidate_token_ceiling: Final = ( + 4 * fusion_router.config.max_candidate_chars * len(fusion_router.config.panel_models) + ) + 1024 + aggregator_input_tokens: Final = ( + original_aggregator_tokens + candidate_token_ceiling if original_aggregator_tokens is not None else None + ) + aggregator_estimate: Final = _estimate_request_input_cost_for_model( + request_body=request_body, + route=route, + model=fusion_router.config.aggregator_model, + llm_router=llm_router, + input_tokens=aggregator_input_tokens, + ) + child_estimates: Final = (*panel_estimates, aggregator_estimate) + known_estimates: Final = tuple(estimate for estimate in child_estimates if estimate is not None) + return sum(known_estimates) if known_estimates else None + + def _estimate_request_input_cost_for_model( request_body: dict, route: str, diff --git a/litellm/router.py b/litellm/router.py index 63f7c571643..b0be104cc8c 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1646,6 +1646,8 @@ class Router: self._base_aanthropic_messages = self.factory_function( litellm.anthropic_messages, call_type="anthropic_messages" ) + # Both public names are intentionally async. Before Fusion, factory_function already + # returned async_wrapper for "anthropic_messages" and assigned it to both aliases. self.aanthropic_messages = self._fusion_aware_aanthropic_messages self.anthropic_messages = self._fusion_aware_aanthropic_messages self.agenerate_content = self.factory_function(litellm.agenerate_content, call_type="agenerate_content") diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index a762929b0e8..a8595353501 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -39,6 +39,7 @@ from litellm.proxy.spend_tracking.budget_reservation import ( TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS, _approximate_input_size, _get_model_access_group_budget_counters, + estimate_request_input_cost, estimate_request_max_cost, get_budget_window_start, invalidate_budget_reservation_counters, @@ -1089,7 +1090,7 @@ def test_fusion_reservation_sums_panels_and_candidate_inflated_aggregator() -> N return 3.0 return {"panel-a": 1.0, "panel-b": 2.0}[model] - with patch( + with patch( # test-quality-ok: isolates child pricing so this test measures Fusion aggregation, not registry prices "litellm.proxy.spend_tracking.budget_reservation._estimate_request_max_cost_for_model", side_effect=child_estimate, ): @@ -1102,6 +1103,60 @@ def test_fusion_reservation_sums_panels_and_candidate_inflated_aggregator() -> N assert estimated == pytest.approx(6.0) +def test_fusion_cancel_floor_sums_child_input_costs() -> None: + router = Router( + model_list=[ + { + "model_name": "panel-a", + "litellm_params": {"model": "openai/panel-a", "api_key": "fake"}, + }, + { + "model_name": "panel-b", + "litellm_params": {"model": "openai/panel-b", "api_key": "fake"}, + }, + { + "model_name": "aggregator", + "litellm_params": {"model": "openai/aggregator", "api_key": "fake"}, + }, + { + "model_name": "fusion/test", + "litellm_params": { + "model": "fusion_router", + "fusion_router_config": { + "panel_models": ["panel-a", "panel-b"], + "aggregator_model": "aggregator", + "max_candidate_chars": 1000, + }, + }, + }, + ] + ) + request_body = { + "model": "fusion/test", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 10, + } + + def child_input_estimate(*, model: str, input_tokens: int | None = None, **_: object) -> float: + if model == "aggregator": + assert input_tokens is not None + assert input_tokens >= 9000 + return 3.0 + return {"panel-a": 1.0, "panel-b": 2.0}[model] + + with patch( # test-quality-ok: isolates child pricing so this test measures Fusion aggregation, not registry prices + "litellm.proxy.spend_tracking.budget_reservation._estimate_request_input_cost_for_model", + side_effect=child_input_estimate, + ): + estimated = estimate_request_input_cost( + request_body=request_body, + route="/chat/completions", + llm_router=router, + ) + + assert estimated == pytest.approx(6.0) + + def test_tiered_reservation_is_all_or_nothing_with_output_tier_from_input_length(): """Dashscope tiered pricing is all-or-nothing: the tier is chosen by the total input tokens and every token (input and output) is billed at that tier's rate. diff --git a/tests/test_litellm/test_fusion_router.py b/tests/test_litellm/test_fusion_router.py index 6cd5ead7e8d..00c4b85d5b7 100644 --- a/tests/test_litellm/test_fusion_router.py +++ b/tests/test_litellm/test_fusion_router.py @@ -1,4 +1,5 @@ import asyncio +import inspect import json from collections.abc import Mapping from typing import Final, cast @@ -220,7 +221,57 @@ async def test_quorum_failure_modes_and_candidate_bound() -> None: await bounded_router.acompletion(messages=[{"role": "user", "content": "Answer"}], stream=False, request_kwargs={}) instruction = str(bounded_completion.calls[-1]["messages"][0]["content"]) payload = json.loads(instruction.split("Candidate responses:\n", 1)[1]) - assert len(payload[0]["content"]) == 1000 + assert len(json.dumps(payload[0], ensure_ascii=False, separators=(",", ":"))) <= 1000 + assert payload[0]["truncated"] is True + + +@pytest.mark.asyncio +async def test_candidate_bound_includes_function_and_custom_tool_payloads() -> None: + oversized_arguments = json.dumps({"patch": "x" * 4000}) + completion = RecordingCompletion( + { + "panel-a": _response( + None, + [ + { + "id": "function-call", + "type": "function", + "function": {"name": "apply_patch", "arguments": oversized_arguments}, + } + ], + ), + "panel-b": _response( + None, + [ + { + "id": "custom-call", + "type": "custom", + "custom": {"name": "research", "input": "漢" * 4000}, + } + ], + ), + "aggregator": _response("bounded"), + } + ) + router = build_fusion_router( + model_name="fusion/bounded-tools", + raw_config={ + "panel_models": ["panel-a", "panel-b"], + "aggregator_model": "aggregator", + "max_candidate_chars": 1000, + }, + completion=completion, + ) + + await router.acompletion(messages=[{"role": "user", "content": "Act"}], stream=False, request_kwargs={}) + + instruction = str(completion.calls[-1]["messages"][0]["content"]) + payload = json.loads(instruction.split("Candidate responses:\n", 1)[1]) + assert len(payload) == 2 + for candidate in payload: + assert len(json.dumps(candidate, ensure_ascii=False, separators=(",", ":"))) <= 1000 + assert candidate["truncated"] is True + assert candidate["tool_proposals"] == [] @pytest.mark.asyncio @@ -365,6 +416,9 @@ async def test_router_responses_api_bridges_through_the_same_fusion_model() -> N async def test_router_anthropic_messages_bridges_through_the_same_fusion_model() -> None: router = Router(model_list=_router_model_list()) + assert inspect.iscoroutinefunction(router.aanthropic_messages) + assert inspect.iscoroutinefunction(router.anthropic_messages) + response = await router.aanthropic_messages( model="fusion/test", messages=[{"role": "user", "content": "Answer"}],