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) <noreply@anthropic.com>
This commit is contained in:
utkarsh240799
2026-03-25 17:35:18 +05:30
parent 211a7570e7
commit 596a624716
2 changed files with 64 additions and 9 deletions
+9 -9
View File
@@ -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")
+55
View File
@@ -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()