Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 596a624716 | |||
| 211a7570e7 |
+39
-11
@@ -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")
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user