From ccd6ef8d44dc447e973ad2ae874eedea7133ed27 Mon Sep 17 00:00:00 2001 From: Darsh Joshi Date: Fri, 2 Oct 2026 19:22:07 -0400 Subject: [PATCH 1/4] fix(websearch_interception): honor allowed_domains and blocked_domains on web_search The pre-request hooks swap the Anthropic web_search tool for litellm_web_search, which dropped the domain fields, and the search ran without search_domain_filter. The hooks now keep the domain filter on the request, allowed_domains is forwarded as search_domain_filter, and results are filtered by URL host for every provider, so providers that ignore the filter and blocked_domains are covered too Fixes #44188 --- .../websearch_interception/handler.py | 64 +++++++++++++++++-- .../websearch_interception/transformation.py | 32 ++++++++++ .../integrations/websearch_interception.py | 8 +++ .../test_websearch_interception_handler.py | 47 ++++++++++++++ 4 files changed, 145 insertions(+), 6 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 1db94e82066..e3aea26193e 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -11,9 +11,11 @@ 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 +from pydantic import TypeAdapter, ValidationError from typing_extensions import Never, ReadOnly import litellm @@ -32,6 +34,7 @@ from litellm.integrations.websearch_interception.tools import ( ) from litellm.integrations.websearch_interception.transformation import ( WebSearchTransformation, + domain_host, ) from litellm.litellm_core_utils.agentic_loop_settings import ( validated_max_agentic_loops, @@ -49,6 +52,7 @@ from litellm.types.integrations.websearch_interception import ( RichWebSearchInput, SearchFailed, SearchOutcome, + WebSearchDomainFilter, WebSearchInterceptionConfig, ) from litellm.types.llms.anthropic import AnthropicThinkingParam @@ -90,6 +94,8 @@ 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" +WEBSEARCH_DOMAIN_FILTER_KEY: Final = "_litellm_websearch_domain_filter" + _RESPONSE_CONTENT_FIELD: Final = "content" _ResponseT: Final = TypeVar("_ResponseT") @@ -126,6 +132,10 @@ class _DeploymentCallKwargsView(TypedDict): model: ReadOnly[str] +class _RequestToolsView(TypedDict): + tools: ReadOnly[object] + + class _AcreateNamedParams(TypedDict, total=False): metadata: ReadOnly[Never] stop_sequences: ReadOnly[Never] @@ -142,7 +152,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] @@ -200,6 +210,31 @@ _NO_ACREATE_NAMED: Final[_AcreateNamedParams] = {} _NO_ASEARCH_NAMED: Final[_AsearchNamedParams] = {} +_TOOLS_ADAPTER: Final = TypeAdapter(tuple[Mapping[str, object], ...]) +_DOMAINS_ADAPTER: Final = TypeAdapter(tuple[str, ...]) + + +def _domain_entries(value: object) -> tuple[str, ...]: + try: + entries: Final = _DOMAINS_ADAPTER.validate_python(value or ()) + except ValidationError: + return () + return tuple(entry.strip() for entry in entries if entry.strip()) + + +def _web_search_domain_filter(tools: object) -> WebSearchDomainFilter | None: + try: + validated_tools: Final = _TOOLS_ADAPTER.validate_python(tools) + except ValidationError: + return None + native_tools: Final = tuple(tool for tool in validated_tools if is_anthropic_native_web_search_tool(tool)) + domain_filter: Final = WebSearchDomainFilter( + allowed=tuple(chain.from_iterable(_domain_entries(tool.get("allowed_domains")) for tool in native_tools)), + blocked=tuple(chain.from_iterable(_domain_entries(tool.get("blocked_domains")) for tool in native_tools)), + ) + return domain_filter if domain_filter.allowed or domain_filter.blocked else None + + def _as_str_mapping(value: object) -> Mapping[str, object] | None: return value if isinstance(value, Mapping) else None # pyright: ignore[reportUnknownVariableType] # str-keyed request metadata is not narrowable from object @@ -467,6 +502,9 @@ class WebSearchInterceptionLogger(CustomLogger): # so flagging here ensures the signal isn't lost regardless of order. if any(is_anthropic_native_web_search_tool(t) for t in tools): kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True + deployment_domain_filter: Final = _web_search_domain_filter(tools) + if deployment_domain_filter is not None: + kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = deployment_domain_filter # Convert native/custom web_search tools to LiteLLM standard converted_tools: Final = [] @@ -642,6 +680,10 @@ class WebSearchInterceptionLogger(CustomLogger): # prefix ensures it is stripped before the follow-up call kwargs. if any(is_anthropic_native_web_search_tool(t) for t in tools): kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True + requested_tools: Final[_RequestToolsView] = {"tools": tools} + domain_filter: Final = _web_search_domain_filter(requested_tools["tools"]) + if domain_filter is not None: + kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filter # Convert native web search tools to LiteLLM standard converted_tools: Final[list[dict[str, object]]] = [] @@ -1598,19 +1640,29 @@ class WebSearchInterceptionLogger(CustomLogger): 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 - ) + 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 + 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 + ) + unfiltered_result: Final = ( + await litellm.asearch(query=query_arg, search_provider=search_provider, **named_params, **search_kwargs) if search_metadata is None else await litellm.asearch( query=query_arg, search_provider=search_provider, litellm_metadata=search_metadata, - **_NO_ASEARCH_NAMED, + **named_params, **search_kwargs, ) ) + result: Final = ( + unfiltered_result + if domains is None + else WebSearchTransformation.filter_search_response(unfiltered_result, domains) + ) # Format using transformation function search_result_text: Final = WebSearchTransformation.format_search_response(result) diff --git a/litellm/integrations/websearch_interception/transformation.py b/litellm/integrations/websearch_interception/transformation.py index 47af73570fc..44821855335 100644 --- a/litellm/integrations/websearch_interception/transformation.py +++ b/litellm/integrations/websearch_interception/transformation.py @@ -7,6 +7,7 @@ Transforms between Anthropic/OpenAI tool_use format and LiteLLM search format. import json from collections.abc import Sequence from typing import Any, Final +from urllib.parse import urlsplit from typing_extensions import assert_never @@ -18,10 +19,34 @@ from litellm.types.integrations.websearch_interception import ( SearchFailed, SearchOutcome, SearchSucceeded, + WebSearchDomainFilter, WebSearchToolResultErrorCode, ) +def _split_domain(domain: str) -> tuple[str, str]: + rule: Final = urlsplit(domain if "://" in domain else f"//{domain}") + return (rule.hostname or "").lower(), rule.path.rstrip("/") + + +def domain_host(domain: str) -> str: + return _split_domain(domain)[0] + + +def _url_matches_domain(url: str, domain: str) -> bool: + """Anthropic web_search domain semantics: subdomains are included and an optional path is a prefix.""" + target: Final = urlsplit(url) + host: Final = (target.hostname or "").lower() + rule_host, rule_path = _split_domain(domain) + host_matches: Final = host == rule_host or host.endswith(f".{rule_host}") + return host_matches and target.path.startswith(rule_path) + + +def _url_passes_domain_filter(url: str, domains: WebSearchDomainFilter) -> bool: + allowed: Final = not domains.allowed or any(_url_matches_domain(url, domain) for domain in domains.allowed) + return allowed and not any(_url_matches_domain(url, domain) for domain in domains.blocked) + + class WebSearchTransformation: """ Transformation class for WebSearch tool interception. @@ -527,6 +552,13 @@ class WebSearchTransformation: case _: assert_never(outcome) + @staticmethod + def filter_search_response(result: SearchResponse, domains: WebSearchDomainFilter) -> SearchResponse: + """Drop results outside ``allowed_domains`` or inside ``blocked_domains``, whatever the provider honored.""" + return result.model_copy( + update={"results": [r for r in result.results if _url_passes_domain_filter(r.url, domains)]} + ) + @staticmethod def format_search_response(result: SearchResponse) -> str: """ diff --git a/litellm/types/integrations/websearch_interception.py b/litellm/types/integrations/websearch_interception.py index 22233404fb3..516c83a5b41 100644 --- a/litellm/types/integrations/websearch_interception.py +++ b/litellm/types/integrations/websearch_interception.py @@ -72,6 +72,14 @@ class SearchFailed: SearchOutcome: TypeAlias = SearchSucceeded | SearchFailed +@dataclass(frozen=True, slots=True) +class WebSearchDomainFilter: + """``allowed_domains`` / ``blocked_domains`` from the client's Anthropic ``web_search`` tool.""" + + allowed: tuple[str, ...] + blocked: tuple[str, ...] + + class WebSearchInterceptionConfig(TypedDict, total=False): """ Configuration parameters for WebSearchInterceptionLogger. 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 a5ab28ba72a..cb8eac91d21 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py @@ -899,3 +899,50 @@ async def test_pre_request_hook_syncs_forced_tool_choice(): "type": "tool", "name": LITELLM_WEB_SEARCH_TOOL_NAME, } + + +_DOMAIN_FILTER_URLS = ("https://www.example.com/a", "https://docs.example.org/b", "https://example.net/c") + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("domain_field", "expected_urls", "expected_provider_filter"), + [ + ("allowed_domains", ["https://docs.example.org/b"], ["example.org"]), + ("blocked_domains", ["https://www.example.com/a", "https://docs.example.org/b"], None), + ], +) +async def test_messages_web_search_honors_the_native_tool_domain_filter( + monkeypatch: pytest.MonkeyPatch, + domain_field: str, + expected_urls: list[str], + expected_provider_filter: list[str] | None, +): + import litellm + from litellm.llms.anthropic.pass_through.messages.handler import anthropic_messages + from litellm.llms.base_llm.search.transformation import SearchResult + from litellm.proxy import proxy_server + + domain = {"allowed_domains": "example.org", "blocked_domains": "example.net"}[domain_field] + mock_asearch = AsyncMock( + return_value=SearchResponse( + object="search", + results=[SearchResult(title=url, url=url, snippet="snippet") for url in _DOMAIN_FILTER_URLS], + ) + ) + monkeypatch.setattr(proxy_server, "llm_router", _perplexity_router()) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + monkeypatch.setattr(litellm, "callbacks", [WebSearchInterceptionLogger(enabled_providers=["bedrock"])]) + + response = 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", domain_field: [domain]}], + ) + + text = next(block["text"] for block in response["content"] if block["type"] == "text") + returned_urls = [url for url in _DOMAIN_FILTER_URLS if url in text] + assert returned_urls == expected_urls + assert mock_asearch.await_args.kwargs.get("search_domain_filter") == expected_provider_filter From 56e8d3031bf2d9109699b2493d7e4c9594a8f93f Mon Sep 17 00:00:00 2001 From: Darsh Joshi Date: Fri, 2 Oct 2026 19:46:06 -0400 Subject: [PATCH 2/4] test(websearch_interception): cover domain filter on the chat completions path --- .../test_websearch_interception_handler.py | 60 +++++++++++++++++++ 1 file changed, 60 insertions(+) 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 cb8eac91d21..69a30c74363 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py @@ -946,3 +946,63 @@ async def test_messages_web_search_honors_the_native_tool_domain_filter( returned_urls = [url for url in _DOMAIN_FILTER_URLS if url in text] assert returned_urls == expected_urls assert mock_asearch.await_args.kwargs.get("search_domain_filter") == expected_provider_filter + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("domain_field", "expected_urls", "expected_provider_filter"), + [ + ("allowed_domains", ["https://docs.example.org/b"], ["example.org"]), + ("blocked_domains", ["https://www.example.com/a", "https://docs.example.org/b"], None), + ], +) +async def test_router_chat_completion_web_search_honors_the_native_tool_domain_filter( + monkeypatch: pytest.MonkeyPatch, + domain_field: str, + expected_urls: list[str], + expected_provider_filter: list[str] | None, +): + import litellm + from litellm import Router + from litellm.integrations.custom_logger import CustomLogger + from litellm.llms.base_llm.search.transformation import SearchResult + from litellm.proxy import proxy_server + + sent_messages: list[list[dict]] = [] + + class _SentMessages(CustomLogger): + def log_pre_api_call(self, model, messages, kwargs): + sent_messages.append(messages) + + domain = {"allowed_domains": "example.org", "blocked_domains": "example.net"}[domain_field] + mock_asearch = AsyncMock( + return_value=SearchResponse( + object="search", + results=[SearchResult(title=url, url=url, snippet="snippet") for url in _DOMAIN_FILTER_URLS], + ) + ) + monkeypatch.setattr(proxy_server, "llm_router", _perplexity_router()) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + monkeypatch.setattr( + litellm, "callbacks", [WebSearchInterceptionLogger(enabled_providers=["openai"]), _SentMessages()] + ) + router = Router( + model_list=[{"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}}] + ) + + await router.acompletion( + model="gpt-4o", + messages=[{"role": "user", "content": "litellm"}], + tools=[{"type": "web_search_20250305", "name": "web_search", domain_field: [domain]}], + mock_tool_calls=[ + { + "id": "call_1", + "type": "function", + "function": {"name": LITELLM_WEB_SEARCH_TOOL_NAME, "arguments": '{"query": "litellm"}'}, + } + ], + ) + + 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 From 758950a432bcd43fad706c2a437d81e3d07607d5 Mon Sep 17 00:00:00 2001 From: Darsh Joshi Date: Sat, 3 Oct 2026 12:06:07 -0400 Subject: [PATCH 3/4] 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 From b7a62f4a71f3885584d6ff6aebb1adeba24c3fb8 Mon Sep 17 00:00:00 2001 From: Darsh Joshi Date: Sat, 3 Oct 2026 12:21:24 -0400 Subject: [PATCH 4/4] fix(websearch_interception): match domain path rules on whole path segments A rule like example.com/docs/ lost its trailing slash and was then matched as a plain prefix, so it also admitted example.com/docs-private. Match the path exactly or at a segment boundary instead --- .../websearch_interception/transformation.py | 5 +- .../test_websearch_interception_handler.py | 52 +++++++++++++++++++ 2 files changed, 55 insertions(+), 2 deletions(-) diff --git a/litellm/integrations/websearch_interception/transformation.py b/litellm/integrations/websearch_interception/transformation.py index 44821855335..634735dbfc4 100644 --- a/litellm/integrations/websearch_interception/transformation.py +++ b/litellm/integrations/websearch_interception/transformation.py @@ -34,12 +34,13 @@ def domain_host(domain: str) -> str: def _url_matches_domain(url: str, domain: str) -> bool: - """Anthropic web_search domain semantics: subdomains are included and an optional path is a prefix.""" + """Anthropic web_search domain semantics: subdomains are included and an optional path matches whole path segments.""" target: Final = urlsplit(url) host: Final = (target.hostname or "").lower() rule_host, rule_path = _split_domain(domain) host_matches: Final = host == rule_host or host.endswith(f".{rule_host}") - return host_matches and target.path.startswith(rule_path) + path_matches: Final = not rule_path or target.path == rule_path or target.path.startswith(f"{rule_path}/") + return host_matches and path_matches def _url_passes_domain_filter(url: str, domains: WebSearchDomainFilter) -> bool: 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 15e7e80660c..8cacbc21726 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py @@ -948,6 +948,58 @@ async def test_messages_web_search_honors_the_native_tool_domain_filter( assert mock_asearch.await_args.kwargs.get("search_domain_filter") == expected_provider_filter +_PATH_RULE_URLS = ( + "https://example.com/docs/intro", + "https://example.com/docs-private/x", + "https://example.com/other/y", +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("domain_field", "rule", "expected_urls"), + [ + ("allowed_domains", "example.com/docs/", ["https://example.com/docs/intro"]), + ("allowed_domains", "example.com/docs", ["https://example.com/docs/intro"]), + ("blocked_domains", "example.com/docs/", ["https://example.com/docs-private/x", "https://example.com/other/y"]), + ], +) +async def test_messages_web_search_domain_path_rules_match_whole_path_segments( + monkeypatch: pytest.MonkeyPatch, + domain_field: str, + rule: str, + expected_urls: list[str], +): + import litellm + from litellm.llms.anthropic.pass_through.messages.handler import anthropic_messages + from litellm.llms.base_llm.search.transformation import SearchResult + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_router", _perplexity_router()) + monkeypatch.setattr( + litellm, + "asearch", + AsyncMock( + return_value=SearchResponse( + object="search", + results=[SearchResult(title=url, url=url, snippet="snippet") for url in _PATH_RULE_URLS], + ) + ), + ) + monkeypatch.setattr(litellm, "callbacks", [WebSearchInterceptionLogger(enabled_providers=["bedrock"])]) + + response = 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", domain_field: [rule]}], + ) + + text = next(block["text"] for block in response["content"] if block["type"] == "text") + assert [url for url in _PATH_RULE_URLS if url in text] == expected_urls + + @pytest.mark.asyncio @pytest.mark.parametrize( ("domain_field", "expected_urls", "expected_provider_filter"),