From 211a7570e754827ade9e9cff8823bf8cefbc5314 Mon Sep 17 00:00:00 2001 From: Varun Chawla Date: Sat, 7 Feb 2026 17:14:44 -0800 Subject: [PATCH] Fix: prevent double embedding in mem0.add (fixes #3723) This fix addresses issue #3723 where mem0.add() was calling the embedding API twice, unnecessarily doubling costs and latency. Changes made: 1. Modified _create_memory() to accept embeddings as either a dict (for caching) or a precomputed vector, preventing redundant calls 2. Updated infer=False path to pass embeddings as a dict 3. Added caching for action_text embeddings in infer=True path for both ADD and UPDATE operations, since the LLM may rephrase facts 4. Applied same fixes to both sync and async Memory classes 5. Added regression test to verify embedding is called only once The root cause was that when infer=False, embeddings were passed directly to _create_memory without a dict wrapper, causing it to re-embed. When infer=True, if the LLM rephrased extracted facts, the action_text wouldn't match the cache key, triggering re-embedding. --- mem0/memory/main.py | 40 ++++++++++++++++++---- tests/test_memory.py | 79 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 113 insertions(+), 6 deletions(-) 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()