mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 265126e6f4 into 461a58c40a
This commit is contained in:
commit
33912e1cee
4 changed files with 328 additions and 10 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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__"):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue