fix(memory): exclude vector-store-rejected records from ADD results (#7066)
This commit is contained in:
committed by
GitHub
parent
f4acc89a29
commit
2c6ff619d1
+41
-15
@@ -22,7 +22,7 @@ from mem0.configs.prompts import (
|
||||
PROCEDURAL_MEMORY_SYSTEM_PROMPT,
|
||||
generate_additive_extraction_prompt,
|
||||
)
|
||||
from mem0.exceptions import LLMError
|
||||
from mem0.exceptions import LLMError, VectorStoreError
|
||||
from mem0.exceptions import ValidationError as Mem0ValidationError
|
||||
from mem0.memory.base import MemoryBase
|
||||
from mem0.memory.notices import (
|
||||
@@ -1049,19 +1049,31 @@ class Memory(MemoryBase):
|
||||
all_ids = [r[0] for r in records]
|
||||
all_payloads = [r[3] for r in records]
|
||||
|
||||
# Only records confirmed to be stored make it into history, entity
|
||||
# links, and the returned results — a record the store rejected must
|
||||
# never be reported back as a successful ADD.
|
||||
persisted_records = []
|
||||
try:
|
||||
self.vector_store.insert(
|
||||
vectors=all_vectors,
|
||||
ids=all_ids,
|
||||
payloads=all_payloads,
|
||||
)
|
||||
persisted_records = records
|
||||
except Exception:
|
||||
# Fallback: insert one by one
|
||||
for mid, vec, pay in zip(all_ids, all_vectors, all_payloads):
|
||||
for rec in records:
|
||||
try:
|
||||
self.vector_store.insert(vectors=[vec], ids=[mid], payloads=[pay])
|
||||
self.vector_store.insert(vectors=[rec[2]], ids=[rec[0]], payloads=[rec[3]])
|
||||
persisted_records.append(rec)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to insert memory {mid}: {e}")
|
||||
logger.error(f"Failed to insert memory {rec[0]}: {e}")
|
||||
|
||||
if not persisted_records:
|
||||
self.db.save_messages(messages, session_scope)
|
||||
raise VectorStoreError(
|
||||
f"Failed to insert any of the {len(records)} extracted memories into the vector store"
|
||||
)
|
||||
|
||||
# Batch history
|
||||
history_records = [
|
||||
@@ -1073,7 +1085,7 @@ class Memory(MemoryBase):
|
||||
"created_at": r[3].get("created_at"),
|
||||
"is_deleted": 0,
|
||||
}
|
||||
for r in records
|
||||
for r in persisted_records
|
||||
]
|
||||
try:
|
||||
self.db.batch_add_history(history_records)
|
||||
@@ -1087,12 +1099,12 @@ class Memory(MemoryBase):
|
||||
|
||||
# Phase 7: Batch entity linking
|
||||
try:
|
||||
all_texts = [r[1] for r in records]
|
||||
all_texts = [r[1] for r in persisted_records]
|
||||
all_entities = extract_entities_batch(all_texts)
|
||||
|
||||
# 7a: Global dedup — collect unique entities across all memories
|
||||
global_entities = {} # normalized_key -> (entity_type, entity_text, set of memory_ids)
|
||||
for idx, (memory_id, text, embedding, payload) in enumerate(records):
|
||||
for idx, (memory_id, text, embedding, payload) in enumerate(persisted_records):
|
||||
entities = all_entities[idx] if idx < len(all_entities) else []
|
||||
for entity_type, entity_text in entities:
|
||||
key = self._normalize_entity_text(entity_text)
|
||||
@@ -1196,7 +1208,7 @@ class Memory(MemoryBase):
|
||||
|
||||
returned_memories = [
|
||||
{"id": r[0], "memory": r[1], "event": "ADD"}
|
||||
for r in records
|
||||
for r in persisted_records
|
||||
]
|
||||
|
||||
keys, encoded_ids = process_telemetry_filters(filters)
|
||||
@@ -2709,6 +2721,10 @@ class AsyncMemory(MemoryBase):
|
||||
all_ids = [r[0] for r in records]
|
||||
all_payloads = [r[3] for r in records]
|
||||
|
||||
# Only records confirmed to be stored make it into history, entity
|
||||
# links, and the returned results — a record the store rejected must
|
||||
# never be reported back as a successful ADD.
|
||||
persisted_records = []
|
||||
try:
|
||||
await asyncio.to_thread(
|
||||
self.vector_store.insert,
|
||||
@@ -2716,12 +2732,22 @@ class AsyncMemory(MemoryBase):
|
||||
ids=all_ids,
|
||||
payloads=all_payloads,
|
||||
)
|
||||
persisted_records = records
|
||||
except Exception:
|
||||
for mid, vec, pay in zip(all_ids, all_vectors, all_payloads):
|
||||
for rec in records:
|
||||
try:
|
||||
await asyncio.to_thread(self.vector_store.insert, vectors=[vec], ids=[mid], payloads=[pay])
|
||||
await asyncio.to_thread(
|
||||
self.vector_store.insert, vectors=[rec[2]], ids=[rec[0]], payloads=[rec[3]]
|
||||
)
|
||||
persisted_records.append(rec)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to insert memory {mid} (async): {e}")
|
||||
logger.error(f"Failed to insert memory {rec[0]} (async): {e}")
|
||||
|
||||
if not persisted_records:
|
||||
await asyncio.to_thread(self.db.save_messages, messages, session_scope)
|
||||
raise VectorStoreError(
|
||||
f"Failed to insert any of the {len(records)} extracted memories into the vector store"
|
||||
)
|
||||
|
||||
# Batch history
|
||||
history_records = [
|
||||
@@ -2733,7 +2759,7 @@ class AsyncMemory(MemoryBase):
|
||||
"created_at": r[3].get("created_at"),
|
||||
"is_deleted": 0,
|
||||
}
|
||||
for r in records
|
||||
for r in persisted_records
|
||||
]
|
||||
try:
|
||||
await asyncio.to_thread(self.db.batch_add_history, history_records)
|
||||
@@ -2749,12 +2775,12 @@ class AsyncMemory(MemoryBase):
|
||||
|
||||
# Phase 7: Batch entity linking
|
||||
try:
|
||||
all_texts = [r[1] for r in records]
|
||||
all_texts = [r[1] for r in persisted_records]
|
||||
all_entities = await asyncio.to_thread(extract_entities_batch, all_texts)
|
||||
|
||||
# 7a: Global dedup
|
||||
global_entities = {}
|
||||
for idx, (memory_id, text, embedding, payload) in enumerate(records):
|
||||
for idx, (memory_id, text, embedding, payload) in enumerate(persisted_records):
|
||||
entities = all_entities[idx] if idx < len(all_entities) else []
|
||||
for entity_type, entity_text in entities:
|
||||
key = self._normalize_entity_text(entity_text)
|
||||
@@ -2856,7 +2882,7 @@ class AsyncMemory(MemoryBase):
|
||||
|
||||
returned_memories = [
|
||||
{"id": r[0], "memory": r[1], "event": "ADD"}
|
||||
for r in records
|
||||
for r in persisted_records
|
||||
]
|
||||
|
||||
keys, encoded_ids = process_telemetry_filters(effective_filters)
|
||||
|
||||
+156
-1
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
@@ -6,7 +7,7 @@ from unittest.mock import MagicMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.exceptions import LLMError
|
||||
from mem0.exceptions import LLMError, VectorStoreError
|
||||
from mem0.memory.main import AsyncMemory, Memory
|
||||
|
||||
|
||||
@@ -1178,3 +1179,157 @@ class TestAddPipelineEntityEmbeddingCountGuard:
|
||||
assert any("padding/truncating" in r.message for r in caplog.records), (
|
||||
"expected count-mismatch warning was not emitted"
|
||||
)
|
||||
|
||||
class TestPartialInsertFailure:
|
||||
"""Records the vector store rejects must never be reported as successful ADDs (#6911)."""
|
||||
|
||||
LLM_RESPONSE = json.dumps(
|
||||
{
|
||||
"memory": [
|
||||
{"text": "User's name is Aryan"},
|
||||
{"text": "User is allergic to penicillin"},
|
||||
{"text": "User works as an engineer"},
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
@pytest.fixture
|
||||
def mock_memory(self, mocker):
|
||||
mock_llm, _ = _setup_mocks(mocker)
|
||||
|
||||
memory = Memory()
|
||||
memory.config = mocker.MagicMock()
|
||||
memory.config.custom_instructions = None
|
||||
memory.config.custom_update_memory_prompt = None
|
||||
memory.custom_instructions = None
|
||||
memory.api_version = "v1.1"
|
||||
memory.db.get_last_messages = MagicMock(return_value=[])
|
||||
memory.db.save_messages = MagicMock()
|
||||
memory.db.batch_add_history = MagicMock()
|
||||
memory.embedding_model.embed_batch = Mock(side_effect=lambda texts, action: [[0.1, 0.2, 0.3] for _ in texts])
|
||||
mocker.patch("mem0.memory.main.extract_entities_batch", return_value=[])
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
return memory
|
||||
|
||||
def test_rejected_records_are_not_reported_as_add(self, mock_memory):
|
||||
"""A record rejected by the vector store must be absent from results and history."""
|
||||
poison = "User is allergic to penicillin"
|
||||
real_insert = mock_memory.vector_store.insert
|
||||
|
||||
def flaky_insert(vectors, ids, payloads):
|
||||
if len(ids) > 1:
|
||||
raise RuntimeError("batch insert rejected")
|
||||
if payloads[0]["data"] == poison:
|
||||
raise RuntimeError("record rejected by vector store")
|
||||
return real_insert(vectors=vectors, ids=ids, payloads=payloads)
|
||||
|
||||
mock_memory.vector_store.insert = Mock(side_effect=flaky_insert)
|
||||
mock_memory.llm.generate_response.return_value = self.LLM_RESPONSE
|
||||
|
||||
result = mock_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "My name is Aryan. I work as an engineer."}],
|
||||
metadata={},
|
||||
filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert [r["memory"] for r in result] == [
|
||||
"User's name is Aryan",
|
||||
"User works as an engineer",
|
||||
]
|
||||
assert all(r["event"] == "ADD" for r in result)
|
||||
|
||||
history = mock_memory.db.batch_add_history.call_args.args[0]
|
||||
assert {h["new_memory"] for h in history} == {
|
||||
"User's name is Aryan",
|
||||
"User works as an engineer",
|
||||
}
|
||||
|
||||
def test_all_inserts_failed_raises_vector_store_error(self, mock_memory):
|
||||
"""If nothing persisted, add() must raise VectorStoreError instead of reporting success."""
|
||||
mock_memory.vector_store.insert = Mock(side_effect=RuntimeError("vector store down"))
|
||||
mock_memory.llm.generate_response.return_value = self.LLM_RESPONSE
|
||||
|
||||
with pytest.raises(VectorStoreError, match="Failed to insert any"):
|
||||
mock_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
metadata={},
|
||||
filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
# Raw messages are still saved so a later retry can re-extract them.
|
||||
mock_memory.db.save_messages.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestAsyncPartialInsertFailure:
|
||||
"""Async mirror of TestPartialInsertFailure (#6911)."""
|
||||
|
||||
LLM_RESPONSE = TestPartialInsertFailure.LLM_RESPONSE
|
||||
|
||||
@pytest.fixture
|
||||
def mock_async_memory(self, mocker):
|
||||
mock_llm, _ = _setup_mocks(mocker)
|
||||
|
||||
memory = AsyncMemory()
|
||||
memory.config = mocker.MagicMock()
|
||||
memory.config.custom_instructions = None
|
||||
memory.config.custom_update_memory_prompt = None
|
||||
memory.custom_instructions = None
|
||||
memory.api_version = "v1.1"
|
||||
memory.db.get_last_messages = MagicMock(return_value=[])
|
||||
memory.db.save_messages = MagicMock()
|
||||
memory.db.batch_add_history = MagicMock()
|
||||
memory.embedding_model.embed_batch = Mock(side_effect=lambda texts, action: [[0.1, 0.2, 0.3] for _ in texts])
|
||||
mocker.patch("mem0.memory.main.extract_entities_batch", return_value=[])
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
return memory
|
||||
|
||||
async def test_rejected_records_are_not_reported_as_add(self, mock_async_memory):
|
||||
poison = "User is allergic to penicillin"
|
||||
real_insert = mock_async_memory.vector_store.insert
|
||||
|
||||
def flaky_insert(vectors, ids, payloads):
|
||||
if len(ids) > 1:
|
||||
raise RuntimeError("batch insert rejected")
|
||||
if payloads[0]["data"] == poison:
|
||||
raise RuntimeError("record rejected by vector store")
|
||||
return real_insert(vectors=vectors, ids=ids, payloads=payloads)
|
||||
|
||||
mock_async_memory.vector_store.insert = Mock(side_effect=flaky_insert)
|
||||
mock_async_memory.llm.generate_response.return_value = self.LLM_RESPONSE
|
||||
|
||||
result = await mock_async_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "My name is Aryan. I work as an engineer."}],
|
||||
metadata={},
|
||||
effective_filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert [r["memory"] for r in result] == [
|
||||
"User's name is Aryan",
|
||||
"User works as an engineer",
|
||||
]
|
||||
|
||||
history = mock_async_memory.db.batch_add_history.call_args.args[0]
|
||||
assert {h["new_memory"] for h in history} == {
|
||||
"User's name is Aryan",
|
||||
"User works as an engineer",
|
||||
}
|
||||
|
||||
async def test_all_inserts_failed_raises_vector_store_error(self, mock_async_memory):
|
||||
mock_async_memory.vector_store.insert = Mock(side_effect=RuntimeError("vector store down"))
|
||||
mock_async_memory.llm.generate_response.return_value = self.LLM_RESPONSE
|
||||
|
||||
with pytest.raises(VectorStoreError, match="Failed to insert any"):
|
||||
await mock_async_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
metadata={},
|
||||
effective_filters={},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
mock_async_memory.db.save_messages.assert_called_once()
|
||||
|
||||
Reference in New Issue
Block a user