mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(websearch): preserve user-defined web_search tools
This commit is contained in:
parent
f2b399d451
commit
a622324182
2 changed files with 41 additions and 5 deletions
|
|
@ -161,7 +161,9 @@ 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 "web_search"
|
||||
- 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"
|
||||
|
||||
Args:
|
||||
tool: Tool dictionary to check
|
||||
|
|
@ -185,8 +187,15 @@ 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 in (LITELLM_WEB_SEARCH_TOOL_NAME, "web_search"):
|
||||
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 set(function_def) == {"name"}:
|
||||
return True
|
||||
|
||||
# Check for LiteLLM standard tool (Anthropic format)
|
||||
|
|
@ -269,8 +278,14 @@ 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 in (LITELLM_WEB_SEARCH_TOOL_NAME, "web_search"):
|
||||
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 set(function_def) == {"name"}:
|
||||
return True
|
||||
|
||||
# Check for LiteLLM standard tool (Anthropic format)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
"""Unit tests for web search tool shape detection."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.websearch_interception.tools import (
|
||||
|
|
@ -15,7 +17,7 @@ from litellm.integrations.websearch_interception.tools import (
|
|||
{"type": "function", "function": {"name": "web_search"}},
|
||||
],
|
||||
)
|
||||
def test_openai_function_web_search_shapes_are_detected(tool):
|
||||
def test_openai_function_web_search_shapes_are_detected(tool: dict[str, Any]):
|
||||
"""Recognize both LiteLLM and conventional web_search function names."""
|
||||
assert is_web_search_tool(tool) is True
|
||||
assert is_web_search_tool_chat_completion(tool) is True
|
||||
|
|
@ -28,7 +30,26 @@ def test_openai_function_web_search_shapes_are_detected(tool):
|
|||
{"type": "function", "function": {"name": "search"}},
|
||||
],
|
||||
)
|
||||
def test_unrelated_openai_function_tools_are_not_detected(tool):
|
||||
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