This commit is contained in:
JingHao-Leon 2026-10-04 12:47:44 -07:00 • committed by GitHub
commit 33912e1cee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 328 additions and 10 deletions

View file

@ -11,6 +11,7 @@ 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
@ -90,11 +91,55 @@ 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 _web_search_domain_strings(tool: Mapping[str, object], key: str) -> tuple[str, ...]:
value = tool.get(key)
if not isinstance(value, list):
return ()
return tuple(item for item in value if isinstance(item, str) and item)
def _extract_web_search_domain_filters(
tools: Sequence[dict[str, object]],
) -> Mapping[str, tuple[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.
"""
web_tools: Final = tuple(tool for tool in tools if is_web_search_tool(tool))
allowed: Final = tuple(
chain.from_iterable(_web_search_domain_strings(tool, "allowed_domains") for tool in web_tools)
)
blocked: Final = tuple(
chain.from_iterable(_web_search_domain_strings(tool, "blocked_domains") for tool in web_tools)
)
if not allowed and not blocked:
return None
if not blocked:
return MappingProxyType({"allowed_domains": allowed})
if not allowed:
return MappingProxyType({"blocked_domains": blocked})
return MappingProxyType({"allowed_domains": allowed, "blocked_domains": blocked})
class _PlanMetadataView(TypedDict):
websearch_native_blocks: Sequence[Mapping[str, object]] | None
@ -142,7 +187,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 +422,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 +519,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 +700,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 +1658,40 @@ 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.
request_domain_filters: Final = kwargs.get(WEBSEARCH_DOMAIN_FILTER_KEY) if kwargs is not None else None
domain_view: Final = (
request_domain_filters if isinstance(request_domain_filters, Mapping) else MappingProxyType({})
)
allowed: Final = tuple(
item for item in domain_view.get("allowed_domains", ()) if isinstance(item, str) and item
)
blocked: Final = tuple(
f"-{item}" for item in domain_view.get("blocked_domains", ()) if isinstance(item, str) and item
)
search_domain_filter: Final = [*allowed, *blocked] or None # mutable-ok: JSON request array, not mutated
if search_domain_filter is not None:
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,

View file

@ -3611,11 +3611,22 @@ class Router:
initial_kwargs["original_function"] = router_self._completion
initial_kwargs["messages"] = messages
router_self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs)
fallback_response = router_self.function_with_fallbacks(
**initial_kwargs,
# Pass the MidStreamFallbackError through the common fallback utils, like the
# async twin does. Calling function_with_fallbacks() here instead re-runs the
# original (failing) group first, and because each nested Router.completion()
# wraps its stream in this same iterator, every retry fails again only when
# the caller iterates it — recursing until the stack runs out.
fallback_response = run_async_function(
router_self.async_function_with_fallbacks_common_utils,
e,
disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs),
fallbacks=fallbacks,
context_window_fallbacks=context_window_fallbacks,
content_policy_fallbacks=content_policy_fallbacks,
model_group=model_group,
args=(),
kwargs=initial_kwargs,
include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True,
)
if hasattr(fallback_response, "__iter__"):

View file

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

View file

@ -3507,7 +3507,7 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste
}
return chunk
with patch.object(router, "function_with_fallbacks", return_value=NestedFallbackStream()):
with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=NestedFallbackStream()):
result = router._completion_streaming_iterator(
model_response=FailedStream(),
messages=[{"role": "user", "content": "hi"}],
@ -3685,7 +3685,7 @@ def test_completion_streaming_iterator_adopts_fallback_response_headers():
def __iter__(self):
return iter([])
with patch.object(router, "function_with_fallbacks", return_value=FallbackStream()):
with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=FallbackStream()):
result = router._completion_streaming_iterator(
model_response=FailedStream(),
messages=[{"role": "user", "content": "hi"}],
@ -3752,7 +3752,7 @@ def test_completion_streaming_iterator_fallback_on_429():
with patch.object(
router,
"function_with_fallbacks",
"async_function_with_fallbacks_common_utils",
return_value=mock_fallback_response,
) as mock_fallback:
result = router._completion_streaming_iterator(
@ -3765,10 +3765,99 @@ def test_completion_streaming_iterator_fallback_on_429():
assert mock_fallback.called
call_kwargs = mock_fallback.call_args
assert mock_fallback.call_args.args[0] is rate_limit_error
# Pre-first-chunk: should use original messages, no continuation prompt
assert call_kwargs.kwargs.get("messages") == messages
assert call_kwargs.kwargs.get("kwargs", {}).get("messages") == messages
# Verify original_function is _completion (sync)
assert call_kwargs.kwargs.get("original_function") == router._completion
assert call_kwargs.kwargs.get("kwargs", {}).get("original_function") == router._completion
def test_completion_streaming_iterator_routes_mid_stream_fallback_through_common_utils():
"""Regression (#43945): the sync mid-stream fallback re-entry must hand the
MidStreamFallbackError to async_function_with_fallbacks_common_utils, like the async
twin does. Calling function_with_fallbacks() instead re-runs the original (failing)
group first, and because each nested Router.completion() wraps its own stream in this
same iterator, every retry fails again only while being iterated — recursing until
the stack runs out (measured: 478 requests to the failing group, then
InternalServerError, with a healthy fallback configured)."""
from unittest.mock import MagicMock
from litellm.exceptions import MidStreamFallbackError
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake-key"},
}
],
)
messages = [{"role": "user", "content": "Test"}]
initial_kwargs = {"model": "gpt-4", "stream": True}
pre_first_chunk_error = MidStreamFallbackError(
message="upstream died before the first chunk",
model="gpt-4",
llm_provider="openai",
generated_content="",
is_pre_first_chunk=True,
)
class SyncIteratorImmediateError:
def __init__(self):
self.model = "gpt-4"
self.custom_llm_provider = "openai"
self.logging_obj = MagicMock()
self.chunks = []
def __iter__(self):
return self
def __next__(self):
raise pre_first_chunk_error
class FallbackStream:
def __init__(self):
self._chunks = iter(
[
litellm.ModelResponseStream(
choices=[{"index": 0, "delta": {"content": "from the fallback"}}]
)
]
)
def __iter__(self):
return self
def __next__(self):
return next(self._chunks)
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
return_value=FallbackStream(),
) as mock_utils:
with patch.object(router, "function_with_fallbacks") as mock_function_with_fallbacks:
result = router._completion_streaming_iterator(
model_response=SyncIteratorImmediateError(),
messages=messages,
initial_kwargs=initial_kwargs,
)
collected_chunks = list(result)
assert mock_utils.called
# the triggering error must reach the common utils, so cooldowns apply and
# the walk starts from the fallback list, not the failing group
assert mock_utils.call_args.args[0] is pre_first_chunk_error
assert mock_utils.call_args.kwargs.get("kwargs", {}).get("messages") == messages
assert not mock_function_with_fallbacks.called, (
"re-running the original group is what recurses; common utils already "
"excludes the deployment that raised"
)
assert len(collected_chunks) == 1
def test_completion_streaming_iterator_preserves_hidden_params():
@ -3980,7 +4069,7 @@ def test_completion_streaming_iterator_reraises_mid_chunk_error_with_no_text_con
mock_response = SyncIteratorNoTextChunkError()
with patch.object(router, "function_with_fallbacks") as mock_fallback:
with patch.object(router, "async_function_with_fallbacks_common_utils") as mock_fallback:
result = router._completion_streaming_iterator(
model_response=mock_response,
messages=messages,