Feat/llm monitoring callback (#2877)
This commit is contained in:
@@ -119,6 +119,7 @@ Here's a comprehensive list of all parameters that can be used across different
|
||||
| `seed` | Seed for deterministic sampling | Sarvam |
|
||||
| `stop` | Stop sequences (max 4) | Sarvam |
|
||||
| `lmstudio_base_url` | Base URL for LM Studio API | LM Studio |
|
||||
| `response_callback` | LLM response callback function | OpenAI |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Provider |
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import List, Optional
|
||||
from typing import Any, Callable, List, Optional
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
|
||||
@@ -28,6 +28,8 @@ class OpenAIConfig(BaseLlmConfig):
|
||||
openrouter_base_url: Optional[str] = None,
|
||||
site_url: Optional[str] = None,
|
||||
app_name: Optional[str] = None,
|
||||
# Response monitoring callback
|
||||
response_callback: Optional[Callable[[Any, dict, dict], None]] = None,
|
||||
):
|
||||
"""
|
||||
Initialize OpenAI configuration.
|
||||
@@ -48,6 +50,7 @@ class OpenAIConfig(BaseLlmConfig):
|
||||
openrouter_base_url: OpenRouter base URL, defaults to None
|
||||
site_url: Site URL for OpenRouter, defaults to None
|
||||
app_name: Application name for OpenRouter, defaults to None
|
||||
response_callback: Optional callback for monitoring LLM responses.
|
||||
"""
|
||||
# Initialize base parameters
|
||||
super().__init__(
|
||||
@@ -69,3 +72,5 @@ class OpenAIConfig(BaseLlmConfig):
|
||||
self.openrouter_base_url = openrouter_base_url
|
||||
self.site_url = site_url
|
||||
self.app_name = app_name
|
||||
# Response monitoring
|
||||
self.response_callback = response_callback
|
||||
|
||||
+10
-2
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
@@ -130,6 +131,13 @@ class OpenAILLM(LLMBase):
|
||||
if tools: # TODO: Remove tools if no issues found with new memory addition logic
|
||||
params["tools"] = tools
|
||||
params["tool_choice"] = tool_choice
|
||||
|
||||
response = self.client.chat.completions.create(**params)
|
||||
return self._parse_response(response, tools)
|
||||
parsed_response = self._parse_response(response, tools)
|
||||
if self.config.response_callback:
|
||||
try:
|
||||
self.config.response_callback(self, response, params)
|
||||
except Exception as e:
|
||||
# Log error but don't propagate
|
||||
logging.error(f"Error due to callback: {e}")
|
||||
pass
|
||||
return parsed_response
|
||||
|
||||
@@ -104,3 +104,106 @@ def test_generate_response_with_tools(mock_openai_client):
|
||||
assert len(response["tool_calls"]) == 1
|
||||
assert response["tool_calls"][0]["name"] == "add_memory"
|
||||
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
|
||||
|
||||
|
||||
def test_response_callback_invocation(mock_openai_client):
|
||||
# Setup mock callback
|
||||
mock_callback = Mock()
|
||||
|
||||
config = OpenAIConfig(model="gpt-4o", response_callback=mock_callback)
|
||||
llm = OpenAILLM(config)
|
||||
messages = [{"role": "user", "content": "Test callback"}]
|
||||
|
||||
# Mock response
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [Mock(message=Mock(content="Response"))]
|
||||
mock_openai_client.chat.completions.create.return_value = mock_response
|
||||
|
||||
# Call method
|
||||
llm.generate_response(messages)
|
||||
|
||||
# Verify callback called with correct arguments
|
||||
mock_callback.assert_called_once()
|
||||
args = mock_callback.call_args[0]
|
||||
assert args[0] is llm # llm_instance
|
||||
assert args[1] == mock_response # raw_response
|
||||
assert "messages" in args[2] # params
|
||||
|
||||
|
||||
def test_no_response_callback(mock_openai_client):
|
||||
config = OpenAIConfig(model="gpt-4o")
|
||||
llm = OpenAILLM(config)
|
||||
messages = [{"role": "user", "content": "Test no callback"}]
|
||||
|
||||
# Mock response
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [Mock(message=Mock(content="Response"))]
|
||||
mock_openai_client.chat.completions.create.return_value = mock_response
|
||||
|
||||
# Should complete without calling any callback
|
||||
response = llm.generate_response(messages)
|
||||
assert response == "Response"
|
||||
|
||||
# Verify no callback is set
|
||||
assert llm.config.response_callback is None
|
||||
|
||||
|
||||
def test_callback_exception_handling(mock_openai_client):
|
||||
# Callback that raises exception
|
||||
def faulty_callback(*args):
|
||||
raise ValueError("Callback error")
|
||||
|
||||
config = OpenAIConfig(model="gpt-4o", response_callback=faulty_callback)
|
||||
llm = OpenAILLM(config)
|
||||
messages = [{"role": "user", "content": "Test exception"}]
|
||||
|
||||
# Mock response
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [Mock(message=Mock(content="Expected response"))]
|
||||
mock_openai_client.chat.completions.create.return_value = mock_response
|
||||
|
||||
# Should complete without raising
|
||||
response = llm.generate_response(messages)
|
||||
assert response == "Expected response"
|
||||
|
||||
# Verify callback was called (even though it raised an exception)
|
||||
assert llm.config.response_callback is faulty_callback
|
||||
|
||||
|
||||
def test_callback_with_tools(mock_openai_client):
|
||||
mock_callback = Mock()
|
||||
config = OpenAIConfig(model="gpt-4o", response_callback=mock_callback)
|
||||
llm = OpenAILLM(config)
|
||||
messages = [{"role": "user", "content": "Test tools"}]
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "test_tool",
|
||||
"description": "A test tool",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"param1": {"type": "string"}},
|
||||
"required": ["param1"],
|
||||
},
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
# Mock tool response
|
||||
mock_response = Mock()
|
||||
mock_message = Mock()
|
||||
mock_message.content = "Tool response"
|
||||
mock_tool_call = Mock()
|
||||
mock_tool_call.function.name = "test_tool"
|
||||
mock_tool_call.function.arguments = '{"param1": "value1"}'
|
||||
mock_message.tool_calls = [mock_tool_call]
|
||||
mock_response.choices = [Mock(message=mock_message)]
|
||||
mock_openai_client.chat.completions.create.return_value = mock_response
|
||||
|
||||
llm.generate_response(messages, tools=tools)
|
||||
|
||||
# Verify callback called with tool response
|
||||
mock_callback.assert_called_once()
|
||||
# Check that tool_calls exists in the message
|
||||
assert hasattr(mock_callback.call_args[0][1].choices[0].message, 'tool_calls')
|
||||
|
||||
Reference in New Issue
Block a user