fix(memory): return attributed_to from get/get_all/search (#5629)
This commit is contained in:
@@ -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}
|
||||
|
||||
|
||||
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user