diff --git a/mem0/memory/main.py b/mem0/memory/main.py index ac75baa04..05a55c876 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -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) diff --git a/tests/memory/test_main.py b/tests/memory/test_main.py index cea7e3842..00d962c34 100644 --- a/tests/memory/test_main.py +++ b/tests/memory/test_main.py @@ -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()