Skip to main content

adk_graph/
deferred.rs

1//! Deferred node (fan-in barrier) support for graph workflows.
2//!
3//! Provides fan-in barrier semantics for nodes that wait on multiple upstream
4//! parallel paths before executing. This enables scatter-gather patterns where
5//! work is distributed across parallel branches and then collected at a single
6//! synchronization point.
7//!
8//! # Overview
9//!
10//! A deferred node is declared with a [`DeferredNodeConfig`] that specifies:
11//! - [`MergeStrategy`]: How upstream outputs are combined (collect, merge maps, first, or custom).
12//! - `fan_in_timeout`: Optional maximum wait duration for all upstream paths.
13//!
14//! The [`FanInTracker`] tracks which upstream paths have completed and merges
15//! their outputs according to the configured strategy.
16//!
17//! # Example
18//!
19//! ```rust
20//! use std::time::Duration;
21//! use adk_graph::deferred::{DeferredNodeConfig, FanInTracker, MergeStrategy};
22//! use serde_json::json;
23//!
24//! // Configure a deferred node that collects all upstream outputs
25//! let config = DeferredNodeConfig {
26//!     merge_strategy: MergeStrategy::Collect,
27//!     fan_in_timeout: Some(Duration::from_secs(30)),
28//!     ..Default::default()
29//! };
30//!
31//! // Track upstream completions
32//! let mut tracker = FanInTracker::new(vec!["branch_a", "branch_b", "branch_c"]);
33//!
34//! tracker.record("branch_a", json!({"result": 1}));
35//! tracker.record("branch_b", json!({"result": 2}));
36//! assert!(!tracker.is_ready());
37//!
38//! tracker.record("branch_c", json!({"result": 3}));
39//! assert!(tracker.is_ready());
40//!
41//! // Merge outputs using the configured strategy
42//! let merged = tracker.merge(&config.merge_strategy);
43//! assert_eq!(merged, json!([{"result": 1}, {"result": 2}, {"result": 3}]));
44//! ```
45
46use std::collections::{HashMap, HashSet};
47use std::fmt;
48use std::sync::Arc;
49use std::time::Duration;
50
51use serde_json::Value;
52
53/// How to combine outputs from multiple upstream parallel paths.
54///
55/// The merge strategy determines how the collected outputs from all upstream
56/// branches are combined into a single value for the deferred node's input.
57///
58/// # Example
59///
60/// ```rust
61/// use adk_graph::deferred::MergeStrategy;
62///
63/// // Default strategy collects all outputs into a Vec
64/// let strategy = MergeStrategy::default();
65/// assert!(matches!(strategy, MergeStrategy::Collect));
66///
67/// // MergeMap combines all output maps with last-write-wins
68/// let strategy = MergeStrategy::MergeMap;
69/// ```
70#[derive(Clone, Default)]
71pub enum MergeStrategy {
72    /// Collect all outputs into a `Vec<Value>`.
73    ///
74    /// Outputs are ordered by the insertion order of source nodes
75    /// (the order in which they were recorded).
76    #[default]
77    Collect,
78
79    /// Merge all output maps into a single map (last-write-wins on key conflict).
80    ///
81    /// Each upstream output is expected to be a JSON object. Non-object outputs
82    /// are skipped. When multiple outputs contain the same key, the value from
83    /// the later-recorded source wins.
84    MergeMap,
85
86    /// Use only the first completed output.
87    ///
88    /// Returns the output from whichever upstream path completed first
89    /// (i.e., was recorded first).
90    First,
91
92    /// Custom merge function.
93    ///
94    /// Accepts a closure that takes all collected outputs and produces a
95    /// single merged value.
96    ///
97    /// # Example
98    ///
99    /// ```rust
100    /// use std::sync::Arc;
101    /// use adk_graph::deferred::MergeStrategy;
102    /// use serde_json::{json, Value};
103    ///
104    /// let strategy = MergeStrategy::Custom(Arc::new(|outputs: Vec<Value>| {
105    ///     json!({ "count": outputs.len() })
106    /// }));
107    /// ```
108    Custom(Arc<dyn Fn(Vec<Value>) -> Value + Send + Sync>),
109}
110
111impl fmt::Debug for MergeStrategy {
112    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
113        match self {
114            Self::Collect => write!(f, "Collect"),
115            Self::MergeMap => write!(f, "MergeMap"),
116            Self::First => write!(f, "First"),
117            Self::Custom(_) => write!(f, "Custom(<fn>)"),
118        }
119    }
120}
121
122/// Configuration for a deferred (fan-in) node.
123///
124/// A deferred node waits for all upstream parallel paths to complete before
125/// executing. The configuration controls how outputs are merged and how long
126/// the node waits.
127///
128/// # Example
129///
130/// ```rust
131/// use std::time::Duration;
132/// use adk_graph::deferred::{DeferredNodeConfig, MergeStrategy};
133///
134/// let config = DeferredNodeConfig {
135///     merge_strategy: MergeStrategy::MergeMap,
136///     fan_in_timeout: Some(Duration::from_secs(60)),
137///     ..Default::default()
138/// };
139/// ```
140#[derive(Debug, Clone, Default)]
141pub struct DeferredNodeConfig {
142    /// Strategy for combining upstream outputs.
143    pub merge_strategy: MergeStrategy,
144
145    /// Maximum time to wait for all upstream paths to complete.
146    ///
147    /// - `None`: Wait indefinitely for all upstream paths.
148    /// - `Some(duration)`: If the timeout expires and some paths have completed,
149    ///   proceed with partial results. If zero paths have completed, return
150    ///   `GraphError::FanInTimedOut`.
151    pub fan_in_timeout: Option<Duration>,
152    /// How many predecessors must have arrived for `fan_in_timeout` to release
153    /// the node with partial results.
154    ///
155    /// `None` means one, which is the behaviour this field replaced. Set it to
156    /// the full predecessor count to require all of them even after a timeout, or
157    /// to any value between for an n-of-m join: release once that many branches
158    /// have answered and abandon the rest.
159    pub min_predecessors: Option<usize>,
160}
161
162/// Tracks which upstream paths have completed for a deferred node.
163///
164/// The tracker maintains a set of expected source nodes and records their
165/// outputs as they arrive. Once all expected sources have reported, the
166/// tracker is ready and outputs can be merged.
167///
168/// # Example
169///
170/// ```rust
171/// use adk_graph::deferred::{FanInTracker, MergeStrategy};
172/// use serde_json::json;
173///
174/// let mut tracker = FanInTracker::new(vec!["node_a", "node_b"]);
175///
176/// assert!(!tracker.is_ready());
177/// assert_eq!(tracker.received_count(), 0);
178/// assert_eq!(tracker.expected_count(), 2);
179///
180/// tracker.record("node_a", json!("output_a"));
181/// assert!(!tracker.is_ready());
182///
183/// tracker.record("node_b", json!("output_b"));
184/// assert!(tracker.is_ready());
185///
186/// let merged = tracker.merge(&MergeStrategy::Collect);
187/// assert_eq!(merged, json!(["output_a", "output_b"]));
188/// ```
189pub struct FanInTracker {
190    /// The set of source node names we expect to receive output from.
191    expected: HashSet<String>,
192    /// Outputs received so far, keyed by source node name.
193    received: HashMap<String, Value>,
194    /// Insertion order of received outputs (for deterministic merge ordering).
195    insertion_order: Vec<String>,
196}
197
198impl FanInTracker {
199    /// Create a new tracker expecting outputs from the given source nodes.
200    ///
201    /// # Arguments
202    ///
203    /// * `expected_sources` - Names of upstream nodes that must complete
204    ///   before this deferred node can execute.
205    ///
206    /// # Example
207    ///
208    /// ```rust
209    /// use adk_graph::deferred::FanInTracker;
210    ///
211    /// let tracker = FanInTracker::new(vec!["branch_1", "branch_2", "branch_3"]);
212    /// assert_eq!(tracker.expected_count(), 3);
213    /// assert!(!tracker.is_ready());
214    /// ```
215    pub fn new(expected_sources: Vec<&str>) -> Self {
216        Self {
217            expected: expected_sources.iter().map(|s| (*s).to_string()).collect(),
218            received: HashMap::new(),
219            insertion_order: Vec::new(),
220        }
221    }
222
223    /// Returns `true` when all expected sources have reported their output.
224    ///
225    /// # Example
226    ///
227    /// ```rust
228    /// use adk_graph::deferred::FanInTracker;
229    /// use serde_json::json;
230    ///
231    /// let mut tracker = FanInTracker::new(vec!["a"]);
232    /// assert!(!tracker.is_ready());
233    ///
234    /// tracker.record("a", json!(42));
235    /// assert!(tracker.is_ready());
236    /// ```
237    pub fn is_ready(&self) -> bool {
238        self.expected.iter().all(|s| self.received.contains_key(s))
239    }
240
241    /// Record the output from a source node.
242    ///
243    /// If the source has already been recorded, the previous value is
244    /// overwritten (last-write-wins). Recording a source that is not in
245    /// the expected set is a no-op for readiness but the value is still stored.
246    ///
247    /// # Arguments
248    ///
249    /// * `source_node` - The name of the upstream node that produced the output.
250    /// * `output` - The output value from the source node.
251    ///
252    /// # Example
253    ///
254    /// ```rust
255    /// use adk_graph::deferred::FanInTracker;
256    /// use serde_json::json;
257    ///
258    /// let mut tracker = FanInTracker::new(vec!["worker_1", "worker_2"]);
259    /// tracker.record("worker_1", json!({"status": "done"}));
260    /// assert_eq!(tracker.received_count(), 1);
261    /// ```
262    pub fn record(&mut self, source_node: &str, output: Value) {
263        let key = source_node.to_string();
264        if !self.received.contains_key(&key) {
265            self.insertion_order.push(key.clone());
266        }
267        self.received.insert(key, output);
268    }
269
270    /// Merge all received outputs according to the given strategy.
271    ///
272    /// The merge operation combines all recorded outputs into a single
273    /// [`Value`] based on the [`MergeStrategy`]:
274    ///
275    /// - [`MergeStrategy::Collect`]: Returns a JSON array of all outputs in
276    ///   insertion order.
277    /// - [`MergeStrategy::MergeMap`]: Merges all JSON object outputs into a
278    ///   single object (last-write-wins). Non-object outputs are skipped.
279    /// - [`MergeStrategy::First`]: Returns the first recorded output.
280    /// - [`MergeStrategy::Custom`]: Invokes the custom function with all outputs.
281    ///
282    /// # Arguments
283    ///
284    /// * `strategy` - The merge strategy to apply.
285    ///
286    /// # Example
287    ///
288    /// ```rust
289    /// use adk_graph::deferred::{FanInTracker, MergeStrategy};
290    /// use serde_json::json;
291    ///
292    /// let mut tracker = FanInTracker::new(vec!["a", "b"]);
293    /// tracker.record("a", json!({"x": 1}));
294    /// tracker.record("b", json!({"y": 2}));
295    ///
296    /// // Collect strategy
297    /// let result = tracker.merge(&MergeStrategy::Collect);
298    /// assert_eq!(result, json!([{"x": 1}, {"y": 2}]));
299    ///
300    /// // MergeMap strategy
301    /// let result = tracker.merge(&MergeStrategy::MergeMap);
302    /// assert_eq!(result, json!({"x": 1, "y": 2}));
303    /// ```
304    pub fn merge(&self, strategy: &MergeStrategy) -> Value {
305        match strategy {
306            MergeStrategy::Collect => {
307                let outputs: Vec<Value> = self
308                    .insertion_order
309                    .iter()
310                    .filter_map(|key| self.received.get(key).cloned())
311                    .collect();
312                Value::Array(outputs)
313            }
314            MergeStrategy::MergeMap => {
315                let mut merged = serde_json::Map::new();
316                for key in &self.insertion_order {
317                    if let Some(Value::Object(map)) = self.received.get(key) {
318                        for (k, v) in map {
319                            merged.insert(k.clone(), v.clone());
320                        }
321                    }
322                }
323                Value::Object(merged)
324            }
325            MergeStrategy::First => self
326                .insertion_order
327                .first()
328                .and_then(|key| self.received.get(key).cloned())
329                .unwrap_or(Value::Null),
330            MergeStrategy::Custom(f) => {
331                let outputs: Vec<Value> = self
332                    .insertion_order
333                    .iter()
334                    .filter_map(|key| self.received.get(key).cloned())
335                    .collect();
336                f(outputs)
337            }
338        }
339    }
340
341    /// Returns the number of outputs received so far.
342    pub fn received_count(&self) -> usize {
343        self.received.len()
344    }
345
346    /// Returns the number of expected source nodes.
347    pub fn expected_count(&self) -> usize {
348        self.expected.len()
349    }
350
351    /// Returns the names of sources that have not yet reported.
352    pub fn pending_sources(&self) -> Vec<&str> {
353        self.expected
354            .iter()
355            .filter(|s| !self.received.contains_key(*s))
356            .map(|s| s.as_str())
357            .collect()
358    }
359
360    /// Returns the names of sources that have reported.
361    pub fn completed_sources(&self) -> Vec<&str> {
362        self.insertion_order.iter().map(|s| s.as_str()).collect()
363    }
364}
365
366impl fmt::Debug for FanInTracker {
367    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
368        f.debug_struct("FanInTracker")
369            .field("expected", &self.expected)
370            .field("received_keys", &self.insertion_order)
371            .field("is_ready", &self.is_ready())
372            .finish()
373    }
374}
375
376#[cfg(test)]
377mod tests {
378    use super::*;
379    use serde_json::json;
380
381    #[test]
382    fn test_tracker_new_empty_not_ready() {
383        let tracker = FanInTracker::new(vec!["a", "b", "c"]);
384        assert!(!tracker.is_ready());
385        assert_eq!(tracker.expected_count(), 3);
386        assert_eq!(tracker.received_count(), 0);
387    }
388
389    #[test]
390    fn test_tracker_ready_when_all_received() {
391        let mut tracker = FanInTracker::new(vec!["a", "b"]);
392        tracker.record("a", json!(1));
393        assert!(!tracker.is_ready());
394        tracker.record("b", json!(2));
395        assert!(tracker.is_ready());
396    }
397
398    #[test]
399    fn test_merge_collect() {
400        let mut tracker = FanInTracker::new(vec!["x", "y", "z"]);
401        tracker.record("x", json!("first"));
402        tracker.record("y", json!("second"));
403        tracker.record("z", json!("third"));
404
405        let result = tracker.merge(&MergeStrategy::Collect);
406        assert_eq!(result, json!(["first", "second", "third"]));
407    }
408
409    #[test]
410    fn test_merge_map_combines_objects() {
411        let mut tracker = FanInTracker::new(vec!["a", "b"]);
412        tracker.record("a", json!({"key1": "val1", "shared": "from_a"}));
413        tracker.record("b", json!({"key2": "val2", "shared": "from_b"}));
414
415        let result = tracker.merge(&MergeStrategy::MergeMap);
416        assert_eq!(result, json!({"key1": "val1", "key2": "val2", "shared": "from_b"}));
417    }
418
419    #[test]
420    fn test_merge_map_skips_non_objects() {
421        let mut tracker = FanInTracker::new(vec!["a", "b"]);
422        tracker.record("a", json!(42)); // Not an object, skipped
423        tracker.record("b", json!({"key": "value"}));
424
425        let result = tracker.merge(&MergeStrategy::MergeMap);
426        assert_eq!(result, json!({"key": "value"}));
427    }
428
429    #[test]
430    fn test_merge_first() {
431        let mut tracker = FanInTracker::new(vec!["a", "b", "c"]);
432        tracker.record("b", json!("first_to_arrive"));
433        tracker.record("a", json!("second_to_arrive"));
434        tracker.record("c", json!("third_to_arrive"));
435
436        let result = tracker.merge(&MergeStrategy::First);
437        assert_eq!(result, json!("first_to_arrive"));
438    }
439
440    #[test]
441    fn test_merge_first_empty() {
442        let tracker = FanInTracker::new(vec!["a"]);
443        let result = tracker.merge(&MergeStrategy::First);
444        assert_eq!(result, Value::Null);
445    }
446
447    #[test]
448    fn test_merge_custom() {
449        let mut tracker = FanInTracker::new(vec!["a", "b"]);
450        tracker.record("a", json!(10));
451        tracker.record("b", json!(20));
452
453        let strategy = MergeStrategy::Custom(Arc::new(|outputs| {
454            let sum: i64 = outputs.iter().filter_map(|v| v.as_i64()).sum();
455            json!(sum)
456        }));
457
458        let result = tracker.merge(&strategy);
459        assert_eq!(result, json!(30));
460    }
461
462    #[test]
463    fn test_record_overwrites_previous() {
464        let mut tracker = FanInTracker::new(vec!["a"]);
465        tracker.record("a", json!("first"));
466        tracker.record("a", json!("second"));
467
468        assert!(tracker.is_ready());
469        assert_eq!(tracker.received_count(), 1);
470
471        let result = tracker.merge(&MergeStrategy::First);
472        assert_eq!(result, json!("second"));
473    }
474
475    #[test]
476    fn test_pending_and_completed_sources() {
477        let mut tracker = FanInTracker::new(vec!["a", "b", "c"]);
478        tracker.record("b", json!(1));
479
480        let mut pending = tracker.pending_sources();
481        pending.sort();
482        assert_eq!(pending, vec!["a", "c"]);
483        assert_eq!(tracker.completed_sources(), vec!["b"]);
484    }
485
486    #[test]
487    fn test_default_config() {
488        let config = DeferredNodeConfig::default();
489        assert!(matches!(config.merge_strategy, MergeStrategy::Collect));
490        assert!(config.fan_in_timeout.is_none());
491    }
492
493    #[test]
494    fn test_merge_strategy_debug() {
495        assert_eq!(format!("{:?}", MergeStrategy::Collect), "Collect");
496        assert_eq!(format!("{:?}", MergeStrategy::MergeMap), "MergeMap");
497        assert_eq!(format!("{:?}", MergeStrategy::First), "First");
498        let custom = MergeStrategy::Custom(Arc::new(Value::Array));
499        assert_eq!(format!("{:?}", custom), "Custom(<fn>)");
500    }
501}