fix(memory): return attributed_to from get/get_all/search (#5629)

This commit is contained in:
Yash Singh
2026-06-23 16:43:30 +05:30
committed by GitHub
parent c2e723352e
commit 879c68555c
2 changed files with 91 additions and 0 deletions
+6
View File
@@ -1102,6 +1102,7 @@ class Memory(MemoryBase):
"run_id",
"actor_id",
"role",
"attributed_to",
]
core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys}
@@ -1215,6 +1216,7 @@ class Memory(MemoryBase):
"run_id",
"actor_id",
"role",
"attributed_to",
]
core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys}
@@ -1550,6 +1552,7 @@ class Memory(MemoryBase):
"run_id",
"actor_id",
"role",
"attributed_to",
]
core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys}
@@ -2638,6 +2641,7 @@ class AsyncMemory(MemoryBase):
"run_id",
"actor_id",
"role",
"attributed_to",
]
core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys}
@@ -2751,6 +2755,7 @@ class AsyncMemory(MemoryBase):
"run_id",
"actor_id",
"role",
"attributed_to",
]
core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys}
@@ -3092,6 +3097,7 @@ class AsyncMemory(MemoryBase):
"run_id",
"actor_id",
"role",
"attributed_to",
]
core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys}
+85
View File
@@ -345,6 +345,91 @@ def test_get_all_handles_flat_list_from_postgres(mock_sqlite, mock_llm_factory,
assert result[1]["memory"] == "Memory 2"
@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_read_apis_surface_attributed_to(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory):
"""
attributed_to is written to the payload on add (and the extraction prompt marks it
required), so get/get_all/search must return it instead of dropping it. It must be a
top-level field, not buried inside metadata.
"""
mock_embedder = MagicMock()
mock_embedder.embed.return_value = [0.1, 0.2, 0.3]
mock_embedder_factory.return_value = mock_embedder
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_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.embedding_model = mock_embedder
payload = {"data": "User likes Python", "attributed_to": "user", "user_id": "u1"}
# get
mock_vector_store.get.return_value = MockVectorMemory("mem_1", payload)
got = memory.get("mem_1")
assert got["attributed_to"] == "user"
assert "attributed_to" not in (got.get("metadata") or {})
# get_all
mock_vector_store.list.return_value = [MockVectorMemory("mem_1", payload)]
listed = memory._get_all_from_vector_store({"user_id": "u1"}, 100)
assert listed[0]["attributed_to"] == "user"
assert "attributed_to" not in (listed[0].get("metadata") or {})
# search
mock_vector_store.search.return_value = [MockVectorMemory("mem_1", payload, score=0.9)]
mock_vector_store.keyword_search.return_value = []
searched = memory._search_vector_store("python", {"user_id": "u1"}, 10)
assert searched[0]["attributed_to"] == "user"
assert "attributed_to" not in (searched[0].get("metadata") or {})
@pytest.mark.asyncio
@patch('mem0.utils.factory.EmbedderFactory.create')
@patch('mem0.utils.factory.VectorStoreFactory.create')
@patch('mem0.utils.factory.LlmFactory.create')
@patch('mem0.memory.storage.SQLiteManager')
async def test_async_read_apis_surface_attributed_to(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory):
"""AsyncMemory get/get_all/search must surface attributed_to, same as the sync path."""
mock_embedder = MagicMock()
mock_embedder.embed.return_value = [0.1, 0.2, 0.3]
mock_embedder_factory.return_value = mock_embedder
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
memory = AsyncMemory(MemoryConfig())
memory.embedding_model = mock_embedder
payload = {"data": "User likes Python", "attributed_to": "user", "user_id": "u1"}
# get
mock_vector_store.get.return_value = MockVectorMemory("mem_1", payload)
got = await memory.get("mem_1")
assert got["attributed_to"] == "user"
assert "attributed_to" not in (got.get("metadata") or {})
# get_all
mock_vector_store.list.return_value = [MockVectorMemory("mem_1", payload)]
listed = await memory._get_all_from_vector_store({"user_id": "u1"}, 100)
assert listed[0]["attributed_to"] == "user"
assert "attributed_to" not in (listed[0].get("metadata") or {})
# search
mock_vector_store.search.return_value = [MockVectorMemory("mem_1", payload, score=0.9)]
mock_vector_store.keyword_search.return_value = []
searched = await memory._search_vector_store("python", {"user_id": "u1"}, 10)
assert searched[0]["attributed_to"] == "user"
assert "attributed_to" not in (searched[0].get("metadata") or {})
@patch('mem0.utils.factory.EmbedderFactory.create')
@patch('mem0.utils.factory.VectorStoreFactory.create')
@patch('mem0.utils.factory.LlmFactory.create')