From 8a470779d8f3dc89ccff2abd331a3154a6219a11 Mon Sep 17 00:00:00 2001 From: Alishahryar1 Date: Thu, 18 Jun 2026 12:53:31 -0700 Subject: [PATCH] Fix messaging clear persisted status mappings --- messaging/session.py | 39 ++++++++-- tests/messaging/test_restart_reply_restore.py | 72 ++++++++++++++++-- .../test_session_store_edge_cases.py | 74 +++++++++++++++++++ 3 files changed, 174 insertions(+), 11 deletions(-) diff --git a/messaging/session.py b/messaging/session.py index bbb0cdc1..89297726 100644 --- a/messaging/session.py +++ b/messaging/session.py @@ -240,6 +240,36 @@ class SessionStore: # ==================== Tree Methods ==================== + @staticmethod + def _tree_lookup_ids(tree_data: dict) -> set[str]: + """Return lookup IDs represented by serialized tree nodes.""" + lookup_ids: set[str] = set() + nodes = tree_data.get("nodes", {}) + if not isinstance(nodes, dict): + return lookup_ids + + for node_key, node_data in nodes.items(): + lookup_ids.add(str(node_key)) + if not isinstance(node_data, dict): + continue + node_id = node_data.get("node_id") + if node_id is not None: + lookup_ids.add(str(node_id)) + status_message_id = node_data.get("status_message_id") + if status_message_id is not None: + lookup_ids.add(str(status_message_id)) + return lookup_ids + + def _remove_tree_lookup_ids_unlocked(self, root_id: str) -> None: + """Remove all lookup IDs currently pointing at a root. Caller holds lock.""" + stale_lookup_ids = [ + lookup_id + for lookup_id, mapped_root_id in self._node_to_tree.items() + if mapped_root_id == root_id + ] + for lookup_id in stale_lookup_ids: + self._node_to_tree.pop(lookup_id, None) + def save_tree(self, root_id: str, tree_data: dict) -> None: """ Save a message tree. @@ -251,9 +281,9 @@ class SessionStore: with self._lock: self._trees[root_id] = tree_data - # Update node-to-tree mapping - for node_id in tree_data.get("nodes", {}): - self._node_to_tree[node_id] = root_id + self._remove_tree_lookup_ids_unlocked(root_id) + for lookup_id in self._tree_lookup_ids(tree_data): + self._node_to_tree[lookup_id] = root_id self._schedule_save() logger.debug(f"Saved tree {root_id}") @@ -281,8 +311,7 @@ class SessionStore: with self._lock: tree_data = self._trees.pop(root_id, None) if tree_data: - for node_id in tree_data.get("nodes", {}): - self._node_to_tree.pop(node_id, None) + self._remove_tree_lookup_ids_unlocked(root_id) self._schedule_save() def get_all_trees(self) -> dict[str, dict]: diff --git a/tests/messaging/test_restart_reply_restore.py b/tests/messaging/test_restart_reply_restore.py index 634d4cea..82342783 100644 --- a/tests/messaging/test_restart_reply_restore.py +++ b/tests/messaging/test_restart_reply_restore.py @@ -69,7 +69,7 @@ async def test_reply_to_old_status_message_after_restore_routes_to_parent( @pytest.mark.asyncio -async def test_reply_to_old_status_message_without_mapping_creates_new_conversation( +async def test_save_tree_persists_status_message_mapping_without_manual_register( tmp_path, mock_platform, mock_cli_manager ): store_path = tmp_path / "sessions.json" @@ -86,7 +86,6 @@ async def test_reply_to_old_status_message_without_mapping_creates_new_conversat tree = await handler1.tree_queue.create_tree( "A", a_incoming, status_message_id="status_A" ) - # Intentionally do NOT register "status_A" mapping. store.save_tree(tree.root_id, tree.to_dict()) store.flush_pending_save() @@ -116,7 +115,68 @@ async def test_reply_to_old_status_message_without_mapping_creates_new_conversat with patch.object(handler2.tree_queue, "enqueue", AsyncMock(return_value=False)): await handler2.handle_message(reply) - # Since the mapping is missing, this should be treated as a new conversation. - new_tree = handler2.tree_queue.get_tree_for_node("R1") - assert new_tree is not None - assert new_tree.root_id == "R1" + restored_tree = handler2.tree_queue.get_tree_for_node("A") + assert restored_tree is not None + node_r1 = restored_tree.get_node("R1") + assert node_r1 is not None + assert node_r1.parent_id == "A" + + +@pytest.mark.asyncio +async def test_reply_clear_purges_removed_status_mapping_from_persisted_store( + tmp_path, mock_platform, mock_cli_manager +): + store_path = tmp_path / "sessions.json" + store = SessionStore(storage_path=str(store_path)) + handler = MessagingWorkflow(mock_platform, mock_cli_manager, store) + + root_incoming = IncomingMessage( + text="root", + chat_id="chat_1", + user_id="user_1", + message_id="root", + platform="telegram", + ) + tree = await handler.tree_queue.create_tree( + "root", root_incoming, status_message_id="root_status" + ) + handler.tree_queue.register_node("root_status", tree.root_id) + + child_incoming = IncomingMessage( + text="child", + chat_id="chat_1", + user_id="user_1", + message_id="child", + platform="telegram", + reply_to_message_id="root", + ) + await handler.tree_queue.add_to_tree( + "root", "child", child_incoming, status_message_id="child_status" + ) + handler.tree_queue.register_node("child_status", tree.root_id) + store.save_tree(tree.root_id, tree.to_dict()) + + clear_reply = IncomingMessage( + text="/clear", + chat_id="chat_1", + user_id="user_1", + message_id="clear_command", + platform="telegram", + reply_to_message_id="child", + ) + + await handler.handle_message(clear_reply) + store.flush_pending_save() + + restored_store = SessionStore(storage_path=str(store_path)) + restored_tree_queue = TreeQueueManager.from_dict( + { + "trees": restored_store.get_all_trees(), + "node_to_tree": restored_store.get_node_mapping(), + } + ) + + assert restored_tree_queue.get_tree_for_node("root") is not None + assert restored_tree_queue.get_tree_for_node("root_status") is not None + assert restored_tree_queue.get_tree_for_node("child") is None + assert restored_tree_queue.get_tree_for_node("child_status") is None diff --git a/tests/messaging/test_session_store_edge_cases.py b/tests/messaging/test_session_store_edge_cases.py index 8656c366..ca3d8b83 100644 --- a/tests/messaging/test_session_store_edge_cases.py +++ b/tests/messaging/test_session_store_edge_cases.py @@ -15,6 +15,13 @@ def tmp_store(tmp_path): return SessionStore(storage_path=path) +def _tree_node(node_id: str, status_message_id: str) -> dict: + return { + "node_id": node_id, + "status_message_id": status_message_id, + } + + class TestSessionStoreLoadEdgeCases: """Tests for loading corrupted/malformed data.""" @@ -90,6 +97,73 @@ class TestSessionStoreSaveEdgeCases: tmp_store._write_data(tmp_store._snapshot()) +class TestSessionStoreTreeMappings: + def test_save_tree_rebuilds_lookup_ids_for_that_root(self, tmp_path): + path = str(tmp_path / "sessions.json") + store = SessionStore(storage_path=path) + store.register_node("unrelated_status", "other_root") + + store.save_tree( + "root", + { + "root_id": "root", + "nodes": { + "root": _tree_node("root", "root_status"), + "child": _tree_node("child", "child_status"), + }, + }, + ) + + mapping = store.get_node_mapping() + assert mapping["root"] == "root" + assert mapping["root_status"] == "root" + assert mapping["child"] == "root" + assert mapping["child_status"] == "root" + + store.save_tree( + "root", + { + "root_id": "root", + "nodes": { + "root": _tree_node("root", "root_status"), + }, + }, + ) + + mapping = store.get_node_mapping() + assert mapping["root"] == "root" + assert mapping["root_status"] == "root" + assert "child" not in mapping + assert "child_status" not in mapping + assert mapping["unrelated_status"] == "other_root" + + def test_remove_tree_removes_all_lookup_ids_for_that_root(self, tmp_path): + path = str(tmp_path / "sessions.json") + store = SessionStore(storage_path=path) + store.register_node("old_status", "root") + store.register_node("unrelated_status", "other_root") + store.save_tree( + "root", + { + "root_id": "root", + "nodes": { + "root": _tree_node("root", "root_status"), + "child": _tree_node("child", "child_status"), + }, + }, + ) + + store.remove_tree("root") + + mapping = store.get_node_mapping() + assert "root" not in mapping + assert "root_status" not in mapping + assert "child" not in mapping + assert "child_status" not in mapping + assert "old_status" not in mapping + assert mapping["unrelated_status"] == "other_root" + + class TestSessionStoreAtomicWrites: """Atomic persistence: failed replace must not truncate the prior file."""