From 596a62471627b68d2f49411deb2e76a0d306d89c Mon Sep 17 00:00:00 2001 From: utkarsh240799 Date: Wed, 25 Mar 2026 17:35:18 +0530 Subject: [PATCH] fix: improve double-embedding fix with type safety and UPDATE path test (#3723) Make isinstance checks numpy-safe by using `not isinstance(dict)` instead of `isinstance(list)`, add type hints to existing_embeddings parameter, and add regression test for the infer=True UPDATE path. Co-Authored-By: Claude Opus 4.6 (1M context) --- mem0/memory/main.py | 18 +++++++-------- tests/test_memory.py | 55 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 64 insertions(+), 9 deletions(-) diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 175a56e08..e5f4456c8 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -9,7 +9,7 @@ import uuid import warnings from copy import deepcopy from datetime import datetime, timezone -from typing import Any, Dict, Optional +from typing import Any, Dict, List, Optional, Union from pydantic import ValidationError @@ -1160,12 +1160,12 @@ class Memory(MemoryBase): capture_event("mem0.history", self, {"memory_id": memory_id, "sync_type": "sync"}) return self.db.get_history(memory_id) - def _create_memory(self, data, existing_embeddings, metadata=None): + def _create_memory(self, data: str, existing_embeddings: Union[Dict[str, List[float]], List[float]], metadata=None): logger.debug(f"Creating memory with {data=}") # 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): + elif not isinstance(existing_embeddings, dict): embeddings = existing_embeddings else: embeddings = self.embedding_model.embed(data, memory_action="add") @@ -1230,7 +1230,7 @@ class Memory(MemoryBase): return result - def _update_memory(self, memory_id, data, existing_embeddings, metadata=None): + def _update_memory(self, memory_id, data: str, existing_embeddings: Union[Dict[str, List[float]], List[float]], metadata=None): logger.info(f"Updating memory with {data=}") try: @@ -1265,7 +1265,7 @@ class Memory(MemoryBase): if isinstance(existing_embeddings, dict) and data in existing_embeddings: embeddings = existing_embeddings[data] - elif isinstance(existing_embeddings, list): + elif not isinstance(existing_embeddings, dict): embeddings = existing_embeddings else: embeddings = self.embedding_model.embed(data, "update") @@ -2252,12 +2252,12 @@ class AsyncMemory(MemoryBase): capture_event("mem0.history", self, {"memory_id": memory_id, "sync_type": "async"}) return await asyncio.to_thread(self.db.get_history, memory_id) - async def _create_memory(self, data, existing_embeddings, metadata=None): + async def _create_memory(self, data: str, existing_embeddings: Union[Dict[str, List[float]], List[float]], metadata=None): logger.debug(f"Creating memory with {data=}") # 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): + elif not isinstance(existing_embeddings, dict): embeddings = existing_embeddings else: embeddings = await asyncio.to_thread(self.embedding_model.embed, data, memory_action="add") @@ -2341,7 +2341,7 @@ class AsyncMemory(MemoryBase): return result - async def _update_memory(self, memory_id, data, existing_embeddings, metadata=None): + async def _update_memory(self, memory_id, data: str, existing_embeddings: Union[Dict[str, List[float]], List[float]], metadata=None): logger.info(f"Updating memory with {data=}") try: @@ -2377,7 +2377,7 @@ class AsyncMemory(MemoryBase): if isinstance(existing_embeddings, dict) and data in existing_embeddings: embeddings = existing_embeddings[data] - elif isinstance(existing_embeddings, list): + elif not isinstance(existing_embeddings, dict): 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 047d02b45..827ed523d 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -517,3 +517,58 @@ def test_add_infer_true_caches_embedding_on_llm_rewrite(mock_sqlite, mock_llm_fa # It should NOT be called a 3rd time inside _create_memory assert embedder.embed.call_count == 2 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_update_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 UPDATE path): when the LLM rewrites a fact during + an UPDATE action, the embedding should be computed once and cached, not computed again inside _update_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 + + # Existing memory that will be matched for update + existing_memory = MockVectorMemory( + memory_id="existing-mem-id", + payload={ + "data": "User likes Python", + "hash": "abc123", + "created_at": "2025-01-01T00:00:00+00:00", + }, + ) + + mock_vector_store = MagicMock() + mock_vector_store.search.return_value = [existing_memory] + mock_vector_store.get.return_value = existing_memory + mock_vector_store.insert.return_value = None + mock_vector_store.update.return_value = None + telemetry_vector_store = MagicMock() + mock_vector_factory.side_effect = [mock_vector_store, telemetry_vector_store] + + # LLM extracts fact "User loves Python now", then UPDATE action rewrites to "The user loves Python" + mock_llm = MagicMock() + mock_llm.generate_response.side_effect = [ + json.dumps({"facts": ["User loves Python now"]}), + json.dumps({"memory": [{"id": "0", "text": "The user loves Python", "event": "UPDATE", "old_memory": "User likes Python"}]}), + ] + 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 love Python now", user_id="test_user", infer=True) + + # embed should be called exactly twice: + # 1. For the extracted fact "User loves Python now" (search) + # 2. For the rewritten text "The user loves Python" (pre-cached before _update_memory) + # It should NOT be called a 3rd time inside _update_memory + assert embedder.embed.call_count == 2 + mock_vector_store.update.assert_called_once()