fix(bedrock): trigger websearch_interception agentic loop for /chat/completions

Signed-off-by: Nishchay Veer <nishchayveer19@gmail.com>
This commit is contained in:
Nishchay Veer 2026-04-06 02:57:20 +05:30
parent d251238bd7
commit b882bec39a
2 changed files with 347 additions and 42 deletions

View file

@ -1,9 +1,10 @@
import json
from typing import Any, Optional, Union
from typing import Any, Dict, List, Optional, Union
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.anthropic_beta_headers_manager import (
update_headers_with_filtered_beta,
)
@ -89,6 +90,83 @@ class BedrockConverseLLM(BaseAWSLLM):
def __init__(self) -> None:
super().__init__()
async def _call_agentic_chat_completion_hooks(
self,
response: Any,
model: str,
messages: List[Dict],
optional_params: Dict,
logging_obj: LiteLLMLoggingObject,
stream: bool,
litellm_params: Dict,
) -> Optional[Any]:
"""
Call agentic chat completion hooks for all custom loggers.
Checks each callback for async_should_run_chat_completion_agentic_loop
and if triggered, runs async_run_chat_completion_agentic_loop.
Returns the response from the agentic loop, or None if no hook runs.
"""
from litellm.integrations.custom_logger import CustomLogger
callbacks = litellm.callbacks + (
getattr(logging_obj, "dynamic_success_callbacks", None) or []
)
tools = optional_params.get("tools", [])
custom_llm_provider = litellm_params.get("custom_llm_provider", "bedrock")
for callback in callbacks:
try:
if isinstance(callback, CustomLogger):
if not hasattr(
callback, "async_should_run_chat_completion_agentic_loop"
):
continue
(
should_run,
tool_calls,
) = await callback.async_should_run_chat_completion_agentic_loop(
response=response,
model=model,
messages=messages,
tools=tools,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=litellm_params,
)
if should_run:
kwargs_with_provider = (
litellm_params.copy() if litellm_params else {}
)
kwargs_with_provider["custom_llm_provider"] = (
custom_llm_provider
)
agentic_response = (
await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
)
)
return agentic_response
except Exception as e:
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in agentic completion hooks for Bedrock: %s",
str(e),
)
return None
async def async_streaming(
self,
model: str,
@ -235,7 +313,7 @@ class BedrockConverseLLM(BaseAWSLLM):
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
return litellm.AmazonConverseConfig()._transform_response(
initial_response = litellm.AmazonConverseConfig()._transform_response(
model=model,
response=response,
model_response=model_response,
@ -248,6 +326,18 @@ class BedrockConverseLLM(BaseAWSLLM):
encoding=encoding,
)
agentic_response = await self._call_agentic_chat_completion_hooks(
response=initial_response,
model=model,
messages=messages,
optional_params=optional_params,
logging_obj=logging_obj,
stream=False,
litellm_params=litellm_params,
)
return agentic_response if agentic_response is not None else initial_response
def completion( # noqa: PLR0915
self,
model: str,
@ -332,9 +422,9 @@ class BedrockConverseLLM(BaseAWSLLM):
aws_external_id = optional_params.pop("aws_external_id", None)
optional_params.pop("aws_region_name", None)
litellm_params[
"aws_region_name"
] = aws_region_name # [DO NOT DELETE] important for async calls
litellm_params["aws_region_name"] = (
aws_region_name # [DO NOT DELETE] important for async calls
)
credentials: Credentials = self.get_credentials(
aws_access_key_id=aws_access_key_id,

View file

@ -4,6 +4,7 @@ Integration tests for WebSearch interception with chat completions API.
Tests the end-to-end flow of websearch_interception callback with
litellm.acompletion() for transparent server-side web search execution.
"""
import os
from unittest.mock import AsyncMock, MagicMock, patch
@ -45,7 +46,7 @@ def websearch_logger():
)
async def test_websearch_chat_completion_with_openai():
"""Test websearch interception with OpenAI chat completions API.
This test verifies that:
1. Model calls litellm_web_search tool
2. Server executes web search automatically
@ -58,12 +59,15 @@ async def test_websearch_chat_completion_with_openai():
enabled_providers=[LlmProviders.OPENAI]
)
litellm.callbacks = [websearch_logger]
try:
response = await litellm.acompletion(
model="gpt-4o-mini", # Use cheaper model for testing
messages=[
{"role": "user", "content": "What's the weather in San Francisco today?"}
{
"role": "user",
"content": "What's the weather in San Francisco today?",
}
],
tools=[
{
@ -85,12 +89,12 @@ async def test_websearch_chat_completion_with_openai():
}
],
)
# Verify response structure
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content is not None
assert len(response.choices[0].message.content) > 0
# If agentic loop worked, we should NOT have tool_calls in final response
# (they should have been executed and replaced with final answer)
if hasattr(response.choices[0].message, "tool_calls"):
@ -99,10 +103,9 @@ async def test_websearch_chat_completion_with_openai():
pytest.skip(
"Agentic loop did not execute - search tool may not be configured"
)
# Verify we got a meaningful response
assert response.choices[0].finish_reason in ["stop", "end_turn"]
finally:
# Restore original callbacks
litellm.callbacks = original_callbacks
@ -117,11 +120,11 @@ async def test_websearch_chat_completion_hook_detection():
Function,
Message,
)
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.OPENAI]
)
# Mock response with litellm_web_search tool call
mock_response = ModelResponse(
id="test-123",
@ -142,14 +145,14 @@ async def test_websearch_chat_completion_hook_detection():
),
)
],
)
),
)
],
model="gpt-4o",
object="chat.completion",
created=1234567890,
)
# Test should_run_chat_completion_agentic_loop
should_run, tools_dict = (
await websearch_logger.async_should_run_chat_completion_agentic_loop(
@ -167,7 +170,7 @@ async def test_websearch_chat_completion_hook_detection():
kwargs={},
)
)
# Verify hook detected the tool call
assert should_run is True
assert "tool_calls" in tools_dict
@ -180,11 +183,11 @@ async def test_websearch_chat_completion_hook_detection():
async def test_websearch_not_triggered_without_tool():
"""Test that websearch hook is NOT triggered when no web search tool in request."""
from litellm.types.utils import Choices, Message
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.OPENAI]
)
mock_response = ModelResponse(
id="test-123",
choices=[
@ -195,14 +198,14 @@ async def test_websearch_not_triggered_without_tool():
role="assistant",
content="Here's the answer",
tool_calls=None,
)
),
)
],
model="gpt-4o",
object="chat.completion",
created=1234567890,
)
# Test without web search tool
should_run, tools_dict = (
await websearch_logger.async_should_run_chat_completion_agentic_loop(
@ -220,7 +223,7 @@ async def test_websearch_not_triggered_without_tool():
kwargs={},
)
)
# Verify hook did NOT trigger
assert should_run is False
assert tools_dict == {}
@ -240,7 +243,7 @@ async def test_websearch_not_triggered_for_disabled_provider():
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.BEDROCK]
)
mock_response = ModelResponse(
id="test-123",
choices=[
@ -260,14 +263,14 @@ async def test_websearch_not_triggered_for_disabled_provider():
),
)
],
)
),
)
],
model="gpt-4o",
object="chat.completion",
created=1234567890,
)
# Test with OpenAI provider (not enabled)
should_run, tools_dict = (
await websearch_logger.async_should_run_chat_completion_agentic_loop(
@ -285,7 +288,7 @@ async def test_websearch_not_triggered_for_disabled_provider():
kwargs={},
)
)
# Verify hook did NOT trigger
assert should_run is False
assert tools_dict == {}
@ -294,7 +297,7 @@ async def test_websearch_not_triggered_for_disabled_provider():
@pytest.mark.asyncio
async def test_websearch_json_serialization_fix():
"""Test that tool call arguments are properly JSON serialized.
Regression test for the bug where arguments were converted to Python
string representation instead of proper JSON, causing providers like
MiniMax to reject requests with 'invalid function arguments json string'.
@ -311,25 +314,25 @@ async def test_websearch_json_serialization_fix():
"input": {"query": "weather in SF"}, # Dict input
}
]
search_results = ["Weather: 65°F, partly cloudy"]
# Transform to OpenAI format
assistant_message, tool_messages = WebSearchTransformation.transform_response(
tool_calls=tool_calls,
search_results=search_results,
response_format="openai",
)
# Verify arguments are properly JSON serialized
import json
arguments_str = assistant_message["tool_calls"][0]["function"]["arguments"]
# Should be valid JSON
parsed_args = json.loads(arguments_str)
assert parsed_args == {"query": "weather in SF"}
# Should NOT be Python string representation like "{'query': 'weather in SF'}"
assert arguments_str == '{"query": "weather in SF"}'
assert arguments_str != "{'query': 'weather in SF'}"
@ -343,7 +346,7 @@ async def test_websearch_json_serialization_fix():
)
async def test_websearch_streaming_conversion():
"""Test that streaming requests are converted to non-streaming for web search.
When stream=True is passed with web search tools, the handler should:
1. Convert stream=True to stream=False for initial request
2. Execute web search
@ -353,13 +356,11 @@ async def test_websearch_streaming_conversion():
enabled_providers=[LlmProviders.OPENAI], search_tool_name="perplexity-search"
)
litellm.callbacks = [websearch_logger]
try:
response = await litellm.acompletion(
model="gpt-4o-mini",
messages=[
{"role": "user", "content": "What's the latest AI news?"}
],
messages=[{"role": "user", "content": "What's the latest AI news?"}],
tools=[
{
"type": "function",
@ -375,24 +376,238 @@ async def test_websearch_streaming_conversion():
],
stream=True,
)
# Response should be a streaming iterator
chunks = []
async for chunk in response:
chunks.append(chunk)
# Verify we got streaming chunks
assert len(chunks) > 0
# Verify chunks have expected structure
for chunk in chunks:
assert hasattr(chunk, "choices")
assert len(chunk.choices) > 0
finally:
litellm.callbacks = []
@pytest.mark.asyncio
async def test_bedrock_converse_agentic_hook_triggers():
"""Test that BedrockConverseLLM._call_agentic_chat_completion_hooks
correctly triggers the websearch interception agentic loop.
Regression test for https://github.com/BerriAI/litellm/issues/25191:
websearch_interception worked via /v1/messages but NOT via
/chat/completions for Bedrock because BedrockConverseLLM.async_completion
did not call agentic chat completion hooks.
"""
from unittest.mock import MagicMock
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Choices,
Function,
Message,
)
handler = BedrockConverseLLM()
# Build a ModelResponse that looks like Bedrock returned a tool call
mock_response = ModelResponse(
id="chatcmpl-bedrock-123",
choices=[
Choices(
finish_reason="tool_calls",
index=0,
message=Message(
role="assistant",
content="",
tool_calls=[
ChatCompletionMessageToolCall(
id="tooluse_abc",
type="function",
function=Function(
name="litellm_web_search",
arguments='{"query": "weather tomorrow"}',
),
)
],
),
)
],
model="us.anthropic.claude-sonnet-4-6",
object="chat.completion",
created=1234567890,
)
# Create a mock callback that simulates the websearch interception logger
mock_callback = AsyncMock(spec=WebSearchInterceptionLogger)
mock_callback.async_should_run_chat_completion_agentic_loop = AsyncMock(
return_value=(
True,
{
"tool_calls": [
{
"id": "tooluse_abc",
"type": "function",
"name": "litellm_web_search",
"function": {
"name": "litellm_web_search",
"arguments": {"query": "weather tomorrow"},
},
"input": {"query": "weather tomorrow"},
}
],
"tool_type": "websearch",
"provider": "bedrock",
"response_format": "openai",
},
)
)
# The agentic loop should return a final answer (no tool_calls)
final_response = ModelResponse(
id="chatcmpl-bedrock-456",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
role="assistant",
content="Tomorrow's weather will be sunny, 25°C.",
),
)
],
model="us.anthropic.claude-sonnet-4-6",
object="chat.completion",
created=1234567891,
)
mock_callback.async_run_chat_completion_agentic_loop = AsyncMock(
return_value=final_response
)
# Patch litellm.callbacks so the handler finds our mock
original_callbacks = litellm.callbacks
litellm.callbacks = [mock_callback]
logging_obj = MagicMock()
logging_obj.dynamic_success_callbacks = []
try:
result = await handler._call_agentic_chat_completion_hooks(
response=mock_response,
model="us.anthropic.claude-sonnet-4-6",
messages=[{"role": "user", "content": "What is the weather tomorrow?"}],
optional_params={
"tools": [
{
"type": "function",
"function": {
"name": "litellm_web_search",
"description": "Search the web",
"parameters": {
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"],
},
},
}
]
},
logging_obj=logging_obj,
stream=False,
litellm_params={"custom_llm_provider": "bedrock"},
)
# Agentic hook should have returned the final response
assert result is not None
assert result.choices[0].finish_reason == "stop"
assert "sunny" in result.choices[0].message.content
# Verify the callback methods were called
mock_callback.async_should_run_chat_completion_agentic_loop.assert_called_once()
mock_callback.async_run_chat_completion_agentic_loop.assert_called_once()
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_bedrock_converse_agentic_hook_skips_when_no_tool_call():
"""Test that the agentic hook returns None when the model response
does not contain a litellm_web_search tool call (normal response)."""
from unittest.mock import MagicMock
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
from litellm.types.utils import Choices, Message
handler = BedrockConverseLLM()
# Build a normal response (no tool calls)
mock_response = ModelResponse(
id="chatcmpl-bedrock-789",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
role="assistant",
content="Hello! How can I help you?",
),
)
],
model="us.anthropic.claude-sonnet-4-6",
object="chat.completion",
created=1234567890,
)
# Use real websearch logger (not mock) so real detection logic runs
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.BEDROCK]
)
original_callbacks = litellm.callbacks
litellm.callbacks = [websearch_logger]
logging_obj = MagicMock()
logging_obj.dynamic_success_callbacks = []
try:
result = await handler._call_agentic_chat_completion_hooks(
response=mock_response,
model="us.anthropic.claude-sonnet-4-6",
messages=[{"role": "user", "content": "Hello"}],
optional_params={
"tools": [
{
"type": "function",
"function": {
"name": "litellm_web_search",
"description": "Search the web",
"parameters": {
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"],
},
},
}
]
},
logging_obj=logging_obj,
stream=False,
litellm_params={"custom_llm_provider": "bedrock"},
)
# No tool calls in response → hook should return None
assert result is None
finally:
litellm.callbacks = original_callbacks
if __name__ == "__main__":
# Run with: pytest test_websearch_chat_completion.py -v -s
pytest.main([__file__, "-v", "-s"])