Tool call support for LangchainLLM (#3542)

This commit is contained in:
Alex Kondratev
2025-10-05 21:33:44 +03:00
committed by GitHub
parent 8ba032e029
commit ee0202764b
2 changed files with 95 additions and 21 deletions
+50 -21
View File
@@ -5,6 +5,7 @@ from mem0.llms.base import LLMBase
try:
from langchain.chat_models.base import BaseChatModel
from langchain_core.messages import AIMessage
except ImportError:
raise ImportError("langchain is not installed. Please install it using `pip install langchain`")
@@ -21,6 +22,35 @@ class LangchainLLM(LLMBase):
self.langchain_model = self.config.model
def _parse_response(self, response: AIMessage, tools: Optional[List[Dict]]):
"""
Process the response based on whether tools are used or not.
Args:
response: AI Message.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if not tools:
return response.content
processed_response = {
"content": response.content,
"tool_calls": [],
}
for tool_call in response.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call["name"],
"arguments": tool_call["args"],
}
)
return processed_response
def generate_response(
self,
messages: List[Dict[str, str]],
@@ -34,32 +64,31 @@ class LangchainLLM(LLMBase):
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Not used in Langchain.
tools (list, optional): List of tools that the model can call. Not used in Langchain.
tool_choice (str, optional): Tool choice method. Not used in Langchain.
tools (list, optional): List of tools that the model can call.
tool_choice (str, optional): Tool choice method.
Returns:
str: The generated response.
"""
try:
# Convert the messages to LangChain's tuple format
langchain_messages = []
for message in messages:
role = message["role"]
content = message["content"]
# Convert the messages to LangChain's tuple format
langchain_messages = []
for message in messages:
role = message["role"]
content = message["content"]
if role == "system":
langchain_messages.append(("system", content))
elif role == "user":
langchain_messages.append(("human", content))
elif role == "assistant":
langchain_messages.append(("ai", content))
if role == "system":
langchain_messages.append(("system", content))
elif role == "user":
langchain_messages.append(("human", content))
elif role == "assistant":
langchain_messages.append(("ai", content))
if not langchain_messages:
raise ValueError("No valid messages found in the messages list")
if not langchain_messages:
raise ValueError("No valid messages found in the messages list")
ai_message = self.langchain_model.invoke(langchain_messages)
langchain_model = self.langchain_model
if tools:
langchain_model = langchain_model.bind_tools(tools=tools, tool_choice=tool_choice)
return ai_message.content
except Exception as e:
raise Exception(f"Error generating response using langchain model: {str(e)}")
response: AIMessage = langchain_model.invoke(langchain_messages)
return self._parse_response(response, tools)
+45
View File
@@ -68,6 +68,51 @@ def test_generate_response(mock_langchain_model):
assert response == "This is a test response"
def test_generate_response_with_tools(mock_langchain_model):
config = BaseLlmConfig(model=mock_langchain_model, temperature=0.7, max_tokens=100, api_key="test-api-key")
llm = LangchainLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_response = Mock()
mock_response.content = "I've added the memory for you."
mock_tool_call = Mock()
mock_tool_call.__getitem__ = Mock(
side_effect={"name": "add_memory", "args": {"data": "Today is a sunny day."}}.__getitem__
)
mock_response.tool_calls = [mock_tool_call]
mock_langchain_model.invoke.return_value = mock_response
mock_langchain_model.bind_tools.return_value = mock_langchain_model
response = llm.generate_response(messages, tools=tools)
mock_langchain_model.invoke.assert_called_once()
assert response["content"] == "I've added the memory for you."
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_invalid_model():
"""Test that LangchainLLM raises an error with an invalid model."""
config = BaseLlmConfig(model="not-a-valid-model-instance", temperature=0.7, max_tokens=100, api_key="test-api-key")