diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 1db94e82066..543908d1043 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -11,6 +11,7 @@ import math import uuid from collections.abc import AsyncIterator, Mapping, Sequence from dataclasses import dataclass +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, cast @@ -90,11 +91,55 @@ 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 _web_search_domain_strings(tool: Mapping[str, object], key: str) -> tuple[str, ...]: + value = tool.get(key) + if not isinstance(value, list): + return () + return tuple(item for item in value if isinstance(item, str) and item) + + +def _extract_web_search_domain_filters( + tools: Sequence[dict[str, object]], +) -> Mapping[str, tuple[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. + """ + web_tools: Final = tuple(tool for tool in tools if is_web_search_tool(tool)) + allowed: Final = tuple( + chain.from_iterable(_web_search_domain_strings(tool, "allowed_domains") for tool in web_tools) + ) + blocked: Final = tuple( + chain.from_iterable(_web_search_domain_strings(tool, "blocked_domains") for tool in web_tools) + ) + if not allowed and not blocked: + return None + if not blocked: + return MappingProxyType({"allowed_domains": allowed}) + if not allowed: + return MappingProxyType({"blocked_domains": blocked}) + return MappingProxyType({"allowed_domains": allowed, "blocked_domains": blocked}) + + class _PlanMetadataView(TypedDict): websearch_native_blocks: Sequence[Mapping[str, object]] | None @@ -142,7 +187,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 +422,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 +519,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 +700,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 +1658,40 @@ 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. + request_domain_filters: Final = kwargs.get(WEBSEARCH_DOMAIN_FILTER_KEY) if kwargs is not None else None + domain_view: Final = ( + request_domain_filters if isinstance(request_domain_filters, Mapping) else MappingProxyType({}) + ) + allowed: Final = tuple( + item for item in domain_view.get("allowed_domains", ()) if isinstance(item, str) and item + ) + blocked: Final = tuple( + f"-{item}" for item in domain_view.get("blocked_domains", ()) if isinstance(item, str) and item + ) + search_domain_filter: Final = [*allowed, *blocked] or None # mutable-ok: JSON request array, not mutated + if search_domain_filter is not None: + 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..2cbffab71a9 --- /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_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,