Skip to main content

tokio_prompt_orchestrator/
conversation_graph.rs

1//! DAG-based conversation branching.
2//!
3//! A [`ConversationGraph`] is a directed acyclic graph (DAG) where each node
4//! represents one conversation turn (a role + content pair).  Edges represent
5//! the "follows from" relationship between turns.  The graph enforces the
6//! acyclicity invariant on every [`ConversationGraph::add_edge`] call via a
7//! depth-first reachability check.
8
9use std::collections::{HashMap, HashSet, VecDeque};
10
11/// Opaque node identifier.
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub struct NodeId(pub u64);
14
15/// A single node in the conversation DAG.
16#[derive(Debug, Clone)]
17pub struct ConversationNode {
18    /// Unique identifier of this node.
19    pub id: NodeId,
20    /// Speaker role, e.g. `"user"`, `"assistant"`, or `"system"`.
21    pub role: String,
22    /// Text content of the turn.
23    pub content: String,
24    /// Token count for this turn.
25    pub tokens: usize,
26    /// Arbitrary key-value metadata.
27    pub metadata: HashMap<String, String>,
28}
29
30/// Errors that can be returned by graph operations.
31#[derive(Debug, thiserror::Error, PartialEq, Eq)]
32pub enum GraphError {
33    /// The requested edge would introduce a cycle into the DAG.
34    #[error("adding this edge would create a cycle")]
35    CycleDetected,
36    /// One of the referenced node IDs does not exist in the graph.
37    #[error("node {0:?} not found")]
38    NodeNotFound(NodeId),
39    /// The edge specification is logically invalid (e.g. self-loop).
40    #[error("invalid edge: {0}")]
41    InvalidEdge(String),
42}
43
44/// Directed acyclic graph of conversation nodes.
45///
46/// # Invariants
47///
48/// - Every `NodeId` returned by `add_node` is unique and stable for the
49///   lifetime of the graph.
50/// - `add_edge` rejects any edge that would create a cycle.
51/// - `linearize` always returns a valid topological ordering.
52#[derive(Debug, Default)]
53pub struct ConversationGraph {
54    /// All nodes indexed by their ID.
55    nodes: HashMap<NodeId, ConversationNode>,
56    /// Forward adjacency list: node -> list of successor nodes.
57    edges: HashMap<NodeId, Vec<NodeId>>,
58    /// Monotonically increasing counter for generating fresh `NodeId`s.
59    next_id: u64,
60}
61
62impl ConversationGraph {
63    /// Create an empty graph.
64    pub fn new() -> Self {
65        Self::default()
66    }
67
68    // -------------------------------------------------------------------------
69    // Internal helpers
70    // -------------------------------------------------------------------------
71
72    fn fresh_id(&mut self) -> NodeId {
73        let id = NodeId(self.next_id);
74        self.next_id += 1;
75        id
76    }
77
78    /// Return `true` if `target` is reachable from `start` by following forward
79    /// edges (DFS).  Used to detect would-be cycles before inserting an edge.
80    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(&current) {
89                    stack.extend_from_slice(succs);
90                }
91            }
92        }
93        false
94    }
95
96    // -------------------------------------------------------------------------
97    // Public API
98    // -------------------------------------------------------------------------
99
100    /// Add a new node to the graph and return its [`NodeId`].
101    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    /// Add a directed edge `from -> to`.
116    ///
117    /// # Errors
118    ///
119    /// - [`GraphError::NodeNotFound`] if either `from` or `to` is unknown.
120    /// - [`GraphError::InvalidEdge`] if `from == to` (self-loop).
121    /// - [`GraphError::CycleDetected`] if the edge would create a cycle.
122    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        // Adding from->to creates a cycle iff `from` is already reachable from `to`.
133        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    /// Clone the node at `node_id`, attach the clone as a child, and return
141    /// the clone's [`NodeId`].
142    ///
143    /// # Errors
144    ///
145    /// Returns [`GraphError::NodeNotFound`] if `node_id` is unknown.
146    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        // Attach original -> clone
163        self.edges.entry(node_id).or_default().push(new_id);
164        Ok(new_id)
165    }
166
167    /// Create a merge node that references both `path_a` and `path_b`.
168    ///
169    /// The merge node has role `"merge"` and empty content.  Edges are added
170    /// from the last node of each path to the merge node.
171    ///
172    /// # Errors
173    ///
174    /// - [`GraphError::InvalidEdge`] if either path is empty.
175    /// - Propagates errors from `add_edge`.
176    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    /// Topological sort of all nodes reachable from `root` using Kahn's algorithm.
193    ///
194    /// Returns nodes in breadth-first topological order.
195    ///
196    /// # Errors
197    ///
198    /// - [`GraphError::NodeNotFound`] if `root` is unknown.
199    /// - [`GraphError::CycleDetected`] if the reachable subgraph is not a DAG
200    ///   (should not happen if all edges were added through `add_edge`).
201    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        // Collect the subgraph reachable from root.
207        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        // Build in-degree counts restricted to reachable nodes.
218        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        // Kahn's BFS starting from nodes with in-degree 0 that are reachable.
230        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    /// Sum the token counts of all nodes in `path`.
259    ///
260    /// Unknown node IDs are silently skipped.
261    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    /// Borrow the node for `id`, if it exists.
269    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        // The clone should have the same content
339        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        // root must come first, merge must come last
379        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}