Compare commits

...

2 Commits

Author SHA1 Message Date
utkarsh240799 596a624716 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>
2026-03-25 17:35:18 +05:30
Varun Chawla 211a7570e7 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.
2026-03-23 14:15:05 +05:30
2 changed files with 173 additions and 11 deletions
+39 -11
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
@@ -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,
@@ -1153,10 +1160,13 @@ 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=}")
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 not isinstance(existing_embeddings, dict):
embeddings = existing_embeddings
else:
embeddings = self.embedding_model.embed(data, memory_action="add")
memory_id = str(uuid.uuid4())
@@ -1220,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:
@@ -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 not isinstance(existing_embeddings, dict):
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"]],
@@ -2229,10 +2252,13 @@ 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=}")
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 not isinstance(existing_embeddings, dict):
embeddings = existing_embeddings
else:
embeddings = await asyncio.to_thread(self.embedding_model.embed, data, memory_action="add")
@@ -2315,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:
@@ -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 not isinstance(existing_embeddings, dict):
embeddings = existing_embeddings
else:
embeddings = await asyncio.to_thread(self.embedding_model.embed, data, "update")
+134
View File
@@ -438,3 +438,137 @@ 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()
@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()