diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 1db94e82066..f7b7205339f 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,11 +34,12 @@ 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, ) -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, @@ -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]]] = [] @@ -1507,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, @@ -1598,19 +1667,32 @@ 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 + provider_domain_filter: Final = self._provider_domain_filter( + search_provider, configured_search_kwargs, domains + ) + named_params: Final[_AsearchNamedParams] = ( + _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) 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..634735dbfc4 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,35 @@ 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 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}") + 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: + 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 +553,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/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/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..8cacbc21726 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,251 @@ 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 + + +_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"), + [ + ("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 + + +_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