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}