refactor: new agentic loop event hook

simplifies how to create logic for tool based multi llm calls
This commit is contained in:
Krrish Dholakia 2026-04-14 09:22:53 -07:00
parent 26c7412339
commit 9c20b8c743
40 changed files with 591 additions and 149 deletions

View file

@ -20,6 +20,7 @@ from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
from litellm.types.integrations.argilla import ArgillaItem
from litellm.types.llms.openai import AllMessageValues, ChatCompletionRequest
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.integrations.custom_logger import AgenticLoopPlan
from litellm.types.utils import (
AdapterCompletionStreamWrapper,
CallTypes,
@ -676,6 +677,26 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
pass
async def async_build_agentic_loop_plan(
self,
tools: Dict,
model: str,
messages: List[Dict],
response: Any,
anthropic_messages_provider_config: Any,
anthropic_messages_optional_request_params: Dict,
logging_obj: "LiteLLMLoggingObj",
stream: bool,
kwargs: Dict,
) -> AgenticLoopPlan:
"""
Build a typed rerun plan for Anthropic Messages agentic loops.
Override this method to separate callback decision/tool execution from
follow-up request execution (handled by BaseLLMHTTPHandler).
"""
return AgenticLoopPlan(run_agentic_loop=False)
async def async_should_run_chat_completion_agentic_loop(
self,
response: Any,
@ -707,6 +728,22 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
pass
async def async_build_chat_completion_agentic_loop_plan(
self,
tools: Dict,
model: str,
messages: List[Dict],
response: Any,
optional_params: Dict,
logging_obj: "LiteLLMLoggingObj",
stream: bool,
kwargs: Dict,
) -> AgenticLoopPlan:
"""
Build a typed rerun plan for chat-completions agentic loops.
"""
return AgenticLoopPlan(run_agentic_loop=False)
# Useful helpers for custom logger classes
def truncate_standard_logging_payload_content(

View file

@ -28,6 +28,10 @@ from litellm.integrations.websearch_interception.transformation import (
from litellm.types.integrations.websearch_interception import (
WebSearchInterceptionConfig,
)
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
@ -573,6 +577,35 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs=kwargs,
)
async def async_build_agentic_loop_plan(
self,
tools: Dict,
model: str,
messages: List[Dict],
response: Any,
anthropic_messages_provider_config: Any,
anthropic_messages_optional_request_params: Dict,
logging_obj: Any,
stream: bool,
kwargs: Dict,
) -> AgenticLoopPlan:
tool_calls = tools["tool_calls"]
thinking_blocks = tools.get("thinking_blocks", [])
request_patch = await self._build_anthropic_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
thinking_blocks=thinking_blocks,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
kwargs=kwargs,
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={"tool_type": "websearch", "response_format": "anthropic"},
)
async def async_run_chat_completion_agentic_loop(
self,
tools: Dict,
@ -608,6 +641,33 @@ class WebSearchInterceptionLogger(CustomLogger):
response_format=response_format,
)
async def async_build_chat_completion_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"]
response_format = tools.get("response_format", "openai")
request_patch = await self._build_chat_completion_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
optional_params=optional_params,
kwargs=kwargs,
response_format=response_format,
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={"tool_type": "websearch", "response_format": response_format},
)
@staticmethod
def _resolve_max_tokens(
optional_params: Dict,
@ -672,7 +732,48 @@ class WebSearchInterceptionLogger(CustomLogger):
stream: bool,
kwargs: Dict,
) -> Any:
"""Execute litellm.search() and make follow-up request"""
"""Legacy path: execute search + build patch + run follow-up call."""
request_patch = await self._build_anthropic_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
thinking_blocks=thinking_blocks,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
kwargs=kwargs,
)
if request_patch.messages is None:
raise ValueError("WebSearchInterception: missing follow-up messages")
optional_params = dict(anthropic_messages_optional_request_params)
optional_params.update(request_patch.optional_params)
max_tokens = request_patch.max_tokens
if max_tokens is None:
max_tokens = cast(Optional[int], optional_params.pop("max_tokens", None))
else:
optional_params.pop("max_tokens", None)
if max_tokens is None:
max_tokens = cast(int, kwargs.get("max_tokens", 1024))
return await anthropic_messages.acreate(
max_tokens=max_tokens,
messages=request_patch.messages,
model=request_patch.model or model,
**optional_params,
**request_patch.kwargs,
)
async def _build_anthropic_request_patch(
self,
model: str,
messages: List[Dict],
tool_calls: List[Dict],
thinking_blocks: List[Dict],
anthropic_messages_optional_request_params: Dict,
logging_obj: Any,
kwargs: Dict,
) -> AgenticLoopRequestPatch:
"""Execute litellm.search() and build follow-up request patch."""
# Extract search queries from tool_use blocks
search_tasks = []
@ -721,20 +822,8 @@ class WebSearchInterceptionLogger(CustomLogger):
thinking_blocks=thinking_blocks,
)
# Make follow-up request with search results
# Type cast: user_message is a Dict for Anthropic format (default response_format)
follow_up_messages = messages + [assistant_message, cast(Dict, user_message)]
verbose_logger.debug(
"WebSearchInterception: Making follow-up request with search results"
)
verbose_logger.debug(
f"WebSearchInterception: Follow-up messages count: {len(follow_up_messages)}"
)
verbose_logger.debug(
f"WebSearchInterception: Last message (tool_result): {user_message}"
)
# Correlation context for structured logging
_call_id = getattr(logging_obj, "litellm_call_id", None) or kwargs.get(
"litellm_call_id", "unknown"
@ -742,61 +831,39 @@ class WebSearchInterceptionLogger(CustomLogger):
full_model_name = model # safe default before try block
# Use anthropic_messages.acreate for follow-up request
try:
max_tokens = self._resolve_max_tokens(
anthropic_messages_optional_request_params, kwargs
)
max_tokens = self._resolve_max_tokens(
anthropic_messages_optional_request_params, kwargs
)
verbose_logger.debug(
f"WebSearchInterception: Using max_tokens={max_tokens} for follow-up request"
)
verbose_logger.debug(
f"WebSearchInterception: Using max_tokens={max_tokens} for follow-up request"
)
# Create a copy of optional params without max_tokens (since we pass it explicitly)
optional_params_without_max_tokens = {
k: v
for k, v in anthropic_messages_optional_request_params.items()
if k != "max_tokens"
}
optional_params_without_max_tokens = {
k: v
for k, v in anthropic_messages_optional_request_params.items()
if k != "max_tokens"
}
kwargs_for_followup = self._prepare_followup_kwargs(kwargs)
kwargs_for_followup = self._prepare_followup_kwargs(kwargs)
# Get model from logging_obj.model_call_details["agentic_loop_params"]
# This preserves the full model name with provider prefix (e.g., "bedrock/invoke/...")
if logging_obj is not None:
agentic_params = logging_obj.model_call_details.get(
"agentic_loop_params", {}
)
full_model_name = agentic_params.get("model", model)
verbose_logger.debug(
f"WebSearchInterception: Using model name: {full_model_name}"
)
final_response = await anthropic_messages.acreate(
max_tokens=max_tokens,
messages=follow_up_messages,
model=full_model_name,
**optional_params_without_max_tokens,
**kwargs_for_followup,
)
verbose_logger.debug(
f"WebSearchInterception: Follow-up request completed, response type: {type(final_response)}"
)
verbose_logger.debug(
f"WebSearchInterception: Final response: {final_response}"
)
return final_response
except Exception as e:
verbose_logger.exception(
"WebSearchInterception: Follow-up request failed "
"[call_id=%s model=%s messages=%d searches=%d]: %s",
_call_id,
full_model_name,
len(follow_up_messages),
len(final_search_results),
str(e),
)
raise
if logging_obj is not None:
agentic_params = logging_obj.model_call_details.get("agentic_loop_params", {})
full_model_name = agentic_params.get("model", model)
verbose_logger.debug(
"WebSearchInterception: Built anthropic request patch "
"[call_id=%s model=%s messages=%d searches=%d]",
_call_id,
full_model_name,
len(follow_up_messages),
len(final_search_results),
)
return AgenticLoopRequestPatch(
model=full_model_name,
messages=follow_up_messages,
max_tokens=max_tokens,
optional_params=optional_params_without_max_tokens,
kwargs=kwargs_for_followup,
)
async def _execute_search(self, query: str) -> str:
"""Execute a single web search using router's search tools"""
@ -883,7 +950,36 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs: Dict,
response_format: str = "openai",
) -> Any:
"""Execute litellm.search() and make follow-up chat completion request"""
"""Legacy path: execute search + build patch + run follow-up call."""
request_patch = await self._build_chat_completion_request_patch(
model=model,
messages=messages,
tool_calls=tool_calls,
optional_params=optional_params,
kwargs=kwargs,
response_format=response_format,
)
if request_patch.messages is None:
raise ValueError("WebSearchInterception: missing follow-up messages")
params = dict(optional_params)
params.update(request_patch.optional_params)
return await litellm.acompletion(
model=request_patch.model or model,
messages=request_patch.messages,
**params,
**request_patch.kwargs,
)
async def _build_chat_completion_request_patch( # noqa: PLR0915
self,
model: str,
messages: List[Dict],
tool_calls: List[Dict],
optional_params: Dict,
kwargs: Dict,
response_format: str = "openai",
) -> AgenticLoopRequestPatch:
"""Execute litellm.search() and build chat-completion rerun patch."""
# Extract search queries from tool_calls
search_tasks = []
@ -963,74 +1059,56 @@ class WebSearchInterceptionLogger(CustomLogger):
f"WebSearchInterception: Follow-up messages count: {len(follow_up_messages)}"
)
# Use litellm.acompletion for follow-up request
try:
# Remove internal parameters that shouldn't be passed to follow-up request
internal_params = {
"_websearch_interception",
"acompletion",
"litellm_logging_obj",
"custom_llm_provider",
# Remove internal parameters that shouldn't be passed to follow-up request
internal_params = {
"_websearch_interception",
"acompletion",
"litellm_logging_obj",
"custom_llm_provider",
"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 internal_params
}
full_model_name = model
if "custom_llm_provider" in kwargs:
custom_llm_provider = kwargs["custom_llm_provider"]
if not model.startswith(custom_llm_provider) and "/" not in model:
full_model_name = f"{custom_llm_provider}/{model}"
verbose_logger.debug(
"WebSearchInterception: Built chat completion request patch model=%s messages=%d",
full_model_name,
len(follow_up_messages),
)
tools_param = optional_params.get("tools")
optional_params_clean = {
k: v
for k, v in optional_params.items()
if k
not in {
"tools",
"extra_body",
"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 internal_params
}
}
if tools_param is not None:
optional_params_clean["tools"] = tools_param
# Get full model name from kwargs
full_model_name = model
if "custom_llm_provider" in kwargs:
custom_llm_provider = kwargs["custom_llm_provider"]
# Reconstruct full model name with provider prefix if needed
if not model.startswith(custom_llm_provider):
# Check if model already has a provider prefix
if "/" not in model:
full_model_name = f"{custom_llm_provider}/{model}"
verbose_logger.debug(
f"WebSearchInterception: Using model name: {full_model_name}"
)
# Prepare tools for follow-up request (same as original)
tools_param = optional_params.get("tools")
# Remove tools and extra_body from optional_params to avoid issues
# extra_body often contains internal LiteLLM params that shouldn't be forwarded
optional_params_clean = {
k: v
for k, v in optional_params.items()
if k
not in {
"tools",
"extra_body",
"model_alias_map",
"stream_response",
"custom_prompt_dict",
}
}
final_response = await litellm.acompletion(
model=full_model_name,
messages=follow_up_messages,
tools=tools_param,
**optional_params_clean,
**kwargs_for_followup,
)
verbose_logger.debug(
f"WebSearchInterception: Follow-up request completed, response type: {type(final_response)}"
)
return final_response
except Exception as e:
verbose_logger.exception(
f"WebSearchInterception: Follow-up request failed: {str(e)}"
)
raise
return AgenticLoopRequestPatch(
model=full_model_name,
messages=follow_up_messages,
optional_params=optional_params_clean,
kwargs=kwargs_for_followup,
)
async def _create_empty_search_result(self) -> str:
"""Create an empty search result for tool calls without queries"""

View file

@ -78,6 +78,10 @@ from litellm.types.containers.main import (
DeleteContainerResult,
)
from litellm.types.files import TwoStepFileUploadConfig
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -4449,6 +4453,129 @@ class BaseLLMHTTPHandler:
return stream, data
return stream, data
@staticmethod
def _get_agentic_loop_settings(kwargs: Dict) -> Tuple[int, int, List[str]]:
depth = int(kwargs.get("_agentic_loop_depth", 0) or 0)
max_loops = int(kwargs.get("max_agentic_loops", 3) or 3)
fingerprints = list(kwargs.get("_agentic_loop_fingerprints", []) or [])
return depth, max(max_loops, 1), fingerprints
@staticmethod
def _fingerprint_agentic_tools(tools: Dict) -> str:
try:
return json.dumps(tools, sort_keys=True, default=str)
except Exception:
return str(tools)
async def _execute_anthropic_agentic_plan(
self,
plan: AgenticLoopPlan,
model: str,
messages: List[Dict],
anthropic_messages_optional_request_params: Dict,
logging_obj: "LiteLLMLoggingObj",
kwargs: Dict,
depth: int,
max_loops: int,
fingerprints: List[str],
fingerprint: str,
) -> Any:
from litellm.anthropic_interface import messages as anthropic_messages
patch = plan.request_patch or AgenticLoopRequestPatch()
if patch.messages is None:
raise ValueError("Agentic loop plan missing patched messages")
full_model_name = model
if logging_obj is not None:
agentic_params = logging_obj.model_call_details.get("agentic_loop_params", {})
full_model_name = cast(str, agentic_params.get("model", model))
optional_params = dict(anthropic_messages_optional_request_params)
optional_params.update(patch.optional_params)
if patch.tools is not None:
optional_params["tools"] = patch.tools
max_tokens = patch.max_tokens
if max_tokens is None:
max_tokens = cast(Optional[int], optional_params.pop("max_tokens", None))
else:
optional_params.pop("max_tokens", None)
if max_tokens is None:
max_tokens = cast(int, kwargs.get("max_tokens", 1024))
internal_keys = {"litellm_logging_obj"}
kwargs_for_followup = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception") and k not in internal_keys
}
kwargs_for_followup.update(patch.kwargs)
kwargs_for_followup["_agentic_loop_depth"] = depth + 1
kwargs_for_followup["max_agentic_loops"] = max_loops
kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint]
return await anthropic_messages.acreate(
max_tokens=max_tokens,
messages=patch.messages,
model=patch.model or full_model_name,
**optional_params,
**kwargs_for_followup,
)
async def _execute_chat_completion_agentic_plan(
self,
plan: AgenticLoopPlan,
model: str,
messages: List[Dict],
optional_params: Dict,
kwargs: Dict,
custom_llm_provider: str,
depth: int,
max_loops: int,
fingerprints: List[str],
fingerprint: str,
) -> Any:
patch = plan.request_patch or AgenticLoopRequestPatch()
if patch.messages is None:
raise ValueError("Agentic loop plan missing patched messages")
full_model_name = patch.model or model
if "/" not in full_model_name:
full_model_name = f"{custom_llm_provider}/{full_model_name}"
optional_params_for_followup = dict(optional_params)
optional_params_for_followup.update(patch.optional_params)
if patch.tools is not None:
optional_params_for_followup["tools"] = patch.tools
internal_params = {
"_websearch_interception",
"acompletion",
"litellm_logging_obj",
"custom_llm_provider",
"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 internal_params
}
kwargs_for_followup.update(patch.kwargs)
kwargs_for_followup["_agentic_loop_depth"] = depth + 1
kwargs_for_followup["max_agentic_loops"] = max_loops
kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint]
return await litellm.acompletion(
model=full_model_name,
messages=patch.messages,
**optional_params_for_followup,
**kwargs_for_followup,
)
async def _call_agentic_completion_hooks(
self,
response: Any,
@ -4474,6 +4601,7 @@ class BaseLLMHTTPHandler:
callbacks = litellm.callbacks + (logging_obj.dynamic_success_callbacks or [])
tools = anthropic_messages_optional_request_params.get("tools", [])
depth, max_loops, fingerprints = self._get_agentic_loop_settings(kwargs=kwargs)
for callback in callbacks:
try:
@ -4493,13 +4621,36 @@ class BaseLLMHTTPHandler:
)
if should_run:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
fingerprint = self._fingerprint_agentic_tools(tool_calls)
if fingerprint in fingerprints:
raise ValueError(
"Agentic loop detected repeated tool-call fingerprint; aborting rerun"
)
if depth >= max_loops:
raise ValueError(
f"Exceeded max_agentic_loops={max_loops} for model={model}"
)
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
agentic_response = await callback.async_run_agentic_loop(
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
)
if not build_plan_overridden:
return await callback.async_run_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
)
plan = await callback.async_build_agentic_loop_plan(
tools=tool_calls,
model=model,
messages=messages,
@ -4510,8 +4661,31 @@ class BaseLLMHTTPHandler:
stream=stream,
kwargs=kwargs_with_provider,
)
# First hook that runs agentic loop wins
return agentic_response
if plan.response_override is not None:
return plan.response_override
if plan.terminate:
verbose_logger.debug(
"Agentic loop terminated by callback=%s reason=%s",
callback.__class__.__name__,
plan.stop_reason,
)
return response
if not plan.run_agentic_loop:
continue
return await self._execute_anthropic_agentic_plan(
plan=plan,
model=model,
messages=messages,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
kwargs=kwargs_with_provider,
depth=depth,
max_loops=max_loops,
fingerprints=fingerprints,
fingerprint=fingerprint,
)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
@ -4586,6 +4760,7 @@ class BaseLLMHTTPHandler:
callbacks = litellm.callbacks + (logging_obj.dynamic_success_callbacks or [])
tools = optional_params.get("tools", [])
depth, max_loops, fingerprints = self._get_agentic_loop_settings(kwargs=kwargs)
for callback in callbacks:
try:
@ -4611,14 +4786,24 @@ class BaseLLMHTTPHandler:
)
if should_run:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
fingerprint = self._fingerprint_agentic_tools(tool_calls)
if fingerprint in fingerprints:
raise ValueError(
"Agentic loop detected repeated tool-call fingerprint; aborting rerun"
)
if depth >= max_loops:
raise ValueError(
f"Exceeded max_agentic_loops={max_loops} for model={model}"
)
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
agentic_response = (
await callback.async_run_chat_completion_agentic_loop(
kwargs_with_provider["custom_llm_provider"] = custom_llm_provider
build_plan_overridden = (
callback.__class__.async_build_chat_completion_agentic_loop_plan
is not CustomLogger.async_build_chat_completion_agentic_loop_plan
)
if not build_plan_overridden:
return await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
@ -4628,9 +4813,42 @@ class BaseLLMHTTPHandler:
stream=stream,
kwargs=kwargs_with_provider,
)
plan = await callback.async_build_chat_completion_agentic_loop_plan(
tools=tool_calls,
model=model,
messages=messages,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
)
if plan.response_override is not None:
return plan.response_override
if plan.terminate:
verbose_logger.debug(
"Agentic chat loop terminated by callback=%s reason=%s",
callback.__class__.__name__,
plan.stop_reason,
)
return response
if not plan.run_agentic_loop:
continue
return await self._execute_chat_completion_agentic_plan(
plan=plan,
model=model,
messages=messages,
optional_params=optional_params,
kwargs=kwargs_with_provider,
custom_llm_provider=custom_llm_provider,
depth=depth,
max_loops=max_loops,
fingerprints=fingerprints,
fingerprint=fingerprint,
)
# First hook that runs agentic loop wins
return agentic_response
except Exception as e:
verbose_logger.exception(

View file

@ -1,6 +1,6 @@
from typing import Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel
from pydantic import BaseModel, Field
class StandardCustomLoggerInitParams(BaseModel):
@ -9,3 +9,29 @@ class StandardCustomLoggerInitParams(BaseModel):
"""
turn_off_message_logging: Optional[bool] = False
class AgenticLoopRequestPatch(BaseModel):
"""
Patch returned by callbacks to request a follow-up LLM call.
"""
model: Optional[str] = None
messages: Optional[List[Dict[str, Any]]] = None
tools: Optional[List[Dict[str, Any]]] = None
max_tokens: Optional[int] = None
optional_params: Dict[str, Any] = Field(default_factory=dict)
kwargs: Dict[str, Any] = Field(default_factory=dict)
class AgenticLoopPlan(BaseModel):
"""
Typed callback response for agentic-loop reruns.
"""
run_agentic_loop: bool = False
request_patch: Optional[AgenticLoopRequestPatch] = None
response_override: Optional[Any] = None
terminate: bool = False
stop_reason: Optional[str] = None
metadata: Dict[str, Any] = Field(default_factory=dict)

View file

@ -4,7 +4,7 @@ Unit tests for WebSearch Interception Handler
Tests the WebSearchInterceptionLogger class and helper functions.
"""
from unittest.mock import MagicMock, Mock
from unittest.mock import AsyncMock, MagicMock, Mock
import pytest
@ -69,6 +69,61 @@ async def test_async_should_run_agentic_loop():
assert tools_dict == {}
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_returns_request_patch():
"""Callback should return a typed patch for base handler reruns."""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
logger._execute_search = AsyncMock( # type: ignore
return_value="Title: LiteLLM\nURL: docs\nSnippet: test"
)
tools_dict = {
"tool_calls": [
{
"id": "toolu_123",
"type": "tool_use",
"name": "litellm_web_search",
"input": {"query": "what is litellm"},
}
],
"response_format": "anthropic",
}
logging_obj = MagicMock()
logging_obj.model_call_details = {
"agentic_loop_params": {"model": "bedrock/invoke/claude-3-5-sonnet"}
}
kwargs = {
"temperature": 0.2,
"_websearch_interception_converted_stream": True,
"litellm_logging_obj": object(),
}
plan = await logger.async_build_agentic_loop_plan(
tools=tools_dict,
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "search LiteLLM"}],
response=None,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={
"max_tokens": 1024,
"tools": [{"name": "litellm_web_search"}],
},
logging_obj=logging_obj,
stream=False,
kwargs=kwargs,
)
assert plan.run_agentic_loop is True
assert plan.request_patch is not None
assert plan.request_patch.model == "bedrock/invoke/claude-3-5-sonnet"
assert plan.request_patch.max_tokens == 1024
assert plan.request_patch.messages is not None
assert len(plan.request_patch.messages) == 3
assert "_websearch_interception_converted_stream" not in plan.request_patch.kwargs
assert "litellm_logging_obj" not in plan.request_patch.kwargs
assert plan.request_patch.kwargs["temperature"] == 0.2
@pytest.mark.asyncio
async def test_internal_flags_filtered_from_followup_kwargs():
"""Test that internal _websearch_interception flags are filtered from follow-up request kwargs.

View file

@ -74,6 +74,34 @@ def test_prepare_fake_stream_request():
assert result_data["messages"] == [{"role": "user", "content": "Hello"}]
def test_get_agentic_loop_settings_defaults_and_overrides():
handler = BaseLLMHTTPHandler()
depth, max_loops, fingerprints = handler._get_agentic_loop_settings(kwargs={})
assert depth == 0
assert max_loops == 3
assert fingerprints == []
depth, max_loops, fingerprints = handler._get_agentic_loop_settings(
kwargs={
"_agentic_loop_depth": 2,
"max_agentic_loops": 7,
"_agentic_loop_fingerprints": ["fp-1", "fp-2"],
}
)
assert depth == 2
assert max_loops == 7
assert fingerprints == ["fp-1", "fp-2"]
def test_fingerprint_agentic_tools_is_deterministic():
handler = BaseLLMHTTPHandler()
tools_a = {"tool_calls": [{"id": "1", "input": {"q": "abc"}, "name": "web_search"}]}
tools_b = {"tool_calls": [{"name": "web_search", "input": {"q": "abc"}, "id": "1"}]}
assert handler._fingerprint_agentic_tools(tools_a) == handler._fingerprint_agentic_tools(tools_b)
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_extra_headers():
"""