diff --git a/mem0/llms/langchain.py b/mem0/llms/langchain.py index 3e722d609..9833cd5bd 100644 --- a/mem0/llms/langchain.py +++ b/mem0/llms/langchain.py @@ -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) diff --git a/tests/llms/test_langchain.py b/tests/llms/test_langchain.py index 11764a6e5..c0ab2701c 100644 --- a/tests/llms/test_langchain.py +++ b/tests/llms/test_langchain.py @@ -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")