From 758950a432bcd43fad706c2a437d81e3d07607d5 Mon Sep 17 00:00:00 2001 From: Darsh Joshi Date: Sat, 3 Oct 2026 12:06:07 -0400 Subject: [PATCH] fix(websearch_interception): keep configured domain filters and single-domain limits allowed_domains was forwarded as search_domain_filter whenever the search tool config did not set that exact key. Providers map it onto their native field and skip a configured includeDomains or include_domains, so a caller could replace the operator's allowlist, and google_pse and dataforseo, which take one domain, got only the first entry or an invalid list. Each search config now declares the params that constrain domains and how many search_domain_filter entries it applies. The filter is forwarded only when none of those params is configured and the provider takes every allowed domain. Results are still filtered by URL in every case --- .../websearch_interception/handler.py | 48 ++++++++-- .../llms/base_llm/search/transformation.py | 22 +++++ .../llms/dataforseo/search/transformation.py | 7 ++ litellm/llms/exa_ai/search/transformation.py | 3 + .../llms/google_pse/search/transformation.py | 7 ++ litellm/llms/linkup/search/transformation.py | 3 + litellm/llms/nimble/search/transformation.py | 3 + .../llms/parallel_ai/search/transformation.py | 3 + litellm/llms/tavily/search/transformation.py | 3 + litellm/llms/you_com/search/transformation.py | 3 + .../test_websearch_interception_handler.py | 89 +++++++++++++++++++ 11 files changed, 182 insertions(+), 9 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index e3aea26193e..f7b7205339f 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -39,7 +39,7 @@ from litellm.integrations.websearch_interception.transformation import ( from litellm.litellm_core_utils.agentic_loop_settings import ( validated_max_agentic_loops, ) -from litellm.llms.base_llm.search.transformation import SearchResponse +from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse from litellm.types.integrations.custom_logger import ( CHAT_COMPLETION_AGENTIC_SURFACE, RESPONSES_AGENTIC_SURFACE, @@ -1549,19 +1549,46 @@ class WebSearchInterceptionLogger(CustomLogger): return None @staticmethod - def _provider_supports_rich_search(search_provider: str | None) -> bool: - """Whether the provider's search config accepts objective + multi-query input.""" + def _search_provider_config(search_provider: str | None) -> BaseSearchConfig | None: if not search_provider: - return False + return None try: from litellm.utils import ProviderConfigManager except ImportError: - return False + return None # SearchProviders is a str enum, so an unknown provider string simply # misses the config map and returns None rather than raising. - config = ProviderConfigManager.get_provider_search_config(search_provider) # pyright: ignore[reportArgumentType] -- SearchProviders is a str enum, so the router's provider string hashes to the matching member; unknown strings miss the map and yield None + return ProviderConfigManager.get_provider_search_config(search_provider) # pyright: ignore[reportArgumentType] -- SearchProviders is a str enum, so the router's provider string hashes to the matching member; unknown strings miss the map and yield None + + @classmethod + def _provider_supports_rich_search(cls, search_provider: str | None) -> bool: + """Whether the provider's search config accepts objective + multi-query input.""" + config: Final = cls._search_provider_config(search_provider) return config is not None and config.supports_rich_search_input() + @classmethod + def _provider_domain_filter( + cls, + search_provider: str | None, + configured_search_kwargs: Mapping[str, object], + domains: WebSearchDomainFilter | None, + ) -> list[str] | None: + """ + The client's ``allowed_domains`` hosts to send as ``search_domain_filter``, or None. + + Only sent when the search tool config sets no domain filter of its own (the + operator's filter must never be replaced) and the provider applies every + entry; results are filtered by URL afterwards either way. + """ + if domains is None or not domains.allowed: + return None + config: Final = cls._search_provider_config(search_provider) + if config is None or not config.domain_filter_params().isdisjoint(configured_search_kwargs): + return None + hosts: Final = list(dict.fromkeys(host for host in map(domain_host, domains.allowed) if host)) + max_entries: Final = config.max_search_domain_filter_entries() + return hosts if hosts and (max_entries is None or len(hosts) <= max_entries) else None + async def _execute_search( self, query: str, @@ -1642,10 +1669,13 @@ class WebSearchInterceptionLogger(CustomLogger): ) requested_domains: Final = None if kwargs is None else kwargs.get(WEBSEARCH_DOMAIN_FILTER_KEY) domains: Final = requested_domains if isinstance(requested_domains, WebSearchDomainFilter) else None + provider_domain_filter: Final = self._provider_domain_filter( + search_provider, configured_search_kwargs, domains + ) named_params: Final[_AsearchNamedParams] = ( - {"search_domain_filter": list(dict.fromkeys(domain_host(domain) for domain in domains.allowed))} - if domains is not None and domains.allowed and "search_domain_filter" not in configured_search_kwargs - else _NO_ASEARCH_NAMED + _NO_ASEARCH_NAMED + if provider_domain_filter is None + else {"search_domain_filter": provider_domain_filter} ) unfiltered_result: Final = ( await litellm.asearch(query=query_arg, search_provider=search_provider, **named_params, **search_kwargs) diff --git a/litellm/llms/base_llm/search/transformation.py b/litellm/llms/base_llm/search/transformation.py index 4edbf260d99..c2286d3e233 100644 --- a/litellm/llms/base_llm/search/transformation.py +++ b/litellm/llms/base_llm/search/transformation.py @@ -107,6 +107,28 @@ class BaseSearchConfig: """ return False + def domain_filter_params(self) -> frozenset[str]: + """ + Request params through which this provider restricts results by domain: + the unified ``search_domain_filter`` plus any provider-native keys + (e.g. ``include_domains``) that reach the request through pass-through. + + Integrations that add their own domain filter (e.g. websearch + interception) leave a request alone when any of these is already set, + so a filter configured on the search tool is never replaced. + """ + return frozenset(("search_domain_filter",)) + + def max_search_domain_filter_entries(self) -> int | None: + """ + How many ``search_domain_filter`` entries this provider's request + applies, or None when it takes the whole list. + + Providers whose API takes a single domain return 1, so callers with + more domains than that know a forwarded filter would drop the rest. + """ + return None + def get_http_method(self) -> Literal["GET", "POST"]: """ Get HTTP method for search requests. diff --git a/litellm/llms/dataforseo/search/transformation.py b/litellm/llms/dataforseo/search/transformation.py index fcd4ae70645..10a987cffc3 100644 --- a/litellm/llms/dataforseo/search/transformation.py +++ b/litellm/llms/dataforseo/search/transformation.py @@ -32,6 +32,13 @@ class DataForSEOSearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "DataForSEO" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "domain", "target")) + + def max_search_domain_filter_entries(self) -> int | None: + # `domain` takes a single domain + return 1 + def get_http_method(self) -> Literal["GET", "POST"]: """ DataForSEO uses POST requests with JSON body. diff --git a/litellm/llms/exa_ai/search/transformation.py b/litellm/llms/exa_ai/search/transformation.py index cd904e493f2..06e3841a9e4 100644 --- a/litellm/llms/exa_ai/search/transformation.py +++ b/litellm/llms/exa_ai/search/transformation.py @@ -53,6 +53,9 @@ class ExaAISearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "Exa AI" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "includeDomains", "excludeDomains")) + def validate_environment( self, headers: dict, diff --git a/litellm/llms/google_pse/search/transformation.py b/litellm/llms/google_pse/search/transformation.py index 0128db9a823..83a55499be1 100644 --- a/litellm/llms/google_pse/search/transformation.py +++ b/litellm/llms/google_pse/search/transformation.py @@ -62,6 +62,13 @@ class GooglePSESearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "Google PSE" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "siteSearch", "siteSearchFilter")) + + def max_search_domain_filter_entries(self) -> int | None: + # `siteSearch` takes a single domain + return 1 + def get_http_method(self) -> Literal["GET", "POST"]: """ Google PSE uses GET requests with query parameters. diff --git a/litellm/llms/linkup/search/transformation.py b/litellm/llms/linkup/search/transformation.py index 8ecd5a24c08..bf8e7968cb2 100644 --- a/litellm/llms/linkup/search/transformation.py +++ b/litellm/llms/linkup/search/transformation.py @@ -49,6 +49,9 @@ class LinkupSearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "Linkup" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "includeDomains", "excludeDomains")) + def validate_environment( self, headers: dict, diff --git a/litellm/llms/nimble/search/transformation.py b/litellm/llms/nimble/search/transformation.py index f40ca60fb6c..656af03d673 100644 --- a/litellm/llms/nimble/search/transformation.py +++ b/litellm/llms/nimble/search/transformation.py @@ -86,6 +86,9 @@ class NimbleSearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "Nimble" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "include_domains", "exclude_domains")) + def validate_environment( self, headers: dict[str, str], # mutable-ok: BaseSearchConfig.validate_environment signature diff --git a/litellm/llms/parallel_ai/search/transformation.py b/litellm/llms/parallel_ai/search/transformation.py index d91e532a2cf..2933bc8f6cd 100644 --- a/litellm/llms/parallel_ai/search/transformation.py +++ b/litellm/llms/parallel_ai/search/transformation.py @@ -90,6 +90,9 @@ class ParallelAISearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "Parallel AI" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "include_domains", "exclude_domains")) + def supports_rich_search_input(self) -> bool: # The v1 search API takes `objective` + multiple `search_queries` # natively; sending both is the documented best practice. diff --git a/litellm/llms/tavily/search/transformation.py b/litellm/llms/tavily/search/transformation.py index 158346051f6..8e0dab724c1 100644 --- a/litellm/llms/tavily/search/transformation.py +++ b/litellm/llms/tavily/search/transformation.py @@ -52,6 +52,9 @@ class TavilySearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "Tavily" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "include_domains", "exclude_domains")) + def validate_environment( self, headers: dict, diff --git a/litellm/llms/you_com/search/transformation.py b/litellm/llms/you_com/search/transformation.py index 8b2df949a93..b339ad0c613 100644 --- a/litellm/llms/you_com/search/transformation.py +++ b/litellm/llms/you_com/search/transformation.py @@ -50,6 +50,9 @@ class YouComSearchConfig(BaseSearchConfig): def ui_friendly_name() -> str: return "You.com" + def domain_filter_params(self) -> frozenset[str]: + return frozenset(("search_domain_filter", "include_domains", "exclude_domains")) + def validate_environment( self, headers: dict, diff --git a/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py index 69a30c74363..15e7e80660c 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py @@ -1006,3 +1006,92 @@ async def test_router_chat_completion_web_search_honors_the_native_tool_domain_f search_result = next(m["content"] for messages in sent_messages for m in messages if m["role"] == "tool") assert [url for url in _DOMAIN_FILTER_URLS if url in search_result] == expected_urls assert mock_asearch.await_args.kwargs.get("search_domain_filter") == expected_provider_filter + + +_PROVIDER_DOMAIN_FIELDS = ("includeDomains", "domain", "siteSearch", "siteSearchFilter") + + +def _sent_domain_fields(request) -> dict: + import json + + sent = dict(request.url.params) if request.method == "GET" else json.loads(request.content) + task = sent[0] if isinstance(sent, list) else sent + return {field: task[field] for field in _PROVIDER_DOMAIN_FIELDS if field in task} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("search_tool_params", "allowed_domains", "expected_sent"), + [ + pytest.param( + {"search_provider": "exa_ai", "includeDomains": ["trusted.example"]}, + ["attacker.example"], + {"includeDomains": ["trusted.example"]}, + id="exa-keeps-operator-includeDomains", + ), + pytest.param( + {"search_provider": "linkup", "includeDomains": ["trusted.example"]}, + ["attacker.example"], + {"includeDomains": ["trusted.example"]}, + id="linkup-keeps-operator-includeDomains", + ), + pytest.param( + {"search_provider": "exa_ai"}, + ["example.org", "https://www.example.com/docs"], + {"includeDomains": ["example.org", "www.example.com"]}, + id="exa-takes-every-allowed-domain", + ), + pytest.param( + {"search_provider": "dataforseo", "api_key": "login:password"}, + ["example.org", "example.com"], + {}, + id="dataforseo-single-domain-gets-no-list", + ), + pytest.param( + {"search_provider": "google_pse"}, + ["example.org", "example.com"], + {}, + id="google-pse-not-narrowed-to-first-domain", + ), + pytest.param( + {"search_provider": "google_pse"}, + ["example.org"], + {"siteSearch": "example.org", "siteSearchFilter": "i"}, + id="google-pse-takes-one-allowed-domain", + ), + ], +) +async def test_messages_web_search_forwards_allowed_domains_only_where_the_provider_applies_them( + monkeypatch: pytest.MonkeyPatch, + search_tool_params: dict, + allowed_domains: list[str], + expected_sent: dict, +): + import httpx + import respx + + import litellm + from litellm.llms.anthropic.pass_through.messages.handler import anthropic_messages + from litellm.proxy import proxy_server + + router = MagicMock() + router.search_tools = [ + {"search_tool_name": "operator-search", "litellm_params": {"api_key": "fake-key", **search_tool_params}} + ] + monkeypatch.setenv("GOOGLE_PSE_ENGINE_ID", "fake-engine") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(litellm, "callbacks", [WebSearchInterceptionLogger(enabled_providers=["bedrock"])]) + + with respx.mock(assert_all_called=False) as provider: + provider.route().mock(return_value=httpx.Response(200, json={"results": [], "tasks": []})) + await anthropic_messages( + max_tokens=512, + messages=[{"role": "user", "content": "litellm"}], + model="bedrock/converse/test-model", + custom_llm_provider="bedrock", + tools=[{"type": "web_search_20250305", "name": "web_search", "allowed_domains": allowed_domains}], + ) + sent = _sent_domain_fields(provider.calls.last.request) + + assert sent == expected_sent