File size: 10,786 Bytes
2415446
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
05e7f80
2415446
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
"""In-memory graph for one messaging conversation tree."""

from loguru import logger

from ..models import MessageScope
from .identity import TreeIdentity
from .node import MessageNode, MessageReferenceKind, MessageState
from .snapshot import (
    TreeSnapshot,
    node_from_snapshot,
    node_scope_from_snapshot,
    node_to_snapshot,
)


class MessageTreeGraph:
    """Own parent/child links, node lookup, and status-message lookup."""

    def __init__(self, root_node: MessageNode) -> None:
        if root_node.status_message_id == root_node.node_id:
            raise ValueError("Prompt and status message IDs must be distinct")
        self.root_id = root_node.node_id
        self.identity = TreeIdentity(
            scope=root_node.scope,
            root_id=root_node.node_id,
        )
        self._nodes: dict[str, MessageNode] = {root_node.node_id: root_node}
        self._status_to_node: dict[str, str] = {}
        if root_node.status_message_id is not None:
            self._status_to_node[root_node.status_message_id] = root_node.node_id

    def add_node(
        self,
        *,
        node_id: str,
        scope: MessageScope,
        prompt: str,
        status_message_id: str,
        parent_id: str,
        parent_reference_id: str,
    ) -> MessageNode:
        if scope != self.identity.scope:
            raise ValueError("A reply cannot cross platform chat boundaries")
        if parent_id not in self._nodes:
            raise ValueError(f"Parent node {parent_id} not found in tree")
        parent_reference = self.resolve_reference(parent_reference_id)
        if parent_reference is None or parent_reference[0].node_id != parent_id:
            raise ValueError("Reply reference does not belong to its logical parent")
        if node_id in self._nodes:
            raise ValueError(f"Node {node_id} already exists in tree")
        if status_message_id == node_id:
            raise ValueError("Prompt and status message IDs must be distinct")
        if node_id in self._status_to_node:
            raise ValueError(f"Message reference {node_id} already exists in tree")
        if status_message_id in self._status_to_node:
            raise ValueError(
                f"Status message {status_message_id} already exists in tree"
            )
        if status_message_id in self._nodes:
            raise ValueError(
                f"Message reference {status_message_id} already exists in tree"
            )

        node = MessageNode(
            node_id=node_id,
            scope=scope,
            prompt=prompt,
            status_message_id=status_message_id,
            parent_id=parent_id,
            parent_reference_id=parent_reference_id,
            state=MessageState.PENDING,
        )
        self._nodes[node_id] = node
        self._status_to_node[status_message_id] = node_id
        self._nodes[parent_id].children_ids.append(node_id)
        logger.debug("Added node {} as child of {}", node_id, parent_id)
        return node

    def get_node(self, node_id: str) -> MessageNode | None:
        return self._nodes.get(node_id)

    def get_root(self) -> MessageNode:
        return self._nodes[self.root_id]

    def get_parent(self, node_id: str) -> MessageNode | None:
        node = self._nodes.get(node_id)
        if not node or not node.parent_id:
            return None
        return self._nodes.get(node.parent_id)

    def get_parent_session_id(self, node_id: str) -> str | None:
        parent = self.get_parent(node_id)
        return parent.session_id if parent else None

    def find_node_by_status_message(self, status_msg_id: str) -> MessageNode | None:
        node_id = self._status_to_node.get(status_msg_id)
        return self._nodes.get(node_id) if node_id else None

    def resolve_reference(
        self, reference_id: str
    ) -> tuple[MessageNode, MessageReferenceKind] | None:
        """Resolve an exact prompt or FCC status reference."""
        node = self._nodes.get(reference_id)
        if node is not None:
            return node, MessageReferenceKind.PROMPT
        node = self.find_node_by_status_message(reference_id)
        if node is not None:
            return node, MessageReferenceKind.STATUS
        return None

    def all_nodes(self) -> list[MessageNode]:
        return list(self._nodes.values())

    def get_descendants(self, node_id: str) -> list[str]:
        if node_id not in self._nodes:
            return []
        result: list[str] = []
        stack = [node_id]
        while stack:
            current_id = stack.pop()
            result.append(current_id)
            node = self._nodes.get(current_id)
            if node:
                stack.extend(node.children_ids)
        return result

    def get_reference_descendants(self, reference_id: str) -> list[str]:
        """Return the literal platform reply subtree rooted at a reference."""
        if self.resolve_reference(reference_id) is None:
            return []

        children: dict[str, list[str]] = {}
        for node in self._nodes.values():
            if node.status_message_id is not None:
                children.setdefault(node.node_id, []).append(node.status_message_id)
            if node.parent_reference_id is not None:
                children.setdefault(node.parent_reference_id, []).append(node.node_id)

        result: list[str] = []
        stack = [reference_id]
        while stack:
            current_id = stack.pop()
            result.append(current_id)
            stack.extend(children.get(current_id, ()))
        return result

    def remove_nodes(self, node_ids: set[str]) -> None:
        """Remove an exact set of nodes after reference-subtree calculation."""
        for node_id in node_ids:
            node = self._nodes.get(node_id)
            if node is None:
                continue
            if node.parent_id is not None:
                parent = self._nodes.get(node.parent_id)
                if parent is not None:
                    parent.children_ids = [
                        child_id
                        for child_id in parent.children_ids
                        if child_id != node_id
                    ]
            if node.status_message_id is not None:
                self._status_to_node.pop(node.status_message_id, None)
            self._nodes.pop(node_id, None)

    def clear_status(self, node_id: str) -> None:
        """Remove one status reference while preserving its prompt node."""
        node = self._nodes.get(node_id)
        if node is None or node.status_message_id is None:
            return
        self._status_to_node.pop(node.status_message_id, None)
        node.clear_status()

    def all_reference_ids(self) -> set[str]:
        """Return every prompt and live FCC status reference in the tree."""
        references = set(self._nodes)
        references.update(self._status_to_node)
        return references

    def snapshot(self) -> TreeSnapshot:
        return TreeSnapshot(
            scope=self.identity.scope,
            root_id=self.root_id,
            nodes={
                node_id: node_to_snapshot(node) for node_id, node in self._nodes.items()
            },
        )

    @classmethod
    def from_snapshot(cls, snapshot: TreeSnapshot) -> "MessageTreeGraph":
        root_data = snapshot.nodes[snapshot.root_id]
        if not isinstance(root_data, dict):
            raise ValueError("Tree snapshot contains an invalid root node")
        if node_scope_from_snapshot(root_data) not in (None, snapshot.scope):
            raise ValueError("Tree snapshot contains a cross-scope node")
        root_node = node_from_snapshot(root_data, snapshot.scope)
        if root_node.node_id != snapshot.root_id:
            raise ValueError("Tree snapshot root key does not match its node ID")
        graph = cls(root_node)
        reference_owner = {root_node.node_id: root_node.node_id}
        if root_node.status_message_id is not None:
            reference_owner[root_node.status_message_id] = root_node.node_id
        for snapshot_node_id, node_data in snapshot.nodes.items():
            if snapshot_node_id == snapshot.root_id:
                continue
            if not isinstance(node_data, dict):
                raise ValueError("Tree snapshot contains an invalid node")
            if node_scope_from_snapshot(node_data) not in (None, snapshot.scope):
                raise ValueError("Tree snapshot contains a cross-scope node")
            node = node_from_snapshot(node_data, snapshot.scope)
            if str(snapshot_node_id) != node.node_id:
                raise ValueError("Tree snapshot node key does not match its node ID")
            if node.status_message_id == node.node_id:
                raise ValueError("Prompt and status message IDs must be distinct")
            if node.node_id in graph._nodes:
                raise ValueError(f"Duplicate node {node.node_id} in tree snapshot")
            references = {node.node_id}
            if node.status_message_id is not None:
                references.add(node.status_message_id)
            for reference in references:
                owner = reference_owner.get(reference)
                if owner is not None and owner != node.node_id:
                    raise ValueError(
                        f"Duplicate message reference {reference} in tree snapshot"
                    )
                reference_owner[reference] = node.node_id
            graph._nodes[node.node_id] = node
            if node.status_message_id is not None:
                graph._status_to_node[node.status_message_id] = node.node_id

        if root_node.parent_id is not None or root_node.parent_reference_id is not None:
            raise ValueError("Tree snapshot root cannot have a parent")
        for node in graph._nodes.values():
            if node.node_id == graph.root_id:
                continue
            if node.parent_id is None or node.parent_id not in graph._nodes:
                raise ValueError(f"Node {node.node_id} has no valid parent")
            if node.parent_reference_id is None:
                raise ValueError(f"Node {node.node_id} has no exact parent reference")
            parent_reference = graph.resolve_reference(node.parent_reference_id)
            if (
                parent_reference is None
                or parent_reference[0].node_id != node.parent_id
            ):
                raise ValueError(
                    f"Node {node.node_id} has an invalid exact parent reference"
                )
            graph._nodes[node.parent_id].children_ids.append(node.node_id)
        if set(graph.get_descendants(graph.root_id)) != set(graph._nodes):
            raise ValueError("Tree snapshot contains a disconnected branch")
        return graph