mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Add callback for websearch in completion method
This commit is contained in:
parent
6207bf8f68
commit
88778a871d
2 changed files with 214 additions and 6 deletions
|
|
@ -302,7 +302,7 @@ class BaseLLMHTTPHandler:
|
|||
logging_obj=logging_obj,
|
||||
signed_json_body=signed_json_body,
|
||||
)
|
||||
return provider_config.transform_response(
|
||||
initial_response = provider_config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
|
|
@ -316,6 +316,20 @@ class BaseLLMHTTPHandler:
|
|||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
# Call agentic chat completion hooks
|
||||
final_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,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=litellm_params,
|
||||
)
|
||||
|
||||
return final_response if final_response is not None else initial_response
|
||||
|
||||
def completion(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -412,6 +426,11 @@ class BaseLLMHTTPHandler:
|
|||
},
|
||||
)
|
||||
|
||||
# Check if stream was converted for WebSearch interception
|
||||
# This is set by the async_pre_request_hook in WebSearchInterceptionLogger
|
||||
if litellm_params.get("_websearch_interception_converted_stream", False):
|
||||
logging_obj.model_call_details["websearch_interception_converted_stream"] = True
|
||||
|
||||
if acompletion is True:
|
||||
if stream is True:
|
||||
data = self._add_stream_param_to_request_body(
|
||||
|
|
@ -419,7 +438,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=provider_config,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
return self.acompletion_stream_function(
|
||||
response = self.acompletion_stream_function(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
|
|
@ -4361,10 +4380,10 @@ class BaseLLMHTTPHandler:
|
|||
kwargs: Dict,
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Call agentic completion hooks for all custom loggers.
|
||||
Call agentic completion hooks for all custom loggers (Anthropic Messages API).
|
||||
|
||||
1. Call async_should_run_agentic_completion to check if agentic loop is needed
|
||||
2. If yes, call async_run_agentic_completion to execute the loop
|
||||
1. Call async_should_run_agentic_loop to check if agentic loop is needed
|
||||
2. If yes, call async_run_agentic_loop to execute the loop
|
||||
|
||||
Returns the response from agentic loop, or None if no hook runs.
|
||||
"""
|
||||
|
|
@ -4453,6 +4472,105 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
return None
|
||||
|
||||
async def _call_agentic_chat_completion_hooks(
|
||||
self,
|
||||
response: Any,
|
||||
model: str,
|
||||
messages: List[Dict],
|
||||
optional_params: Dict,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
stream: bool,
|
||||
custom_llm_provider: str,
|
||||
kwargs: Dict,
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Call agentic chat completion hooks for all custom loggers (Chat Completions API).
|
||||
|
||||
1. Call async_should_run_chat_completion_agentic_loop to check if agentic loop is needed
|
||||
2. If yes, call async_run_chat_completion_agentic_loop to execute the loop
|
||||
|
||||
Returns the response from agentic loop, or None if no hook runs.
|
||||
"""
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
callbacks = litellm.callbacks + (
|
||||
logging_obj.dynamic_success_callbacks or []
|
||||
)
|
||||
tools = optional_params.get("tools", [])
|
||||
|
||||
for callback in callbacks:
|
||||
try:
|
||||
if isinstance(callback, CustomLogger):
|
||||
# Check if callback has the chat completion agentic loop method
|
||||
if not hasattr(callback, "async_should_run_chat_completion_agentic_loop"):
|
||||
continue
|
||||
|
||||
# First: Check if agentic loop should run
|
||||
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=kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
if should_run:
|
||||
# Second: Execute agentic loop
|
||||
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
|
||||
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(
|
||||
tools=tool_calls,
|
||||
model=model,
|
||||
messages=messages,
|
||||
response=response,
|
||||
optional_params=optional_params,
|
||||
logging_obj=logging_obj,
|
||||
stream=stream,
|
||||
kwargs=kwargs_with_provider,
|
||||
)
|
||||
# First hook that runs agentic loop wins
|
||||
return agentic_response
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: {str(e)}"
|
||||
)
|
||||
|
||||
# Check if we need to convert response to fake stream for chat completions
|
||||
# This happens when:
|
||||
# 1. Stream was originally True but converted to False for WebSearch interception
|
||||
# 2. No agentic loop ran (LLM didn't use the tool)
|
||||
# 3. We have a non-streaming response that needs to be converted to streaming
|
||||
websearch_converted_stream = (
|
||||
logging_obj.model_call_details.get("websearch_interception_converted_stream", False)
|
||||
if logging_obj is not None
|
||||
else False
|
||||
)
|
||||
|
||||
if websearch_converted_stream:
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.base_model_iterator import (
|
||||
convert_model_response_to_streaming,
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
"WebSearchInterception: No tool call made, converting non-streaming chat completion to fake stream"
|
||||
)
|
||||
|
||||
# Convert the non-streaming ModelResponse to a fake stream
|
||||
if hasattr(response, "choices"):
|
||||
# Use the existing converter for ModelResponse
|
||||
fake_stream = convert_model_response_to_streaming(response)
|
||||
return fake_stream
|
||||
|
||||
return None
|
||||
|
||||
def _handle_error(
|
||||
self,
|
||||
e: Exception,
|
||||
|
|
|
|||
|
|
@ -501,6 +501,82 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
else:
|
||||
raise e
|
||||
|
||||
async def _call_agentic_completion_hooks_openai(
|
||||
self,
|
||||
response: Any,
|
||||
model: str,
|
||||
messages: List[Dict],
|
||||
optional_params: Dict,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
stream: bool,
|
||||
litellm_params: Dict,
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Call agentic completion hooks for all custom loggers (OpenAI Chat Completions API).
|
||||
|
||||
1. Call async_should_run_chat_completion_agentic_loop to check if agentic loop is needed
|
||||
2. If yes, call async_run_chat_completion_agentic_loop to execute the loop
|
||||
|
||||
Returns the response from agentic loop, or None if no hook runs.
|
||||
"""
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
callbacks = litellm.callbacks + (
|
||||
logging_obj.dynamic_success_callbacks or []
|
||||
)
|
||||
print(f"🔥callbacks: {callbacks}")
|
||||
tools = optional_params.get("tools", [])
|
||||
print(f"🔥tools: {tools}")
|
||||
# Get custom_llm_provider from litellm_params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider", "openai")
|
||||
|
||||
for callback in callbacks:
|
||||
try:
|
||||
if isinstance(callback, CustomLogger):
|
||||
# Check if the callback has the chat completion agentic loop methods
|
||||
if not hasattr(callback, 'async_should_run_chat_completion_agentic_loop'):
|
||||
continue
|
||||
|
||||
# First: Check if agentic loop should run (using chat completion method)
|
||||
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:
|
||||
# Second: Execute agentic loop
|
||||
kwargs_with_provider = litellm_params.copy() if litellm_params else {}
|
||||
kwargs_with_provider["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
# For OpenAI Chat Completions, use the chat completion agentic loop method
|
||||
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,
|
||||
)
|
||||
# First hook that runs agentic loop wins
|
||||
return agentic_response
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"LiteLLM.AgenticHookError: Exception in agentic completion hooks for OpenAI: {str(e)}"
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def mock_streaming(
|
||||
self,
|
||||
response: ModelResponse,
|
||||
|
|
@ -844,7 +920,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
logging_obj=logging_obj,
|
||||
)
|
||||
stringified_response = response.model_dump()
|
||||
|
||||
print(f"🔥stringified_response: {stringified_response}")
|
||||
logging_obj.post_call(
|
||||
input=data["messages"],
|
||||
api_key=api_key,
|
||||
|
|
@ -859,6 +935,20 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
_response_headers=headers,
|
||||
)
|
||||
|
||||
# Call agentic completion hooks (e.g., for websearch_interception)
|
||||
agentic_response = await self._call_agentic_completion_hooks_openai(
|
||||
response=final_response_obj,
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
logging_obj=logging_obj,
|
||||
stream=False,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
if agentic_response is not None:
|
||||
final_response_obj = agentic_response
|
||||
|
||||
if fake_stream is True:
|
||||
return self.mock_streaming(
|
||||
response=cast(ModelResponse, final_response_obj),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue