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:
+9
-9
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user