Spaces:
Sleeping
Sleeping
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
|