fix: reset() only drops history table, leaving stale messages (#5541)
This commit is contained in:
+4
-8
@@ -1934,10 +1934,8 @@ class Memory(MemoryBase):
|
|||||||
"""
|
"""
|
||||||
logger.warning("Resetting all memories")
|
logger.warning("Resetting all memories")
|
||||||
|
|
||||||
if hasattr(self.db, "connection") and self.db.connection:
|
self.db.reset()
|
||||||
self.db.connection.execute("DROP TABLE IF EXISTS history")
|
self.db.close()
|
||||||
self.db.connection.close()
|
|
||||||
|
|
||||||
self.db = SQLiteManager(self.config.history_db_path)
|
self.db = SQLiteManager(self.config.history_db_path)
|
||||||
|
|
||||||
if hasattr(self.vector_store, "reset"):
|
if hasattr(self.vector_store, "reset"):
|
||||||
@@ -3509,10 +3507,8 @@ class AsyncMemory(MemoryBase):
|
|||||||
if hasattr(self.vector_store, "client") and hasattr(self.vector_store.client, "close"):
|
if hasattr(self.vector_store, "client") and hasattr(self.vector_store.client, "close"):
|
||||||
await asyncio.to_thread(self.vector_store.client.close)
|
await asyncio.to_thread(self.vector_store.client.close)
|
||||||
|
|
||||||
if hasattr(self.db, "connection") and self.db.connection:
|
await asyncio.to_thread(self.db.reset)
|
||||||
await asyncio.to_thread(lambda: self.db.connection.execute("DROP TABLE IF EXISTS history"))
|
await asyncio.to_thread(self.db.close)
|
||||||
await asyncio.to_thread(self.db.connection.close)
|
|
||||||
|
|
||||||
self.db = SQLiteManager(self.config.history_db_path)
|
self.db = SQLiteManager(self.config.history_db_path)
|
||||||
|
|
||||||
self.vector_store = VectorStoreFactory.create(
|
self.vector_store = VectorStoreFactory.create(
|
||||||
|
|||||||
@@ -324,7 +324,9 @@ class SQLiteManager:
|
|||||||
]
|
]
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
"""Drop and recreate the history and messages tables."""
|
"""Drop both tables. Caller is expected to replace this instance."""
|
||||||
|
if not self.connection:
|
||||||
|
raise RuntimeError("Cannot reset a closed SQLiteManager")
|
||||||
with self._lock:
|
with self._lock:
|
||||||
try:
|
try:
|
||||||
self.connection.execute("BEGIN")
|
self.connection.execute("BEGIN")
|
||||||
@@ -335,8 +337,6 @@ class SQLiteManager:
|
|||||||
self.connection.execute("ROLLBACK")
|
self.connection.execute("ROLLBACK")
|
||||||
logger.error(f"Failed to reset tables: {e}")
|
logger.error(f"Failed to reset tables: {e}")
|
||||||
raise
|
raise
|
||||||
self._create_history_table()
|
|
||||||
self._create_messages_table()
|
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
if self.connection:
|
if self.connection:
|
||||||
|
|||||||
@@ -280,3 +280,22 @@ class TestSQLiteManager:
|
|||||||
assert history[0]["actor_id"] is None
|
assert history[0]["actor_id"] is None
|
||||||
assert history[0]["is_deleted"] is False
|
assert history[0]["is_deleted"] is False
|
||||||
mgr.close()
|
mgr.close()
|
||||||
|
|
||||||
|
def test_reset_drops_tables(self, temp_db_path):
|
||||||
|
"""reset() must drop both history and messages tables."""
|
||||||
|
mgr = SQLiteManager(temp_db_path)
|
||||||
|
mgr.add_history(memory_id="m1", old_memory=None, new_memory="new", event="ADD")
|
||||||
|
mgr.save_messages([{"role": "user", "content": "hello", "name": None}], "sess1")
|
||||||
|
mgr.reset()
|
||||||
|
tables = mgr.connection.execute(
|
||||||
|
"SELECT name FROM sqlite_master WHERE type='table' AND name IN ('history','messages')"
|
||||||
|
).fetchall()
|
||||||
|
assert tables == [], "both tables should be dropped after reset"
|
||||||
|
mgr.close()
|
||||||
|
|
||||||
|
mgr2 = SQLiteManager(temp_db_path)
|
||||||
|
msg_count = mgr2.connection.execute("SELECT COUNT(*) FROM messages").fetchone()[0]
|
||||||
|
hist_count = mgr2.connection.execute("SELECT COUNT(*) FROM history").fetchone()[0]
|
||||||
|
assert msg_count == 0
|
||||||
|
assert hist_count == 0
|
||||||
|
mgr2.close()
|
||||||
|
|||||||
@@ -118,6 +118,57 @@ def test_collection_name_preserved_after_reset(mock_sqlite, mock_llm_factory, mo
|
|||||||
assert reset_config.collection_name == test_collection_name, f"Reset used wrong collection name: {reset_config.collection_name}"
|
assert reset_config.collection_name == test_collection_name, f"Reset used wrong collection name: {reset_config.collection_name}"
|
||||||
|
|
||||||
|
|
||||||
|
@patch('mem0.utils.factory.EmbedderFactory.create')
|
||||||
|
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
||||||
|
@patch('mem0.utils.factory.LlmFactory.create')
|
||||||
|
def test_memory_reset_clears_messages_table(mock_llm_factory, mock_vector_factory, mock_embedder_factory, tmp_path):
|
||||||
|
"""Regression: Memory.reset() must clear the messages table, not just history."""
|
||||||
|
mock_embedder_factory.return_value = MagicMock()
|
||||||
|
mock_vector_factory.return_value = MagicMock()
|
||||||
|
mock_llm_factory.return_value = MagicMock()
|
||||||
|
|
||||||
|
config = MemoryConfig()
|
||||||
|
config.history_db_path = str(tmp_path / "test.db")
|
||||||
|
memory = Memory(config)
|
||||||
|
|
||||||
|
memory.db.save_messages([{"role": "user", "content": "hello", "name": None}], "sess1")
|
||||||
|
memory.db.add_history(memory_id="m1", old_memory=None, new_memory="x", event="ADD")
|
||||||
|
|
||||||
|
memory.reset()
|
||||||
|
|
||||||
|
msg_count = memory.db.connection.execute("SELECT COUNT(*) FROM messages").fetchone()[0]
|
||||||
|
hist_count = memory.db.connection.execute("SELECT COUNT(*) FROM history").fetchone()[0]
|
||||||
|
assert msg_count == 0, "messages table must be empty after Memory.reset()"
|
||||||
|
assert hist_count == 0, "history table must be empty after Memory.reset()"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@patch('mem0.utils.factory.EmbedderFactory.create')
|
||||||
|
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
||||||
|
@patch('mem0.utils.factory.LlmFactory.create')
|
||||||
|
async def test_async_memory_reset_clears_messages_table(mock_llm_factory, mock_vector_factory, mock_embedder_factory, tmp_path):
|
||||||
|
"""Regression: AsyncMemory.reset() must clear the messages table, not just history."""
|
||||||
|
mock_embedder_factory.return_value = MagicMock()
|
||||||
|
mock_vector_factory.return_value = MagicMock()
|
||||||
|
mock_llm_factory.return_value = MagicMock()
|
||||||
|
|
||||||
|
from mem0 import AsyncMemory
|
||||||
|
|
||||||
|
config = MemoryConfig()
|
||||||
|
config.history_db_path = str(tmp_path / "test.db")
|
||||||
|
memory = AsyncMemory(config)
|
||||||
|
|
||||||
|
memory.db.save_messages([{"role": "user", "content": "hi", "name": None}], "s1")
|
||||||
|
memory.db.add_history(memory_id="m1", old_memory=None, new_memory="x", event="ADD")
|
||||||
|
|
||||||
|
await memory.reset()
|
||||||
|
|
||||||
|
msg_count = memory.db.connection.execute("SELECT COUNT(*) FROM messages").fetchone()[0]
|
||||||
|
hist_count = memory.db.connection.execute("SELECT COUNT(*) FROM history").fetchone()[0]
|
||||||
|
assert msg_count == 0, "messages must be empty after AsyncMemory.reset()"
|
||||||
|
assert hist_count == 0, "history must be empty after AsyncMemory.reset()"
|
||||||
|
|
||||||
|
|
||||||
@patch('mem0.utils.factory.EmbedderFactory.create')
|
@patch('mem0.utils.factory.EmbedderFactory.create')
|
||||||
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
||||||
@patch('mem0.utils.factory.LlmFactory.create')
|
@patch('mem0.utils.factory.LlmFactory.create')
|
||||||
|
|||||||
Reference in New Issue
Block a user