mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(websearch_interception): apply allowed_domains / blocked_domains to intercepted searches
With websearch_interception enabled, the Anthropic-native web_search tool's allowed_domains / blocked_domains were silently dropped: the tool is replaced by the standard LiteLLM web search tool (which has no domain fields) and _execute_search() called litellm.asearch() without search_domain_filter. A /v1/messages request that limited the search to one domain got results from every domain (issue #44188). - Pre-call hooks now stash the domain limits on the request kwargs before the native tool is replaced (same side-channel pattern as the native-blocks flag), and the short-circuit path applies them the same way. - _execute_search() forwards them to litellm.asearch() as search_domain_filter: allowed_domains pass through as an allowlist and blocked_domains as '-'-prefixed exclusions, the convention litellm.asearch's providers (e.g. Perplexity) use. - _AsearchNamedParams now permits search_domain_filter. - New unit tests cover extraction, hook stashing, and the asearch call wiring.
This commit is contained in:
parent
aeb88738e8
commit
3cf5bb2753
2 changed files with 214 additions and 2 deletions
|
|
@ -90,11 +90,51 @@ 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"
|
||||
|
||||
# Key used to flag, on per-request kwargs, that the originating client sent
|
||||
# domain filters (``allowed_domains`` / ``blocked_domains``) on an
|
||||
# Anthropic-native ``web_search_*`` tool. The standard LiteLLM tool drops
|
||||
# them (the model must not see client policy), so they are stashed here and
|
||||
# applied to the downstream ``litellm.asearch()`` call as
|
||||
# ``search_domain_filter``.
|
||||
WEBSEARCH_DOMAIN_FILTER_KEY: Final = "_websearch_interception_domain_filter"
|
||||
|
||||
_RESPONSE_CONTENT_FIELD: Final = "content"
|
||||
|
||||
_ResponseT: Final = TypeVar("_ResponseT")
|
||||
|
||||
|
||||
def _extract_web_search_domain_filters(
|
||||
tools: Sequence[dict[str, object]],
|
||||
) -> dict[str, list[str]] | None:
|
||||
"""Collect ``allowed_domains`` / ``blocked_domains`` from web search tools.
|
||||
|
||||
Anthropic-native ``web_search_*`` tools carry optional domain limits. The
|
||||
standard LiteLLM tool deliberately drops them (client policy, not model
|
||||
input), so they are collected here, stashed on the request kwargs, and
|
||||
applied to the downstream ``litellm.asearch()`` call as
|
||||
``search_domain_filter``.
|
||||
|
||||
Returns None when no web search tool carries a domain limit.
|
||||
"""
|
||||
allowed: list[str] = []
|
||||
blocked: list[str] = []
|
||||
for tool in tools:
|
||||
if not is_web_search_tool(tool):
|
||||
continue
|
||||
for key, bucket in (("allowed_domains", allowed), ("blocked_domains", blocked)):
|
||||
value = tool.get(key)
|
||||
if isinstance(value, list):
|
||||
bucket.extend(item for item in value if isinstance(item, str) and item)
|
||||
if not allowed and not blocked:
|
||||
return None
|
||||
domain_filters: dict[str, list[str]] = {}
|
||||
if allowed:
|
||||
domain_filters["allowed_domains"] = allowed
|
||||
if blocked:
|
||||
domain_filters["blocked_domains"] = blocked
|
||||
return domain_filters
|
||||
|
||||
|
||||
class _PlanMetadataView(TypedDict):
|
||||
websearch_native_blocks: Sequence[Mapping[str, object]] | None
|
||||
|
||||
|
|
@ -142,7 +182,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]
|
||||
|
|
@ -377,6 +417,12 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
None,
|
||||
)
|
||||
|
||||
# Apply any domain limits the client set on the native web search
|
||||
# tool before the search executes.
|
||||
domain_filters: Final = _extract_web_search_domain_filters(tools)
|
||||
if domain_filters is not None and isinstance(kwargs, dict):
|
||||
kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters
|
||||
|
||||
outcome: Final = await self._short_circuit_search_outcome(query, kwargs=kwargs)
|
||||
search_result_text: Final = WebSearchTransformation.search_outcome_text(outcome)
|
||||
|
||||
|
|
@ -468,6 +514,12 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if any(is_anthropic_native_web_search_tool(t) for t in tools):
|
||||
kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True
|
||||
|
||||
# Same for domain limits: stash them before the native tool is
|
||||
# replaced, so the downstream search can apply them.
|
||||
domain_filters: Final = _extract_web_search_domain_filters(tools)
|
||||
if domain_filters is not None:
|
||||
kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters
|
||||
|
||||
# Convert native/custom web_search tools to LiteLLM standard
|
||||
converted_tools: Final = []
|
||||
for tool in tools:
|
||||
|
|
@ -643,6 +695,12 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if any(is_anthropic_native_web_search_tool(t) for t in tools):
|
||||
kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True
|
||||
|
||||
# Same for domain limits: stash them before the native tool is
|
||||
# replaced, so the downstream search can apply them.
|
||||
domain_filters: Final = _extract_web_search_domain_filters(tools)
|
||||
if domain_filters is not None:
|
||||
kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] = domain_filters
|
||||
|
||||
# Convert native web search tools to LiteLLM standard
|
||||
converted_tools: Final[list[dict[str, object]]] = []
|
||||
for tool in tools:
|
||||
|
|
@ -1595,17 +1653,39 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
rich_objective = rich.get("objective")
|
||||
if rich_objective and "objective" not in configured_search_kwargs:
|
||||
configured_search_kwargs["objective"] = rich_objective
|
||||
# Domain limits stashed by the pre-call hooks from the client's
|
||||
# native web_search tool. ``allowed_domains`` pass through as an
|
||||
# allowlist and ``blocked_domains`` as '-'-prefixed exclusions —
|
||||
# the convention litellm.asearch()'s providers (e.g. Perplexity)
|
||||
# use for search_domain_filter.
|
||||
search_domain_filter: list[str] | None = None
|
||||
request_domain_filters: Final = kwargs.get(WEBSEARCH_DOMAIN_FILTER_KEY) if kwargs is not None else None
|
||||
if isinstance(request_domain_filters, dict):
|
||||
allowed: Final[list[str]] = [
|
||||
item for item in request_domain_filters.get("allowed_domains", []) if isinstance(item, str) and item
|
||||
]
|
||||
blocked: Final[list[str]] = [
|
||||
item for item in request_domain_filters.get("blocked_domains", []) if isinstance(item, str) and item
|
||||
]
|
||||
if allowed or blocked:
|
||||
search_domain_filter = allowed + [f"-{item}" for item in blocked]
|
||||
verbose_logger.debug("WebSearchInterception: Applying domain filter %s", search_domain_filter)
|
||||
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
|
||||
query=query_arg,
|
||||
search_provider=search_provider,
|
||||
search_domain_filter=search_domain_filter,
|
||||
**_NO_ASEARCH_NAMED,
|
||||
**search_kwargs,
|
||||
)
|
||||
if search_metadata is None
|
||||
else await litellm.asearch(
|
||||
query=query_arg,
|
||||
search_provider=search_provider,
|
||||
search_domain_filter=search_domain_filter,
|
||||
litellm_metadata=search_metadata,
|
||||
**_NO_ASEARCH_NAMED,
|
||||
**search_kwargs,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,132 @@
|
|||
"""
|
||||
Tests for domain-limit passthrough in web search interception.
|
||||
|
||||
Covers bug #44188: ``allowed_domains`` / ``blocked_domains`` set on an
|
||||
Anthropic-native ``web_search_*`` tool must survive the conversion to the
|
||||
standard LiteLLM web search tool and be applied to the downstream
|
||||
``litellm.asearch()`` call as ``search_domain_filter``.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.websearch_interception.handler import (
|
||||
WEBSEARCH_DOMAIN_FILTER_KEY,
|
||||
WebSearchInterceptionLogger,
|
||||
_extract_web_search_domain_filters,
|
||||
)
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse, SearchResult
|
||||
|
||||
|
||||
def _make_search_response() -> SearchResponse:
|
||||
return SearchResponse(
|
||||
results=[
|
||||
SearchResult(
|
||||
title="LiteLLM Docs",
|
||||
url="https://docs.litellm.ai/",
|
||||
snippet="Unified interface for LLMs.",
|
||||
date=None,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
class TestExtractWebSearchDomainFilters:
|
||||
def test_collects_allowed_and_blocked(self):
|
||||
tools = [
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com", "x.com"],
|
||||
}
|
||||
]
|
||||
assert _extract_web_search_domain_filters(tools) == {
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com", "x.com"],
|
||||
}
|
||||
|
||||
def test_returns_none_without_domain_limits(self):
|
||||
tools = [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}]
|
||||
assert _extract_web_search_domain_filters(tools) is None
|
||||
|
||||
def test_ignores_domains_on_non_web_search_tools(self):
|
||||
tools = [
|
||||
{"name": "bash", "allowed_domains": ["example.com"]},
|
||||
]
|
||||
assert _extract_web_search_domain_filters(tools) is None
|
||||
|
||||
def test_ignores_non_string_entries(self):
|
||||
tools = [
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"allowed_domains": ["docs.litellm.ai", 42, None],
|
||||
}
|
||||
]
|
||||
assert _extract_web_search_domain_filters(tools) == {"allowed_domains": ["docs.litellm.ai"]}
|
||||
|
||||
|
||||
class TestDeploymentHookStashesDomainFilters:
|
||||
@pytest.mark.asyncio
|
||||
async def test_stashes_filters_for_native_tool(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
kwargs = {
|
||||
"tools": [
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com"],
|
||||
}
|
||||
],
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
}
|
||||
out = await logger.async_pre_call_deployment_hook(kwargs, None)
|
||||
assert out is not None
|
||||
assert kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] == {
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com"],
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_stash_without_domain_limits(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
kwargs = {
|
||||
"tools": [{"type": "web_search_20250305", "name": "web_search"}],
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
}
|
||||
await logger.async_pre_call_deployment_hook(kwargs, None)
|
||||
assert WEBSEARCH_DOMAIN_FILTER_KEY not in kwargs
|
||||
|
||||
|
||||
class TestExecuteSearchAppliesDomainFilter:
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_search_domain_filter(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
kwargs = {
|
||||
WEBSEARCH_DOMAIN_FILTER_KEY: {
|
||||
"allowed_domains": ["docs.litellm.ai"],
|
||||
"blocked_domains": ["twitter.com"],
|
||||
}
|
||||
}
|
||||
asearch = AsyncMock(return_value=_make_search_response())
|
||||
with patch("litellm.asearch", asearch):
|
||||
await logger._execute_search("what is litellm", kwargs=kwargs)
|
||||
|
||||
assert asearch.await_count == 1
|
||||
assert asearch.await_args.kwargs.get("search_domain_filter") == [
|
||||
"docs.litellm.ai",
|
||||
"-twitter.com",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_filter_when_kwargs_empty(self):
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
asearch = AsyncMock(return_value=_make_search_response())
|
||||
with patch("litellm.asearch", asearch):
|
||||
await logger._execute_search("what is litellm", kwargs={})
|
||||
|
||||
assert asearch.await_count == 1
|
||||
assert asearch.await_args.kwargs.get("search_domain_filter") is None
|
||||
Loading…
Add table
Reference in a new issue