mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
56e8d3031b
commit
758950a432
11 changed files with 182 additions and 9 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue