From aeb88738e8ff6c34766e7b80c211476bfceafbbb Mon Sep 17 00:00:00 2001 From: JingHao-Leon Date: Thu, 1 Oct 2026 11:15:09 +0800 Subject: [PATCH 1/3] fix(router): sync streaming fallback re-ran the failing group instead of falling back MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _completion_streaming_iterator.stream_with_fallbacks() re-entered the fallback chain via function_with_fallbacks(), which tries the original group first. The retry returns a fresh stream that succeeds at creation time (HTTP 200) and only fails while being iterated; each nested Router.completion() wraps its own stream in this same iterator, so every retry fails again during iteration and re-enters the chain — recursing until the stack runs out (measured: 478 requests to the failing group, ~116 s, then InternalServerError, with a healthy fallback configured). Mirror the async twin: hand the MidStreamFallbackError to async_function_with_fallbacks_common_utils() via run_async_function, which cools down the failed deployment and walks the fallback list directly. Sync tests that patched function_with_fallbacks are updated to the new re-entry point, plus a regression test asserting the triggering error reaches the common utils and the original group is not re-run. Fixes #43945 --- litellm/router.py | 15 +++- tests/unit/test_router/test_router.py | 101 ++++++++++++++++++++++++-- 2 files changed, 108 insertions(+), 8 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index bb118639839..26e1ff16152 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3582,11 +3582,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/test_router/test_router.py b/tests/unit/test_router/test_router.py index 96dddf15869..34eb200e7cd 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -3438,7 +3438,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"}], @@ -3616,7 +3616,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"}], @@ -3683,7 +3683,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( @@ -3696,10 +3696,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(): @@ -3911,7 +4000,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, From 3cf5bb27535ed3e2b6ce18cf586a24ce027ab464 Mon Sep 17 00:00:00 2001 From: JingHao-Leon Date: Sat, 3 Oct 2026 01:22:33 +0800 Subject: [PATCH 2/3] fix(websearch_interception): apply allowed_domains / blocked_domains to intercepted searches With websearch_interception enabled, the Anthropic-native web_search tool's allowed_domains / blocked_domains were silently dropped: the tool is replaced by the standard LiteLLM web search tool (which has no domain fields) and _execute_search() called litellm.asearch() without search_domain_filter. A /v1/messages request that limited the search to one domain got results from every domain (issue #44188). - Pre-call hooks now stash the domain limits on the request kwargs before the native tool is replaced (same side-channel pattern as the native-blocks flag), and the short-circuit path applies them the same way. - _execute_search() forwards them to litellm.asearch() as search_domain_filter: allowed_domains pass through as an allowlist and blocked_domains as '-'-prefixed exclusions, the convention litellm.asearch's providers (e.g. Perplexity) use. - _AsearchNamedParams now permits search_domain_filter. - New unit tests cover extraction, hook stashing, and the asearch call wiring. --- .../websearch_interception/handler.py | 84 ++++++++++- .../test_websearch_domain_filters.py | 132 ++++++++++++++++++ 2 files changed, 214 insertions(+), 2 deletions(-) create mode 100644 tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 6ebf485d717..2394ae4a43b 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/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 From 265126e6f4ed988327032822d38e699f21f747b3 Mon Sep 17 00:00:00 2001 From: JingHao-Leon <102573344+JingHao-Leon@users.noreply.github.com> Date: Sat, 3 Oct 2026 13:33:34 +0800 Subject: [PATCH 3/3] refactor(websearch_interception): immutable domain-filter plumbing for the type-discipline gate The LIT002 ceiling rejected the mutable list/dict plumbing used to collect and stash the domain filters. Collect them as tuples, stash behind MappingProxyType, and rebuild the wire list at the single asearch() call site (original allowed-first order preserved, one construction carries a '# mutable-ok:' reason). No behavior change. --- .../websearch_interception/handler.py | 60 ++++++++++--------- .../test_websearch_domain_filters.py | 10 ++-- 2 files changed, 38 insertions(+), 32 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 2394ae4a43b..470a0c03a0d 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 @@ -103,9 +104,16 @@ _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]], -) -> dict[str, list[str]] | None: +) -> Mapping[str, tuple[str, ...]] | None: """Collect ``allowed_domains`` / ``blocked_domains`` from web search tools. Anthropic-native ``web_search_*`` tools carry optional domain limits. The @@ -116,23 +124,20 @@ def _extract_web_search_domain_filters( 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) + 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 - domain_filters: dict[str, list[str]] = {} - if allowed: - domain_filters["allowed_domains"] = allowed - if blocked: - domain_filters["blocked_domains"] = blocked - return domain_filters + 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): @@ -1658,18 +1663,19 @@ class WebSearchInterceptionLogger(CustomLogger): # 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) + 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()} ) diff --git a/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py b/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py index 70a082f83a3..2cbffab71a9 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py @@ -43,8 +43,8 @@ class TestExtractWebSearchDomainFilters: } ] assert _extract_web_search_domain_filters(tools) == { - "allowed_domains": ["docs.litellm.ai"], - "blocked_domains": ["twitter.com", "x.com"], + "allowed_domains": ("docs.litellm.ai",), + "blocked_domains": ("twitter.com", "x.com"), } def test_returns_none_without_domain_limits(self): @@ -65,7 +65,7 @@ class TestExtractWebSearchDomainFilters: "allowed_domains": ["docs.litellm.ai", 42, None], } ] - assert _extract_web_search_domain_filters(tools) == {"allowed_domains": ["docs.litellm.ai"]} + assert _extract_web_search_domain_filters(tools) == {"allowed_domains": ("docs.litellm.ai",)} class TestDeploymentHookStashesDomainFilters: @@ -86,8 +86,8 @@ class TestDeploymentHookStashesDomainFilters: 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"], + "allowed_domains": ("docs.litellm.ai",), + "blocked_domains": ("twitter.com",), } @pytest.mark.asyncio