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