mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge f92341acf6 into f4308bc124
This commit is contained in:
commit
a9f8ee38d6
4 changed files with 121 additions and 10 deletions
|
|
@ -258,6 +258,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:
|
||||
|
|
@ -270,6 +271,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
|
||||
|
|
@ -279,8 +284,17 @@ 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:
|
||||
"""
|
||||
|
|
@ -350,7 +364,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
|
||||
|
|
@ -453,7 +467,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
|
||||
|
|
@ -471,7 +485,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)
|
||||
|
|
@ -543,6 +557,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
|
||||
|
|
@ -561,6 +578,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
|
||||
|
|
@ -625,7 +643,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
|
||||
|
||||
|
|
@ -646,7 +664,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(
|
||||
|
|
@ -721,7 +739,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, {}
|
||||
|
|
@ -817,7 +835,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, {}
|
||||
|
|
|
|||
|
|
@ -160,16 +160,17 @@ 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).
|
||||
|
||||
This is a stricter version that ONLY checks for the exact LiteLLM web search tool name.
|
||||
This is a stricter version that checks the recognized web search tool names.
|
||||
Use this for Chat Completions API to avoid false positives with user-defined tools.
|
||||
|
||||
Detects ONLY:
|
||||
- LiteLLM standard: name == "litellm_web_search" (Anthropic format)
|
||||
- OpenAI format: type == "function" with function.name == "litellm_web_search"
|
||||
- Optionally (recognize_conventional_name=True): name-only function ``web_search``
|
||||
|
||||
Args:
|
||||
tool: Tool dictionary to check
|
||||
|
|
@ -193,9 +194,20 @@ def is_web_search_tool_chat_completion(tool: dict[str, Any]) -> bool:
|
|||
# Check for OpenAI format: {"type": "function", "function": {"name": "litellm_web_search"}}
|
||||
if tool_type == "function" and "function" in tool:
|
||||
function_def: Final = tool.get("function", {})
|
||||
if not isinstance(function_def, dict):
|
||||
return False
|
||||
function_name: Final = function_def.get("name", "")
|
||||
if function_name == LITELLM_WEB_SEARCH_TOOL_NAME:
|
||||
return True
|
||||
# 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)
|
||||
if tool_name == LITELLM_WEB_SEARCH_TOOL_NAME:
|
||||
|
|
@ -225,7 +237,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).
|
||||
|
||||
|
|
@ -277,9 +289,20 @@ def is_web_search_tool(tool: dict[str, Any]) -> bool:
|
|||
# Check for OpenAI format: {"type": "function", "function": {"name": "..."}}
|
||||
if tool_type == "function" and "function" in tool:
|
||||
function_def: Final = tool.get("function", {})
|
||||
if not isinstance(function_def, dict):
|
||||
return False
|
||||
function_name: Final = function_def.get("name", "")
|
||||
if function_name == LITELLM_WEB_SEARCH_TOOL_NAME:
|
||||
return True
|
||||
# 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)
|
||||
if tool_name == LITELLM_WEB_SEARCH_TOOL_NAME:
|
||||
|
|
|
|||
|
|
@ -92,3 +92,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: ReadOnly[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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,64 @@
|
|||
"""Unit tests for web search tool shape detection."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.websearch_interception.tools import (
|
||||
is_web_search_tool,
|
||||
is_web_search_tool_chat_completion,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"tool",
|
||||
[
|
||||
{"type": "function", "function": {"name": "litellm_web_search"}},
|
||||
],
|
||||
)
|
||||
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",
|
||||
[
|
||||
{"type": "function", "function": {"name": "web_search_helper"}},
|
||||
{"type": "function", "function": {"name": "search"}},
|
||||
{"type": "function", "function": None},
|
||||
],
|
||||
)
|
||||
def test_unrelated_openai_function_tools_are_not_detected(tool: dict[str, Any]):
|
||||
"""Do not classify similarly named user tools as web search."""
|
||||
assert is_web_search_tool(tool) is False
|
||||
assert is_web_search_tool_chat_completion(tool) is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"tool",
|
||||
[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "web_search",
|
||||
"description": "Search a private index",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
def test_user_defined_web_search_functions_are_not_detected(tool: dict[str, Any]):
|
||||
"""Preserve user-defined schemas even when their name is web_search."""
|
||||
assert is_web_search_tool(tool) is False
|
||||
assert is_web_search_tool_chat_completion(tool) is False
|
||||
Loading…
Add table
Reference in a new issue