tokio_prompt_orchestrator/
conversation_graph.rs1use std::collections::{HashMap, HashSet, VecDeque};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub struct NodeId(pub u64);
14
15#[derive(Debug, Clone)]
17pub struct ConversationNode {
18 pub id: NodeId,
20 pub role: String,
22 pub content: String,
24 pub tokens: usize,
26 pub metadata: HashMap<String, String>,
28}
29
30#[derive(Debug, thiserror::Error, PartialEq, Eq)]
32pub enum GraphError {
33 #[error("adding this edge would create a cycle")]
35 CycleDetected,
36 #[error("node {0:?} not found")]
38 NodeNotFound(NodeId),
39 #[error("invalid edge: {0}")]
41 InvalidEdge(String),
42}
43
44#[derive(Debug, Default)]
53pub struct ConversationGraph {
54 nodes: HashMap<NodeId, ConversationNode>,
56 edges: HashMap<NodeId, Vec<NodeId>>,
58 next_id: u64,
60}
61
62impl ConversationGraph {
63 pub fn new() -> Self {
65 Self::default()
66 }
67
68 fn fresh_id(&mut self) -> NodeId {
73 let id = NodeId(self.next_id);
74 self.next_id += 1;
75 id
76 }
77
78 fn is_reachable(&self, start: NodeId, target: NodeId) -> bool {
81 let mut visited: HashSet<NodeId> = HashSet::new();
82 let mut stack: Vec<NodeId> = vec![start];
83 while let Some(current) = stack.pop() {
84 if current == target {
85 return true;
86 }
87 if visited.insert(current) {
88 if let Some(succs) = self.edges.get(¤t) {
89 stack.extend_from_slice(succs);
90 }
91 }
92 }
93 false
94 }
95
96 pub fn add_node(&mut self, role: impl Into<String>, content: impl Into<String>, tokens: usize) -> NodeId {
102 let id = self.fresh_id();
103 let node = ConversationNode {
104 id,
105 role: role.into(),
106 content: content.into(),
107 tokens,
108 metadata: HashMap::new(),
109 };
110 self.nodes.insert(id, node);
111 self.edges.entry(id).or_default();
112 id
113 }
114
115 pub fn add_edge(&mut self, from: NodeId, to: NodeId) -> Result<(), GraphError> {
123 if !self.nodes.contains_key(&from) {
124 return Err(GraphError::NodeNotFound(from));
125 }
126 if !self.nodes.contains_key(&to) {
127 return Err(GraphError::NodeNotFound(to));
128 }
129 if from == to {
130 return Err(GraphError::InvalidEdge("self-loop".into()));
131 }
132 if self.is_reachable(to, from) {
134 return Err(GraphError::CycleDetected);
135 }
136 self.edges.entry(from).or_default().push(to);
137 Ok(())
138 }
139
140 pub fn branch_from(&mut self, node_id: NodeId) -> Result<NodeId, GraphError> {
147 let original = self
148 .nodes
149 .get(&node_id)
150 .ok_or(GraphError::NodeNotFound(node_id))?
151 .clone();
152 let new_id = self.fresh_id();
153 let clone = ConversationNode {
154 id: new_id,
155 role: original.role,
156 content: original.content,
157 tokens: original.tokens,
158 metadata: original.metadata,
159 };
160 self.nodes.insert(new_id, clone);
161 self.edges.entry(new_id).or_default();
162 self.edges.entry(node_id).or_default().push(new_id);
164 Ok(new_id)
165 }
166
167 pub fn merge_paths(
177 &mut self,
178 path_a: &[NodeId],
179 path_b: &[NodeId],
180 ) -> Result<NodeId, GraphError> {
181 let (Some(&tail_a), Some(&tail_b)) = (path_a.last(), path_b.last()) else {
182 return Err(GraphError::InvalidEdge("paths must be non-empty".into()));
183 };
184 let merge_id = self.add_node("merge", "", 0);
185 self.add_edge(tail_a, merge_id)?;
186 if tail_b != tail_a {
187 self.add_edge(tail_b, merge_id)?;
188 }
189 Ok(merge_id)
190 }
191
192 pub fn linearize(&self, root: NodeId) -> Result<Vec<NodeId>, GraphError> {
202 if !self.nodes.contains_key(&root) {
203 return Err(GraphError::NodeNotFound(root));
204 }
205
206 let mut reachable: HashSet<NodeId> = HashSet::new();
208 let mut dfs_stack = vec![root];
209 while let Some(n) = dfs_stack.pop() {
210 if reachable.insert(n) {
211 if let Some(succs) = self.edges.get(&n) {
212 dfs_stack.extend_from_slice(succs);
213 }
214 }
215 }
216
217 let mut in_degree: HashMap<NodeId, usize> = reachable.iter().map(|&n| (n, 0)).collect();
219 for &n in &reachable {
220 if let Some(succs) = self.edges.get(&n) {
221 for &s in succs {
222 if reachable.contains(&s) {
223 *in_degree.entry(s).or_insert(0) += 1;
224 }
225 }
226 }
227 }
228
229 let mut queue: VecDeque<NodeId> = in_degree
231 .iter()
232 .filter(|(_, &d)| d == 0)
233 .map(|(&n, _)| n)
234 .collect();
235 let mut order: Vec<NodeId> = Vec::with_capacity(reachable.len());
236
237 while let Some(n) = queue.pop_front() {
238 order.push(n);
239 if let Some(succs) = self.edges.get(&n) {
240 for &s in succs {
241 if reachable.contains(&s) {
242 let d = in_degree.entry(s).or_insert(0);
243 *d = d.saturating_sub(1);
244 if *d == 0 {
245 queue.push_back(s);
246 }
247 }
248 }
249 }
250 }
251
252 if order.len() != reachable.len() {
253 return Err(GraphError::CycleDetected);
254 }
255 Ok(order)
256 }
257
258 pub fn path_tokens(&self, path: &[NodeId]) -> usize {
262 path.iter()
263 .filter_map(|id| self.nodes.get(id))
264 .map(|n| n.tokens)
265 .sum()
266 }
267
268 pub fn get_node(&self, id: NodeId) -> Option<&ConversationNode> {
270 self.nodes.get(&id)
271 }
272}
273
274#[cfg(test)]
275mod tests {
276 use super::*;
277
278 #[test]
279 fn add_and_retrieve_nodes() {
280 let mut g = ConversationGraph::new();
281 let a = g.add_node("user", "Hello", 5);
282 let b = g.add_node("assistant", "Hi there", 10);
283 assert_ne!(a, b);
284 assert_eq!(g.get_node(a).unwrap().tokens, 5);
285 assert_eq!(g.get_node(b).unwrap().role, "assistant");
286 }
287
288 #[test]
289 fn add_edge_success() {
290 let mut g = ConversationGraph::new();
291 let a = g.add_node("user", "A", 1);
292 let b = g.add_node("assistant", "B", 2);
293 assert!(g.add_edge(a, b).is_ok());
294 }
295
296 #[test]
297 fn add_edge_self_loop_rejected() {
298 let mut g = ConversationGraph::new();
299 let a = g.add_node("user", "A", 1);
300 assert_eq!(g.add_edge(a, a), Err(GraphError::InvalidEdge("self-loop".into())));
301 }
302
303 #[test]
304 fn add_edge_cycle_rejected() {
305 let mut g = ConversationGraph::new();
306 let a = g.add_node("user", "A", 1);
307 let b = g.add_node("assistant", "B", 2);
308 g.add_edge(a, b).unwrap();
309 assert_eq!(g.add_edge(b, a), Err(GraphError::CycleDetected));
310 }
311
312 #[test]
313 fn add_edge_longer_cycle_rejected() {
314 let mut g = ConversationGraph::new();
315 let a = g.add_node("user", "A", 1);
316 let b = g.add_node("assistant", "B", 2);
317 let c = g.add_node("user", "C", 3);
318 g.add_edge(a, b).unwrap();
319 g.add_edge(b, c).unwrap();
320 assert_eq!(g.add_edge(c, a), Err(GraphError::CycleDetected));
321 }
322
323 #[test]
324 fn add_edge_unknown_node() {
325 let mut g = ConversationGraph::new();
326 let a = g.add_node("user", "A", 1);
327 let phantom = NodeId(999);
328 assert_eq!(g.add_edge(a, phantom), Err(GraphError::NodeNotFound(phantom)));
329 assert_eq!(g.add_edge(phantom, a), Err(GraphError::NodeNotFound(phantom)));
330 }
331
332 #[test]
333 fn branch_from_creates_child() {
334 let mut g = ConversationGraph::new();
335 let a = g.add_node("user", "Hello", 5);
336 let b = g.branch_from(a).unwrap();
337 assert_ne!(a, b);
338 assert_eq!(g.get_node(b).unwrap().content, "Hello");
340 }
341
342 #[test]
343 fn merge_paths_creates_merge_node() {
344 let mut g = ConversationGraph::new();
345 let a = g.add_node("user", "A", 1);
346 let b = g.add_node("assistant", "B", 2);
347 let c = g.add_node("user", "C", 3);
348 let merge = g.merge_paths(&[a], &[b, c]).unwrap();
349 let merge_node = g.get_node(merge).unwrap();
350 assert_eq!(merge_node.role, "merge");
351 }
352
353 #[test]
354 fn linearize_simple_chain() {
355 let mut g = ConversationGraph::new();
356 let a = g.add_node("user", "A", 1);
357 let b = g.add_node("assistant", "B", 2);
358 let c = g.add_node("user", "C", 3);
359 g.add_edge(a, b).unwrap();
360 g.add_edge(b, c).unwrap();
361 let order = g.linearize(a).unwrap();
362 assert_eq!(order, vec![a, b, c]);
363 }
364
365 #[test]
366 fn linearize_diamond() {
367 let mut g = ConversationGraph::new();
368 let root = g.add_node("system", "Root", 0);
369 let left = g.add_node("user", "Left", 1);
370 let right = g.add_node("user", "Right", 2);
371 let merge = g.add_node("assistant", "Merge", 3);
372 g.add_edge(root, left).unwrap();
373 g.add_edge(root, right).unwrap();
374 g.add_edge(left, merge).unwrap();
375 g.add_edge(right, merge).unwrap();
376 let order = g.linearize(root).unwrap();
377 assert_eq!(order.len(), 4);
378 assert_eq!(order[0], root);
380 assert_eq!(order[3], merge);
381 }
382
383 #[test]
384 fn path_tokens_sums_correctly() {
385 let mut g = ConversationGraph::new();
386 let a = g.add_node("user", "A", 10);
387 let b = g.add_node("assistant", "B", 20);
388 let c = g.add_node("user", "C", 5);
389 assert_eq!(g.path_tokens(&[a, b, c]), 35);
390 }
391
392 #[test]
393 fn linearize_unknown_root() {
394 let g = ConversationGraph::new();
395 assert_eq!(
396 g.linearize(NodeId(42)),
397 Err(GraphError::NodeNotFound(NodeId(42)))
398 );
399 }
400}