Skip to main content

pe_graph/
pending_writes.rs

1//! Pending writes tracking for fault tolerance within a superstep.
2//!
3//! When multiple nodes execute in parallel, some may succeed while others fail.
4//! `PendingWrites` records which nodes produced valid updates so that on resume,
5//! successful nodes don't re-execute.
6
7use pe_core::error::PeError;
8use pe_core::state::StateUpdate;
9
10/// Tracks node write outcomes within a single superstep.
11///
12/// # Example
13///
14/// ```ignore
15/// let mut writes = PendingWrites::new();
16/// writes.record_success("chat", update);
17/// writes.record_failure("tools", &PeError::ToolExecution { .. });
18///
19/// if writes.has_failures() {
20///     // Only re-run failed nodes on resume
21/// }
22/// ```
23#[derive(Debug, Clone)]
24pub struct PendingWrites<U: StateUpdate> {
25    successes: Vec<(String, U)>,
26    failures: Vec<(String, String)>,
27}
28
29impl<U: StateUpdate> PendingWrites<U> {
30    /// Create an empty write tracker.
31    pub fn new() -> Self {
32        Self {
33            successes: Vec::new(),
34            failures: Vec::new(),
35        }
36    }
37
38    /// Record a successful node execution with its update.
39    pub fn record_success(&mut self, node_name: &str, update: U) {
40        self.successes.push((node_name.to_string(), update));
41    }
42
43    /// Record a failed node execution.
44    pub fn record_failure(&mut self, node_name: &str, error: &PeError) {
45        self.failures
46            .push((node_name.to_string(), error.to_string()));
47    }
48
49    /// Returns true if any node failed in this superstep.
50    pub fn has_failures(&self) -> bool {
51        !self.failures.is_empty()
52    }
53
54    /// View successful node writes.
55    pub fn successes(&self) -> &[(String, U)] {
56        &self.successes
57    }
58
59    /// View failed nodes and their error messages.
60    pub fn failures(&self) -> &[(String, String)] {
61        &self.failures
62    }
63
64    /// Drain all successful writes, leaving the tracker empty.
65    pub fn drain_successes(&mut self) -> Vec<(String, U)> {
66        std::mem::take(&mut self.successes)
67    }
68}
69
70impl<U: StateUpdate> Default for PendingWrites<U> {
71    fn default() -> Self {
72        Self::new()
73    }
74}
75
76#[cfg(test)]
77mod tests {
78    use super::*;
79    use serde::{Deserialize, Serialize};
80
81    #[derive(Debug, Clone, Default, Serialize, Deserialize)]
82    struct FakeUpdate {
83        value: i32,
84    }
85    impl StateUpdate for FakeUpdate {}
86
87    #[test]
88    fn test_record_and_access() {
89        let mut writes = PendingWrites::new();
90        writes.record_success("node_a", FakeUpdate { value: 1 });
91        writes.record_success("node_b", FakeUpdate { value: 2 });
92
93        assert_eq!(writes.successes().len(), 2);
94        assert_eq!(writes.successes()[0].0, "node_a");
95        assert_eq!(writes.successes()[1].1.value, 2);
96        assert!(!writes.has_failures());
97    }
98
99    #[test]
100    fn test_record_failure() {
101        let mut writes: PendingWrites<FakeUpdate> = PendingWrites::new();
102        writes.record_failure(
103            "bad_node",
104            &PeError::Internal {
105                details: "boom".into(),
106            },
107        );
108
109        assert!(writes.has_failures());
110        assert_eq!(writes.failures().len(), 1);
111        assert_eq!(writes.failures()[0].0, "bad_node");
112    }
113
114    #[test]
115    fn test_drain_successes() {
116        let mut writes = PendingWrites::new();
117        writes.record_success("a", FakeUpdate { value: 10 });
118        writes.record_success("b", FakeUpdate { value: 20 });
119
120        let drained = writes.drain_successes();
121        assert_eq!(drained.len(), 2);
122        assert!(writes.successes().is_empty());
123    }
124}