diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 6ebf485d717..5daaca1343f 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -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, {} diff --git a/litellm/integrations/websearch_interception/tools.py b/litellm/integrations/websearch_interception/tools.py index 2e1ae07eb68..c56f650b2e1 100644 --- a/litellm/integrations/websearch_interception/tools.py +++ b/litellm/integrations/websearch_interception/tools.py @@ -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: diff --git a/litellm/types/integrations/websearch_interception.py b/litellm/types/integrations/websearch_interception.py index 22233404fb3..97f72ff1590 100644 --- a/litellm/types/integrations/websearch_interception.py +++ b/litellm/types/integrations/websearch_interception.py @@ -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. + """ diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_tools.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_tools.py new file mode 100644 index 00000000000..b9a74e99be9 --- /dev/null +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_tools.py @@ -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