fix(websearch): preserve Anthropic domain filters

Keep allowed and blocked domains through tool conversion and search execution, then filter results before building text and citations

Fixes #44188
This commit is contained in:
Vrajesh Sulakhe 2026-10-03 07:03:01 +05:30
parent 9e31afbc6d
commit f5bc46e08e
5 changed files with 416 additions and 18 deletions

View file

@ -14,6 +14,7 @@ from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, cast
from pydantic import TypeAdapter
from typing_extensions import Never, ReadOnly
import litellm
@ -25,10 +26,12 @@ from litellm.integrations.websearch_interception.tools import (
get_litellm_web_search_tool,
get_litellm_web_search_tool_openai,
get_litellm_web_search_tool_responses,
get_web_search_domain_filters,
is_anthropic_native_web_search_tool,
is_web_search_tool,
is_web_search_tool_chat_completion,
is_web_search_tool_responses,
resolve_web_search_domain_filters,
)
from litellm.integrations.websearch_interception.transformation import (
WebSearchTransformation,
@ -49,6 +52,7 @@ from litellm.types.integrations.websearch_interception import (
RichWebSearchInput,
SearchFailed,
SearchOutcome,
WebSearchDomainFilters,
WebSearchInterceptionConfig,
)
from litellm.types.llms.anthropic import AnthropicThinkingParam
@ -91,6 +95,7 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
_RESPONSE_CONTENT_FIELD: Final = "content"
_WEB_SEARCH_REQUEST: Final = TypeAdapter(dict[str, object])
_ResponseT: Final = TypeVar("_ResponseT")
@ -377,7 +382,11 @@ class WebSearchInterceptionLogger(CustomLogger):
None,
)
outcome: Final = await self._short_circuit_search_outcome(query, kwargs=kwargs)
outcome: Final = await self._short_circuit_search_outcome(
query,
kwargs=kwargs,
domains=resolve_web_search_domain_filters(_WEB_SEARCH_REQUEST.validate_python({"tools": tools})),
)
search_result_text: Final = WebSearchTransformation.search_outcome_text(outcome)
content: Final[list[dict[str, object]]] = []
@ -473,7 +482,7 @@ class WebSearchInterceptionLogger(CustomLogger):
for tool in tools:
if is_web_search_tool(tool):
# Convert to LiteLLM standard web search tool
converted_tool = get_litellm_web_search_tool_openai()
converted_tool = get_litellm_web_search_tool_openai(**get_web_search_domain_filters(tool))
converted_tools.append(converted_tool)
verbose_logger.debug(
"WebSearchInterception: Converted %s (type=%s) to %s",
@ -647,7 +656,9 @@ class WebSearchInterceptionLogger(CustomLogger):
converted_tools: Final[list[dict[str, object]]] = []
for tool in tools:
if is_web_search_tool(tool):
standard_tool = get_litellm_web_search_tool()
standard_tool = get_litellm_web_search_tool(
**get_web_search_domain_filters(_WEB_SEARCH_REQUEST.validate_python(tool))
)
converted_tools.append(standard_tool)
verbose_logger.debug(
"WebSearchInterception: Converted %s (type=%s) to %s",
@ -1188,7 +1199,19 @@ class WebSearchInterceptionLogger(CustomLogger):
search_tasks: Final = [
(
self._execute_search(
tool_call["input"]["query"], kwargs=kwargs, rich=self._rich_search_input(tool_call["input"])
tool_call["input"]["query"],
kwargs=kwargs,
rich=self._rich_search_input(tool_call["input"]),
**resolve_web_search_domain_filters(
_WEB_SEARCH_REQUEST.validate_python(
{
"tools": optional_params["tools"]
if "tools" in optional_params
else kwargs.get("tools"),
"input": tool_call["input"],
}
)
),
)
if isinstance(tool_call.get("input"), dict) and tool_call["input"].get("query")
else self._create_empty_search_result()
@ -1408,7 +1431,23 @@ class WebSearchInterceptionLogger(CustomLogger):
if query:
verbose_logger.debug("WebSearchInterception: Queuing search for query='%s'", query)
search_tasks.append(
self._execute_search(query, kwargs=kwargs, rich=self._rich_search_input(tool_call["input"]))
self._execute_search(
query,
kwargs=kwargs,
rich=self._rich_search_input(tool_call["input"]),
**resolve_web_search_domain_filters(
_WEB_SEARCH_REQUEST.validate_python(
{
"tools": (
anthropic_messages_optional_request_params["tools"]
if "tools" in anthropic_messages_optional_request_params
else kwargs.get("tools")
),
"input": tool_call["input"],
}
)
),
)
)
else:
verbose_logger.debug("WebSearchInterception: Tool call %s has no query", tool_call["id"])
@ -1467,12 +1506,14 @@ class WebSearchInterceptionLogger(CustomLogger):
)
return patch, search_outcomes
async def _short_circuit_search_outcome(self, query: str, kwargs: Mapping[str, object] | None) -> SearchOutcome:
async def _short_circuit_search_outcome(
self, query: str, kwargs: Mapping[str, object] | None, domains: WebSearchDomainFilters
) -> SearchOutcome:
try:
result: Final = (
await self._execute_search(query)
await self._execute_search(query, **domains)
if kwargs is None
else await self._execute_search(query, kwargs=kwargs)
else await self._execute_search(query, kwargs=kwargs, **domains)
)
except Exception as e:
return WebSearchTransformation.search_outcome(e)
@ -1525,6 +1566,8 @@ class WebSearchInterceptionLogger(CustomLogger):
query: str,
kwargs: Mapping[str, object] | None = None,
rich: RichWebSearchInput | None = None,
allowed_domains: Sequence[str] | None = None,
blocked_domains: Sequence[str] | None = None,
) -> tuple[str, SearchResponse | None]:
"""
Execute a single web search using router's search tools.
@ -1596,7 +1639,11 @@ class WebSearchInterceptionLogger(CustomLogger):
if rich_objective and "objective" not in configured_search_kwargs:
configured_search_kwargs["objective"] = rich_objective
search_kwargs: Final = MappingProxyType(
{**configured_search_kwargs, **parent_correlation.as_search_kwargs()}
{
**configured_search_kwargs,
**parent_correlation.as_search_kwargs(),
**({"search_domain_filter": list(allowed_domains)} if allowed_domains is not None else {}),
}
)
result: Final = (
await litellm.asearch(
@ -1613,12 +1660,19 @@ class WebSearchInterceptionLogger(CustomLogger):
)
# Format using transformation function
search_result_text: Final = WebSearchTransformation.format_search_response(result)
filtered_result: Final = WebSearchTransformation.filter_search_response(
result, allowed_domains=allowed_domains, blocked_domains=blocked_domains
)
search_result_text: Final = (
"No search results found."
if filtered_result is not result and not filtered_result.results
else WebSearchTransformation.format_search_response(filtered_result)
)
verbose_logger.debug(
"WebSearchInterception: Search completed for '%s', got %s chars", query, len(search_result_text)
)
return search_result_text, result
return search_result_text, filtered_result
except Exception as e:
verbose_logger.error("WebSearchInterception: Search failed for '%s': %s", query, e)
raise
@ -1863,7 +1917,23 @@ class WebSearchInterceptionLogger(CustomLogger):
if query:
verbose_logger.debug("WebSearchInterception: Queuing search for query='%s'", query)
search_tasks.append(self._execute_search(query, kwargs=kwargs, rich=self._rich_search_input(tool_args)))
search_tasks.append(
self._execute_search(
query,
kwargs=kwargs,
rich=self._rich_search_input(tool_args),
**resolve_web_search_domain_filters(
_WEB_SEARCH_REQUEST.validate_python(
{
"tools": optional_params["tools"]
if "tools" in optional_params
else kwargs.get("tools"),
"input": tool_args,
}
)
),
)
)
else:
verbose_logger.debug("WebSearchInterception: Tool call %s has no query", tool_call.get("id"))
# Add empty result for tools without query

View file

@ -6,10 +6,17 @@ Native provider tools (like Anthropic's web_search_20250305) are converted
to this format for consistent interception and execution.
"""
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import Any, Final
from pydantic import TypeAdapter
from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME
from litellm.types.integrations.websearch_interception import WebSearchDomainFilters
_DOMAINS: Final = TypeAdapter(tuple[str, ...])
_TOOL_MAPPING: Final = TypeAdapter(dict[str, object])
_TOOLS: Final = TypeAdapter(list[dict[str, object]])
_WEB_SEARCH_TOOL_DESCRIPTION: Final = (
"Search the web for information. Use this when you need current "
@ -17,7 +24,10 @@ _WEB_SEARCH_TOOL_DESCRIPTION: Final = (
)
def _web_search_input_schema() -> dict[str, object]: # mutable-ok: plain-dict tool shape, as the get_* builders
def _web_search_input_schema(
allowed_domains: Sequence[str] | None = None,
blocked_domains: Sequence[str] | None = None,
) -> Mapping[str, object]:
"""
JSON schema for the web search tool's input, shared by every tool format.
@ -51,12 +61,20 @@ def _web_search_input_schema() -> dict[str, object]: # mutable-ok: plain-dict t
"objective for the best results."
),
},
**{
name: {"type": "array", "items": {"type": "string"}, "default": list(domains)}
for name, domains in (("allowed_domains", allowed_domains), ("blocked_domains", blocked_domains))
if domains is not None
},
},
"required": ["query"],
}
def get_litellm_web_search_tool() -> dict[str, object]:
def get_litellm_web_search_tool(
allowed_domains: Sequence[str] | None = None,
blocked_domains: Sequence[str] | None = None,
) -> dict[str, object]:
"""
Get the standard LiteLLM web search tool definition.
@ -78,11 +96,14 @@ def get_litellm_web_search_tool() -> dict[str, object]:
return {
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
"description": _WEB_SEARCH_TOOL_DESCRIPTION,
"input_schema": _web_search_input_schema(),
"input_schema": _web_search_input_schema(allowed_domains, blocked_domains),
}
def get_litellm_web_search_tool_openai() -> dict[str, object]:
def get_litellm_web_search_tool_openai(
allowed_domains: Sequence[str] | None = None,
blocked_domains: Sequence[str] | None = None,
) -> dict[str, object]:
"""
Get the standard LiteLLM web search tool definition in OpenAI format.
@ -98,7 +119,7 @@ def get_litellm_web_search_tool_openai() -> dict[str, object]:
"function": {
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
"description": _WEB_SEARCH_TOOL_DESCRIPTION,
"parameters": _web_search_input_schema(),
"parameters": _web_search_input_schema(allowed_domains, blocked_domains),
},
}
@ -123,6 +144,39 @@ def get_litellm_web_search_tool_responses() -> dict[str, object]:
}
def get_web_search_domain_filters(tool: Mapping[str, object]) -> WebSearchDomainFilters:
function: Final = tool.get("function")
if isinstance(function, Mapping):
return get_web_search_domain_filters(_TOOL_MAPPING.validate_python(function))
schema: Final = _TOOL_MAPPING.validate_python(tool.get("input_schema", tool.get("parameters", {})))
properties: Final = _TOOL_MAPPING.validate_python(schema.get("properties", {}))
defaults: Final = {
name: _TOOL_MAPPING.validate_python(value).get("default")
for name in ("allowed_domains", "blocked_domains")
if isinstance(value := properties.get(name), Mapping)
}
values: Final = {**defaults, **tool}
allowed: Final = values.get("allowed_domains")
blocked: Final = values.get("blocked_domains")
allowed_filter: Final[WebSearchDomainFilters] = (
{"allowed_domains": _DOMAINS.validate_python(allowed)} if allowed is not None else {}
)
blocked_filter: Final[WebSearchDomainFilters] = (
{"blocked_domains": _DOMAINS.validate_python(blocked)} if blocked is not None else {}
)
return {**allowed_filter, **blocked_filter}
def resolve_web_search_domain_filters(request: Mapping[str, object]) -> WebSearchDomainFilters:
definitions: Final = _TOOLS.validate_python(request.get("tools") or [])
configured: Final[WebSearchDomainFilters] = next(
(get_web_search_domain_filters(tool) for tool in definitions if is_web_search_tool(tool)),
WebSearchDomainFilters(),
)
arguments: Final = get_web_search_domain_filters(_TOOL_MAPPING.validate_python(request.get("input") or {}))
return {**arguments, **configured}
def is_web_search_tool_responses(tool: Mapping[str, object]) -> bool:
"""
Check if a tool is a web search tool for the Responses API.

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
@ -22,6 +23,18 @@ from litellm.types.integrations.websearch_interception import (
)
def _search_hostname(url: str) -> str:
try:
parsed: Final = urlsplit(url if "://" in url or url.startswith("//") else f"//{url}")
return (parsed.hostname or "").lower().rstrip(".")
except ValueError:
return ""
def _matches_search_domain(host: str, domains: Sequence[str]) -> bool:
return any(host == domain or host.endswith("." + domain) for domain in domains if domain)
class WebSearchTransformation:
"""
Transformation class for WebSearch tool interception.
@ -527,6 +540,28 @@ class WebSearchTransformation:
case _:
assert_never(outcome)
@staticmethod
def filter_search_response(
result: SearchResponse,
allowed_domains: Sequence[str] | None = None,
blocked_domains: Sequence[str] | None = None,
) -> SearchResponse:
if allowed_domains is None and not blocked_domains:
return result
allowed: Final = tuple(_search_hostname(domain) for domain in allowed_domains or ())
blocked: Final = tuple(_search_hostname(domain) for domain in blocked_domains or ())
return result.model_copy(
update={
"results": [
item
for item in result.results
if (host := _search_hostname(item.url))
and (allowed_domains is None or _matches_search_domain(host, allowed))
and not _matches_search_domain(host, blocked)
]
}
)
@staticmethod
def format_search_response(result: SearchResponse) -> str:
"""

View file

@ -47,6 +47,11 @@ class RichWebSearchInput(TypedDict, total=False):
"""Two to five short keyword queries covering different angles."""
class WebSearchDomainFilters(TypedDict, total=False):
allowed_domains: ReadOnly[tuple[str, ...]]
blocked_domains: ReadOnly[tuple[str, ...]]
WebSearchToolResultErrorCode: TypeAlias = Literal[
"invalid_tool_input",
"unavailable",

View file

@ -0,0 +1,234 @@
from copy import deepcopy
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock
import pytest
import litellm
from litellm.integrations.websearch_interception.handler import WebSearchInterceptionLogger
from litellm.integrations.websearch_interception.tools import get_litellm_web_search_tool
from litellm.integrations.websearch_interception.transformation import WebSearchTransformation
from litellm.llms.base_llm.search.transformation import SearchResponse, SearchResult
from litellm.types.integrations.websearch_interception import SearchSucceeded
from litellm.types.utils import CallTypes
@pytest.fixture
def search_response() -> SearchResponse:
return SearchResponse(
results=[
SearchResult(title=f"Result {index}", url=url, snippet=f"Snippet {index}")
for index, url in enumerate(
(
"https://example.org/article",
"https://docs.example.org/guide",
"https://example.com/article",
"https://example.net/article",
"https://docs.example.net/guide",
"https://notexample.org/article",
"https://example.org.evil.test/article",
)
)
]
)
@pytest.fixture
def search_mock(monkeypatch: pytest.MonkeyPatch, search_response: SearchResponse) -> AsyncMock:
from litellm.proxy import proxy_server
mock: Final = AsyncMock(return_value=search_response)
monkeypatch.setattr(litellm, "asearch", mock)
monkeypatch.setattr(
proxy_server,
"llm_router",
SimpleNamespace(search_tools=[{"litellm_params": {"search_provider": "firecrawl"}}]),
)
return mock
@pytest.mark.asyncio
@pytest.mark.parametrize(
("filters", "expected_indices"),
[
({"allowed_domains": ["example.org"]}, (0, 1)),
({"blocked_domains": ["example.net"]}, (0, 1, 2, 5, 6)),
({"allowed_domains": ["example.org"], "blocked_domains": ["docs.example.org"]}, (0,)),
({"allowed_domains": ["example.org"], "blocked_domains": ["example.org"]}, ()),
({"allowed_domains": []}, ()),
({}, (0, 1, 2, 3, 4, 5, 6)),
],
)
async def test_short_circuit_filters_citations_and_text(
filters: dict[str, list[str]],
expected_indices: tuple[int, ...],
search_response: SearchResponse,
search_mock: AsyncMock,
) -> None:
logger: Final = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
response: Final = await logger.try_short_circuit_search(
model="interception-test-model",
messages=[{"role": "user", "content": "Find domain documentation"}],
tools=[{"type": "web_search_20250305", "name": "web_search", **filters}],
custom_llm_provider="github_copilot",
)
assert response is not None
expected: Final = [search_response.results[index] for index in expected_indices]
content: Final = response["content"]
assert [item["url"] for item in content[1]["content"]] == [item.url for item in expected]
assert content[2]["text"] == (
"\n\n".join(f"Title: {item.title}\nURL: {item.url}\nSnippet: {item.snippet}" for item in expected)
if expected
else "No search results found."
)
assert response["stop_reason"] == "end_turn"
search_mock.assert_awaited_once_with(
query="Find domain documentation",
search_provider="firecrawl",
**({"search_domain_filter": filters["allowed_domains"]} if "allowed_domains" in filters else {}),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["pre_request", "deployment", "both_hooks", "chat_completion", "responses"])
@pytest.mark.parametrize("source", ["definition", "arguments", "attempted_override"])
async def test_agentic_paths_preserve_domain_constraints(
path: str,
source: str,
search_response: SearchResponse,
search_mock: AsyncMock,
) -> None:
logger: Final = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
filters: Final = {"allowed_domains": ["example.org"], "blocked_domains": ["docs.example.org"]}
kwargs: Final = {
"custom_llm_provider": "bedrock",
"litellm_params": {"custom_llm_provider": "bedrock"},
"tools": [{"type": "web_search_20250305", "name": "web_search", **(filters if source != "arguments" else {})}],
}
if path in ("deployment", "both_hooks", "chat_completion"):
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
if path in ("pre_request", "both_hooks", "responses"):
await logger.async_pre_request_hook(model="interception-test-model", messages=[], kwargs=kwargs)
tool_input: Final = {
"query": "Find domain documentation",
**(filters if source == "arguments" else {}),
**({"allowed_domains": ["example.com"], "blocked_domains": []} if source == "attempted_override" else {}),
}
tool_calls: Final = [{"id": "toolu_test", "name": "litellm_web_search", "input": tool_input}]
optional_params: Final = {
"max_tokens": 1024,
**({"tools": kwargs["tools"]} if path != "pre_request" else {}),
}
expected_text: Final = "Title: Result 0\nURL: https://example.org/article\nSnippet: Snippet 0"
if path == "chat_completion":
chat_patch: Final = await logger._build_chat_completion_request_patch(
model="interception-test-model",
messages=[],
tool_calls=tool_calls,
optional_params=optional_params,
kwargs={},
)
assert chat_patch.messages[-1]["content"] == expected_text
elif path == "responses":
responses_patch: Final = await logger._build_responses_request_patch(
model="interception-test-model",
messages=[],
tool_calls=tool_calls,
optional_params=optional_params,
kwargs={},
)
assert responses_patch.messages[-1]["output"] == expected_text
else:
patch, outcomes = await logger._build_anthropic_request_patch(
model="interception-test-model",
messages=[],
tool_calls=tool_calls,
thinking_blocks=[],
anthropic_messages_optional_request_params=optional_params,
logging_obj=None,
kwargs={"tools": kwargs["tools"]},
)
assert patch.messages[-1]["content"][0]["content"] == expected_text
assert isinstance(outcomes[0], SearchSucceeded)
assert outcomes[0].response.results == [search_response.results[0]]
search_mock.assert_awaited_once_with(
query="Find domain documentation", search_provider="firecrawl", search_domain_filter=["example.org"]
)
@pytest.mark.parametrize(
("url", "retained"),
[
("https://EXAMPLE.ORG:8443/article", True),
("https://docs.EXAMPLE.ORG./article", True),
("//docs.example.org/article", True),
("https://example.org@evil.test/article", False),
("https://evil.test/example.org", False),
("https://evil.test/?url=https://example.org", False),
("https://notexample.org/article", False),
("https://example.org.evil.test/article", False),
("https://[invalid/article", False),
("/relative/article", False),
],
)
def test_domain_matching_uses_normalized_hostname(url: str, retained: bool) -> None:
item: Final = SearchResult(title="Title", url=url, snippet="Snippet")
response: Final = SearchResponse(results=[item])
filtered: Final = WebSearchTransformation.filter_search_response(
response, allowed_domains=["https://EXAMPLE.ORG:443"]
)
assert filtered.results == ([item] if retained else [])
assert response.results == [item]
@pytest.mark.asyncio
async def test_tool_builder_preserves_filters_for_execution(search_mock: AsyncMock) -> None:
tool: Final = get_litellm_web_search_tool(allowed_domains=["example.org"], blocked_domains=["docs.example.org"])
logger: Final = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
response: Final = await logger.try_short_circuit_search(
model="interception-test-model",
messages=[{"role": "user", "content": "Find domain documentation"}],
tools=[tool],
custom_llm_provider="github_copilot",
)
assert response["content"] == [
{"type": "text", "text": "Title: Result 0\nURL: https://example.org/article\nSnippet: Snippet 0"}
]
@pytest.mark.asyncio
@pytest.mark.parametrize("hook", ["request", "deployment"])
async def test_disabled_interception_leaves_native_domains_untouched(hook: str) -> None:
logger: Final = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
tool: Final = {
"type": "web_search_20250305",
"name": "web_search",
"allowed_domains": ["example.org"],
"blocked_domains": ["docs.example.org"],
}
kwargs: Final = {
"tools": [tool],
"custom_llm_provider": "anthropic",
"litellm_params": {"custom_llm_provider": "anthropic"},
}
original: Final = deepcopy(kwargs)
result: Final = (
await logger.async_pre_request_hook(model="interception-test-model", messages=[], kwargs=kwargs)
if hook == "request"
else await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
)
assert result is None
assert kwargs == original
@pytest.mark.parametrize("blocked_domains", [None, []])
def test_unconstrained_search_preserves_response(
search_response: SearchResponse, blocked_domains: list[str] | None
) -> None:
assert (
WebSearchTransformation.filter_search_response(search_response, blocked_domains=blocked_domains)
is search_response
)