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/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/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/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/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,