mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(bedrock): trigger websearch_interception agentic loop for /chat/completions
Signed-off-by: Nishchay Veer <nishchayveer19@gmail.com>
This commit is contained in:
parent
d251238bd7
commit
b882bec39a
2 changed files with 347 additions and 42 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue