mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(websearch): preserve Anthropic domain filters
Keep allowed and blocked domains through tool conversion and search execution, then filter results before building text and citations Fixes #44188
This commit is contained in:
parent
9e31afbc6d
commit
f5bc46e08e
5 changed files with 416 additions and 18 deletions
|
|
@ -14,6 +14,7 @@ from dataclasses import dataclass
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import Never, ReadOnly
|
||||
|
||||
import litellm
|
||||
|
|
@ -25,10 +26,12 @@ from litellm.integrations.websearch_interception.tools import (
|
|||
get_litellm_web_search_tool,
|
||||
get_litellm_web_search_tool_openai,
|
||||
get_litellm_web_search_tool_responses,
|
||||
get_web_search_domain_filters,
|
||||
is_anthropic_native_web_search_tool,
|
||||
is_web_search_tool,
|
||||
is_web_search_tool_chat_completion,
|
||||
is_web_search_tool_responses,
|
||||
resolve_web_search_domain_filters,
|
||||
)
|
||||
from litellm.integrations.websearch_interception.transformation import (
|
||||
WebSearchTransformation,
|
||||
|
|
@ -49,6 +52,7 @@ from litellm.types.integrations.websearch_interception import (
|
|||
RichWebSearchInput,
|
||||
SearchFailed,
|
||||
SearchOutcome,
|
||||
WebSearchDomainFilters,
|
||||
WebSearchInterceptionConfig,
|
||||
)
|
||||
from litellm.types.llms.anthropic import AnthropicThinkingParam
|
||||
|
|
@ -91,6 +95,7 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
|
|||
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
|
||||
|
||||
_RESPONSE_CONTENT_FIELD: Final = "content"
|
||||
_WEB_SEARCH_REQUEST: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
_ResponseT: Final = TypeVar("_ResponseT")
|
||||
|
||||
|
|
@ -377,7 +382,11 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
None,
|
||||
)
|
||||
|
||||
outcome: Final = await self._short_circuit_search_outcome(query, kwargs=kwargs)
|
||||
outcome: Final = await self._short_circuit_search_outcome(
|
||||
query,
|
||||
kwargs=kwargs,
|
||||
domains=resolve_web_search_domain_filters(_WEB_SEARCH_REQUEST.validate_python({"tools": tools})),
|
||||
)
|
||||
search_result_text: Final = WebSearchTransformation.search_outcome_text(outcome)
|
||||
|
||||
content: Final[list[dict[str, object]]] = []
|
||||
|
|
@ -473,7 +482,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
for tool in tools:
|
||||
if is_web_search_tool(tool):
|
||||
# Convert to LiteLLM standard web search tool
|
||||
converted_tool = get_litellm_web_search_tool_openai()
|
||||
converted_tool = get_litellm_web_search_tool_openai(**get_web_search_domain_filters(tool))
|
||||
converted_tools.append(converted_tool)
|
||||
verbose_logger.debug(
|
||||
"WebSearchInterception: Converted %s (type=%s) to %s",
|
||||
|
|
@ -647,7 +656,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
converted_tools: Final[list[dict[str, object]]] = []
|
||||
for tool in tools:
|
||||
if is_web_search_tool(tool):
|
||||
standard_tool = get_litellm_web_search_tool()
|
||||
standard_tool = get_litellm_web_search_tool(
|
||||
**get_web_search_domain_filters(_WEB_SEARCH_REQUEST.validate_python(tool))
|
||||
)
|
||||
converted_tools.append(standard_tool)
|
||||
verbose_logger.debug(
|
||||
"WebSearchInterception: Converted %s (type=%s) to %s",
|
||||
|
|
@ -1188,7 +1199,19 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
search_tasks: Final = [
|
||||
(
|
||||
self._execute_search(
|
||||
tool_call["input"]["query"], kwargs=kwargs, rich=self._rich_search_input(tool_call["input"])
|
||||
tool_call["input"]["query"],
|
||||
kwargs=kwargs,
|
||||
rich=self._rich_search_input(tool_call["input"]),
|
||||
**resolve_web_search_domain_filters(
|
||||
_WEB_SEARCH_REQUEST.validate_python(
|
||||
{
|
||||
"tools": optional_params["tools"]
|
||||
if "tools" in optional_params
|
||||
else kwargs.get("tools"),
|
||||
"input": tool_call["input"],
|
||||
}
|
||||
)
|
||||
),
|
||||
)
|
||||
if isinstance(tool_call.get("input"), dict) and tool_call["input"].get("query")
|
||||
else self._create_empty_search_result()
|
||||
|
|
@ -1408,7 +1431,23 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if query:
|
||||
verbose_logger.debug("WebSearchInterception: Queuing search for query='%s'", query)
|
||||
search_tasks.append(
|
||||
self._execute_search(query, kwargs=kwargs, rich=self._rich_search_input(tool_call["input"]))
|
||||
self._execute_search(
|
||||
query,
|
||||
kwargs=kwargs,
|
||||
rich=self._rich_search_input(tool_call["input"]),
|
||||
**resolve_web_search_domain_filters(
|
||||
_WEB_SEARCH_REQUEST.validate_python(
|
||||
{
|
||||
"tools": (
|
||||
anthropic_messages_optional_request_params["tools"]
|
||||
if "tools" in anthropic_messages_optional_request_params
|
||||
else kwargs.get("tools")
|
||||
),
|
||||
"input": tool_call["input"],
|
||||
}
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
else:
|
||||
verbose_logger.debug("WebSearchInterception: Tool call %s has no query", tool_call["id"])
|
||||
|
|
@ -1467,12 +1506,14 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
)
|
||||
return patch, search_outcomes
|
||||
|
||||
async def _short_circuit_search_outcome(self, query: str, kwargs: Mapping[str, object] | None) -> SearchOutcome:
|
||||
async def _short_circuit_search_outcome(
|
||||
self, query: str, kwargs: Mapping[str, object] | None, domains: WebSearchDomainFilters
|
||||
) -> SearchOutcome:
|
||||
try:
|
||||
result: Final = (
|
||||
await self._execute_search(query)
|
||||
await self._execute_search(query, **domains)
|
||||
if kwargs is None
|
||||
else await self._execute_search(query, kwargs=kwargs)
|
||||
else await self._execute_search(query, kwargs=kwargs, **domains)
|
||||
)
|
||||
except Exception as e:
|
||||
return WebSearchTransformation.search_outcome(e)
|
||||
|
|
@ -1525,6 +1566,8 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
query: str,
|
||||
kwargs: Mapping[str, object] | None = None,
|
||||
rich: RichWebSearchInput | None = None,
|
||||
allowed_domains: Sequence[str] | None = None,
|
||||
blocked_domains: Sequence[str] | None = None,
|
||||
) -> tuple[str, SearchResponse | None]:
|
||||
"""
|
||||
Execute a single web search using router's search tools.
|
||||
|
|
@ -1596,7 +1639,11 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if rich_objective and "objective" not in configured_search_kwargs:
|
||||
configured_search_kwargs["objective"] = rich_objective
|
||||
search_kwargs: Final = MappingProxyType(
|
||||
{**configured_search_kwargs, **parent_correlation.as_search_kwargs()}
|
||||
{
|
||||
**configured_search_kwargs,
|
||||
**parent_correlation.as_search_kwargs(),
|
||||
**({"search_domain_filter": list(allowed_domains)} if allowed_domains is not None else {}),
|
||||
}
|
||||
)
|
||||
result: Final = (
|
||||
await litellm.asearch(
|
||||
|
|
@ -1613,12 +1660,19 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
)
|
||||
|
||||
# Format using transformation function
|
||||
search_result_text: Final = WebSearchTransformation.format_search_response(result)
|
||||
filtered_result: Final = WebSearchTransformation.filter_search_response(
|
||||
result, allowed_domains=allowed_domains, blocked_domains=blocked_domains
|
||||
)
|
||||
search_result_text: Final = (
|
||||
"No search results found."
|
||||
if filtered_result is not result and not filtered_result.results
|
||||
else WebSearchTransformation.format_search_response(filtered_result)
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
"WebSearchInterception: Search completed for '%s', got %s chars", query, len(search_result_text)
|
||||
)
|
||||
return search_result_text, result
|
||||
return search_result_text, filtered_result
|
||||
except Exception as e:
|
||||
verbose_logger.error("WebSearchInterception: Search failed for '%s': %s", query, e)
|
||||
raise
|
||||
|
|
@ -1863,7 +1917,23 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
|
||||
if query:
|
||||
verbose_logger.debug("WebSearchInterception: Queuing search for query='%s'", query)
|
||||
search_tasks.append(self._execute_search(query, kwargs=kwargs, rich=self._rich_search_input(tool_args)))
|
||||
search_tasks.append(
|
||||
self._execute_search(
|
||||
query,
|
||||
kwargs=kwargs,
|
||||
rich=self._rich_search_input(tool_args),
|
||||
**resolve_web_search_domain_filters(
|
||||
_WEB_SEARCH_REQUEST.validate_python(
|
||||
{
|
||||
"tools": optional_params["tools"]
|
||||
if "tools" in optional_params
|
||||
else kwargs.get("tools"),
|
||||
"input": tool_args,
|
||||
}
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
else:
|
||||
verbose_logger.debug("WebSearchInterception: Tool call %s has no query", tool_call.get("id"))
|
||||
# Add empty result for tools without query
|
||||
|
|
|
|||
|
|
@ -6,10 +6,17 @@ Native provider tools (like Anthropic's web_search_20250305) are converted
|
|||
to this format for consistent interception and execution.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME
|
||||
from litellm.types.integrations.websearch_interception import WebSearchDomainFilters
|
||||
|
||||
_DOMAINS: Final = TypeAdapter(tuple[str, ...])
|
||||
_TOOL_MAPPING: Final = TypeAdapter(dict[str, object])
|
||||
_TOOLS: Final = TypeAdapter(list[dict[str, object]])
|
||||
|
||||
_WEB_SEARCH_TOOL_DESCRIPTION: Final = (
|
||||
"Search the web for information. Use this when you need current "
|
||||
|
|
@ -17,7 +24,10 @@ _WEB_SEARCH_TOOL_DESCRIPTION: Final = (
|
|||
)
|
||||
|
||||
|
||||
def _web_search_input_schema() -> dict[str, object]: # mutable-ok: plain-dict tool shape, as the get_* builders
|
||||
def _web_search_input_schema(
|
||||
allowed_domains: Sequence[str] | None = None,
|
||||
blocked_domains: Sequence[str] | None = None,
|
||||
) -> Mapping[str, object]:
|
||||
"""
|
||||
JSON schema for the web search tool's input, shared by every tool format.
|
||||
|
||||
|
|
@ -51,12 +61,20 @@ def _web_search_input_schema() -> dict[str, object]: # mutable-ok: plain-dict t
|
|||
"objective for the best results."
|
||||
),
|
||||
},
|
||||
**{
|
||||
name: {"type": "array", "items": {"type": "string"}, "default": list(domains)}
|
||||
for name, domains in (("allowed_domains", allowed_domains), ("blocked_domains", blocked_domains))
|
||||
if domains is not None
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
}
|
||||
|
||||
|
||||
def get_litellm_web_search_tool() -> dict[str, object]:
|
||||
def get_litellm_web_search_tool(
|
||||
allowed_domains: Sequence[str] | None = None,
|
||||
blocked_domains: Sequence[str] | None = None,
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Get the standard LiteLLM web search tool definition.
|
||||
|
||||
|
|
@ -78,11 +96,14 @@ def get_litellm_web_search_tool() -> dict[str, object]:
|
|||
return {
|
||||
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"description": _WEB_SEARCH_TOOL_DESCRIPTION,
|
||||
"input_schema": _web_search_input_schema(),
|
||||
"input_schema": _web_search_input_schema(allowed_domains, blocked_domains),
|
||||
}
|
||||
|
||||
|
||||
def get_litellm_web_search_tool_openai() -> dict[str, object]:
|
||||
def get_litellm_web_search_tool_openai(
|
||||
allowed_domains: Sequence[str] | None = None,
|
||||
blocked_domains: Sequence[str] | None = None,
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Get the standard LiteLLM web search tool definition in OpenAI format.
|
||||
|
||||
|
|
@ -98,7 +119,7 @@ def get_litellm_web_search_tool_openai() -> dict[str, object]:
|
|||
"function": {
|
||||
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"description": _WEB_SEARCH_TOOL_DESCRIPTION,
|
||||
"parameters": _web_search_input_schema(),
|
||||
"parameters": _web_search_input_schema(allowed_domains, blocked_domains),
|
||||
},
|
||||
}
|
||||
|
||||
|
|
@ -123,6 +144,39 @@ def get_litellm_web_search_tool_responses() -> dict[str, object]:
|
|||
}
|
||||
|
||||
|
||||
def get_web_search_domain_filters(tool: Mapping[str, object]) -> WebSearchDomainFilters:
|
||||
function: Final = tool.get("function")
|
||||
if isinstance(function, Mapping):
|
||||
return get_web_search_domain_filters(_TOOL_MAPPING.validate_python(function))
|
||||
schema: Final = _TOOL_MAPPING.validate_python(tool.get("input_schema", tool.get("parameters", {})))
|
||||
properties: Final = _TOOL_MAPPING.validate_python(schema.get("properties", {}))
|
||||
defaults: Final = {
|
||||
name: _TOOL_MAPPING.validate_python(value).get("default")
|
||||
for name in ("allowed_domains", "blocked_domains")
|
||||
if isinstance(value := properties.get(name), Mapping)
|
||||
}
|
||||
values: Final = {**defaults, **tool}
|
||||
allowed: Final = values.get("allowed_domains")
|
||||
blocked: Final = values.get("blocked_domains")
|
||||
allowed_filter: Final[WebSearchDomainFilters] = (
|
||||
{"allowed_domains": _DOMAINS.validate_python(allowed)} if allowed is not None else {}
|
||||
)
|
||||
blocked_filter: Final[WebSearchDomainFilters] = (
|
||||
{"blocked_domains": _DOMAINS.validate_python(blocked)} if blocked is not None else {}
|
||||
)
|
||||
return {**allowed_filter, **blocked_filter}
|
||||
|
||||
|
||||
def resolve_web_search_domain_filters(request: Mapping[str, object]) -> WebSearchDomainFilters:
|
||||
definitions: Final = _TOOLS.validate_python(request.get("tools") or [])
|
||||
configured: Final[WebSearchDomainFilters] = next(
|
||||
(get_web_search_domain_filters(tool) for tool in definitions if is_web_search_tool(tool)),
|
||||
WebSearchDomainFilters(),
|
||||
)
|
||||
arguments: Final = get_web_search_domain_filters(_TOOL_MAPPING.validate_python(request.get("input") or {}))
|
||||
return {**arguments, **configured}
|
||||
|
||||
|
||||
def is_web_search_tool_responses(tool: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
Check if a tool is a web search tool for the Responses API.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
@ -22,6 +23,18 @@ from litellm.types.integrations.websearch_interception import (
|
|||
)
|
||||
|
||||
|
||||
def _search_hostname(url: str) -> str:
|
||||
try:
|
||||
parsed: Final = urlsplit(url if "://" in url or url.startswith("//") else f"//{url}")
|
||||
return (parsed.hostname or "").lower().rstrip(".")
|
||||
except ValueError:
|
||||
return ""
|
||||
|
||||
|
||||
def _matches_search_domain(host: str, domains: Sequence[str]) -> bool:
|
||||
return any(host == domain or host.endswith("." + domain) for domain in domains if domain)
|
||||
|
||||
|
||||
class WebSearchTransformation:
|
||||
"""
|
||||
Transformation class for WebSearch tool interception.
|
||||
|
|
@ -527,6 +540,28 @@ class WebSearchTransformation:
|
|||
case _:
|
||||
assert_never(outcome)
|
||||
|
||||
@staticmethod
|
||||
def filter_search_response(
|
||||
result: SearchResponse,
|
||||
allowed_domains: Sequence[str] | None = None,
|
||||
blocked_domains: Sequence[str] | None = None,
|
||||
) -> SearchResponse:
|
||||
if allowed_domains is None and not blocked_domains:
|
||||
return result
|
||||
allowed: Final = tuple(_search_hostname(domain) for domain in allowed_domains or ())
|
||||
blocked: Final = tuple(_search_hostname(domain) for domain in blocked_domains or ())
|
||||
return result.model_copy(
|
||||
update={
|
||||
"results": [
|
||||
item
|
||||
for item in result.results
|
||||
if (host := _search_hostname(item.url))
|
||||
and (allowed_domains is None or _matches_search_domain(host, allowed))
|
||||
and not _matches_search_domain(host, blocked)
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def format_search_response(result: SearchResponse) -> str:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -47,6 +47,11 @@ class RichWebSearchInput(TypedDict, total=False):
|
|||
"""Two to five short keyword queries covering different angles."""
|
||||
|
||||
|
||||
class WebSearchDomainFilters(TypedDict, total=False):
|
||||
allowed_domains: ReadOnly[tuple[str, ...]]
|
||||
blocked_domains: ReadOnly[tuple[str, ...]]
|
||||
|
||||
|
||||
WebSearchToolResultErrorCode: TypeAlias = Literal[
|
||||
"invalid_tool_input",
|
||||
"unavailable",
|
||||
|
|
|
|||
234
tests/test_litellm/test_websearch_interception_domains_44188.py
Normal file
234
tests/test_litellm/test_websearch_interception_domains_44188.py
Normal file
|
|
@ -0,0 +1,234 @@
|
|||
from copy import deepcopy
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.websearch_interception.handler import WebSearchInterceptionLogger
|
||||
from litellm.integrations.websearch_interception.tools import get_litellm_web_search_tool
|
||||
from litellm.integrations.websearch_interception.transformation import WebSearchTransformation
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse, SearchResult
|
||||
from litellm.types.integrations.websearch_interception import SearchSucceeded
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def search_response() -> SearchResponse:
|
||||
return SearchResponse(
|
||||
results=[
|
||||
SearchResult(title=f"Result {index}", url=url, snippet=f"Snippet {index}")
|
||||
for index, url in enumerate(
|
||||
(
|
||||
"https://example.org/article",
|
||||
"https://docs.example.org/guide",
|
||||
"https://example.com/article",
|
||||
"https://example.net/article",
|
||||
"https://docs.example.net/guide",
|
||||
"https://notexample.org/article",
|
||||
"https://example.org.evil.test/article",
|
||||
)
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def search_mock(monkeypatch: pytest.MonkeyPatch, search_response: SearchResponse) -> AsyncMock:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
mock: Final = AsyncMock(return_value=search_response)
|
||||
monkeypatch.setattr(litellm, "asearch", mock)
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
SimpleNamespace(search_tools=[{"litellm_params": {"search_provider": "firecrawl"}}]),
|
||||
)
|
||||
return mock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("filters", "expected_indices"),
|
||||
[
|
||||
({"allowed_domains": ["example.org"]}, (0, 1)),
|
||||
({"blocked_domains": ["example.net"]}, (0, 1, 2, 5, 6)),
|
||||
({"allowed_domains": ["example.org"], "blocked_domains": ["docs.example.org"]}, (0,)),
|
||||
({"allowed_domains": ["example.org"], "blocked_domains": ["example.org"]}, ()),
|
||||
({"allowed_domains": []}, ()),
|
||||
({}, (0, 1, 2, 3, 4, 5, 6)),
|
||||
],
|
||||
)
|
||||
async def test_short_circuit_filters_citations_and_text(
|
||||
filters: dict[str, list[str]],
|
||||
expected_indices: tuple[int, ...],
|
||||
search_response: SearchResponse,
|
||||
search_mock: AsyncMock,
|
||||
) -> None:
|
||||
logger: Final = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
response: Final = await logger.try_short_circuit_search(
|
||||
model="interception-test-model",
|
||||
messages=[{"role": "user", "content": "Find domain documentation"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search", **filters}],
|
||||
custom_llm_provider="github_copilot",
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
expected: Final = [search_response.results[index] for index in expected_indices]
|
||||
content: Final = response["content"]
|
||||
assert [item["url"] for item in content[1]["content"]] == [item.url for item in expected]
|
||||
assert content[2]["text"] == (
|
||||
"\n\n".join(f"Title: {item.title}\nURL: {item.url}\nSnippet: {item.snippet}" for item in expected)
|
||||
if expected
|
||||
else "No search results found."
|
||||
)
|
||||
assert response["stop_reason"] == "end_turn"
|
||||
search_mock.assert_awaited_once_with(
|
||||
query="Find domain documentation",
|
||||
search_provider="firecrawl",
|
||||
**({"search_domain_filter": filters["allowed_domains"]} if "allowed_domains" in filters else {}),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ["pre_request", "deployment", "both_hooks", "chat_completion", "responses"])
|
||||
@pytest.mark.parametrize("source", ["definition", "arguments", "attempted_override"])
|
||||
async def test_agentic_paths_preserve_domain_constraints(
|
||||
path: str,
|
||||
source: str,
|
||||
search_response: SearchResponse,
|
||||
search_mock: AsyncMock,
|
||||
) -> None:
|
||||
logger: Final = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
filters: Final = {"allowed_domains": ["example.org"], "blocked_domains": ["docs.example.org"]}
|
||||
kwargs: Final = {
|
||||
"custom_llm_provider": "bedrock",
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"tools": [{"type": "web_search_20250305", "name": "web_search", **(filters if source != "arguments" else {})}],
|
||||
}
|
||||
if path in ("deployment", "both_hooks", "chat_completion"):
|
||||
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
|
||||
if path in ("pre_request", "both_hooks", "responses"):
|
||||
await logger.async_pre_request_hook(model="interception-test-model", messages=[], kwargs=kwargs)
|
||||
tool_input: Final = {
|
||||
"query": "Find domain documentation",
|
||||
**(filters if source == "arguments" else {}),
|
||||
**({"allowed_domains": ["example.com"], "blocked_domains": []} if source == "attempted_override" else {}),
|
||||
}
|
||||
tool_calls: Final = [{"id": "toolu_test", "name": "litellm_web_search", "input": tool_input}]
|
||||
optional_params: Final = {
|
||||
"max_tokens": 1024,
|
||||
**({"tools": kwargs["tools"]} if path != "pre_request" else {}),
|
||||
}
|
||||
expected_text: Final = "Title: Result 0\nURL: https://example.org/article\nSnippet: Snippet 0"
|
||||
|
||||
if path == "chat_completion":
|
||||
chat_patch: Final = await logger._build_chat_completion_request_patch(
|
||||
model="interception-test-model",
|
||||
messages=[],
|
||||
tool_calls=tool_calls,
|
||||
optional_params=optional_params,
|
||||
kwargs={},
|
||||
)
|
||||
assert chat_patch.messages[-1]["content"] == expected_text
|
||||
elif path == "responses":
|
||||
responses_patch: Final = await logger._build_responses_request_patch(
|
||||
model="interception-test-model",
|
||||
messages=[],
|
||||
tool_calls=tool_calls,
|
||||
optional_params=optional_params,
|
||||
kwargs={},
|
||||
)
|
||||
assert responses_patch.messages[-1]["output"] == expected_text
|
||||
else:
|
||||
patch, outcomes = await logger._build_anthropic_request_patch(
|
||||
model="interception-test-model",
|
||||
messages=[],
|
||||
tool_calls=tool_calls,
|
||||
thinking_blocks=[],
|
||||
anthropic_messages_optional_request_params=optional_params,
|
||||
logging_obj=None,
|
||||
kwargs={"tools": kwargs["tools"]},
|
||||
)
|
||||
assert patch.messages[-1]["content"][0]["content"] == expected_text
|
||||
assert isinstance(outcomes[0], SearchSucceeded)
|
||||
assert outcomes[0].response.results == [search_response.results[0]]
|
||||
search_mock.assert_awaited_once_with(
|
||||
query="Find domain documentation", search_provider="firecrawl", search_domain_filter=["example.org"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("url", "retained"),
|
||||
[
|
||||
("https://EXAMPLE.ORG:8443/article", True),
|
||||
("https://docs.EXAMPLE.ORG./article", True),
|
||||
("//docs.example.org/article", True),
|
||||
("https://example.org@evil.test/article", False),
|
||||
("https://evil.test/example.org", False),
|
||||
("https://evil.test/?url=https://example.org", False),
|
||||
("https://notexample.org/article", False),
|
||||
("https://example.org.evil.test/article", False),
|
||||
("https://[invalid/article", False),
|
||||
("/relative/article", False),
|
||||
],
|
||||
)
|
||||
def test_domain_matching_uses_normalized_hostname(url: str, retained: bool) -> None:
|
||||
item: Final = SearchResult(title="Title", url=url, snippet="Snippet")
|
||||
response: Final = SearchResponse(results=[item])
|
||||
filtered: Final = WebSearchTransformation.filter_search_response(
|
||||
response, allowed_domains=["https://EXAMPLE.ORG:443"]
|
||||
)
|
||||
assert filtered.results == ([item] if retained else [])
|
||||
assert response.results == [item]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_builder_preserves_filters_for_execution(search_mock: AsyncMock) -> None:
|
||||
tool: Final = get_litellm_web_search_tool(allowed_domains=["example.org"], blocked_domains=["docs.example.org"])
|
||||
logger: Final = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
response: Final = await logger.try_short_circuit_search(
|
||||
model="interception-test-model",
|
||||
messages=[{"role": "user", "content": "Find domain documentation"}],
|
||||
tools=[tool],
|
||||
custom_llm_provider="github_copilot",
|
||||
)
|
||||
assert response["content"] == [
|
||||
{"type": "text", "text": "Title: Result 0\nURL: https://example.org/article\nSnippet: Snippet 0"}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("hook", ["request", "deployment"])
|
||||
async def test_disabled_interception_leaves_native_domains_untouched(hook: str) -> None:
|
||||
logger: Final = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
tool: Final = {
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"allowed_domains": ["example.org"],
|
||||
"blocked_domains": ["docs.example.org"],
|
||||
}
|
||||
kwargs: Final = {
|
||||
"tools": [tool],
|
||||
"custom_llm_provider": "anthropic",
|
||||
"litellm_params": {"custom_llm_provider": "anthropic"},
|
||||
}
|
||||
original: Final = deepcopy(kwargs)
|
||||
result: Final = (
|
||||
await logger.async_pre_request_hook(model="interception-test-model", messages=[], kwargs=kwargs)
|
||||
if hook == "request"
|
||||
else await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
|
||||
)
|
||||
assert result is None
|
||||
assert kwargs == original
|
||||
|
||||
|
||||
@pytest.mark.parametrize("blocked_domains", [None, []])
|
||||
def test_unconstrained_search_preserves_response(
|
||||
search_response: SearchResponse, blocked_domains: list[str] | None
|
||||
) -> None:
|
||||
assert (
|
||||
WebSearchTransformation.filter_search_response(search_response, blocked_domains=blocked_domains)
|
||||
is search_response
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue