diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 947e63530..5adc2c5e1 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -3464,7 +3464,7 @@ class AsyncMemory(MemoryBase): if llm is not None: parsed_messages = convert_to_messages(parsed_messages) response = await asyncio.to_thread(llm.invoke, input=parsed_messages) - procedural_memory = response.content + procedural_memory = remove_code_blocks(response.content) else: procedural_memory = await asyncio.to_thread(self.llm.generate_response, messages=parsed_messages) procedural_memory = remove_code_blocks(procedural_memory) diff --git a/tests/test_memory.py b/tests/test_memory.py index 5c277f50b..551acef5a 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -1496,3 +1496,36 @@ class TestAsyncDeleteAllEntityRace: mock_entity_store.delete.assert_called_once_with(vector_id="entity-alice") assert mock_vector_store.delete.call_count == 2 + + +@pytest.mark.asyncio +@patch("mem0.memory.main.VectorStoreFactory") +@patch("mem0.memory.main.EmbedderFactory") +@patch("mem0.memory.main.LlmFactory") +async def test_async_procedural_memory_langchain_strips_code_blocks(mock_llm_factory, mock_emb, mock_vs): + """Regression #5710: async LangChain path must call remove_code_blocks().""" + mock_vs.return_value = MagicMock() + mock_emb.return_value = MagicMock() + mock_emb.return_value.embed.return_value = [0.1] * 1536 + mock_llm_factory.return_value = MagicMock() + + from mem0.memory.main import AsyncMemory + + config = MemoryConfig() + memory = AsyncMemory(config) + memory.vector_store = MagicMock() + memory.vector_store.insert = MagicMock() + + mock_langchain_llm = MagicMock() + mock_response = MagicMock() + mock_response.content = '```json\n{"key": "value"}\n```' + mock_langchain_llm.invoke.return_value = mock_response + + messages = [{"role": "user", "content": "test"}] + metadata = {"user_id": "test_user"} + + await memory._create_procedural_memory(messages, metadata=metadata, llm=mock_langchain_llm) + + insert_call = memory.vector_store.insert.call_args + stored_data = insert_call[1]["payloads"][0]["data"] + assert "```" not in stored_data