diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 84b760790..175a56e08 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -479,7 +479,8 @@ class Memory(MemoryBase): msg_content = message_dict["content"] msg_embeddings = self.embedding_model.embed(msg_content, "add") - mem_id = self._create_memory(msg_content, msg_embeddings, per_msg_meta) + # Pass embeddings as a dict so _create_memory can reuse the cached embedding + mem_id = self._create_memory(msg_content, {msg_content: msg_embeddings}, per_msg_meta) returned_memories.append( { @@ -608,6 +609,9 @@ class Memory(MemoryBase): event_type = resp.get("event") if event_type == "ADD": + # Ensure action_text has an embedding cached to avoid redundant API calls + if action_text not in new_message_embeddings: + new_message_embeddings[action_text] = self.embedding_model.embed(action_text, "add") memory_id = self._create_memory( data=action_text, existing_embeddings=new_message_embeddings, @@ -615,6 +619,9 @@ class Memory(MemoryBase): ) returned_memories.append({"id": memory_id, "memory": action_text, "event": event_type}) elif event_type == "UPDATE": + # Ensure action_text has an embedding cached to avoid redundant API calls + if action_text not in new_message_embeddings: + new_message_embeddings[action_text] = self.embedding_model.embed(action_text, "update") self._update_memory( memory_id=temp_uuid_mapping[resp.get("id")], data=action_text, @@ -1155,8 +1162,11 @@ class Memory(MemoryBase): def _create_memory(self, data, existing_embeddings, metadata=None): logger.debug(f"Creating memory with {data=}") - if data in existing_embeddings: + # existing_embeddings may be a dict (preferred) or a precomputed vector + if isinstance(existing_embeddings, dict) and data in existing_embeddings: embeddings = existing_embeddings[data] + elif isinstance(existing_embeddings, list): + embeddings = existing_embeddings else: embeddings = self.embedding_model.embed(data, memory_action="add") memory_id = str(uuid.uuid4()) @@ -1253,8 +1263,10 @@ class Memory(MemoryBase): if "role" not in new_metadata and "role" in existing_memory.payload: new_metadata["role"] = existing_memory.payload["role"] - if data in existing_embeddings: + if isinstance(existing_embeddings, dict) and data in existing_embeddings: embeddings = existing_embeddings[data] + elif isinstance(existing_embeddings, list): + embeddings = existing_embeddings else: embeddings = self.embedding_model.embed(data, "update") @@ -1523,7 +1535,8 @@ class AsyncMemory(MemoryBase): msg_content = message_dict["content"] msg_embeddings = await asyncio.to_thread(self.embedding_model.embed, msg_content, "add") - mem_id = await self._create_memory(msg_content, msg_embeddings, per_msg_meta) + # Pass embeddings as a dict so _create_memory can reuse the cached embedding + mem_id = await self._create_memory(msg_content, {msg_content: msg_embeddings}, per_msg_meta) returned_memories.append( { @@ -1651,6 +1664,11 @@ class AsyncMemory(MemoryBase): event_type = resp.get("event") if event_type == "ADD": + # Ensure action_text has an embedding cached to avoid redundant API calls + if action_text not in new_message_embeddings: + new_message_embeddings[action_text] = await asyncio.to_thread( + self.embedding_model.embed, action_text, "add" + ) task = asyncio.create_task( self._create_memory( data=action_text, @@ -1660,6 +1678,11 @@ class AsyncMemory(MemoryBase): ) memory_tasks.append((task, resp, "ADD", None)) elif event_type == "UPDATE": + # Ensure action_text has an embedding cached to avoid redundant API calls + if action_text not in new_message_embeddings: + new_message_embeddings[action_text] = await asyncio.to_thread( + self.embedding_model.embed, action_text, "update" + ) task = asyncio.create_task( self._update_memory( memory_id=temp_uuid_mapping[resp["id"]], @@ -2231,8 +2254,11 @@ class AsyncMemory(MemoryBase): async def _create_memory(self, data, existing_embeddings, metadata=None): logger.debug(f"Creating memory with {data=}") - if data in existing_embeddings: + # existing_embeddings may be a dict (preferred) or a precomputed vector + if isinstance(existing_embeddings, dict) and data in existing_embeddings: embeddings = existing_embeddings[data] + elif isinstance(existing_embeddings, list): + embeddings = existing_embeddings else: embeddings = await asyncio.to_thread(self.embedding_model.embed, data, memory_action="add") @@ -2349,8 +2375,10 @@ class AsyncMemory(MemoryBase): if "role" not in new_metadata and "role" in existing_memory.payload: new_metadata["role"] = existing_memory.payload["role"] - if data in existing_embeddings: + if isinstance(existing_embeddings, dict) and data in existing_embeddings: embeddings = existing_embeddings[data] + elif isinstance(existing_embeddings, list): + embeddings = existing_embeddings else: embeddings = await asyncio.to_thread(self.embedding_model.embed, data, "update") diff --git a/tests/test_memory.py b/tests/test_memory.py index 14214dcb5..047d02b45 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -438,3 +438,82 @@ async def test_async_update_nonexistent_memory_raises_error(mock_sqlite, mock_ll await memory._update_memory("non-existent-id", "new data", {"new data": [0.1, 0.2]}) mock_vector_store.update.assert_not_called() + + +@patch('mem0.utils.factory.EmbedderFactory.create') +@patch('mem0.utils.factory.VectorStoreFactory.create') +@patch('mem0.utils.factory.LlmFactory.create') +@patch('mem0.memory.storage.SQLiteManager') +def test_add_infer_false_embeds_once(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """ + Regression test for issue #3723: adding with infer=False should not trigger duplicate embedding calls. + + Root cause: _create_memory expected a dict for existing_embeddings but received a raw list[float], + causing the cache check `data in existing_embeddings` to always fail and trigger a redundant embed. + """ + embedder = MagicMock() + embedder.embed.return_value = [0.1, 0.2, 0.3] + embedder.config = MagicMock(embedding_dims=3) + mock_embedder_factory.return_value = embedder + + mock_vector_store = MagicMock() + mock_vector_store.search.return_value = [] + mock_vector_store.insert.return_value = None + mock_vector_store.get.return_value = None + telemetry_vector_store = MagicMock() + mock_vector_factory.side_effect = [mock_vector_store, telemetry_vector_store] + + mock_llm_factory.return_value = MagicMock() + mock_sqlite.return_value = MagicMock() + + from mem0.memory.main import Memory as MemoryClass + memory = MemoryClass(MemoryConfig()) + + memory.add("foo", user_id="test_user", infer=False) + + assert embedder.embed.call_count == 1 + mock_vector_store.insert.assert_called_once() + + +@patch('mem0.utils.factory.EmbedderFactory.create') +@patch('mem0.utils.factory.VectorStoreFactory.create') +@patch('mem0.utils.factory.LlmFactory.create') +@patch('mem0.memory.storage.SQLiteManager') +def test_add_infer_true_caches_embedding_on_llm_rewrite(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """ + Regression test for issue #3723 (infer=True path): when the LLM rewrites a fact during the + ADD action, the embedding should be computed once and cached, not computed again inside _create_memory. + """ + embedder = MagicMock() + embedder.embed.return_value = [0.1, 0.2, 0.3] + embedder.config = MagicMock(embedding_dims=3) + mock_embedder_factory.return_value = embedder + + mock_vector_store = MagicMock() + mock_vector_store.search.return_value = [] + mock_vector_store.insert.return_value = None + mock_vector_store.get.return_value = None + telemetry_vector_store = MagicMock() + mock_vector_factory.side_effect = [mock_vector_store, telemetry_vector_store] + + # LLM extracts fact "User likes Python", then ADD action rewrites to "The user enjoys Python" + mock_llm = MagicMock() + mock_llm.generate_response.side_effect = [ + json.dumps({"facts": ["User likes Python"]}), + json.dumps({"memory": [{"id": "0", "text": "The user enjoys Python", "event": "ADD", "old_memory": None}]}), + ] + mock_llm_factory.return_value = mock_llm + + mock_sqlite.return_value = MagicMock() + + from mem0.memory.main import Memory as MemoryClass + memory = MemoryClass(MemoryConfig()) + + memory.add("I like Python", user_id="test_user", infer=True) + + # embed should be called exactly twice: + # 1. For the extracted fact "User likes Python" (search) + # 2. For the rewritten text "The user enjoys Python" (pre-cached before _create_memory) + # It should NOT be called a 3rd time inside _create_memory + assert embedder.embed.call_count == 2 + mock_vector_store.insert.assert_called_once()