fix(memory): guard entity embed_batch count mismatch in v3 add pipeline (#5604)
Co-authored-by: Kartik <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -993,6 +993,17 @@ class Memory(MemoryBase):
|
||||
except Exception:
|
||||
entity_embeddings.append(None)
|
||||
|
||||
|
||||
if len(entity_embeddings) != len(ordered_keys):
|
||||
logger.warning(
|
||||
"embed_batch returned %d vectors for %d entity texts — "
|
||||
"padding/truncating to avoid dropping entity links",
|
||||
len(entity_embeddings),
|
||||
len(ordered_keys),
|
||||
)
|
||||
entity_embeddings = list(entity_embeddings[: len(ordered_keys)])
|
||||
entity_embeddings += [None] * (len(ordered_keys) - len(entity_embeddings))
|
||||
|
||||
# Filter out entities with failed embeddings
|
||||
valid = [(i, k) for i, k in enumerate(ordered_keys) if entity_embeddings[i] is not None]
|
||||
if valid:
|
||||
@@ -2520,6 +2531,16 @@ class AsyncMemory(MemoryBase):
|
||||
except Exception:
|
||||
entity_embeddings.append(None)
|
||||
|
||||
if len(entity_embeddings) != len(ordered_keys):
|
||||
logger.warning(
|
||||
"embed_batch returned %d vectors for %d entity texts — "
|
||||
"padding/truncating to avoid dropping entity links",
|
||||
len(entity_embeddings),
|
||||
len(ordered_keys),
|
||||
)
|
||||
entity_embeddings = list(entity_embeddings[: len(ordered_keys)])
|
||||
entity_embeddings += [None] * (len(ordered_keys) - len(entity_embeddings))
|
||||
|
||||
valid = [(i, k) for i, k in enumerate(ordered_keys) if entity_embeddings[i] is not None]
|
||||
if valid:
|
||||
valid_indices, valid_keys = zip(*valid)
|
||||
|
||||
@@ -857,3 +857,134 @@ class TestEntityBoostParallelism:
|
||||
assert elapsed < 0.75, f"searches did not run concurrently (took {elapsed:.2f}s)"
|
||||
assert concurrent_count["peak"] >= 2, "no overlap observed between entity searches"
|
||||
assert len(boosts) == 4
|
||||
|
||||
|
||||
class TestAddPipelineEntityEmbeddingCountGuard:
|
||||
"""A misbehaving embedder returning fewer (or more) vectors than entity
|
||||
texts must not silently drop ALL entity links via a swallowed IndexError.
|
||||
|
||||
Before the fix, `entity_embeddings[i]` (indexed by ordered_keys length) raised
|
||||
IndexError on a short return; the outer `except Exception` in the Phase 7
|
||||
entity-linking block swallowed it, so no entity was ever searched/inserted and
|
||||
no error surfaced. The search path (_compute_entity_boosts) already guarded
|
||||
this; the add path did not.
|
||||
"""
|
||||
|
||||
@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()
|
||||
return memory
|
||||
|
||||
@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()
|
||||
return memory
|
||||
|
||||
@staticmethod
|
||||
def _short_entity_embed_batch(texts, memory_action="add"):
|
||||
# Memory-text embeddings are well-behaved; entity embeddings come back short.
|
||||
if memory_action == "add" and any(t in ("Alice", "Bob") for t in texts):
|
||||
return [[0.1] * 10] # 1 vector for 2 entity texts
|
||||
return [[0.1] * 10 for _ in texts]
|
||||
|
||||
def test_sync_short_entity_embeddings_still_link_valid_entity(self, mock_memory, mocker, caplog):
|
||||
mock_memory.llm.generate_response.return_value = (
|
||||
'{"memory": [{"text": "Alice met Bob"}, {"text": "Bob called Alice"}]}'
|
||||
)
|
||||
mock_memory.embedding_model = Mock()
|
||||
mock_memory.embedding_model.embed_batch = Mock(side_effect=self._short_entity_embed_batch)
|
||||
mock_memory.embedding_model.embed = Mock(return_value=[0.1] * 10)
|
||||
|
||||
mock_memory._entity_store = Mock()
|
||||
mock_memory._entity_store.search_batch = Mock(return_value=[[]])
|
||||
mock_memory._entity_store.insert = Mock()
|
||||
mock_memory._entity_store.update = Mock()
|
||||
|
||||
mocker.patch(
|
||||
"mem0.memory.main.extract_entities_batch",
|
||||
return_value=[
|
||||
[("person", "Alice"), ("person", "Bob")],
|
||||
[("person", "Bob"), ("person", "Alice")],
|
||||
],
|
||||
)
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = mock_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "Alice met Bob; Bob called Alice"}],
|
||||
metadata={},
|
||||
filters={"user_id": "u1"},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
# Both memories persist regardless.
|
||||
assert len(result) == 2
|
||||
# The entity block did NOT abort: it searched + inserted the one valid entity
|
||||
# instead of swallowing an IndexError and linking nothing.
|
||||
assert mock_memory._entity_store.search_batch.call_count == 1
|
||||
assert mock_memory._entity_store.insert.call_count == 1
|
||||
assert not any("Batch entity linking failed" in r.message for r in caplog.records), (
|
||||
"entity linking aborted on a swallowed IndexError"
|
||||
)
|
||||
assert any("padding/truncating" in r.message for r in caplog.records), (
|
||||
"expected count-mismatch warning was not emitted"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_short_entity_embeddings_still_link_valid_entity(self, mock_async_memory, mocker, caplog):
|
||||
mock_async_memory.llm.generate_response.return_value = (
|
||||
'{"memory": [{"text": "Alice met Bob"}, {"text": "Bob called Alice"}]}'
|
||||
)
|
||||
mock_async_memory.embedding_model = Mock()
|
||||
mock_async_memory.embedding_model.embed_batch = Mock(side_effect=self._short_entity_embed_batch)
|
||||
mock_async_memory.embedding_model.embed = Mock(return_value=[0.1] * 10)
|
||||
|
||||
mock_async_memory._entity_store = Mock()
|
||||
mock_async_memory._entity_store.search_batch = Mock(return_value=[[]])
|
||||
mock_async_memory._entity_store.insert = Mock()
|
||||
mock_async_memory._entity_store.update = Mock()
|
||||
|
||||
mocker.patch(
|
||||
"mem0.memory.main.extract_entities_batch",
|
||||
return_value=[
|
||||
[("person", "Alice"), ("person", "Bob")],
|
||||
[("person", "Bob"), ("person", "Alice")],
|
||||
],
|
||||
)
|
||||
mocker.patch("mem0.memory.main.capture_event")
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = await mock_async_memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "Alice met Bob; Bob called Alice"}],
|
||||
metadata={},
|
||||
effective_filters={"user_id": "u1"},
|
||||
infer=True,
|
||||
)
|
||||
|
||||
assert len(result) == 2
|
||||
assert mock_async_memory._entity_store.search_batch.call_count == 1
|
||||
assert mock_async_memory._entity_store.insert.call_count == 1
|
||||
assert not any("Batch entity linking failed" in r.message for r in caplog.records), (
|
||||
"async entity linking aborted on a swallowed IndexError"
|
||||
)
|
||||
assert any("padding/truncating" in r.message for r in caplog.records), (
|
||||
"expected count-mismatch warning was not emitted"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user