diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index e358636c105..8de7801babb 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -801,7 +801,15 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None _NON_TOKEN_RATE_FIELDS: Final = frozenset( - {"cost_per_second", "input_cost_per_second", "output_cost_per_second", "input_cost_per_query", "tiered_pricing"} + { + "cost_per_second", + "input_cost_per_second", + "output_cost_per_second", + "input_cost_per_query", + "input_cost_per_character", + "output_cost_per_character", + "tiered_pricing", + } ) diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 0271cf1e03c..7421e5e5035 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -897,6 +897,14 @@ def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None: cost: Final = standard_logging_payload.get("response_cost") if cost is None: return + # A 0 cost is ambiguous between "genuinely free" and "pricing failed". + # When the pricing-failure debug info is present, omit the attributes + # rather than reporting a fake $0 that cost backends cannot distinguish + # from a free model (issue #44186) — e.g. Langfuse prefers an ingested + # cost over its own price knowledge, so a fake 0 hides real spend there. + failure_debug: Final = standard_logging_payload.get("response_cost_failure_debug_info") + if failure_debug: + return try: cost_value: Final = float(cost) except (TypeError, ValueError): diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index efd7f6c7dd9..289d8a2628d 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -479,7 +479,14 @@ class LLMCallSpanData: usage=LLMUsage.from_standard_logging_payload(payload), finish_reasons=finish_reasons, error=_parse_error(payload), - response_cost=as_float(payload.get("response_cost")), + # response_cost = 0 is ambiguous between a free model and a + # pricing failure; when the pricing-failure debug info is + # present, report None (unknown) instead of a fake 0 (issue + # #44186) — e.g. Langfuse prefers an ingested cost over its own + # price knowledge, so a fake 0 hides real spend there. + response_cost=( + as_float(payload.get("response_cost")) if not payload.get("response_cost_failure_debug_info") else None + ), cost=LLMCost.from_breakdown(cast("Mapping[str, object] | None", payload.get("cost_breakdown"))), server=ServerInfo.from_api_base(context.api_base), identity=context.identity, @@ -574,7 +581,10 @@ class MCPToolCallSpanData: _json_or_none(meta.get("result")) if capture_content and meta.get("result") is not None else None ), error=_parse_error(payload), - response_cost=as_float(payload.get("response_cost")), + # Same 0-vs-pricing-failure ambiguity as the LLM span above. + response_cost=( + as_float(payload.get("response_cost")) if not payload.get("response_cost_failure_debug_info") else None + ), identity=RequestContext.from_standard_logging_payload(payload).identity, ) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 1db94e82066..c74e7d7c488 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -90,11 +90,51 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b # ``web_search_tool_result`` blocks to inject into the final response. WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks" +# Key used to flag, on per-request kwargs, that the originating client sent +# domain filters (``allowed_domains`` / ``blocked_domains``) on an +# Anthropic-native ``web_search_*`` tool. The standard LiteLLM tool drops +# them (the model must not see client policy), so they are stashed here and +# applied to the downstream ``litellm.asearch()`` call as +# ``search_domain_filter``. +WEBSEARCH_DOMAIN_FILTER_KEY: Final = "_websearch_interception_domain_filter" + _RESPONSE_CONTENT_FIELD: Final = "content" _ResponseT: Final = TypeVar("_ResponseT") +def _extract_web_search_domain_filters( + tools: Sequence[dict[str, object]], +) -> dict[str, list[str]] | None: + """Collect ``allowed_domains`` / ``blocked_domains`` from web search tools. + + Anthropic-native ``web_search_*`` tools carry optional domain limits. The + standard LiteLLM tool deliberately drops them (client policy, not model + input), so they are collected here, stashed on the request kwargs, and + applied to the downstream ``litellm.asearch()`` call as + ``search_domain_filter``. + + Returns None when no web search tool carries a domain limit. + """ + allowed: list[str] = [] + blocked: list[str] = [] + for tool in tools: + if not is_web_search_tool(tool): + continue + for key, bucket in (("allowed_domains", allowed), ("blocked_domains", blocked)): + value = tool.get(key) + if isinstance(value, list): + bucket.extend(item for item in value if isinstance(item, str) and item) + if not allowed and not blocked: + return None + domain_filters: dict[str, list[str]] = {} + if allowed: + domain_filters["allowed_domains"] = allowed + if blocked: + domain_filters["blocked_domains"] = blocked + return domain_filters + + class _PlanMetadataView(TypedDict): websearch_native_blocks: Sequence[Mapping[str, object]] | None @@ -142,7 +182,7 @@ class _AcreateNamedParams(TypedDict, total=False): class _AsearchNamedParams(TypedDict, total=False): max_results: ReadOnly[int | None] - search_domain_filter: ReadOnly[Never] + search_domain_filter: ReadOnly[list[str] | None] max_tokens_per_page: ReadOnly[int | None] country: ReadOnly[str | None] api_key: ReadOnly[str | None] @@ -377,6 +417,12 @@ class WebSearchInterceptionLogger(CustomLogger): None, ) + # Apply any domain limits the client set on the native web search + # tool before the search executes. + domain_filters: Final = _extract_web_search_domain_filters(tools) + if domain_filters is not None and isinstance(kwargs, dict): + kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters + outcome: Final = await self._short_circuit_search_outcome(query, kwargs=kwargs) search_result_text: Final = WebSearchTransformation.search_outcome_text(outcome) @@ -468,6 +514,12 @@ class WebSearchInterceptionLogger(CustomLogger): if any(is_anthropic_native_web_search_tool(t) for t in tools): kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True + # Same for domain limits: stash them before the native tool is + # replaced, so the downstream search can apply them. + domain_filters: Final = _extract_web_search_domain_filters(tools) + if domain_filters is not None: + kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters + # Convert native/custom web_search tools to LiteLLM standard converted_tools: Final = [] for tool in tools: @@ -643,6 +695,12 @@ class WebSearchInterceptionLogger(CustomLogger): if any(is_anthropic_native_web_search_tool(t) for t in tools): kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True + # Same for domain limits: stash them before the native tool is + # replaced, so the downstream search can apply them. + domain_filters: Final = _extract_web_search_domain_filters(tools) + if domain_filters is not None: + kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters + # Convert native web search tools to LiteLLM standard converted_tools: Final[list[dict[str, object]]] = [] for tool in tools: @@ -1595,17 +1653,39 @@ class WebSearchInterceptionLogger(CustomLogger): rich_objective = rich.get("objective") if rich_objective and "objective" not in configured_search_kwargs: configured_search_kwargs["objective"] = rich_objective + # Domain limits stashed by the pre-call hooks from the client's + # native web_search tool. ``allowed_domains`` pass through as an + # allowlist and ``blocked_domains`` as '-'-prefixed exclusions — + # the convention litellm.asearch()'s providers (e.g. Perplexity) + # use for search_domain_filter. + search_domain_filter: list[str] | None = None + request_domain_filters: Final = kwargs.get(WEBSEARCH_DOMAIN_FILTER_KEY) if kwargs is not None else None + if isinstance(request_domain_filters, dict): + allowed: Final[list[str]] = [ + item for item in request_domain_filters.get("allowed_domains", []) if isinstance(item, str) and item + ] + blocked: Final[list[str]] = [ + item for item in request_domain_filters.get("blocked_domains", []) if isinstance(item, str) and item + ] + if allowed or blocked: + search_domain_filter = allowed + [f"-{item}" for item in blocked] + verbose_logger.debug("WebSearchInterception: Applying domain filter %s", search_domain_filter) search_kwargs: Final = MappingProxyType( {**configured_search_kwargs, **parent_correlation.as_search_kwargs()} ) result: Final = ( await litellm.asearch( - query=query_arg, search_provider=search_provider, **_NO_ASEARCH_NAMED, **search_kwargs + query=query_arg, + search_provider=search_provider, + search_domain_filter=search_domain_filter, + **_NO_ASEARCH_NAMED, + **search_kwargs, ) if search_metadata is None else await litellm.asearch( query=query_arg, search_provider=search_provider, + search_domain_filter=search_domain_filter, litellm_metadata=search_metadata, **_NO_ASEARCH_NAMED, **search_kwargs, diff --git a/litellm/llms/deepseek/chat/transformation.py b/litellm/llms/deepseek/chat/transformation.py index 4e428a23392..457e2603e63 100644 --- a/litellm/llms/deepseek/chat/transformation.py +++ b/litellm/llms/deepseek/chat/transformation.py @@ -160,13 +160,18 @@ class DeepSeekChatConfig(OpenAIGPTConfig): def _is_vision_forwardable_content(self, message: AllMessageValues, content: Sequence[object]) -> bool: """ - True only for a user message whose content list holds well-formed - text and image_url blocks with at least one image; a block missing - its payload falls back to the string collapse instead of crashing - or reaching the wire malformed. The model capability gate lives in - the caller. + True only for a user or tool message whose content list holds + well-formed text and image_url blocks with at least one image; a block + missing its payload falls back to the string collapse instead of + crashing or reaching the wire malformed. The model capability gate + lives in the caller. + + ``role="tool"`` is forwardable: the DeepSeek Chat Completions API + accepts and reads image_url blocks in tool results (verified directly + against the API; bug #44211) — agent loops that screenshot inside a + tool rely on it. """ - if message.get("role") != "user": + if message.get("role") not in ("user", "tool"): return False if not all(self._is_forwardable_block(block) for block in content): return False diff --git a/litellm/router.py b/litellm/router.py index 0b9f12c8da3..018053a11b2 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3611,11 +3611,22 @@ class Router: initial_kwargs["original_function"] = router_self._completion initial_kwargs["messages"] = messages router_self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) - fallback_response = router_self.function_with_fallbacks( - **initial_kwargs, + # Pass the MidStreamFallbackError through the common fallback utils, like the + # async twin does. Calling function_with_fallbacks() here instead re-runs the + # original (failing) group first, and because each nested Router.completion() + # wraps its stream in this same iterator, every retry fails again only when + # the caller iterates it — recursing until the stack runs out. + fallback_response = run_async_function( + router_self.async_function_with_fallbacks_common_utils, + e, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, + model_group=model_group, + args=(), + kwargs=initial_kwargs, + include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, ) if hasattr(fallback_response, "__iter__"): diff --git a/tests/unit/integrations/arize/test_arize_utils.py b/tests/unit/integrations/arize/test_arize_utils.py index 167b083e147..4b2f3ef4ee0 100644 --- a/tests/unit/integrations/arize/test_arize_utils.py +++ b/tests/unit/integrations/arize/test_arize_utils.py @@ -70,9 +70,7 @@ def test_arize_set_attributes(): # Simulated LLM response object response_obj = ModelResponse( usage={"total_tokens": 100, "completion_tokens": 60, "prompt_tokens": 40}, - choices=[ - Choices(message={"role": "assistant", "content": "Basic Response Content"}) - ], + choices=[Choices(message={"role": "assistant", "content": "Basic Response Content"})], model="gpt-4o", id="chatcmpl-ID", ) @@ -89,9 +87,7 @@ def test_arize_set_attributes(): assert span.set_attribute.call_count == 26 # Metadata attached to the span - span.set_attribute.assert_any_call( - SpanAttributes.METADATA, json.dumps({"key_1": "value_1", "key_2": None}) - ) + span.set_attribute.assert_any_call(SpanAttributes.METADATA, json.dumps({"key_1": "value_1", "key_2": None})) # Basic LLM information span.set_attribute.assert_any_call(SpanAttributes.LLM_MODEL_NAME, "gpt-4o") @@ -114,16 +110,12 @@ def test_arize_set_attributes(): span.set_attribute.assert_any_call(SpanAttributes.OPENINFERENCE_SPAN_KIND, "LLM") # And TOOL must never be written for an LLM chat completion call. span_kind_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND ] assert "TOOL" not in span_kind_writes # Request message content and metadata - span.set_attribute.assert_any_call( - SpanAttributes.INPUT_VALUE, "Basic Request Content" - ) + span.set_attribute.assert_any_call(SpanAttributes.INPUT_VALUE, "Basic Request Content") span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_INPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_ROLE}", "user", @@ -134,9 +126,7 @@ def test_arize_set_attributes(): ) # Tool call definitions and function names - span.set_attribute.assert_any_call( - f"{SpanAttributes.LLM_TOOLS}.0.name", "get_weather" - ) + span.set_attribute.assert_any_call(f"{SpanAttributes.LLM_TOOLS}.0.name", "get_weather") span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_TOOLS}.0.description", "Fetches weather details.", @@ -146,26 +136,20 @@ def test_arize_set_attributes(): json.dumps( { "type": "object", - "properties": { - "location": {"type": "string", "description": "City name"} - }, + "properties": {"location": {"type": "string", "description": "City name"}}, "required": ["location"], } ), ) # Invocation parameters - span.set_attribute.assert_any_call( - SpanAttributes.LLM_INVOCATION_PARAMETERS, '{"user": "test_user"}' - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_INVOCATION_PARAMETERS, '{"user": "test_user"}') # User ID span.set_attribute.assert_any_call(SpanAttributes.USER_ID, "test_user") # Output message content - span.set_attribute.assert_any_call( - SpanAttributes.OUTPUT_VALUE, "Basic Response Content" - ) + span.set_attribute.assert_any_call(SpanAttributes.OUTPUT_VALUE, "Basic Response Content") span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_ROLE}", "assistant", @@ -228,9 +212,7 @@ def test_arize_set_attributes_responses_api(): ResponseReasoningItem( id="reasoning-001", type="reasoning", - summary=[ - Summary(text="First, I need to analyze...", type="summary_text") - ], + summary=[Summary(text="First, I need to analyze...", type="summary_text")], ), ResponseOutputMessage( id="msg-001", @@ -277,9 +259,7 @@ def test_arize_set_attributes_responses_api(): span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120) - span.set_attribute.assert_any_call( - SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180 - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180) def test_set_usage_outputs_pydantic_completion_usage(): @@ -327,9 +307,7 @@ def test_set_usage_outputs_pydantic_completion_usage(): span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 40) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 60) # reasoning_tokens for chat completions live in completion_tokens_details - span.set_attribute.assert_any_call( - SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 25 - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 25) def test_set_usage_outputs_pydantic_response_api_usage(): @@ -362,9 +340,7 @@ def test_set_usage_outputs_pydantic_response_api_usage(): span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250) - span.set_attribute.assert_any_call( - SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180 - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180) class TestArizeLogger(CustomLogger): @@ -375,16 +351,12 @@ class TestArizeLogger(CustomLogger): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.standard_callback_dynamic_params: Optional[ - StandardCallbackDynamicParams - ] = None + self.standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = None async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): # Capture dynamic params and print them for verification print("logged kwargs", json.dumps(kwargs, indent=4, default=str)) - self.standard_callback_dynamic_params = kwargs.get( - "standard_callback_dynamic_params" - ) + self.standard_callback_dynamic_params = kwargs.get("standard_callback_dynamic_params") @pytest.mark.asyncio @@ -410,14 +382,8 @@ async def test_arize_dynamic_params(): # Assert dynamic parameters were received in the callback assert test_arize_logger.standard_callback_dynamic_params is not None - assert ( - test_arize_logger.standard_callback_dynamic_params.get("arize_api_key") - == "test_api_key_dynamic" - ) - assert ( - test_arize_logger.standard_callback_dynamic_params.get("arize_space_key") - == "test_space_key_dynamic" - ) + assert test_arize_logger.standard_callback_dynamic_params.get("arize_api_key") == "test_api_key_dynamic" + assert test_arize_logger.standard_callback_dynamic_params.get("arize_space_key") == "test_space_key_dynamic" def test_construct_dynamic_arize_headers(): @@ -428,9 +394,7 @@ def test_construct_dynamic_arize_headers(): from litellm.types.utils import StandardCallbackDynamicParams # Test with all parameters present - dynamic_params_full = StandardCallbackDynamicParams( - arize_api_key="test_api_key", arize_space_id="test_space_id" - ) + dynamic_params_full = StandardCallbackDynamicParams(arize_api_key="test_api_key", arize_space_id="test_space_id") arize_logger = ArizeLogger() headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_full) @@ -438,9 +402,7 @@ def test_construct_dynamic_arize_headers(): assert headers == expected_headers # Test with only space_id - dynamic_params_space_id_only = StandardCallbackDynamicParams( - arize_space_id="test_space_id" - ) + dynamic_params_space_id_only = StandardCallbackDynamicParams(arize_space_id="test_space_id") headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_id_only) expected_headers = {"arize-space-id": "test_space_id"} @@ -456,9 +418,7 @@ def test_construct_dynamic_arize_headers(): dynamic_params_space_key_and_api_key = StandardCallbackDynamicParams( arize_space_key="test_space_key", arize_api_key="test_api_key" ) - headers = arize_logger.construct_dynamic_otel_headers( - dynamic_params_space_key_and_api_key - ) + headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_key_and_api_key) expected_headers = {"arize-space-id": "test_space_key", "api_key": "test_api_key"} @@ -528,9 +488,7 @@ def test_arize_emits_no_cache_tokens_when_absent(): from litellm.integrations.arize._utils import _set_usage_outputs span = MagicMock() - response_obj = { - "usage": {"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6} - } + response_obj = {"usage": {"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}} _set_usage_outputs(span, response_obj, SpanAttributes) attrs = _collect_calls(span) assert SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ not in attrs @@ -542,14 +500,8 @@ def test_passthrough_call_type_resolves_to_llm_span_kind(): from litellm.integrations._types.open_inference import OpenInferenceSpanKindValues from litellm.integrations.arize._utils import _infer_open_inference_span_kind - assert ( - _infer_open_inference_span_kind("allm_passthrough_route") - == OpenInferenceSpanKindValues.LLM.value - ) - assert ( - _infer_open_inference_span_kind("llm_passthrough_route") - == OpenInferenceSpanKindValues.LLM.value - ) + assert _infer_open_inference_span_kind("allm_passthrough_route") == OpenInferenceSpanKindValues.LLM.value + assert _infer_open_inference_span_kind("llm_passthrough_route") == OpenInferenceSpanKindValues.LLM.value def test_arize_chat_completion_with_tools_stays_llm_span_kind(): @@ -605,9 +557,7 @@ def test_arize_chat_completion_with_tools_stays_llm_span_kind(): ArizeLogger.set_arize_attributes(span, kwargs, response_obj) span_kind_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND ] assert span_kind_writes, "span.kind must be written" assert all(v == "LLM" for v in span_kind_writes) @@ -659,13 +609,8 @@ def test_arize_emits_assistant_tool_calls_on_output_message(): attrs = _collect_calls(span) base = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_TOOL_CALLS}.0" assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc" - assert ( - attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}"] == "get_weather" - ) - assert ( - attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}"] - == '{"location": "SF"}' - ) + assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}"] == "get_weather" + assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}"] == '{"location": "SF"}' def test_arize_output_value_falls_back_to_tool_calls_summary(): @@ -818,9 +763,7 @@ def test_arize_emits_tool_call_id_and_name_on_input_tool_message(): assert attrs[f"{assistant_base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc" # Tool message at index 2 tool_prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.2" - assert ( - attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}"] == "call_abc" - ) + assert attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}"] == "call_abc" assert attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_NAME}"] == "get_weather" @@ -866,10 +809,7 @@ def test_arize_emits_multimodal_input_contents(): assert attrs[f"{base}.0.message_content.type"] == "text" assert attrs[f"{base}.0.message_content.text"] == "What is in this image?" assert attrs[f"{base}.1.message_content.type"] == "image" - assert ( - attrs[f"{base}.1.message_content.image.image.url"] - == "https://example.com/cat.png" - ) + assert attrs[f"{base}.1.message_content.image.image.url"] == "https://example.com/cat.png" def test_arize_emits_session_and_user_attrs_from_metadata(): @@ -974,11 +914,7 @@ def test_arize_does_not_overwrite_user_id_from_optional_params(): id="r2", ) ArizeLogger.set_arize_attributes(span, kwargs, response_obj) - user_id_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.USER_ID - ] + user_id_writes = [c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.USER_ID] assert "from_metadata" not in user_id_writes @@ -1013,6 +949,72 @@ def test_arize_emits_response_cost(): assert attrs["llm.response.cost"] == 0.0012345 # legacy key still emitted +def test_arize_omits_cost_when_pricing_failed(): + """response_cost = 0 plus pricing-failure debug info means LiteLLM could + not price the call — omit the cost attributes entirely: backends cannot + tell a fake 0 from a free model (issue #44186).""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + "response_cost": 0.0, + "response_cost_failure_debug_info": {"error_str": "model not in cost map"}, + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2}, + choices=[Choices(message={"role": "assistant", "content": "hello"})], + model="gpt-4o", + id="r4", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + assert "llm.cost.total" not in attrs + assert "llm.response.cost" not in attrs + + +def test_arize_emits_zero_cost_for_free_model(): + """A genuine 0 (free model, no pricing-failure info) must still be + emitted — the guard only fires on the pricing-failure marker.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + "response_cost": 0.0, + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2}, + choices=[Choices(message={"role": "assistant", "content": "hello"})], + model="gpt-4o", + id="r5", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + assert attrs["llm.cost.total"] == 0.0 + assert attrs["llm.response.cost"] == 0.0 + + def test_arize_passthrough_bedrock_anthropic_normalization(): """Bedrock-Anthropic passthrough: input/output text must be set so the span renders something other than raw provider attrs.""" @@ -1048,9 +1050,7 @@ def test_arize_passthrough_bedrock_anthropic_normalization(): "complete_input_dict": { "anthropic_version": "bedrock-2023-05-31", "max_tokens": 64, - "messages": [ - {"role": "user", "content": "What is the capital of France?"} - ], + "messages": [{"role": "user", "content": "What is the capital of France?"}], } }, "standard_logging_object": { @@ -1068,19 +1068,13 @@ def test_arize_passthrough_bedrock_anthropic_normalization(): assert attrs[SpanAttributes.INPUT_VALUE] == "What is the capital of France?" msg0 = f"{SpanAttributes.LLM_INPUT_MESSAGES}.0" assert attrs[f"{msg0}.{MessageAttributes.MESSAGE_ROLE}"] == "user" - assert ( - attrs[f"{msg0}.{MessageAttributes.MESSAGE_CONTENT}"] - == "What is the capital of France?" - ) + assert attrs[f"{msg0}.{MessageAttributes.MESSAGE_CONTENT}"] == "What is the capital of France?" # Output rendering (Anthropic content[].text) assert attrs[SpanAttributes.OUTPUT_VALUE] == "The capital of France is Paris." out0 = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0" assert attrs[f"{out0}.{MessageAttributes.MESSAGE_ROLE}"] == "assistant" - assert ( - attrs[f"{out0}.{MessageAttributes.MESSAGE_CONTENT}"] - == "The capital of France is Paris." - ) + assert attrs[f"{out0}.{MessageAttributes.MESSAGE_CONTENT}"] == "The capital of France is Paris." # Token counts (Bedrock input_tokens/output_tokens) — extracted via # coercion of the non-dict response. @@ -1089,9 +1083,7 @@ def test_arize_passthrough_bedrock_anthropic_normalization(): # Span kind defended even though the call_type is a passthrough variant. span_kind_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND ] assert span_kind_writes # at least one assert all(v == "LLM" for v in span_kind_writes) @@ -1109,11 +1101,7 @@ def test_arize_passthrough_call_type_does_not_run_on_chat_completion(): span = MagicMock() _maybe_normalize_passthrough( span, - { - "additional_args": { - "complete_input_dict": {"messages": [{"role": "user", "content": "x"}]} - } - }, + {"additional_args": {"complete_input_dict": {"messages": [{"role": "user", "content": "x"}]}}}, {"choices": [{"message": {"role": "assistant", "content": "y"}}]}, {"choices": [{"message": {"role": "assistant", "content": "y"}}]}, {"call_type": "completion"}, @@ -1133,11 +1121,7 @@ def test_arize_passthrough_skipped_when_message_redaction_enabled(): span = MagicMock() kwargs = { "additional_args": { - "complete_input_dict": { - "messages": [ - {"role": "user", "content": "Patient John Doe, SSN 123-45-6789"} - ] - } + "complete_input_dict": {"messages": [{"role": "user", "content": "Patient John Doe, SSN 123-45-6789"}]} }, # Enables redaction via the dynamic-param path inside # should_redact_message_logging(), without touching globals. @@ -1211,9 +1195,7 @@ def test_arize_mcp_call_tool_result_does_not_break_attribute_setting(): "optional_params": {}, "litellm_params": {"custom_llm_provider": "mcp"}, } - response_obj = CallToolResult( - content=[TextContent(type="text", text="sunny, 21C")], isError=False - ) + response_obj = CallToolResult(content=[TextContent(type="text", text="sunny, 21C")], isError=False) ArizeLogger.set_arize_attributes(span, kwargs, response_obj) @@ -1295,9 +1277,7 @@ def test_arize_mcp_tool_span_renders_name_input_and_output(): from mcp.types import CallToolResult, TextContent span = MagicMock() - response_obj = CallToolResult( - content=[TextContent(type="text", text="sunny, 21C")], isError=False - ) + response_obj = CallToolResult(content=[TextContent(type="text", text="sunny, 21C")], isError=False) ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) @@ -1336,9 +1316,7 @@ def test_arize_mcp_tool_span_respects_message_redaction(): from mcp.types import CallToolResult, TextContent span = MagicMock() - response_obj = CallToolResult( - content=[TextContent(type="text", text="SSN 123-45-6789")], isError=False - ) + response_obj = CallToolResult(content=[TextContent(type="text", text="SSN 123-45-6789")], isError=False) ArizeLogger.set_arize_attributes( span, diff --git a/tests/unit/integrations/otel/test_otel_v2_emitter.py b/tests/unit/integrations/otel/test_otel_v2_emitter.py index 11b2aa5fd67..d7a0be5bc53 100644 --- a/tests/unit/integrations/otel/test_otel_v2_emitter.py +++ b/tests/unit/integrations/otel/test_otel_v2_emitter.py @@ -90,6 +90,22 @@ def test_llm_call_span_cost_breakdown(): assert f"{LiteLLM.COST_PREFIX}margin_total_amount" not in a +def test_llm_call_span_pricing_failure_reports_unknown_cost(): + """A 0 response_cost alongside pricing-failure debug info means LiteLLM + could not price the call — report unknown (None), not a fake 0 that cost + backends cannot distinguish from a free model (issue #44186).""" + data = LLMCallSpanData.from_standard_logging_payload( + _payload(response_cost=0.0, response_cost_failure_debug_info={"error_str": "model not in cost map"}) + ) + assert data.response_cost is None + + +def test_llm_call_span_zero_cost_without_failure_info_is_kept(): + """A genuine 0 (free model) without pricing-failure info stays 0.""" + data = LLMCallSpanData.from_standard_logging_payload(_payload(response_cost=0.0)) + assert data.response_cost == 0.0 + + def test_tracer_scope_carries_litellm_version(): from litellm._version import version as litellm_version diff --git a/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py b/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py new file mode 100644 index 00000000000..70a082f83a3 --- /dev/null +++ b/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py @@ -0,0 +1,132 @@ +""" +Tests for domain-limit passthrough in web search interception. + +Covers bug #44188: ``allowed_domains`` / ``blocked_domains`` set on an +Anthropic-native ``web_search_*`` tool must survive the conversion to the +standard LiteLLM web search tool and be applied to the downstream +``litellm.asearch()`` call as ``search_domain_filter``. +""" + +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.integrations.websearch_interception.handler import ( + WEBSEARCH_DOMAIN_FILTER_KEY, + WebSearchInterceptionLogger, + _extract_web_search_domain_filters, +) +from litellm.llms.base_llm.search.transformation import SearchResponse, SearchResult + + +def _make_search_response() -> SearchResponse: + return SearchResponse( + results=[ + SearchResult( + title="LiteLLM Docs", + url="https://docs.litellm.ai/", + snippet="Unified interface for LLMs.", + date=None, + ) + ] + ) + + +class TestExtractWebSearchDomainFilters: + def test_collects_allowed_and_blocked(self): + tools = [ + { + "type": "web_search_20250305", + "name": "web_search", + "allowed_domains": ["docs.litellm.ai"], + "blocked_domains": ["twitter.com", "x.com"], + } + ] + assert _extract_web_search_domain_filters(tools) == { + "allowed_domains": ["docs.litellm.ai"], + "blocked_domains": ["twitter.com", "x.com"], + } + + def test_returns_none_without_domain_limits(self): + tools = [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}] + assert _extract_web_search_domain_filters(tools) is None + + def test_ignores_domains_on_non_web_search_tools(self): + tools = [ + {"name": "bash", "allowed_domains": ["example.com"]}, + ] + assert _extract_web_search_domain_filters(tools) is None + + def test_ignores_non_string_entries(self): + tools = [ + { + "type": "web_search_20250305", + "name": "web_search", + "allowed_domains": ["docs.litellm.ai", 42, None], + } + ] + assert _extract_web_search_domain_filters(tools) == {"allowed_domains": ["docs.litellm.ai"]} + + +class TestDeploymentHookStashesDomainFilters: + @pytest.mark.asyncio + async def test_stashes_filters_for_native_tool(self): + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + kwargs = { + "tools": [ + { + "type": "web_search_20250305", + "name": "web_search", + "allowed_domains": ["docs.litellm.ai"], + "blocked_domains": ["twitter.com"], + } + ], + "litellm_params": {"custom_llm_provider": "bedrock"}, + } + out = await logger.async_pre_call_deployment_hook(kwargs, None) + assert out is not None + assert kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] == { + "allowed_domains": ["docs.litellm.ai"], + "blocked_domains": ["twitter.com"], + } + + @pytest.mark.asyncio + async def test_no_stash_without_domain_limits(self): + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + kwargs = { + "tools": [{"type": "web_search_20250305", "name": "web_search"}], + "litellm_params": {"custom_llm_provider": "bedrock"}, + } + await logger.async_pre_call_deployment_hook(kwargs, None) + assert WEBSEARCH_DOMAIN_FILTER_KEY not in kwargs + + +class TestExecuteSearchAppliesDomainFilter: + @pytest.mark.asyncio + async def test_forwards_search_domain_filter(self): + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + kwargs = { + WEBSEARCH_DOMAIN_FILTER_KEY: { + "allowed_domains": ["docs.litellm.ai"], + "blocked_domains": ["twitter.com"], + } + } + asearch = AsyncMock(return_value=_make_search_response()) + with patch("litellm.asearch", asearch): + await logger._execute_search("what is litellm", kwargs=kwargs) + + assert asearch.await_count == 1 + assert asearch.await_args.kwargs.get("search_domain_filter") == [ + "docs.litellm.ai", + "-twitter.com", + ] + + @pytest.mark.asyncio + async def test_no_filter_when_kwargs_empty(self): + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + asearch = AsyncMock(return_value=_make_search_response()) + with patch("litellm.asearch", asearch): + await logger._execute_search("what is litellm", kwargs={}) + + assert asearch.await_count == 1 + assert asearch.await_args.kwargs.get("search_domain_filter") is None diff --git a/tests/unit/llms/deepseek/test_deepseek_vision_tool_message.py b/tests/unit/llms/deepseek/test_deepseek_vision_tool_message.py new file mode 100644 index 00000000000..90771cb1194 --- /dev/null +++ b/tests/unit/llms/deepseek/test_deepseek_vision_tool_message.py @@ -0,0 +1,67 @@ +""" +Tests for DeepSeek vision-content forwarding in role=tool messages. + +Bug #44211: ``_is_vision_forwardable_content`` rejected every non-user role, +so a ``role=tool`` message carrying image_url blocks was silently collapsed +to text before the request left LiteLLM — while the DeepSeek API itself +accepts and reads tool-result images. +""" + +import pytest + +from litellm.llms.deepseek.chat.transformation import DeepSeekChatConfig + + +TOOL_MESSAGE_WITH_IMAGE = { + "role": "tool", + "tool_call_id": "abc", + "content": [ + {"type": "text", "text": "screenshot:"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,aGVsbG8="}}, + ], +} + + +class TestVisionForwardableContent: + def test_tool_message_with_image_is_forwardable(self): + config = DeepSeekChatConfig() + assert config._is_vision_forwardable_content( + message=TOOL_MESSAGE_WITH_IMAGE, + content=TOOL_MESSAGE_WITH_IMAGE["content"], + ) + + def test_user_message_with_image_stays_forwardable(self): + config = DeepSeekChatConfig() + message = {"role": "user", "content": TOOL_MESSAGE_WITH_IMAGE["content"]} + assert config._is_vision_forwardable_content(message=message, content=message["content"]) + + def test_assistant_message_stays_collapsed(self): + config = DeepSeekChatConfig() + message = {"role": "assistant", "content": TOOL_MESSAGE_WITH_IMAGE["content"]} + assert not config._is_vision_forwardable_content(message=message, content=message["content"]) + + def test_image_missing_payload_falls_back(self): + config = DeepSeekChatConfig() + message = { + "role": "tool", + "content": [ + {"type": "text", "text": "screenshot:"}, + {"type": "image_url", "image_url": {"url": ""}}, + ], + } + assert not config._is_vision_forwardable_content(message=message, content=message["content"]) + + +class TestForwardOrCollapseContent: + def test_tool_message_image_content_is_not_collapsed(self): + config = DeepSeekChatConfig() + out = config._forward_or_collapse_content(message=TOOL_MESSAGE_WITH_IMAGE, forward_images=True) + assert isinstance(out.get("content"), list) + blocks = out["content"] + assert any(isinstance(block, dict) and block.get("type") == "image_url" for block in blocks) + + def test_tool_message_text_only_still_collapsed(self): + config = DeepSeekChatConfig() + message = {"role": "tool", "content": [{"type": "text", "text": "plain"}]} + out = config._forward_or_collapse_content(message=message, forward_images=True) + assert out.get("content") == "plain" diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 9a7c22f6f82..55d85073021 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -187,7 +187,9 @@ def test_response_cost_calculator_keeps_optional_params_out_of_hidden_params(): assert optional_params["aws_session_token"] == "session-secret" -def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials(monkeypatch: pytest.MonkeyPatch) -> None: +def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials( + monkeypatch: pytest.MonkeyPatch, +) -> None: from litellm.proxy import proxy_server from litellm.proxy.spend_tracking.spend_tracking_utils import _get_proxy_server_request_for_spend_logs_payload @@ -236,10 +238,6 @@ def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials( assert logging_obj.optional_params["extra_headers"] == {"x-goog-api-key": "goog-secret"} - - - - def test_realtime_stream_combines_text_and_audio_token_details(): """Realtime response.done usage with input_token_details / output_token_details.""" from litellm.cost_calculator import RealtimeAPITokenUsageProcessor @@ -1354,8 +1352,6 @@ def test_bedrock_cost_calculator_comparison_with_without_cache(): print(f"Cost with cache: {cost_with_cache}") - - def test_gemini_25_explicit_caching_cost_direct_usage(): """ Test that Gemini 2.5 models correctly calculate costs with explicit caching. @@ -1924,8 +1920,6 @@ def test_cost_margin_with_discount(monkeypatch): print(f" - Expected: ${expected_cost:.6f}") - - def test_completion_cost_extracts_service_tier_from_response(_local_model_cost_map): """Test that completion_cost extracts service_tier from completion_response object.""" from litellm import completion_cost @@ -2675,8 +2669,6 @@ def test_gemini_without_cache_tokens_details(): print("✅ Gemini without cacheTokensDetails works correctly") - - def test_additional_costs_only_for_azure_ai(_local_model_cost_map): """ Test that _get_additional_costs is only called for azure_ai provider. @@ -3221,9 +3213,7 @@ def test_cost_per_token_resolves_per_second_rate_precedence( model: Final = "test-chat-per-second-rate-precedence" entry: Final = {**pricing_fields, "litellm_provider": "together_ai", "mode": "chat"} - litellm.register_model( - model_cost={model: entry} - ) + litellm.register_model(model_cost={model: entry}) assert cost_per_token( model=model, @@ -3648,6 +3638,42 @@ def test_combine_usage_objects_sums_mirrored_cache_write_fields_once(): assert combined_pair.prompt_tokens_details.cache_creation_tokens == 100 +def test_select_model_name_selects_character_priced_deployment(_local_model_cost_map): + """ + A deployment whose only rate is input_cost_per_character must be selected + by router_model_id: the aspeech cost path resolves its price through the + deployment entry, and a character-only entry failing the "prices anything" + check silently produced spend = 0 (issue #44200). + """ + from litellm.cost_calculator import _select_model_name_for_cost_calc + + router_model_id = "openai/qwen-audio-3.1-tts-flash-uuid" + litellm.model_cost[router_model_id] = { + "input_cost_per_character": 1e-8, + "output_cost_per_character": 0.0, + "litellm_provider": "openai", + } + + selected = _select_model_name_for_cost_calc( + model="qwen-audio-3.1-tts-flash", + completion_response=None, + custom_pricing=True, + custom_llm_provider="openai", + router_model_id=router_model_id, + ) + + assert selected == router_model_id + + +def test_cost_map_entry_prices_anything_recognizes_character_rates(): + from litellm.cost_calculator import _cost_map_entry_prices_anything + + assert _cost_map_entry_prices_anything({"input_cost_per_character": 1e-8}) is True + assert _cost_map_entry_prices_anything({"output_cost_per_character": 0.0}) is True + assert _cost_map_entry_prices_anything({"input_cost_per_token": 1e-6}) is True + assert _cost_map_entry_prices_anything({"mode": "audio_speech"}) is False + + def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_map): """A router-facing model_name alias containing "/" whose leading segment is NOT a registered provider must not be double-prefixed into a non-existent cost key. @@ -4815,9 +4841,7 @@ def test_xai_batch_tier_discounts_the_long_context_rate_like_the_flat_batch_rate assert info[f"{prefix}_above_200k_tokens_batches"] < info[f"{prefix}_above_200k_tokens"] -@pytest.mark.parametrize( - ("prompt_tokens", "tier"), [(200_000, "_above_200k_tokens_batches"), (199_999, "_batches")] -) +@pytest.mark.parametrize(("prompt_tokens", "tier"), [(200_000, "_above_200k_tokens_batches"), (199_999, "_batches")]) def test_xai_batch_cost_calculator_bills_the_200k_batch_tier_inclusively( _local_model_cost_map: None, prompt_tokens: int, tier: str ) -> None: diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index d0115593e46..362ef6ddb2f 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -3507,7 +3507,7 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste } return chunk - with patch.object(router, "function_with_fallbacks", return_value=NestedFallbackStream()): + with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=NestedFallbackStream()): result = router._completion_streaming_iterator( model_response=FailedStream(), messages=[{"role": "user", "content": "hi"}], @@ -3685,7 +3685,7 @@ def test_completion_streaming_iterator_adopts_fallback_response_headers(): def __iter__(self): return iter([]) - with patch.object(router, "function_with_fallbacks", return_value=FallbackStream()): + with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=FallbackStream()): result = router._completion_streaming_iterator( model_response=FailedStream(), messages=[{"role": "user", "content": "hi"}], @@ -3752,7 +3752,7 @@ def test_completion_streaming_iterator_fallback_on_429(): with patch.object( router, - "function_with_fallbacks", + "async_function_with_fallbacks_common_utils", return_value=mock_fallback_response, ) as mock_fallback: result = router._completion_streaming_iterator( @@ -3765,10 +3765,99 @@ def test_completion_streaming_iterator_fallback_on_429(): assert mock_fallback.called call_kwargs = mock_fallback.call_args + assert mock_fallback.call_args.args[0] is rate_limit_error # Pre-first-chunk: should use original messages, no continuation prompt - assert call_kwargs.kwargs.get("messages") == messages + assert call_kwargs.kwargs.get("kwargs", {}).get("messages") == messages # Verify original_function is _completion (sync) - assert call_kwargs.kwargs.get("original_function") == router._completion + assert call_kwargs.kwargs.get("kwargs", {}).get("original_function") == router._completion + + +def test_completion_streaming_iterator_routes_mid_stream_fallback_through_common_utils(): + """Regression (#43945): the sync mid-stream fallback re-entry must hand the + MidStreamFallbackError to async_function_with_fallbacks_common_utils, like the async + twin does. Calling function_with_fallbacks() instead re-runs the original (failing) + group first, and because each nested Router.completion() wraps its own stream in this + same iterator, every retry fails again only while being iterated — recursing until + the stack runs out (measured: 478 requests to the failing group, then + InternalServerError, with a healthy fallback configured).""" + from unittest.mock import MagicMock + + from litellm.exceptions import MidStreamFallbackError + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, + } + ], + ) + + messages = [{"role": "user", "content": "Test"}] + initial_kwargs = {"model": "gpt-4", "stream": True} + + pre_first_chunk_error = MidStreamFallbackError( + message="upstream died before the first chunk", + model="gpt-4", + llm_provider="openai", + generated_content="", + is_pre_first_chunk=True, + ) + + class SyncIteratorImmediateError: + def __init__(self): + self.model = "gpt-4" + self.custom_llm_provider = "openai" + self.logging_obj = MagicMock() + self.chunks = [] + + def __iter__(self): + return self + + def __next__(self): + raise pre_first_chunk_error + + class FallbackStream: + def __init__(self): + self._chunks = iter( + [ + litellm.ModelResponseStream( + choices=[{"index": 0, "delta": {"content": "from the fallback"}}] + ) + ] + ) + + def __iter__(self): + return self + + def __next__(self): + return next(self._chunks) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=FallbackStream(), + ) as mock_utils: + with patch.object(router, "function_with_fallbacks") as mock_function_with_fallbacks: + result = router._completion_streaming_iterator( + model_response=SyncIteratorImmediateError(), + messages=messages, + initial_kwargs=initial_kwargs, + ) + + collected_chunks = list(result) + + assert mock_utils.called + # the triggering error must reach the common utils, so cooldowns apply and + # the walk starts from the fallback list, not the failing group + assert mock_utils.call_args.args[0] is pre_first_chunk_error + assert mock_utils.call_args.kwargs.get("kwargs", {}).get("messages") == messages + assert not mock_function_with_fallbacks.called, ( + "re-running the original group is what recurses; common utils already " + "excludes the deployment that raised" + ) + + assert len(collected_chunks) == 1 def test_completion_streaming_iterator_preserves_hidden_params(): @@ -3980,7 +4069,7 @@ def test_completion_streaming_iterator_reraises_mid_chunk_error_with_no_text_con mock_response = SyncIteratorNoTextChunkError() - with patch.object(router, "function_with_fallbacks") as mock_fallback: + with patch.object(router, "async_function_with_fallbacks_common_utils") as mock_fallback: result = router._completion_streaming_iterator( model_response=mock_response, messages=messages,