fix(websearch): gate conventional web_search name behind opt-in

Name-only OpenAI function tools named web_search are valid user tools.
Recognize that conventional envelope only when
websearch_interception_params.recognize_conventional_web_search_name is
true, so enabling interception no longer hijacks those handlers by default.

Addresses greptile P1 on #40951.
This commit is contained in:
leilei3167 2026-09-16 15:10:18 +00:00
parent fb9551064f
commit 7cb01e0795
4 changed files with 63 additions and 21 deletions

View file

@ -213,6 +213,7 @@ class WebSearchInterceptionLogger(CustomLogger):
enabled_providers: list[LlmProviders | str] | None = None,
search_tool_name: str | None = None,
max_agentic_loops: int | None = None,
recognize_conventional_web_search_name: bool = False,
):
"""
Args:
@ -225,6 +226,10 @@ class WebSearchInterceptionLogger(CustomLogger):
max_agentic_loops: How many follow-up model calls one intercepted request
may chain before the loop is refused and the turn ends.
If None, LiteLLM's default of 3 applies.
recognize_conventional_web_search_name: When True, treat name-only OpenAI
function tools named ``web_search`` as server search.
Default False so a user-defined name-only ``web_search``
still reaches the client handler.
"""
super().__init__()
# Convert enum values to strings for comparison
@ -234,8 +239,19 @@ class WebSearchInterceptionLogger(CustomLogger):
self.enabled_providers = [p.value if isinstance(p, LlmProviders) else p for p in enabled_providers]
self.search_tool_name = search_tool_name
self.max_agentic_loops = self._validated_max_agentic_loops(max_agentic_loops)
self.recognize_conventional_web_search_name = bool(recognize_conventional_web_search_name)
self._request_has_websearch = False # Track if current request has web search
def _is_web_search_tool(self, tool: dict[str, Any]) -> bool:
return is_web_search_tool(
tool, recognize_conventional_name=self.recognize_conventional_web_search_name
)
def _is_web_search_tool_chat_completion(self, tool: dict[str, Any]) -> bool:
return is_web_search_tool_chat_completion(
tool, recognize_conventional_name=self.recognize_conventional_web_search_name
)
@staticmethod
def _validated_max_agentic_loops(max_agentic_loops: object) -> int | None:
"""
@ -305,7 +321,7 @@ class WebSearchInterceptionLogger(CustomLogger):
pass # unknown provider enum → safe to short-circuit
# All tools must be web search tools
if not all(is_web_search_tool(t) for t in tools):
if not all(self._is_web_search_tool(t) for t in tools):
return None
# Extract search query from the last user message
@ -419,7 +435,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if call_type in (CallTypes.responses, CallTypes.aresponses):
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
has_websearch: Final = any(is_web_search_tool(t) for t in tools)
has_websearch: Final = any(self._is_web_search_tool(t) for t in tools)
if not has_websearch:
return None
@ -437,7 +453,7 @@ class WebSearchInterceptionLogger(CustomLogger):
# Convert native/custom web_search tools to LiteLLM standard
converted_tools: Final = []
for tool in tools:
if is_web_search_tool(tool):
if self._is_web_search_tool(tool):
# Convert to LiteLLM standard web search tool
converted_tool = get_litellm_web_search_tool_openai()
converted_tools.append(converted_tool)
@ -509,6 +525,9 @@ class WebSearchInterceptionLogger(CustomLogger):
enabled_providers_str: Final = config.get("enabled_providers", None)
search_tool_name: Final = config.get("search_tool_name", None)
max_agentic_loops: Final = config.get("max_agentic_loops", None)
recognize_conventional_web_search_name: Final = bool(
config.get("recognize_conventional_web_search_name", False)
)
# Convert string provider names to LlmProviders enum values
enabled_providers: list[LlmProviders | str] | None = None
@ -527,6 +546,7 @@ class WebSearchInterceptionLogger(CustomLogger):
enabled_providers=enabled_providers,
search_tool_name=search_tool_name,
max_agentic_loops=max_agentic_loops,
recognize_conventional_web_search_name=recognize_conventional_web_search_name,
)
@staticmethod
@ -591,7 +611,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return None
# Check if any tool is a web search tool
has_websearch: Final = any(is_web_search_tool(t) for t in tools)
has_websearch: Final = any(self._is_web_search_tool(t) for t in tools)
if not has_websearch:
return None
@ -612,7 +632,7 @@ class WebSearchInterceptionLogger(CustomLogger):
# Convert native web search tools to LiteLLM standard
converted_tools: Final[list[dict[str, object]]] = []
for tool in tools:
if is_web_search_tool(tool):
if self._is_web_search_tool(tool):
standard_tool = get_litellm_web_search_tool()
converted_tools.append(standard_tool)
verbose_logger.debug(
@ -687,7 +707,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return False, {}
# Check if tools include any web search tool (LiteLLM standard or native)
has_websearch_tool: Final = any(is_web_search_tool(t) for t in (tools or []))
has_websearch_tool: Final = any(self._is_web_search_tool(t) for t in (tools or []))
if not has_websearch_tool:
verbose_logger.debug("WebSearchInterception: No web search tool in request")
return False, {}
@ -783,7 +803,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return False, {}
# Check if tools include any web search tool (strict check for chat completions)
has_websearch_tool: Final = any(is_web_search_tool_chat_completion(t) for t in (tools or []))
has_websearch_tool: Final = any(self._is_web_search_tool_chat_completion(t) for t in (tools or []))
if not has_websearch_tool:
verbose_logger.debug("WebSearchInterception: No litellm_web_search tool in request")
return False, {}

View file

@ -152,7 +152,7 @@ def is_web_search_tool_responses(tool: Mapping[str, object]) -> bool:
return tool_type == "web_search" or tool_type.startswith("web_search_")
def is_web_search_tool_chat_completion(tool: dict[str, Any]) -> bool:
def is_web_search_tool_chat_completion(tool: dict[str, Any], *, recognize_conventional_name: bool = False) -> bool:
"""
Check if a tool is a web search tool for Chat Completions API (strict check).
@ -162,8 +162,7 @@ def is_web_search_tool_chat_completion(tool: dict[str, Any]) -> bool:
Detects ONLY:
- LiteLLM standard: name == "litellm_web_search" (Anthropic format)
- OpenAI format: type == "function" with function.name == "litellm_web_search"
or a bare function definition whose name is "web_search"
or a bare function definition whose name is "web_search"
- Optionally (recognize_conventional_name=True): name-only function ``web_search``
Args:
tool: Tool dictionary to check
@ -192,10 +191,14 @@ def is_web_search_tool_chat_completion(tool: dict[str, Any]) -> bool:
function_name: Final = function_def.get("name", "")
if function_name == LITELLM_WEB_SEARCH_TOOL_NAME:
return True
# The conventional name is only recognized for the name-only shape.
# A user-defined function may also be named ``web_search``; preserving
# any additional definition fields lets that tool pass through intact.
if function_name == "web_search" and len(function_def) == 1 and "name" in function_def:
# Name-only ``web_search`` is opt-in: without the flag, a valid
# user-defined function with only a name still reaches the client.
if (
recognize_conventional_name
and function_name == "web_search"
and len(function_def) == 1
and "name" in function_def
):
return True
# Check for LiteLLM standard tool (Anthropic format)
@ -226,7 +229,7 @@ def is_anthropic_native_web_search_tool(tool: Mapping[str, object]) -> bool:
return tool_type.startswith("web_search_") and tool_type != "function"
def is_web_search_tool(tool: dict[str, Any]) -> bool:
def is_web_search_tool(tool: dict[str, Any], *, recognize_conventional_name: bool = False) -> bool:
"""
Check if a tool is a web search tool (native or LiteLLM standard).
@ -283,9 +286,14 @@ def is_web_search_tool(tool: dict[str, Any]) -> bool:
function_name: Final = function_def.get("name", "")
if function_name == LITELLM_WEB_SEARCH_TOOL_NAME:
return True
# Do not hijack a user-defined function that happens to use the
# conventional name and carries its own schema or description.
if function_name == "web_search" and len(function_def) == 1 and "name" in function_def:
# Name-only ``web_search`` is opt-in: without the flag, a valid
# user-defined function with only a name still reaches the client.
if (
recognize_conventional_name
and function_name == "web_search"
and len(function_def) == 1
and "name" in function_def
):
return True
# Check for LiteLLM standard tool (Anthropic format)

View file

@ -47,3 +47,9 @@ class WebSearchInterceptionConfig(TypedDict, total=False):
max_agentic_loops: ReadOnly[int | None]
"""How many follow-up model calls one intercepted request may chain. If None, LiteLLM's default of 3 applies."""
recognize_conventional_web_search_name: bool
"""When True, treat name-only OpenAI function tools named ``web_search`` as
server search tools. Default is off so a user-defined name-only ``web_search``
function still reaches the client handler.
"""

View file

@ -14,15 +14,23 @@ from litellm.integrations.websearch_interception.tools import (
"tool",
[
{"type": "function", "function": {"name": "litellm_web_search"}},
{"type": "function", "function": {"name": "web_search"}},
],
)
def test_openai_function_web_search_shapes_are_detected(tool: dict[str, Any]):
"""Recognize both LiteLLM and conventional web_search function names."""
def test_openai_function_litellm_web_search_is_detected(tool: dict[str, Any]):
"""Recognize the LiteLLM standard web_search function name."""
assert is_web_search_tool(tool) is True
assert is_web_search_tool_chat_completion(tool) is True
def test_conventional_web_search_requires_opt_in():
"""Name-only web_search stays a user tool unless explicitly opted in."""
tool = {"type": "function", "function": {"name": "web_search"}}
assert is_web_search_tool(tool) is False
assert is_web_search_tool_chat_completion(tool) is False
assert is_web_search_tool(tool, recognize_conventional_name=True) is True
assert is_web_search_tool_chat_completion(tool, recognize_conventional_name=True) is True
@pytest.mark.parametrize(
"tool",
[