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")
|
||||
|
||||
if hasattr(self.db, "connection") and self.db.connection:
|
||||
self.db.connection.execute("DROP TABLE IF EXISTS history")
|
||||
self.db.connection.close()
|
||||
|
||||
self.db.reset()
|
||||
self.db.close()
|
||||
self.db = SQLiteManager(self.config.history_db_path)
|
||||
|
||||
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"):
|
||||
await asyncio.to_thread(self.vector_store.client.close)
|
||||
|
||||
if hasattr(self.db, "connection") and self.db.connection:
|
||||
await asyncio.to_thread(lambda: self.db.connection.execute("DROP TABLE IF EXISTS history"))
|
||||
await asyncio.to_thread(self.db.connection.close)
|
||||
|
||||
await asyncio.to_thread(self.db.reset)
|
||||
await asyncio.to_thread(self.db.close)
|
||||
self.db = SQLiteManager(self.config.history_db_path)
|
||||
|
||||
self.vector_store = VectorStoreFactory.create(
|
||||
|
||||
@@ -324,7 +324,9 @@ class SQLiteManager:
|
||||
]
|
||||
|
||||
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:
|
||||
try:
|
||||
self.connection.execute("BEGIN")
|
||||
@@ -335,8 +337,6 @@ class SQLiteManager:
|
||||
self.connection.execute("ROLLBACK")
|
||||
logger.error(f"Failed to reset tables: {e}")
|
||||
raise
|
||||
self._create_history_table()
|
||||
self._create_messages_table()
|
||||
|
||||
def close(self) -> None:
|
||||
if self.connection:
|
||||
|
||||
@@ -280,3 +280,22 @@ class TestSQLiteManager:
|
||||
assert history[0]["actor_id"] is None
|
||||
assert history[0]["is_deleted"] is False
|
||||
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}"
|
||||
|
||||
|
||||
@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.VectorStoreFactory.create')
|
||||
@patch('mem0.utils.factory.LlmFactory.create')
|
||||
|
||||
Reference in New Issue
Block a user