Tool call support for LangchainLLM (#3542)
This commit is contained in:
+50
-21
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user