diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 1db94e82066..61492202863 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -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 diff --git a/litellm/integrations/websearch_interception/tools.py b/litellm/integrations/websearch_interception/tools.py index 2e1ae07eb68..a65a267871d 100644 --- a/litellm/integrations/websearch_interception/tools.py +++ b/litellm/integrations/websearch_interception/tools.py @@ -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. diff --git a/litellm/integrations/websearch_interception/transformation.py b/litellm/integrations/websearch_interception/transformation.py index 47af73570fc..dbd20510973 100644 --- a/litellm/integrations/websearch_interception/transformation.py +++ b/litellm/integrations/websearch_interception/transformation.py @@ -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: """ diff --git a/litellm/types/integrations/websearch_interception.py b/litellm/types/integrations/websearch_interception.py index 22233404fb3..3a264a543fe 100644 --- a/litellm/types/integrations/websearch_interception.py +++ b/litellm/types/integrations/websearch_interception.py @@ -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", diff --git a/tests/test_litellm/test_websearch_interception_domains_44188.py b/tests/test_litellm/test_websearch_interception_domains_44188.py new file mode 100644 index 00000000000..79c023e589e --- /dev/null +++ b/tests/test_litellm/test_websearch_interception_domains_44188.py @@ -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 + )