mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(websearch): add Responses API surface to websearch interception
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
39e0efa11d
commit
aa7b480f4c
6 changed files with 657 additions and 6 deletions
|
|
@ -19,9 +19,11 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.integrations.websearch_interception.tools import (
|
||||
get_litellm_web_search_tool,
|
||||
get_litellm_web_search_tool_openai,
|
||||
get_litellm_web_search_tool_responses,
|
||||
is_anthropic_native_web_search_tool,
|
||||
is_web_search_tool,
|
||||
is_web_search_tool_chat_completion,
|
||||
is_web_search_tool_responses,
|
||||
)
|
||||
from litellm.integrations.websearch_interception.transformation import (
|
||||
WebSearchTransformation,
|
||||
|
|
@ -32,11 +34,12 @@ from litellm.types.integrations.websearch_interception import (
|
|||
)
|
||||
from litellm.types.integrations.custom_logger import (
|
||||
CHAT_COMPLETION_AGENTIC_SURFACE,
|
||||
RESPONSES_AGENTIC_SURFACE,
|
||||
AgenticLoopPlan,
|
||||
AgenticLoopRequestPatch,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.types.utils import CallTypes, LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
# Key used to flag, on per-request kwargs, that the originating client sent
|
||||
|
|
@ -251,6 +254,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if not tools:
|
||||
return None
|
||||
|
||||
if call_type in (CallTypes.responses, CallTypes.aresponses):
|
||||
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
|
||||
|
||||
# Check if any tool is a web search tool (native or already LiteLLM standard)
|
||||
has_websearch = any(is_web_search_tool(t) for t in tools)
|
||||
|
||||
|
|
@ -291,6 +297,26 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
|
||||
return kwargs
|
||||
|
||||
def _convert_responses_tools(self, kwargs: dict[str, Any], tools: list[dict[str, Any]]) -> dict | None:
|
||||
"""Convert Responses API web search tools to the LiteLLM standard function tool."""
|
||||
if not any(is_web_search_tool_responses(tool) for tool in tools):
|
||||
return None
|
||||
|
||||
verbose_logger.debug("WebSearchInterception: Converting Responses web_search tools to LiteLLM standard")
|
||||
|
||||
converted_tools = [
|
||||
get_litellm_web_search_tool_responses() if is_web_search_tool_responses(tool) else tool for tool in tools
|
||||
]
|
||||
|
||||
converted_kwargs = {**kwargs, "tools": converted_tools}
|
||||
|
||||
if kwargs.get("stream"):
|
||||
verbose_logger.debug("WebSearchInterception: deployment hook converting stream=True to stream=False")
|
||||
converted_kwargs["stream"] = False
|
||||
converted_kwargs["_websearch_interception_converted_stream"] = True
|
||||
|
||||
return converted_kwargs
|
||||
|
||||
@classmethod
|
||||
def from_config_yaml(cls, config: WebSearchInterceptionConfig) -> "WebSearchInterceptionLogger":
|
||||
"""
|
||||
|
|
@ -461,6 +487,17 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
if kwargs.get("_agentic_loop_api_surface") == RESPONSES_AGENTIC_SURFACE:
|
||||
return await self.async_should_run_responses_agentic_loop(
|
||||
response=response,
|
||||
model=model,
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
stream=stream,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"WebSearchInterception: Hook called! provider={custom_llm_provider}, stream={stream}")
|
||||
verbose_logger.debug(f"WebSearchInterception: Response type: {type(response)}")
|
||||
|
||||
|
|
@ -597,6 +634,54 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
}
|
||||
return True, tools_dict
|
||||
|
||||
async def async_should_run_responses_agentic_loop(
|
||||
self,
|
||||
response: Any,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
tools: list[dict] | None,
|
||||
stream: bool,
|
||||
custom_llm_provider: str,
|
||||
kwargs: dict,
|
||||
) -> tuple[bool, dict]:
|
||||
"""Check if WebSearch interception is needed for the Responses API."""
|
||||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Responses hook called! provider={custom_llm_provider}, stream={stream}"
|
||||
)
|
||||
|
||||
if self.enabled_providers is not None and custom_llm_provider not in self.enabled_providers:
|
||||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Skipping provider {custom_llm_provider} (not in enabled list: {self.enabled_providers})"
|
||||
)
|
||||
return False, {}
|
||||
|
||||
has_websearch_tool = any(is_web_search_tool_responses(t) for t in (tools or []))
|
||||
if not has_websearch_tool:
|
||||
verbose_logger.debug("WebSearchInterception: No litellm_web_search tool in responses request")
|
||||
return False, {}
|
||||
|
||||
should_intercept, tool_calls = WebSearchTransformation.transform_request(
|
||||
response=response,
|
||||
stream=stream,
|
||||
response_format="responses",
|
||||
)
|
||||
|
||||
if not should_intercept:
|
||||
verbose_logger.debug("WebSearchInterception: No WebSearch function_call detected in responses output")
|
||||
return False, {}
|
||||
|
||||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Detected {len(tool_calls)} WebSearch function_call(s), executing agentic loop"
|
||||
)
|
||||
|
||||
tools_dict = {
|
||||
"tool_calls": tool_calls,
|
||||
"tool_type": "websearch",
|
||||
"provider": custom_llm_provider,
|
||||
"response_format": "responses",
|
||||
}
|
||||
return True, tools_dict
|
||||
|
||||
async def async_run_agentic_loop(
|
||||
self,
|
||||
tools: Dict,
|
||||
|
|
@ -655,6 +740,18 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
if kwargs.get("_agentic_loop_api_surface") == RESPONSES_AGENTIC_SURFACE:
|
||||
return await self.async_build_responses_agentic_loop_plan(
|
||||
tools=tools,
|
||||
model=model,
|
||||
messages=messages,
|
||||
response=response,
|
||||
optional_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
stream=stream,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
tool_calls = tools["tool_calls"]
|
||||
thinking_blocks = tools.get("thinking_blocks", [])
|
||||
request_patch, structured_results = await self._build_anthropic_request_patch(
|
||||
|
|
@ -809,6 +906,126 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
metadata={"tool_type": "websearch", "response_format": response_format},
|
||||
)
|
||||
|
||||
async def async_build_responses_agentic_loop_plan(
|
||||
self,
|
||||
tools: dict,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
response: Any,
|
||||
optional_params: dict,
|
||||
logging_obj: Any,
|
||||
stream: bool,
|
||||
kwargs: dict,
|
||||
) -> AgenticLoopPlan:
|
||||
tool_calls = tools["tool_calls"]
|
||||
request_patch = await self._build_responses_request_patch(
|
||||
model=model,
|
||||
messages=messages,
|
||||
tool_calls=tool_calls,
|
||||
optional_params=optional_params,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
return AgenticLoopPlan(
|
||||
run_agentic_loop=True,
|
||||
request_patch=request_patch,
|
||||
metadata={"tool_type": "websearch", "response_format": "responses"},
|
||||
)
|
||||
|
||||
async def _build_responses_request_patch(
|
||||
self,
|
||||
model: str,
|
||||
messages: Union[str, list[dict]],
|
||||
tool_calls: list[dict],
|
||||
optional_params: dict,
|
||||
kwargs: dict,
|
||||
) -> AgenticLoopRequestPatch:
|
||||
"""Execute litellm.asearch() and build a Responses API rerun patch."""
|
||||
search_tasks = [
|
||||
(
|
||||
self._execute_search(tool_call["input"]["query"], kwargs=kwargs)
|
||||
if isinstance(tool_call.get("input"), dict) and tool_call["input"].get("query")
|
||||
else self._create_empty_search_result()
|
||||
)
|
||||
for tool_call in tool_calls
|
||||
]
|
||||
|
||||
verbose_logger.debug(f"WebSearchInterception: Executing {len(search_tasks)} responses search(es) in parallel")
|
||||
search_results = await asyncio.gather(*search_tasks, return_exceptions=True)
|
||||
|
||||
search_texts = [self._extract_search_text(result) for result in search_results]
|
||||
|
||||
followup_items = [
|
||||
item
|
||||
for tool_call, search_text in zip(tool_calls, search_texts)
|
||||
for item in (
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": tool_call.get("call_id"),
|
||||
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"arguments": tool_call.get("arguments", ""),
|
||||
},
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": tool_call.get("call_id"),
|
||||
"output": search_text,
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
input_list = self._normalize_responses_input(messages) + followup_items
|
||||
|
||||
tools_param = optional_params.get("tools")
|
||||
optional_params_clean = {
|
||||
k: v
|
||||
for k, v in optional_params.items()
|
||||
if k not in {"tools", "tool_choice", "stream", "model_alias_map", "stream_response", "custom_prompt_dict"}
|
||||
}
|
||||
|
||||
kwargs_for_followup = {
|
||||
k: v
|
||||
for k, v in kwargs.items()
|
||||
if not k.startswith("_websearch_interception")
|
||||
and k not in {"litellm_logging_obj", "acompletion", "custom_llm_provider", "model_alias_map"}
|
||||
}
|
||||
|
||||
full_model_name = model
|
||||
if "/" not in model and isinstance(kwargs.get("custom_llm_provider"), str):
|
||||
full_model_name = f"{kwargs['custom_llm_provider']}/{model}"
|
||||
|
||||
verbose_logger.debug(
|
||||
"WebSearchInterception: Built responses request patch model=%s input_items=%d searches=%d",
|
||||
full_model_name,
|
||||
len(input_list),
|
||||
len(search_texts),
|
||||
)
|
||||
|
||||
return AgenticLoopRequestPatch(
|
||||
model=full_model_name,
|
||||
messages=input_list,
|
||||
tools=tools_param if isinstance(tools_param, list) else None,
|
||||
optional_params=optional_params_clean,
|
||||
kwargs=kwargs_for_followup,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_responses_input(messages: Union[str, list[dict]]) -> list[dict]:
|
||||
if isinstance(messages, str):
|
||||
return [{"role": "user", "content": messages}]
|
||||
if isinstance(messages, list):
|
||||
return list(messages)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _extract_search_text(result: Any) -> str:
|
||||
if isinstance(result, Exception):
|
||||
verbose_logger.error(f"WebSearchInterception: Responses search failed with error: {str(result)}")
|
||||
return f"Search failed: {str(result)}"
|
||||
if isinstance(result, tuple) and len(result) == 2:
|
||||
text_value, _ = result
|
||||
return text_value if isinstance(text_value, str) else str(text_value)
|
||||
verbose_logger.debug(f"WebSearchInterception: Unexpected search result type {type(result)}")
|
||||
return str(result)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_max_tokens(
|
||||
optional_params: Dict,
|
||||
|
|
|
|||
|
|
@ -82,6 +82,75 @@ def get_litellm_web_search_tool_openai() -> Dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
def get_litellm_web_search_tool_responses() -> dict[str, Any]:
|
||||
"""
|
||||
Get the standard LiteLLM web search tool definition in Responses API format.
|
||||
|
||||
Used by async_pre_call_deployment_hook on the Responses API path, where a
|
||||
function tool is a flat object (``type: "function"`` with a top-level
|
||||
``name`` and ``parameters``) rather than the nested ``function`` wrapper
|
||||
used by Chat Completions.
|
||||
|
||||
Returns:
|
||||
Dict containing the Responses-style function tool definition.
|
||||
"""
|
||||
return {
|
||||
"type": "function",
|
||||
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"description": (
|
||||
"Search the web for information. Use this when you need current "
|
||||
"information or answers to questions that require up-to-date data."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "The search query to execute",
|
||||
}
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def is_web_search_tool_responses(tool: dict[str, Any]) -> bool:
|
||||
"""
|
||||
Check if a tool is a web search tool for the Responses API.
|
||||
|
||||
Detects:
|
||||
- OpenAI native Responses web search tools, whose ``type`` is one of
|
||||
``web_search``, ``web_search_2025_08_26``, ``web_search_preview``,
|
||||
``web_search_preview_2025_03_11`` (matched by the ``web_search`` prefix)
|
||||
- The LiteLLM standard function tool in Responses shape:
|
||||
``{"type": "function", "name": "litellm_web_search"}``
|
||||
|
||||
Args:
|
||||
tool: Tool dictionary to check
|
||||
|
||||
Returns:
|
||||
True if tool is a Responses-API web search tool
|
||||
|
||||
Example:
|
||||
>>> is_web_search_tool_responses({"type": "web_search"})
|
||||
True
|
||||
>>> is_web_search_tool_responses({"type": "web_search_preview"})
|
||||
True
|
||||
>>> is_web_search_tool_responses({"type": "function", "name": "litellm_web_search"})
|
||||
True
|
||||
>>> is_web_search_tool_responses({"type": "function", "name": "get_weather"})
|
||||
False
|
||||
"""
|
||||
tool_type = tool.get("type", "")
|
||||
if not isinstance(tool_type, str):
|
||||
return False
|
||||
|
||||
if tool_type == "function":
|
||||
return tool.get("name") == LITELLM_WEB_SEARCH_TOOL_NAME
|
||||
|
||||
return tool_type == "web_search" or tool_type.startswith("web_search_")
|
||||
|
||||
|
||||
def is_web_search_tool_chat_completion(tool: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
Check if a tool is a web search tool for Chat Completions API (strict check).
|
||||
|
|
|
|||
|
|
@ -59,9 +59,76 @@ class WebSearchTransformation:
|
|||
# Parse non-streaming response based on format
|
||||
if response_format == "openai":
|
||||
return WebSearchTransformation._detect_from_openai_response(response)
|
||||
elif response_format == "responses":
|
||||
return WebSearchTransformation._detect_from_responses_response(response)
|
||||
else:
|
||||
return WebSearchTransformation._detect_from_non_streaming_response(response)
|
||||
|
||||
@staticmethod
|
||||
def _detect_from_responses_response(
|
||||
response: Any,
|
||||
) -> tuple[bool, list[dict]]:
|
||||
"""Parse a Responses API response for ``litellm_web_search`` function calls.
|
||||
|
||||
After pre-request conversion the native web search tool is replaced by a
|
||||
``litellm_web_search`` function tool, so the model emits ``function_call``
|
||||
items in ``response.output`` instead of a native ``web_search_call``.
|
||||
"""
|
||||
if isinstance(response, dict):
|
||||
output = response.get("output", [])
|
||||
else:
|
||||
output = getattr(response, "output", None) or []
|
||||
|
||||
if not isinstance(output, list):
|
||||
return False, []
|
||||
|
||||
tool_calls: list[dict] = []
|
||||
for item in output:
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
item_name = item.get("name")
|
||||
call_id = item.get("call_id")
|
||||
arguments = item.get("arguments", "")
|
||||
else:
|
||||
item_type = getattr(item, "type", None)
|
||||
item_name = getattr(item, "name", None)
|
||||
call_id = getattr(item, "call_id", None)
|
||||
arguments = getattr(item, "arguments", "")
|
||||
|
||||
if item_type != "function_call" or item_name not in (
|
||||
LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"web_search",
|
||||
):
|
||||
continue
|
||||
|
||||
if isinstance(arguments, str):
|
||||
try:
|
||||
parsed_input = json.loads(arguments) if arguments else {}
|
||||
except json.JSONDecodeError:
|
||||
verbose_logger.warning(
|
||||
f"WebSearchInterception: Failed to parse function_call arguments: {arguments}"
|
||||
)
|
||||
parsed_input = {}
|
||||
elif isinstance(arguments, dict):
|
||||
parsed_input = arguments
|
||||
else:
|
||||
parsed_input = {}
|
||||
|
||||
arguments_str = arguments if isinstance(arguments, str) else json.dumps(parsed_input)
|
||||
tool_calls.append(
|
||||
{
|
||||
"id": call_id,
|
||||
"call_id": call_id,
|
||||
"type": "function_call",
|
||||
"name": item_name,
|
||||
"arguments": arguments_str,
|
||||
"input": parsed_input,
|
||||
}
|
||||
)
|
||||
verbose_logger.debug(f"WebSearchInterception: Found {item_name} function_call with call_id={call_id}")
|
||||
|
||||
return len(tool_calls) > 0, tool_calls
|
||||
|
||||
@staticmethod
|
||||
def _detect_from_non_streaming_response(
|
||||
response: Any,
|
||||
|
|
|
|||
|
|
@ -2657,9 +2657,10 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
result = final_response if final_response is not None else initial_response
|
||||
if litellm_params.get("_code_interpreter_interception_converted_stream") and not litellm_params.get(
|
||||
"_agentic_loop_depth"
|
||||
):
|
||||
interception_converted_stream = litellm_params.get(
|
||||
"_code_interpreter_interception_converted_stream"
|
||||
) or litellm_params.get("_websearch_interception_converted_stream")
|
||||
if interception_converted_stream and not litellm_params.get("_agentic_loop_depth"):
|
||||
return self._wrap_responses_response_as_fake_stream(
|
||||
result=result,
|
||||
model=model,
|
||||
|
|
@ -5224,6 +5225,8 @@ class BaseLLMHTTPHandler:
|
|||
tools = anthropic_messages_optional_request_params.get("tools", [])
|
||||
depth, max_loops, fingerprints = self._get_agentic_loop_settings(kwargs=kwargs)
|
||||
|
||||
hook_kwargs = {**kwargs, "_agentic_loop_api_surface": api_surface}
|
||||
|
||||
for callback in callbacks:
|
||||
if not isinstance(callback, CustomLogger):
|
||||
continue
|
||||
|
|
@ -5244,7 +5247,7 @@ class BaseLLMHTTPHandler:
|
|||
tools=tools,
|
||||
stream=stream,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
kwargs=hook_kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
|
||||
|
|
@ -5270,7 +5273,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
try:
|
||||
kwargs_with_provider = kwargs.copy() if kwargs else {}
|
||||
kwargs_with_provider = hook_kwargs.copy()
|
||||
kwargs_with_provider["custom_llm_provider"] = custom_llm_provider
|
||||
build_plan_overridden = (
|
||||
callback.__class__.async_build_agentic_loop_plan is not CustomLogger.async_build_agentic_loop_plan
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from typing import Any, Dict, List, Optional
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
CHAT_COMPLETION_AGENTIC_SURFACE = "chat_completions"
|
||||
RESPONSES_AGENTIC_SURFACE = "responses"
|
||||
CODE_INTERPRETER_INTERCEPTION_PREFIX = "_code_interpreter_interception"
|
||||
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES = frozenset(
|
||||
("_websearch_interception", "_compression_interception")
|
||||
|
|
|
|||
|
|
@ -0,0 +1,294 @@
|
|||
"""
|
||||
Integration tests for WebSearch interception with the Responses API.
|
||||
|
||||
Tests that the websearch_interception callback intercepts litellm_web_search
|
||||
tool calls returned by /v1/responses, executes the search server-side, and
|
||||
builds a Responses-format follow-up request.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.websearch_interception.handler import (
|
||||
WebSearchInterceptionLogger,
|
||||
)
|
||||
from litellm.types.integrations.custom_logger import (
|
||||
CHAT_COMPLETION_AGENTIC_SURFACE,
|
||||
RESPONSES_AGENTIC_SURFACE,
|
||||
)
|
||||
from litellm.types.utils import CallTypes, LlmProviders
|
||||
|
||||
|
||||
def _responses_output_with_web_search(call_id: str = "fc_1", query: str = "latest ai news"):
|
||||
return SimpleNamespace(
|
||||
output=[
|
||||
SimpleNamespace(
|
||||
type="function_call",
|
||||
name="litellm_web_search",
|
||||
call_id=call_id,
|
||||
arguments='{"query": "%s"}' % query,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_hook_detects_function_call():
|
||||
"""async_should_run_responses_agentic_loop detects a litellm_web_search function_call."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
should_run, tools_dict = await logger.async_should_run_responses_agentic_loop(
|
||||
response=_responses_output_with_web_search(),
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "What's the latest AI news?"}],
|
||||
tools=[{"type": "function", "name": "litellm_web_search"}],
|
||||
stream=False,
|
||||
custom_llm_provider="openai",
|
||||
kwargs={},
|
||||
)
|
||||
|
||||
assert should_run is True
|
||||
assert tools_dict["response_format"] == "responses"
|
||||
assert len(tools_dict["tool_calls"]) == 1
|
||||
assert tools_dict["tool_calls"][0]["name"] == "litellm_web_search"
|
||||
assert tools_dict["tool_calls"][0]["call_id"] == "fc_1"
|
||||
assert tools_dict["tool_calls"][0]["input"] == {"query": "latest ai news"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_hook_not_triggered_without_tool():
|
||||
"""No web search tool in the request -> hook must not run."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
should_run, tools_dict = await logger.async_should_run_responses_agentic_loop(
|
||||
response=_responses_output_with_web_search(),
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=[{"type": "function", "name": "get_weather"}],
|
||||
stream=False,
|
||||
custom_llm_provider="openai",
|
||||
kwargs={},
|
||||
)
|
||||
|
||||
assert should_run is False
|
||||
assert tools_dict == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_hook_not_triggered_for_disabled_provider():
|
||||
"""Provider not in enabled_providers -> hook must not run."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.BEDROCK])
|
||||
|
||||
should_run, tools_dict = await logger.async_should_run_responses_agentic_loop(
|
||||
response=_responses_output_with_web_search(),
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=[{"type": "function", "name": "litellm_web_search"}],
|
||||
stream=False,
|
||||
custom_llm_provider="openai",
|
||||
kwargs={},
|
||||
)
|
||||
|
||||
assert should_run is False
|
||||
assert tools_dict == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_hook_ignores_non_websearch_function_call():
|
||||
"""A function_call for a different tool must not be intercepted."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
response = SimpleNamespace(
|
||||
output=[SimpleNamespace(type="function_call", name="get_weather", call_id="c1", arguments="{}")]
|
||||
)
|
||||
|
||||
should_run, tools_dict = await logger.async_should_run_responses_agentic_loop(
|
||||
response=response,
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=[{"type": "function", "name": "litellm_web_search"}],
|
||||
stream=False,
|
||||
custom_llm_provider="openai",
|
||||
kwargs={},
|
||||
)
|
||||
|
||||
assert should_run is False
|
||||
assert tools_dict == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_surface_marker_routes_should_run_to_responses_branch():
|
||||
"""async_should_run_agentic_loop must dispatch to the responses branch when the
|
||||
surface marker says responses.
|
||||
|
||||
Without the marker the default anthropic branch runs and never detects the
|
||||
Responses-format function_call, so interception silently no-ops on /v1/responses.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
should_run, tools_dict = await logger.async_should_run_agentic_loop(
|
||||
response=_responses_output_with_web_search(),
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "What's the latest AI news?"}],
|
||||
tools=[{"type": "function", "name": "litellm_web_search"}],
|
||||
stream=False,
|
||||
custom_llm_provider="openai",
|
||||
kwargs={"_agentic_loop_api_surface": RESPONSES_AGENTIC_SURFACE},
|
||||
)
|
||||
|
||||
assert should_run is True
|
||||
assert tools_dict["response_format"] == "responses"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_branch_does_not_detect_responses_output():
|
||||
"""Regression guard: the default (anthropic) branch must not detect a
|
||||
Responses-format function_call, proving the responses branch is required.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
should_run, tools_dict = await logger.async_should_run_agentic_loop(
|
||||
response=_responses_output_with_web_search(),
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
tools=[{"type": "function", "name": "litellm_web_search"}],
|
||||
stream=False,
|
||||
custom_llm_provider="openai",
|
||||
kwargs={},
|
||||
)
|
||||
|
||||
assert should_run is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_responses_plan_produces_responses_input():
|
||||
"""async_build_responses_agentic_loop_plan builds a Responses-format follow-up:
|
||||
the user input followed by function_call + function_call_output items, with
|
||||
the web search tool preserved and tool_choice stripped.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
tools_dict = {
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "fc_1",
|
||||
"call_id": "fc_1",
|
||||
"type": "function_call",
|
||||
"name": "litellm_web_search",
|
||||
"arguments": '{"query": "latest ai news"}',
|
||||
"input": {"query": "latest ai news"},
|
||||
}
|
||||
],
|
||||
"tool_type": "websearch",
|
||||
"provider": "openai",
|
||||
"response_format": "responses",
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
logger,
|
||||
"_execute_search",
|
||||
new=AsyncMock(return_value=("OpenAI shipped a new model", None)),
|
||||
):
|
||||
plan = await logger.async_build_responses_agentic_loop_plan(
|
||||
tools=tools_dict,
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "What's the latest AI news?"}],
|
||||
response=_responses_output_with_web_search(),
|
||||
optional_params={
|
||||
"tools": [{"type": "function", "name": "litellm_web_search"}],
|
||||
"tool_choice": {"type": "function", "name": "litellm_web_search"},
|
||||
},
|
||||
logging_obj=MagicMock(),
|
||||
stream=False,
|
||||
kwargs={"custom_llm_provider": "openai"},
|
||||
)
|
||||
|
||||
assert plan.run_agentic_loop is True
|
||||
patch_obj = plan.request_patch
|
||||
assert patch_obj is not None
|
||||
input_items = patch_obj.messages
|
||||
assert input_items is not None
|
||||
|
||||
assert input_items[0] == {"role": "user", "content": "What's the latest AI news?"}
|
||||
assert input_items[1] == {
|
||||
"type": "function_call",
|
||||
"call_id": "fc_1",
|
||||
"name": "litellm_web_search",
|
||||
"arguments": '{"query": "latest ai news"}',
|
||||
}
|
||||
assert input_items[2] == {
|
||||
"type": "function_call_output",
|
||||
"call_id": "fc_1",
|
||||
"output": "OpenAI shipped a new model",
|
||||
}
|
||||
|
||||
assert patch_obj.tools == [{"type": "function", "name": "litellm_web_search"}]
|
||||
assert "tool_choice" not in patch_obj.optional_params
|
||||
assert patch_obj.model == "openai/gpt-4o"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_hook_converts_native_responses_web_search_tool():
|
||||
"""async_pre_call_deployment_hook converts a native Responses web_search tool
|
||||
into the flat litellm_web_search function tool (Responses shape, not the
|
||||
nested Chat Completions {"function": {...}} wrapper).
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(
|
||||
kwargs={
|
||||
"model": "gpt-4o",
|
||||
"custom_llm_provider": "openai",
|
||||
"tools": [{"type": "web_search"}],
|
||||
},
|
||||
call_type=CallTypes.aresponses,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
converted_tools = result["tools"]
|
||||
assert len(converted_tools) == 1
|
||||
tool = converted_tools[0]
|
||||
assert tool["type"] == "function"
|
||||
assert tool["name"] == "litellm_web_search"
|
||||
assert "function" not in tool
|
||||
assert tool["parameters"]["required"] == ["query"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_hook_responses_returns_none_without_web_search():
|
||||
"""No web search tool in a responses request -> deployment hook makes no change."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(
|
||||
kwargs={
|
||||
"model": "gpt-4o",
|
||||
"custom_llm_provider": "openai",
|
||||
"tools": [{"type": "function", "name": "get_weather"}],
|
||||
},
|
||||
call_type=CallTypes.aresponses,
|
||||
)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_hook_responses_converts_stream_to_non_stream():
|
||||
"""Streaming responses requests are converted to non-streaming so the agentic
|
||||
loop can run, and flagged for re-wrapping afterwards.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(
|
||||
kwargs={
|
||||
"model": "gpt-4o",
|
||||
"custom_llm_provider": "openai",
|
||||
"tools": [{"type": "web_search_preview"}],
|
||||
"stream": True,
|
||||
},
|
||||
call_type=CallTypes.aresponses,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["stream"] is False
|
||||
assert result["_websearch_interception_converted_stream"] is True
|
||||
Loading…
Add table
Reference in a new issue