mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
a292fd409f
commit
ccd6ef8d44
4 changed files with 145 additions and 6 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue