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
This commit is contained in:
Darsh Joshi 2026-10-03 12:06:07 -04:00
parent 56e8d3031b
commit 758950a432
11 changed files with 182 additions and 9 deletions

View file

@ -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)

View file

@ -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.

View file

@ -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.

View file

@ -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,

View file

@ -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.

View file

@ -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,

View file

@ -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

View file

@ -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.

View file

@ -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,

View file

@ -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,

View file

@ -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