diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 481cf7fce8e..e4c759936ad 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -765,6 +765,56 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac """ return AgenticLoopPlan(run_agentic_loop=False) + async def async_should_run_responses_api_agentic_loop( + self, + response: Any, + model: str, + input: Any, + tools: Optional[List[Dict]], + stream: bool, + custom_llm_provider: str, + kwargs: Dict, + original_stream: Optional[bool] = None, + ) -> Tuple[bool, Dict]: + """ + Hook to determine if the Responses-API agentic loop should be executed. + + ``input`` is the OpenAI Responses-API ``input`` (string or list of items), + not the chat-style ``messages``. + + ``stream`` is the post-conversion value seen by the underlying call; + ``original_stream`` is the pre-conversion value the caller requested + (callbacks earlier in the pipeline may have rewritten ``stream`` to + ``False`` so they could consume the response). Use ``original_stream`` + when behavior should depend on what the client asked for, ``stream`` + when it should depend on what the wire actually carried. + """ + return False, {} + + async def async_run_responses_api_agentic_loop( + self, + tools: Dict, + model: str, + input: Any, + response: Any, + response_api_optional_request_params: Dict, + litellm_params: Dict, + logging_obj: "LiteLLMLoggingObj", + stream: bool, + kwargs: Dict, + original_stream: Optional[bool] = None, + ) -> Any: + """ + Hook to execute the Responses-API agentic loop. + + Implementations should run any local tool execution (e.g. web search) + and return a ``ResponsesAPIResponse`` matching the original request shape. + + See ``async_should_run_responses_api_agentic_loop`` for the + ``stream`` / ``original_stream`` distinction. + """ + return response + # Useful helpers for custom logger classes def truncate_standard_logging_payload_content( diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 37528e7dcd5..35dcd93de1b 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -7,6 +7,7 @@ server-side using litellm router's search tools. """ import asyncio +import json import math import uuid from typing import Any, Dict, List, Optional, Tuple, Union, cast @@ -19,9 +20,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_api, is_anthropic_native_web_search_tool, is_web_search_tool, is_web_search_tool_chat_completion, + is_web_search_tool_responses_api, ) from litellm.integrations.websearch_interception.transformation import ( WebSearchTransformation, @@ -255,14 +258,32 @@ class WebSearchInterceptionLogger(CustomLogger): if not tools: return None - # 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) + # Branch on call surface. Chat Completions tools use the OpenAI-nested + # ``{"type": "function", "function": {...}}`` shape. Responses-API tools + # are flat (``{"type": "function", "name": "..."}``) and the server-hosted + # web search variant is ``{"type": "web_search"}`` — neither of which + # ``is_web_search_tool`` matches. + if call_type is None: + call_type_str = "" + else: + call_type_str = getattr(call_type, "value", None) or str(call_type) + is_responses_api_call = call_type_str == "aresponses" + + if is_responses_api_call: + tool_predicate = is_web_search_tool_responses_api + standard_tool_factory = get_litellm_web_search_tool_responses_api + else: + tool_predicate = is_web_search_tool + standard_tool_factory = get_litellm_web_search_tool_openai + + has_websearch = any(tool_predicate(t) for t in tools) if not has_websearch: return None verbose_logger.debug( - "WebSearchInterception: Converting native web_search tools to LiteLLM standard" + "WebSearchInterception: Converting native web_search tools to LiteLLM standard " + f"(call_type={call_type_str or 'unknown'})" ) # If the client sent an Anthropic-native web_search_* tool, mark the @@ -276,9 +297,8 @@ class WebSearchInterceptionLogger(CustomLogger): # Convert native/custom web_search tools to LiteLLM standard converted_tools = [] for tool in tools: - if is_web_search_tool(tool): - # Convert to LiteLLM standard web search tool - converted_tool = get_litellm_web_search_tool_openai() + if tool_predicate(tool): + converted_tool = standard_tool_factory() converted_tools.append(converted_tool) verbose_logger.debug( f"WebSearchInterception: Converted {tool.get('name', 'unknown')} " @@ -810,6 +830,285 @@ class WebSearchInterceptionLogger(CustomLogger): metadata={"tool_type": "websearch", "response_format": response_format}, ) + async def async_should_run_responses_api_agentic_loop( + self, + response: Any, + model: str, + input: Any, + tools: Optional[List[Dict]], + stream: bool, + custom_llm_provider: str, + kwargs: Dict, + original_stream: Optional[bool] = None, + ) -> Tuple[bool, Dict]: + """Detect ``litellm_web_search`` ``function_call`` items in a Responses-API response.""" + verbose_logger.debug( + f"WebSearchInterception: Responses-API 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} " + f"(not in enabled list: {self.enabled_providers})" + ) + return False, {} + + has_websearch_tool = any( + is_web_search_tool_responses_api(t) for t in (tools or []) + ) + if not has_websearch_tool: + verbose_logger.debug( + "WebSearchInterception: No web_search tool in Responses-API 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 litellm_web_search function_call in Responses-API output" + ) + return False, {} + + verbose_logger.debug( + f"WebSearchInterception: Detected {len(tool_calls)} Responses-API function_call(s), " + "executing agentic loop" + ) + return True, { + "tool_calls": tool_calls, + "tool_type": "websearch", + "provider": custom_llm_provider, + "response_format": "responses", + } + + async def async_run_responses_api_agentic_loop( # noqa: PLR0915 + self, + tools: Dict, + model: str, + input: Any, + response: Any, + response_api_optional_request_params: Dict, + litellm_params: Dict, + logging_obj: Any, + stream: bool, + kwargs: Dict, + original_stream: Optional[bool] = None, + ) -> Any: + """Execute searches and re-run the Responses-API call with ``function_call_output`` items.""" + tool_calls = tools["tool_calls"] + + # Check depth + fingerprint cycle BEFORE issuing any work. Otherwise a + # client that drives the model into emitting ``litellm_web_search`` + # calls past the cap still gets ``len(tool_calls)`` parallel Tavily + # requests per iteration before the loop aborts. Mirrors + # ``_check_agentic_loop_safety`` in the chat-completion path, which + # also runs before any rerun work. + depth = int(kwargs.get("_agentic_loop_depth", 0) or 0) + max_loops = max(int(kwargs.get("max_agentic_loops", 3) or 3), 1) + fingerprints = list(kwargs.get("_agentic_loop_fingerprints", []) or []) + try: + fingerprint = json.dumps(tool_calls, sort_keys=True, default=str) + except (TypeError, ValueError): + fingerprint = str(tool_calls) + if fingerprint in fingerprints: + raise ValueError( + "Responses-API agentic loop detected repeated tool-call " + "fingerprint; aborting rerun" + ) + if depth >= max_loops: + raise ValueError( + f"Responses-API agentic loop exceeded max_agentic_loops={max_loops} for model={model}" + ) + + verbose_logger.debug( + f"WebSearchInterception: Executing Responses-API agentic loop for {len(tool_calls)} search(es)" + ) + + # Run searches in parallel. + search_tasks = [] + for tool_call in tool_calls: + query = (tool_call.get("input") or {}).get("query") + if query: + search_tasks.append(self._execute_search(query)) + else: + search_tasks.append(self._create_empty_search_result()) + search_results = await asyncio.gather(*search_tasks, return_exceptions=True) + + result_texts: List[str] = [] + for i, result in enumerate(search_results): + if isinstance(result, Exception): + verbose_logger.error( + f"WebSearchInterception: Search {i} failed: {str(result)}" + ) + result_texts.append(f"Search failed: {str(result)}") + elif isinstance(result, tuple) and len(result) == 2: + text_value, _ = result + result_texts.append( + cast(str, text_value) + if isinstance(text_value, str) + else str(text_value) + ) + else: + result_texts.append(str(result)) + + # Build follow-up ``input`` chain. Responses-API ``input`` accepts + # mixed message + function_call + function_call_output items. Strict + # providers (OpenAI-native) validate that every ``function_call_output`` + # is preceded in the conversation by its assistant ``function_call``, + # and reject the follow-up otherwise. Forward the first response's + # full ``output`` (which already includes the function_call items + # alongside reasoning / text / message blocks) so the assistant turn + # stays intact, then append our paired ``function_call_output`` items. + if isinstance(input, str): + follow_up_input: List[Dict[str, Any]] = [ + {"role": "user", "content": input}, + ] + elif isinstance(input, list): + follow_up_input = list(input) + else: + follow_up_input = [] + + first_response_output = ( + response.get("output", []) or [] + if isinstance(response, dict) + else (getattr(response, "output", None) or []) + ) + for item in first_response_output: + follow_up_input.append( + item if isinstance(item, dict) else self._dump_output_item(item) + ) + + for tool_call, result_text in zip(tool_calls, result_texts): + call_id = tool_call.get("call_id") or tool_call.get("id") or "" + follow_up_input.append( + { + "type": "function_call_output", + "call_id": call_id, + "output": result_text, + } + ) + + # Re-run the Responses-API call. ``previous_response_id`` would be a + # one-line shortcut, but only works when ``store=True`` and isn't + # supported by all backends — rebuilding the input chain works + # universally and matches how the chat-completion path handles this. + followup_kwargs = self._prepare_followup_kwargs(kwargs) + # Strip Responses-API-only flags and the converted-stream marker. + followup_kwargs.pop("response_id", None) + + followup_params = dict(response_api_optional_request_params or {}) + followup_params.pop("stream", None) + followup_params["tools"] = followup_kwargs.pop( + "tools", followup_params.get("tools") + ) + + # Reconstruct ``provider/model`` for the follow-up call. ``model`` here + # is the bare backend ID (e.g. ``openai.gpt-5.5``) because the + # litellm.aresponses dispatcher already stripped the ``bedrock_mantle/`` + # prefix before reaching the HTTP handler. Without the prefix, + # ``get_llm_provider`` raises ``LLM Provider NOT provided``. + custom_llm_provider = ( + (litellm_params or {}).get("custom_llm_provider") + or kwargs.get("custom_llm_provider") + or "" + ) + full_model_name = model + if ( + custom_llm_provider + and "/" not in model + and not model.startswith(f"{custom_llm_provider}/") + ): + full_model_name = f"{custom_llm_provider}/{model}" + + # Forward provider-specific params from the original ``litellm_params`` + # (e.g. ``aws_region_name``, ``api_base``) so the follow-up call lands + # on the same backend region as the initial call. Without this, a + # bedrock_mantle request whose model registration sets + # ``aws_region_name=us-east-2`` falls back to whatever the global + # default points at (``BEDROCK_MANTLE_REGION`` env, otherwise + # us-east-1) and 404s on the follow-up. + _PROVIDER_PARAM_KEYS = ( + "aws_region_name", + "aws_access_key_id", + "aws_secret_access_key", + "aws_session_token", + "aws_role_name", + "aws_session_name", + "aws_profile_name", + "aws_web_identity_token", + "aws_sts_endpoint", + "aws_bedrock_runtime_endpoint", + "api_base", + "api_key", + "api_version", + ) + for k in _PROVIDER_PARAM_KEYS: + v = (litellm_params or {}).get(k) + if v is not None and k not in followup_params: + followup_params[k] = v + + # Forward caller-context fields the proxy uses for budget / spend + # attribution. Without ``metadata`` / ``litellm_metadata`` / ``user``, + # the internal follow-up call is logged against an empty key/team and + # bypasses budgets configured on the original API key. + # ``litellm_logging_obj`` is intentionally excluded — see + # ``_prepare_followup_kwargs`` — so the follow-up creates its own. + _ATTRIBUTION_KEYS = ( + "metadata", + "litellm_metadata", + "user", + "user_api_key", + "user_api_key_alias", + "user_api_key_user_id", + "user_api_key_team_id", + "user_api_key_team_alias", + "user_api_key_org_id", + "proxy_server_request", + ) + for k in _ATTRIBUTION_KEYS: + v = (litellm_params or {}).get(k) + if v is not None and k not in followup_kwargs: + followup_kwargs[k] = v + + # Loop-state propagation. ``max_agentic_loops`` must travel with the + # follow-up call so the cap a deployment configured up-front isn't + # silently reset to the default on each hop. + followup_kwargs["_agentic_loop_depth"] = depth + 1 + followup_kwargs["_agentic_loop_fingerprints"] = fingerprints + [fingerprint] + followup_kwargs["max_agentic_loops"] = max_loops + + verbose_logger.debug( + "WebSearchInterception: Responses-API follow-up call " + f"[items={len(follow_up_input)} model={full_model_name} " + f"depth={depth + 1}/{max_loops}]" + ) + + return await litellm.aresponses( + model=full_model_name, + input=follow_up_input, + **followup_params, + **followup_kwargs, + ) + + @staticmethod + def _dump_output_item(item: Any) -> Dict[str, Any]: + """Best-effort conversion of a Responses-API output item to a dict.""" + if hasattr(item, "model_dump"): + return cast(Dict[str, Any], item.model_dump()) + if hasattr(item, "dict"): + try: + return cast(Dict[str, Any], item.dict()) + except Exception: + pass + return {k: getattr(item, k) for k in dir(item) if not k.startswith("_")} + @staticmethod def _resolve_max_tokens( optional_params: Dict, diff --git a/litellm/integrations/websearch_interception/tools.py b/litellm/integrations/websearch_interception/tools.py index b29372af9ed..a55e4fcf8f6 100644 --- a/litellm/integrations/websearch_interception/tools.py +++ b/litellm/integrations/websearch_interception/tools.py @@ -82,6 +82,60 @@ def get_litellm_web_search_tool_openai() -> Dict[str, Any]: } +def get_litellm_web_search_tool_responses_api() -> Dict[str, Any]: + """ + Get the standard LiteLLM web search tool definition in OpenAI Responses API format. + + Responses-API function tools are flat (no nested ``function`` key): + ``{"type": "function", "name": "...", "description": "...", "parameters": {...}}``. + + Used by ``WebSearchInterceptionLogger.async_pre_call_deployment_hook`` for + ``call_type == "aresponses"`` to convert server-hosted ``web_search`` tools + (which providers like Bedrock Mantle reject) into a function-typed tool we + can intercept and execute server-side. + """ + 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_api(tool: Dict[str, Any]) -> bool: + """ + Check if a tool is a web search tool in OpenAI Responses API shape. + + Detects: + - Server-hosted web search: ``{"type": "web_search"}`` / + ``{"type": "web_search_preview"}`` (sent by Codex CLI and the OpenAI SDK + when ``web_search`` is enabled). + - LiteLLM standard, flat: ``{"type": "function", "name": "litellm_web_search"}`` + - Anthropic-native variants: ``{"type": "web_search_*"}`` (forwarded + verbatim by some clients). + """ + tool_type = tool.get("type", "") + if not isinstance(tool_type, str): + return False + if tool_type == "web_search" or tool_type.startswith("web_search_"): + return True + if tool_type == "function" and tool.get("name") == LITELLM_WEB_SEARCH_TOOL_NAME: + return True + return False + + 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). diff --git a/litellm/integrations/websearch_interception/transformation.py b/litellm/integrations/websearch_interception/transformation.py index 9c20a3f6c77..7cca1ca1042 100644 --- a/litellm/integrations/websearch_interception/transformation.py +++ b/litellm/integrations/websearch_interception/transformation.py @@ -59,10 +59,79 @@ class WebSearchTransformation: return False, [] # Parse non-streaming response based on format + if response_format == "responses": + return WebSearchTransformation._detect_from_responses_api_response(response) if response_format == "openai": return WebSearchTransformation._detect_from_openai_response(response) + return WebSearchTransformation._detect_from_non_streaming_response(response) + + @staticmethod + def _detect_from_responses_api_response( + response: Any, + ) -> Tuple[bool, List[Dict]]: + """Parse OpenAI Responses-API response for ``litellm_web_search`` function calls. + + The Responses API returns ``output`` as a list of items; tool calls + appear as ``{"type": "function_call", "call_id": "...", "name": "...", + "arguments": ""}``. Pre-request conversion replaces all + web-search tools with the LiteLLM standard ``litellm_web_search`` + function tool, so the only name we need to recognize here is the + standard one. ``call_id`` is preserved as ``id`` so the agentic loop + can pair it with a ``function_call_output`` item. + """ + if isinstance(response, dict): + output = response.get("output", []) or [] else: - return WebSearchTransformation._detect_from_non_streaming_response(response) + output = getattr(response, "output", None) or [] + + 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") or item.get("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) or getattr(item, "id", None) + arguments = getattr(item, "arguments", None) + + if item_type != "function_call": + continue + if item_name != LITELLM_WEB_SEARCH_TOOL_NAME: + continue + + if isinstance(arguments, str): + try: + parsed_arguments = json.loads(arguments) + except json.JSONDecodeError: + verbose_logger.warning( + "WebSearchInterception: Failed to parse Responses-API " + f"function_call arguments: {arguments}" + ) + parsed_arguments = {} + elif isinstance(arguments, dict): + parsed_arguments = arguments + else: + parsed_arguments = {} + + tool_calls.append( + { + "id": call_id, + "call_id": call_id, + "type": "function_call", + "name": item_name, + "input": parsed_arguments, + "arguments": parsed_arguments, + } + ) + verbose_logger.debug( + f"WebSearchInterception: Found Responses-API function_call " + f"name={item_name} call_id={call_id}" + ) + + return len(tool_calls) > 0, tool_calls @staticmethod def _detect_from_non_streaming_response( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index eedab7fc36c..6f5a1d4f4d9 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2531,12 +2531,178 @@ class BaseLLMHTTPHandler: provider_config=responses_api_provider_config, ) - return responses_api_provider_config.transform_response_api_response( - model=model, - raw_response=response, - logging_obj=logging_obj, + transformed_response = ( + responses_api_provider_config.transform_response_api_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) ) + # Agentic-loop hook dispatch (e.g. websearch interception). Mirrors + # ``_call_agentic_chat_completion_hooks`` and ``_call_agentic_completion_hooks`` + # used by the Chat Completions and Anthropic Messages paths. The hook + # has access to the full transformed response and can re-run the + # Responses-API call with ``function_call_output`` items spliced into + # ``input``. + # Stash custom_llm_provider on litellm_params so the hook can + # reconstruct the ``provider/model`` string for follow-up calls. + # GenericLiteLLMParams declares ``custom_llm_provider`` as an + # optional field, so the dict often holds it as None — meaning + # ``setdefault`` would skip the assignment. Overwrite explicitly. + agentic_litellm_params = dict(litellm_params) + if not agentic_litellm_params.get("custom_llm_provider"): + agentic_litellm_params["custom_llm_provider"] = custom_llm_provider + + agentic_response: Optional[Any] = None + try: + agentic_response = await self._call_agentic_responses_api_hooks( + response=transformed_response, + model=model, + input=input, + response_api_optional_request_params=response_api_optional_request_params, + litellm_params=agentic_litellm_params, + logging_obj=logging_obj, + stream=stream, + custom_llm_provider=custom_llm_provider, + ) + except Exception as e: + verbose_logger.exception( + f"LiteLLM.AgenticHookError: Exception in Responses-API agentic hooks: {str(e)}" + ) + + final_response = ( + agentic_response if agentic_response is not None else transformed_response + ) + + # If a callback (e.g. websearch interception) silently converted a + # client-requested stream=True call to stream=False so it could + # consume the response, the proxy SSE layer still expects an async + # iterator on the way out. Wrap a completed ``ResponsesAPIResponse`` + # in ``CachedResponsesAPIStreamingIterator`` (the same wrapper the + # cache hit path uses) so ``async for chunk in stream_iterator`` + # works downstream. Mirrors the chat-completion path that sets + # ``model_call_details["websearch_interception_converted_stream"]`` + # from the same flag for the same reason. + converted_stream = bool( + agentic_litellm_params.get( + "_websearch_interception_converted_stream", False + ) + ) + if ( + converted_stream + and getattr(logging_obj, "model_call_details", None) is not None + ): + logging_obj.model_call_details[ + "websearch_interception_converted_stream" + ] = True + if converted_stream and not isinstance( + final_response, BaseResponsesAPIStreamingIterator + ): + from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + ) + + return CachedResponsesAPIStreamingIterator( + response=final_response, + logging_obj=logging_obj, + request_data=request_context, + call_type=CallTypes.responses.value, + ) + + return final_response + + async def _call_agentic_responses_api_hooks( + self, + response: Any, + model: str, + input: Union[str, ResponseInputParam], + response_api_optional_request_params: Dict, + litellm_params: Dict, + logging_obj: "LiteLLMLoggingObj", + stream: bool, + custom_llm_provider: str, + ) -> Optional[Any]: + """Dispatch ``async_should_run_responses_api_agentic_loop`` / + ``async_run_responses_api_agentic_loop`` for any ``CustomLogger`` in + ``litellm.callbacks`` that overrides them. Returns the agentic-loop + response if any callback ran, else ``None``.""" + from litellm._logging import verbose_logger + from litellm.integrations.custom_logger import CustomLogger + + callbacks = litellm.callbacks + (logging_obj.dynamic_success_callbacks or []) + tools = response_api_optional_request_params.get("tools", []) or [] + + # Surface the original (pre-conversion) stream value so custom loggers + # can distinguish "client requested streaming, we converted internally" + # from "client requested non-streaming". By the time this dispatcher + # runs, ``stream`` itself reflects the post-conversion value. + original_stream = stream or bool( + litellm_params.get("_websearch_interception_converted_stream", False) + ) + + # Plumb agentic-loop safety state into the hook so re-entrant calls + # can bound depth + detect cycles. Defaults match + # ``_get_agentic_loop_settings`` so the chat-completion and Responses- + # API paths share the same semantics. + loop_state = { + "_agentic_loop_depth": int( + litellm_params.get("_agentic_loop_depth", 0) or 0 + ), + "max_agentic_loops": int(litellm_params.get("max_agentic_loops", 3) or 3), + "_agentic_loop_fingerprints": list( + litellm_params.get("_agentic_loop_fingerprints", []) or [] + ), + } + + for callback in callbacks: + if not isinstance(callback, CustomLogger): + continue + + try: + should_run, hook_tools = ( + await callback.async_should_run_responses_api_agentic_loop( + response=response, + model=model, + input=input, + tools=tools, + stream=stream, + original_stream=original_stream, + custom_llm_provider=custom_llm_provider, + kwargs=dict(loop_state), + ) + ) + except Exception as e: + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in " + f"async_should_run_responses_api_agentic_loop: {str(e)}" + ) + continue + + if not should_run: + continue + + try: + return await callback.async_run_responses_api_agentic_loop( + tools=hook_tools, + model=model, + input=input, + response=response, + response_api_optional_request_params=response_api_optional_request_params, + litellm_params=litellm_params, + logging_obj=logging_obj, + stream=stream, + original_stream=original_stream, + kwargs=dict(loop_state), + ) + except Exception as e: + verbose_logger.exception( + "LiteLLM.AgenticHookError: Exception in " + f"async_run_responses_api_agentic_loop: {str(e)}" + ) + + return None + async def async_delete_response_api_handler( self, response_id: str, diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_responses_api.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_responses_api.py new file mode 100644 index 00000000000..9edfac6c8e2 --- /dev/null +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_responses_api.py @@ -0,0 +1,445 @@ +""" +Unit tests for WebSearchInterceptionLogger on the OpenAI Responses API path. + +Covers the call_type=="aresponses" branch of async_pre_call_deployment_hook, +the new Responses-API output parser, the should-run hook, and the agentic +loop that rebuilds the input chain with function_call_output items. +""" + +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm.integrations.websearch_interception.handler import ( + WebSearchInterceptionLogger, +) +from litellm.integrations.websearch_interception.tools import ( + get_litellm_web_search_tool_responses_api, + is_web_search_tool_responses_api, +) +from litellm.integrations.websearch_interception.transformation import ( + WebSearchTransformation, +) +from litellm.types.utils import LlmProviders + + +@pytest.fixture +def logger() -> WebSearchInterceptionLogger: + return WebSearchInterceptionLogger( + enabled_providers=[LlmProviders.BEDROCK_MANTLE, LlmProviders.OPENAI], + search_tool_name="tavily-search", + ) + + +def test_is_web_search_tool_responses_api_detects_codex_shape() -> None: + assert is_web_search_tool_responses_api({"type": "web_search"}) + assert is_web_search_tool_responses_api({"type": "web_search_preview"}) + assert is_web_search_tool_responses_api({"type": "web_search_20250305"}) + assert is_web_search_tool_responses_api( + {"type": "function", "name": "litellm_web_search"} + ) + assert not is_web_search_tool_responses_api({"type": "function", "name": "foo"}) + assert not is_web_search_tool_responses_api({"type": "image_generation"}) + assert not is_web_search_tool_responses_api({}) + + +def test_get_litellm_web_search_tool_responses_api_is_flat() -> None: + tool = get_litellm_web_search_tool_responses_api() + # Flat shape — no nested ``function`` key (which would be Chat-Completions style). + assert tool["type"] == "function" + assert tool["name"] == "litellm_web_search" + assert "parameters" in tool + assert "function" not in tool + + +def test_responses_api_response_parser_finds_function_call() -> None: + response = { + "output": [ + {"type": "reasoning", "id": "rs_1", "summary": []}, + { + "type": "function_call", + "call_id": "call_abc", + "name": "litellm_web_search", + "arguments": json.dumps({"query": "weather hong kong"}), + }, + ] + } + should, calls = WebSearchTransformation.transform_request( + response=response, stream=False, response_format="responses" + ) + assert should is True + assert len(calls) == 1 + assert calls[0]["call_id"] == "call_abc" + assert calls[0]["input"]["query"] == "weather hong kong" + + +def test_responses_api_response_parser_ignores_other_function_calls() -> None: + response = { + "output": [ + { + "type": "function_call", + "call_id": "call_xyz", + "name": "calculator", + "arguments": json.dumps({"x": 1}), + } + ] + } + should, calls = WebSearchTransformation.transform_request( + response=response, stream=False, response_format="responses" + ) + assert should is False + assert calls == [] + + +@pytest.mark.asyncio +async def test_pre_call_hook_aresponses_uses_flat_shape( + logger: WebSearchInterceptionLogger, +) -> None: + kwargs = { + "model": "openai.gpt-5.5", + "custom_llm_provider": "bedrock_mantle", + "tools": [{"type": "web_search"}], + "stream": True, + } + out = await logger.async_pre_call_deployment_hook(kwargs, call_type="aresponses") + assert out is not None + converted = out["tools"][0] + # Flat — no nested ``function`` key. + assert converted["type"] == "function" + assert converted["name"] == "litellm_web_search" + assert "function" not in converted + # Streaming converted to non-streaming for interception. + assert out["stream"] is False + assert out["_websearch_interception_converted_stream"] is True + + +@pytest.mark.asyncio +async def test_pre_call_hook_aresponses_skips_disabled_provider( + logger: WebSearchInterceptionLogger, +) -> None: + # vertex_ai isn't in the logger's enabled list. + kwargs = { + "model": "publishers/google/models/gemini-1.5-pro", + "custom_llm_provider": "vertex_ai", + "tools": [{"type": "web_search"}], + } + out = await logger.async_pre_call_deployment_hook(kwargs, call_type="aresponses") + assert out is None + + +@pytest.mark.asyncio +async def test_should_run_responses_api_hook_returns_tool_calls( + logger: WebSearchInterceptionLogger, +) -> None: + response = { + "output": [ + { + "type": "function_call", + "call_id": "call_1", + "name": "litellm_web_search", + "arguments": json.dumps({"query": "btc price"}), + } + ] + } + should, hook_tools = await logger.async_should_run_responses_api_agentic_loop( + response=response, + model="openai.gpt-5.5", + input="hi", + tools=[{"type": "web_search"}], + stream=False, + custom_llm_provider="bedrock_mantle", + kwargs={}, + ) + assert should is True + assert hook_tools["response_format"] == "responses" + assert len(hook_tools["tool_calls"]) == 1 + + +@pytest.mark.asyncio +async def test_run_responses_api_hook_preserves_assistant_turn_then_appends_outputs( + logger: WebSearchInterceptionLogger, +) -> None: + """The follow-up ``input`` must carry the user message, then the entire + first-turn assistant ``output`` (reasoning + function_call), then the + paired ``function_call_output`` items. Strict providers (OpenAI-native) + 400 if a ``function_call_output`` is not preceded by its matching + ``function_call`` from the same conversation turn — see + https://platform.openai.com/docs/api-reference/responses.""" + + fake_search = MagicMock() + fake_search.results = [ + MagicMock(title="t", url="https://example.com", snippet="snippet"), + ] + + captured = {} + + async def fake_aresponses(**kwargs): + captured.update(kwargs) + return {"id": "resp_final", "output": [{"type": "message"}]} + + initial_response = { + "id": "resp_initial", + "output": [ + {"type": "reasoning", "id": "rs_1", "summary": []}, + { + "type": "function_call", + "call_id": "call_1", + "name": "litellm_web_search", + "arguments": '{"query": "btc price"}', + }, + ], + } + + with ( + patch.object(litellm, "asearch", new=AsyncMock(return_value=fake_search)), + patch.object(litellm, "aresponses", new=AsyncMock(side_effect=fake_aresponses)), + ): + await logger.async_run_responses_api_agentic_loop( + tools={ + "tool_calls": [ + { + "id": "call_1", + "call_id": "call_1", + "type": "function_call", + "name": "litellm_web_search", + "input": {"query": "btc price"}, + "arguments": {"query": "btc price"}, + } + ], + "response_format": "responses", + }, + model="openai.gpt-5.5", + input="What's the BTC price?", + response=initial_response, + response_api_optional_request_params={"tools": [{"type": "function"}]}, + litellm_params={}, + logging_obj=MagicMock(), + stream=False, + kwargs={}, + ) + + follow_up_input = captured["input"] + assert follow_up_input[0] == {"role": "user", "content": "What's the BTC price?"} + assert follow_up_input[1]["type"] == "reasoning" + assert follow_up_input[2]["type"] == "function_call" + assert follow_up_input[2]["call_id"] == "call_1" + assert follow_up_input[2]["name"] == "litellm_web_search" + assert follow_up_input[3]["type"] == "function_call_output" + assert follow_up_input[3]["call_id"] == "call_1" + assert "snippet" in follow_up_input[3]["output"] + + +@pytest.mark.asyncio +async def test_run_responses_api_hook_aborts_on_max_loops( + logger: WebSearchInterceptionLogger, +) -> None: + """Depth guard prevents a model from inducing unbounded recursion by + keeping it returning the same ``litellm_web_search`` tool call.""" + + asearch_mock = AsyncMock() + aresponses_mock = AsyncMock() + + with ( + patch.object(litellm, "asearch", new=asearch_mock), + patch.object(litellm, "aresponses", new=aresponses_mock), + ): + with pytest.raises(ValueError, match="exceeded max_agentic_loops"): + await logger.async_run_responses_api_agentic_loop( + tools={ + "tool_calls": [ + { + "id": "c", + "call_id": "c", + "type": "function_call", + "name": "litellm_web_search", + "input": {"query": "q"}, + "arguments": {"query": "q"}, + } + ], + "response_format": "responses", + }, + model="openai.gpt-5.5", + input="hi", + response={"id": "r", "output": []}, + response_api_optional_request_params={}, + litellm_params={}, + logging_obj=MagicMock(), + stream=False, + kwargs={ + "_agentic_loop_depth": 3, + "max_agentic_loops": 3, + "_agentic_loop_fingerprints": [], + }, + ) + + # Guard MUST run before any work — capped-out clients shouldn't burn + # parallel Tavily calls per iteration. + asearch_mock.assert_not_called() + aresponses_mock.assert_not_called() + + +@pytest.mark.asyncio +async def test_run_responses_api_hook_propagates_max_agentic_loops( + logger: WebSearchInterceptionLogger, +) -> None: + """Re-entrant call must carry ``max_agentic_loops`` forward; otherwise a + deployment-configured cap silently resets to the default on each hop.""" + + fake_search = MagicMock() + fake_search.results = [MagicMock(title="t", url="u", snippet="s")] + captured = {} + + async def fake_aresponses(**kwargs): + captured.update(kwargs) + return {"id": "r2", "output": []} + + with ( + patch.object(litellm, "asearch", new=AsyncMock(return_value=fake_search)), + patch.object(litellm, "aresponses", new=AsyncMock(side_effect=fake_aresponses)), + ): + await logger.async_run_responses_api_agentic_loop( + tools={ + "tool_calls": [ + { + "id": "c", + "call_id": "c", + "type": "function_call", + "name": "litellm_web_search", + "input": {"query": "q"}, + "arguments": {"query": "q"}, + } + ], + "response_format": "responses", + }, + model="openai.gpt-5.5", + input="hi", + response={"id": "r1", "output": []}, + response_api_optional_request_params={}, + litellm_params={}, + logging_obj=MagicMock(), + stream=False, + kwargs={ + "_agentic_loop_depth": 1, + "max_agentic_loops": 7, + "_agentic_loop_fingerprints": [], + }, + ) + + assert captured["max_agentic_loops"] == 7 + assert captured["_agentic_loop_depth"] == 2 + assert len(captured["_agentic_loop_fingerprints"]) == 1 + + +@pytest.mark.asyncio +async def test_run_responses_api_hook_propagates_caller_attribution( + logger: WebSearchInterceptionLogger, +) -> None: + """Internal follow-up call must carry caller-attribution fields + (``metadata``, ``litellm_metadata``, ``user``, ``user_api_key_*``) so + proxy budget / spend logging accounts the call against the original + API key and team rather than an empty owner.""" + + fake_search = MagicMock() + fake_search.results = [MagicMock(title="t", url="u", snippet="s")] + captured = {} + + async def fake_aresponses(**kwargs): + captured.update(kwargs) + return {"id": "r2", "output": []} + + with ( + patch.object(litellm, "asearch", new=AsyncMock(return_value=fake_search)), + patch.object(litellm, "aresponses", new=AsyncMock(side_effect=fake_aresponses)), + ): + await logger.async_run_responses_api_agentic_loop( + tools={ + "tool_calls": [ + { + "id": "c", + "call_id": "c", + "type": "function_call", + "name": "litellm_web_search", + "input": {"query": "q"}, + "arguments": {"query": "q"}, + } + ], + "response_format": "responses", + }, + model="openai.gpt-5.5", + input="hi", + response={"id": "r1", "output": []}, + response_api_optional_request_params={}, + litellm_params={ + "metadata": {"trace_id": "abc"}, + "litellm_metadata": {"user_api_key_alias": "qmachu"}, + "user": "u-1", + "user_api_key": "sk-test", + "user_api_key_user_id": "u-1", + "user_api_key_team_id": "t-1", + "user_api_key_team_alias": "team-a", + "user_api_key_org_id": "o-1", + "proxy_server_request": {"arrival_time": "now"}, + }, + logging_obj=MagicMock(), + stream=False, + kwargs={}, + ) + + assert captured["metadata"] == {"trace_id": "abc"} + assert captured["litellm_metadata"] == {"user_api_key_alias": "qmachu"} + assert captured["user"] == "u-1" + assert captured["user_api_key"] == "sk-test" + assert captured["user_api_key_team_id"] == "t-1" + assert captured["user_api_key_team_alias"] == "team-a" + assert captured["user_api_key_org_id"] == "o-1" + assert captured["proxy_server_request"] == {"arrival_time": "now"} + + +@pytest.mark.asyncio +async def test_run_responses_api_hook_aborts_on_repeated_fingerprint( + logger: WebSearchInterceptionLogger, +) -> None: + """Cycle break: same tool_calls fingerprint twice in a row aborts.""" + + fake_search = MagicMock() + fake_search.results = [MagicMock(title="t", url="u", snippet="s")] + + tool_calls = [ + { + "id": "c", + "call_id": "c", + "type": "function_call", + "name": "litellm_web_search", + "input": {"query": "q"}, + "arguments": {"query": "q"}, + } + ] + import json as _json + + seen_fingerprint = _json.dumps(tool_calls, sort_keys=True, default=str) + + with ( + patch.object(litellm, "asearch", new=AsyncMock(return_value=fake_search)), + patch.object(litellm, "aresponses", new=AsyncMock()), + ): + with pytest.raises(ValueError, match="repeated tool-call fingerprint"): + await logger.async_run_responses_api_agentic_loop( + tools={ + "tool_calls": tool_calls, + "response_format": "responses", + }, + model="openai.gpt-5.5", + input="hi", + response={"id": "r", "output": []}, + response_api_optional_request_params={}, + litellm_params={}, + logging_obj=MagicMock(), + stream=False, + kwargs={ + "_agentic_loop_depth": 0, + "max_agentic_loops": 5, + "_agentic_loop_fingerprints": [seen_fingerprint], + }, + )