mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 3941b5ce7f into 635085ac14
This commit is contained in:
commit
f9645adfe7
13 changed files with 437 additions and 12 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,11 +34,12 @@ 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,
|
||||
)
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse
|
||||
from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse
|
||||
from litellm.types.integrations.custom_logger import (
|
||||
CHAT_COMPLETION_AGENTIC_SURFACE,
|
||||
RESPONSES_AGENTIC_SURFACE,
|
||||
|
|
@ -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]]] = []
|
||||
|
|
@ -1507,19 +1549,46 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _provider_supports_rich_search(search_provider: str | None) -> bool:
|
||||
"""Whether the provider's search config accepts objective + multi-query input."""
|
||||
def _search_provider_config(search_provider: str | None) -> BaseSearchConfig | None:
|
||||
if not search_provider:
|
||||
return False
|
||||
return None
|
||||
try:
|
||||
from litellm.utils import ProviderConfigManager
|
||||
except ImportError:
|
||||
return False
|
||||
return None
|
||||
# SearchProviders is a str enum, so an unknown provider string simply
|
||||
# misses the config map and returns None rather than raising.
|
||||
config = ProviderConfigManager.get_provider_search_config(search_provider) # pyright: ignore[reportArgumentType] -- SearchProviders is a str enum, so the router's provider string hashes to the matching member; unknown strings miss the map and yield None
|
||||
return ProviderConfigManager.get_provider_search_config(search_provider) # pyright: ignore[reportArgumentType] -- SearchProviders is a str enum, so the router's provider string hashes to the matching member; unknown strings miss the map and yield None
|
||||
|
||||
@classmethod
|
||||
def _provider_supports_rich_search(cls, search_provider: str | None) -> bool:
|
||||
"""Whether the provider's search config accepts objective + multi-query input."""
|
||||
config: Final = cls._search_provider_config(search_provider)
|
||||
return config is not None and config.supports_rich_search_input()
|
||||
|
||||
@classmethod
|
||||
def _provider_domain_filter(
|
||||
cls,
|
||||
search_provider: str | None,
|
||||
configured_search_kwargs: Mapping[str, object],
|
||||
domains: WebSearchDomainFilter | None,
|
||||
) -> list[str] | None:
|
||||
"""
|
||||
The client's ``allowed_domains`` hosts to send as ``search_domain_filter``, or None.
|
||||
|
||||
Only sent when the search tool config sets no domain filter of its own (the
|
||||
operator's filter must never be replaced) and the provider applies every
|
||||
entry; results are filtered by URL afterwards either way.
|
||||
"""
|
||||
if domains is None or not domains.allowed:
|
||||
return None
|
||||
config: Final = cls._search_provider_config(search_provider)
|
||||
if config is None or not config.domain_filter_params().isdisjoint(configured_search_kwargs):
|
||||
return None
|
||||
hosts: Final = list(dict.fromkeys(host for host in map(domain_host, domains.allowed) if host))
|
||||
max_entries: Final = config.max_search_domain_filter_entries()
|
||||
return hosts if hosts and (max_entries is None or len(hosts) <= max_entries) else None
|
||||
|
||||
async def _execute_search(
|
||||
self,
|
||||
query: str,
|
||||
|
|
@ -1598,19 +1667,32 @@ 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
|
||||
provider_domain_filter: Final = self._provider_domain_filter(
|
||||
search_provider, configured_search_kwargs, domains
|
||||
)
|
||||
named_params: Final[_AsearchNamedParams] = (
|
||||
_NO_ASEARCH_NAMED
|
||||
if provider_domain_filter is None
|
||||
else {"search_domain_filter": provider_domain_filter}
|
||||
)
|
||||
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,35 @@ 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 matches whole path segments."""
|
||||
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}")
|
||||
path_matches: Final = not rule_path or target.path == rule_path or target.path.startswith(f"{rule_path}/")
|
||||
return host_matches and path_matches
|
||||
|
||||
|
||||
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 +553,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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -107,6 +107,28 @@ class BaseSearchConfig:
|
|||
"""
|
||||
return False
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
"""
|
||||
Request params through which this provider restricts results by domain:
|
||||
the unified ``search_domain_filter`` plus any provider-native keys
|
||||
(e.g. ``include_domains``) that reach the request through pass-through.
|
||||
|
||||
Integrations that add their own domain filter (e.g. websearch
|
||||
interception) leave a request alone when any of these is already set,
|
||||
so a filter configured on the search tool is never replaced.
|
||||
"""
|
||||
return frozenset(("search_domain_filter",))
|
||||
|
||||
def max_search_domain_filter_entries(self) -> int | None:
|
||||
"""
|
||||
How many ``search_domain_filter`` entries this provider's request
|
||||
applies, or None when it takes the whole list.
|
||||
|
||||
Providers whose API takes a single domain return 1, so callers with
|
||||
more domains than that know a forwarded filter would drop the rest.
|
||||
"""
|
||||
return None
|
||||
|
||||
def get_http_method(self) -> Literal["GET", "POST"]:
|
||||
"""
|
||||
Get HTTP method for search requests.
|
||||
|
|
|
|||
|
|
@ -32,6 +32,13 @@ class DataForSEOSearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "DataForSEO"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "domain", "target"))
|
||||
|
||||
def max_search_domain_filter_entries(self) -> int | None:
|
||||
# `domain` takes a single domain
|
||||
return 1
|
||||
|
||||
def get_http_method(self) -> Literal["GET", "POST"]:
|
||||
"""
|
||||
DataForSEO uses POST requests with JSON body.
|
||||
|
|
|
|||
|
|
@ -53,6 +53,9 @@ class ExaAISearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "Exa AI"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "includeDomains", "excludeDomains"))
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -62,6 +62,13 @@ class GooglePSESearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "Google PSE"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "siteSearch", "siteSearchFilter"))
|
||||
|
||||
def max_search_domain_filter_entries(self) -> int | None:
|
||||
# `siteSearch` takes a single domain
|
||||
return 1
|
||||
|
||||
def get_http_method(self) -> Literal["GET", "POST"]:
|
||||
"""
|
||||
Google PSE uses GET requests with query parameters.
|
||||
|
|
|
|||
|
|
@ -49,6 +49,9 @@ class LinkupSearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "Linkup"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "includeDomains", "excludeDomains"))
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -86,6 +86,9 @@ class NimbleSearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "Nimble"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "include_domains", "exclude_domains"))
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict[str, str], # mutable-ok: BaseSearchConfig.validate_environment signature
|
||||
|
|
|
|||
|
|
@ -90,6 +90,9 @@ class ParallelAISearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "Parallel AI"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "include_domains", "exclude_domains"))
|
||||
|
||||
def supports_rich_search_input(self) -> bool:
|
||||
# The v1 search API takes `objective` + multiple `search_queries`
|
||||
# natively; sending both is the documented best practice.
|
||||
|
|
|
|||
|
|
@ -52,6 +52,9 @@ class TavilySearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "Tavily"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "include_domains", "exclude_domains"))
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -50,6 +50,9 @@ class YouComSearchConfig(BaseSearchConfig):
|
|||
def ui_friendly_name() -> str:
|
||||
return "You.com"
|
||||
|
||||
def domain_filter_params(self) -> frozenset[str]:
|
||||
return frozenset(("search_domain_filter", "include_domains", "exclude_domains"))
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -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,251 @@ 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
|
||||
|
||||
|
||||
_PATH_RULE_URLS = (
|
||||
"https://example.com/docs/intro",
|
||||
"https://example.com/docs-private/x",
|
||||
"https://example.com/other/y",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("domain_field", "rule", "expected_urls"),
|
||||
[
|
||||
("allowed_domains", "example.com/docs/", ["https://example.com/docs/intro"]),
|
||||
("allowed_domains", "example.com/docs", ["https://example.com/docs/intro"]),
|
||||
("blocked_domains", "example.com/docs/", ["https://example.com/docs-private/x", "https://example.com/other/y"]),
|
||||
],
|
||||
)
|
||||
async def test_messages_web_search_domain_path_rules_match_whole_path_segments(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
domain_field: str,
|
||||
rule: str,
|
||||
expected_urls: list[str],
|
||||
):
|
||||
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
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _perplexity_router())
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"asearch",
|
||||
AsyncMock(
|
||||
return_value=SearchResponse(
|
||||
object="search",
|
||||
results=[SearchResult(title=url, url=url, snippet="snippet") for url in _PATH_RULE_URLS],
|
||||
)
|
||||
),
|
||||
)
|
||||
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: [rule]}],
|
||||
)
|
||||
|
||||
text = next(block["text"] for block in response["content"] if block["type"] == "text")
|
||||
assert [url for url in _PATH_RULE_URLS if url in text] == expected_urls
|
||||
|
||||
|
||||
@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_router_chat_completion_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 import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.base_llm.search.transformation import SearchResult
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
sent_messages: list[list[dict]] = []
|
||||
|
||||
class _SentMessages(CustomLogger):
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
sent_messages.append(messages)
|
||||
|
||||
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=["openai"]), _SentMessages()]
|
||||
)
|
||||
router = Router(
|
||||
model_list=[{"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake-key"}}]
|
||||
)
|
||||
|
||||
await router.acompletion(
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "litellm"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search", domain_field: [domain]}],
|
||||
mock_tool_calls=[
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": LITELLM_WEB_SEARCH_TOOL_NAME, "arguments": '{"query": "litellm"}'},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
search_result = next(m["content"] for messages in sent_messages for m in messages if m["role"] == "tool")
|
||||
assert [url for url in _DOMAIN_FILTER_URLS if url in search_result] == expected_urls
|
||||
assert mock_asearch.await_args.kwargs.get("search_domain_filter") == expected_provider_filter
|
||||
|
||||
|
||||
_PROVIDER_DOMAIN_FIELDS = ("includeDomains", "domain", "siteSearch", "siteSearchFilter")
|
||||
|
||||
|
||||
def _sent_domain_fields(request) -> dict:
|
||||
import json
|
||||
|
||||
sent = dict(request.url.params) if request.method == "GET" else json.loads(request.content)
|
||||
task = sent[0] if isinstance(sent, list) else sent
|
||||
return {field: task[field] for field in _PROVIDER_DOMAIN_FIELDS if field in task}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("search_tool_params", "allowed_domains", "expected_sent"),
|
||||
[
|
||||
pytest.param(
|
||||
{"search_provider": "exa_ai", "includeDomains": ["trusted.example"]},
|
||||
["attacker.example"],
|
||||
{"includeDomains": ["trusted.example"]},
|
||||
id="exa-keeps-operator-includeDomains",
|
||||
),
|
||||
pytest.param(
|
||||
{"search_provider": "linkup", "includeDomains": ["trusted.example"]},
|
||||
["attacker.example"],
|
||||
{"includeDomains": ["trusted.example"]},
|
||||
id="linkup-keeps-operator-includeDomains",
|
||||
),
|
||||
pytest.param(
|
||||
{"search_provider": "exa_ai"},
|
||||
["example.org", "https://www.example.com/docs"],
|
||||
{"includeDomains": ["example.org", "www.example.com"]},
|
||||
id="exa-takes-every-allowed-domain",
|
||||
),
|
||||
pytest.param(
|
||||
{"search_provider": "dataforseo", "api_key": "login:password"},
|
||||
["example.org", "example.com"],
|
||||
{},
|
||||
id="dataforseo-single-domain-gets-no-list",
|
||||
),
|
||||
pytest.param(
|
||||
{"search_provider": "google_pse"},
|
||||
["example.org", "example.com"],
|
||||
{},
|
||||
id="google-pse-not-narrowed-to-first-domain",
|
||||
),
|
||||
pytest.param(
|
||||
{"search_provider": "google_pse"},
|
||||
["example.org"],
|
||||
{"siteSearch": "example.org", "siteSearchFilter": "i"},
|
||||
id="google-pse-takes-one-allowed-domain",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_messages_web_search_forwards_allowed_domains_only_where_the_provider_applies_them(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
search_tool_params: dict,
|
||||
allowed_domains: list[str],
|
||||
expected_sent: dict,
|
||||
):
|
||||
import httpx
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.pass_through.messages.handler import anthropic_messages
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.search_tools = [
|
||||
{"search_tool_name": "operator-search", "litellm_params": {"api_key": "fake-key", **search_tool_params}}
|
||||
]
|
||||
monkeypatch.setenv("GOOGLE_PSE_ENGINE_ID", "fake-engine")
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(litellm, "callbacks", [WebSearchInterceptionLogger(enabled_providers=["bedrock"])])
|
||||
|
||||
with respx.mock(assert_all_called=False) as provider:
|
||||
provider.route().mock(return_value=httpx.Response(200, json={"results": [], "tasks": []}))
|
||||
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", "allowed_domains": allowed_domains}],
|
||||
)
|
||||
sent = _sent_domain_fields(provider.calls.last.request)
|
||||
|
||||
assert sent == expected_sent
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue