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
This commit is contained in:
Darsh Joshi 2026-10-02 19:22:07 -04:00
parent a292fd409f
commit ccd6ef8d44
4 changed files with 145 additions and 6 deletions

View file

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

View file

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

View file

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

View file

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