Skip to main content

ironflow_engine/context/steps/
parallel.rs

1//! Parallel step wave for [`WorkflowContext`].
2
3use std::collections::HashMap;
4
5use chrono::Utc;
6use rust_decimal::Decimal;
7use serde_json::to_value;
8use tokio::task::{Id, JoinSet};
9use tokio::time::timeout;
10use tracing::{error, info};
11use uuid::Uuid;
12
13use ironflow_store::models::{NewStep, StepStatus, StepUpdate, step_trace_id};
14
15use crate::budget::step_budget_usd;
16use crate::config::StepConfig;
17use crate::context::WorkflowContext;
18use crate::context::failure::{
19    allowed_failure_output, extract_debug_messages_from_error, extract_partial_usage_from_error,
20    extract_raw_response_from_error,
21};
22use crate::error::EngineError;
23use crate::executor::{ParallelStepResult, StepOutput, StepResult, execute_step_config};
24use crate::guard::WorkflowRejection;
25use crate::log_sender::StepLogSender;
26use crate::notify::{WorkflowAgentStepTokensUsedEvent, WorkflowEvent};
27
28impl WorkflowContext {
29    /// Execute multiple steps concurrently (wait-all model).
30    ///
31    /// All steps in the batch execute in parallel via `tokio::JoinSet`.
32    /// Each step is recorded with the same `position` (execution wave).
33    /// Dependencies on previous steps are recorded automatically.
34    ///
35    /// When `fail_fast` is true, remaining steps are aborted on the first
36    /// failure. When false, all steps run to completion and the first
37    /// error is returned.
38    ///
39    /// # Errors
40    ///
41    /// Returns [`EngineError`] if any step fails.
42    ///
43    /// # Examples
44    ///
45    /// ```no_run
46    /// use ironflow_engine::context::WorkflowContext;
47    /// use ironflow_engine::config::{StepConfig, ShellConfig};
48    /// use ironflow_engine::error::EngineError;
49    ///
50    /// # async fn example(ctx: &mut WorkflowContext) -> Result<(), EngineError> {
51    /// let results = ctx.parallel(
52    ///     vec![
53    ///         ("test-unit", StepConfig::Shell(ShellConfig::new("cargo test --lib"))),
54    ///         ("lint", StepConfig::Shell(ShellConfig::new("cargo clippy"))),
55    ///     ],
56    ///     true,
57    /// ).await?;
58    ///
59    /// for r in &results {
60    ///     println!("{}: {:?}", r.name, r.output.output);
61    /// }
62    /// # Ok(())
63    /// # }
64    /// ```
65    pub async fn parallel(
66        &mut self,
67        steps: Vec<(&str, StepConfig)>,
68        fail_fast: bool,
69    ) -> Result<Vec<ParallelStepResult>, EngineError> {
70        if steps.is_empty() {
71            return Ok(Vec::new());
72        }
73
74        // Guard timeout: checked before launching the wave.
75        self.check_guard_timeout()?;
76
77        // Cost cap: the whole wave is charged at once. Refused before any step
78        // record is created, so nothing in the wave starts.
79        let wave_budget: Decimal = steps
80            .iter()
81            .filter_map(|(_, config)| match config {
82                StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
83                _ => None,
84            })
85            .map(step_budget_usd)
86            .sum();
87        self.check_run_budget(wave_budget)?;
88
89        let wave_position = self.position;
90        self.position += 1;
91
92        let now = Utc::now();
93        let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
94            Vec::with_capacity(steps.len());
95
96        for (name, config) in &steps {
97            let kind = config.kind();
98            let trace_id = step_trace_id(self.run_id, name, wave_position);
99            let step = self
100                .store
101                .create_step(NewStep {
102                    run_id: self.run_id,
103                    trace_id,
104                    name: name.to_string(),
105                    kind,
106                    position: wave_position,
107                    input: Some(to_value(config)?),
108                    is_error_handler: false,
109                })
110                .await?;
111
112            self.start_step(step.id, now).await?;
113
114            // Inputs are materialized before any step in the wave starts, so a
115            // missing one fails the wave rather than a half-run command.
116            if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
117                self.fail_step(step.id, &err).await;
118                if !config.allow_failure() {
119                    return Err(err);
120                }
121                self.has_allowed_failure = true;
122                info!(
123                    run_id = %self.run_id,
124                    step = %name,
125                    error = %err,
126                    "parallel step input preparation failed but allow_failure is set, skipping"
127                );
128                continue;
129            }
130
131            let mut config_with_trace = config.clone();
132            let step_trace = self.trace_context.child();
133            match config_with_trace {
134                StepConfig::Agent(ref mut agent_config) => {
135                    agent_config.trace_context = Some(step_trace);
136                }
137                StepConfig::Http(ref mut http_config) => {
138                    http_config.trace_context = Some(step_trace);
139                }
140                _ => {}
141            }
142            step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
143        }
144
145        let mut join_set = JoinSet::new();
146        let mut task_index: HashMap<Id, usize> = HashMap::new();
147        let parallel_timeout = self.guard_remaining_timeout();
148        for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
149            let provider = self.provider.clone();
150            let config = config.clone();
151            let step_log_sender = self
152                .log_sender
153                .as_ref()
154                .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
155            let handle = join_set.spawn(async move {
156                let result = match parallel_timeout {
157                    Some(dur) => {
158                        match timeout(
159                            dur,
160                            execute_step_config(&config, &provider, step_log_sender),
161                        )
162                        .await
163                        {
164                            Ok(r) => r,
165                            Err(_elapsed) => {
166                                Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
167                                    elapsed_secs: 0,
168                                    max: 0,
169                                }))
170                            }
171                        }
172                    }
173                    None => execute_step_config(&config, &provider, step_log_sender).await,
174                };
175                (idx, result)
176            });
177            task_index.insert(handle.id(), idx);
178        }
179
180        // JoinSet returns in completion order; indexed_results restores input order.
181        let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
182            vec![None; step_records.len()];
183        let mut first_error: Option<EngineError> = None;
184
185        while let Some(join_result) = join_set.join_next().await {
186            let (idx, step_result) = match join_result {
187                Ok(r) => r,
188                Err(e) => {
189                    let error_msg = format!("join error: {e}");
190                    if let Some(&idx) = task_index.get(&e.id()) {
191                        let (step_id, _, step_name, _) = &step_records[idx];
192                        let completed_at = Utc::now();
193                        error!(
194                            run_id = %self.run_id,
195                            step = %step_name,
196                            error = %error_msg,
197                            "parallel step panicked or was cancelled"
198                        );
199                        if let Err(store_err) = self
200                            .store
201                            .update_step(
202                                *step_id,
203                                StepUpdate {
204                                    status: Some(StepStatus::Failed),
205                                    error: Some(error_msg.clone()),
206                                    completed_at: Some(completed_at),
207                                    ..StepUpdate::default()
208                                },
209                            )
210                            .await
211                        {
212                            error!(
213                                run_id = %self.run_id,
214                                step_id = %step_id,
215                                error = %store_err,
216                                "failed to persist JoinError for step"
217                            );
218                        }
219                        indexed_results[idx] = Some(Err(error_msg.clone()));
220                    }
221                    if first_error.is_none() {
222                        first_error = Some(EngineError::StepConfig(error_msg));
223                    }
224                    if fail_fast {
225                        join_set.abort_all();
226                    }
227                    continue;
228                }
229            };
230
231            let (step_id, step_trace, step_name, step_config) = &step_records[idx];
232            let completed_at = Utc::now();
233
234            if let Err(err) = self
235                .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
236                .await
237            {
238                self.fail_step(*step_id, &err).await;
239                indexed_results[idx] = Some(Err(err.to_string()));
240                if first_error.is_none() {
241                    first_error = Some(err);
242                }
243                if fail_fast {
244                    join_set.abort_all();
245                }
246                continue;
247            }
248
249            match step_result {
250                Ok(output) => {
251                    self.total_cost_usd += output.cost_usd;
252                    self.total_duration_ms += output.duration_ms;
253
254                    // Record token usage in the guard for agent steps.
255                    if matches!(step_config, StepConfig::Agent(_)) {
256                        let tokens = output
257                            .input_tokens
258                            .unwrap_or(0)
259                            .saturating_add(output.output_tokens.unwrap_or(0));
260                        if tokens > 0
261                            && let Err(guard_err) = self.guard_record_tokens(tokens)
262                        {
263                            if first_error.is_none() {
264                                first_error = Some(guard_err);
265                            }
266                            if fail_fast {
267                                join_set.abort_all();
268                            }
269                        }
270                    }
271
272                    let debug_messages_json = output.debug_messages_json();
273
274                    self.store
275                        .update_step(
276                            *step_id,
277                            StepUpdate {
278                                status: Some(StepStatus::Completed),
279                                output: Some(output.output.clone()),
280                                duration_ms: Some(output.duration_ms),
281                                cost_usd: Some(output.cost_usd),
282                                input_tokens: output.input_tokens,
283                                output_tokens: output.output_tokens,
284                                completed_at: Some(completed_at),
285                                debug_messages: debug_messages_json,
286                                ..StepUpdate::default()
287                            },
288                        )
289                        .await?;
290
291                    self.step_results.push(StepResult::from_success(
292                        *step_trace,
293                        step_name,
294                        &output,
295                    ));
296
297                    if let Some(ref bus) = self.event_bus
298                        && matches!(step_config, StepConfig::Agent(_))
299                    {
300                        let tokens = output
301                            .input_tokens
302                            .unwrap_or(0)
303                            .saturating_add(output.output_tokens.unwrap_or(0));
304                        bus.publish(
305                            self.run_id,
306                            WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
307                                step_name: step_name.clone(),
308                                tokens,
309                                cost_usd: output.cost_usd,
310                            }),
311                        );
312                    }
313
314                    info!(
315                        run_id = %self.run_id,
316                        step = %step_name,
317                        trace_id = %step_trace,
318                        duration_ms = output.duration_ms,
319                        "parallel step completed"
320                    );
321
322                    indexed_results[idx] = Some(Ok(output));
323                }
324                Err(err) => {
325                    let err_msg = err.to_string();
326                    let debug_messages_json = extract_debug_messages_from_error(&err);
327                    let partial = extract_partial_usage_from_error(&err);
328                    let raw_response_output = extract_raw_response_from_error(&err);
329
330                    if let Some(ref usage) = partial {
331                        if let Some(cost) = usage.cost_usd {
332                            self.total_cost_usd += cost;
333                        }
334                        if let Some(dur) = usage.duration_ms {
335                            self.total_duration_ms += dur;
336                        }
337                    }
338
339                    if let Err(store_err) = self
340                        .store
341                        .update_step(
342                            *step_id,
343                            StepUpdate {
344                                status: Some(StepStatus::Failed),
345                                error: Some(err_msg.clone()),
346                                output: raw_response_output.clone(),
347                                completed_at: Some(completed_at),
348                                debug_messages: debug_messages_json,
349                                duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
350                                cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
351                                input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
352                                output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
353                                ..StepUpdate::default()
354                            },
355                        )
356                        .await
357                    {
358                        error!(
359                            step_id = %step_id,
360                            error = %store_err,
361                            "failed to persist parallel step failure"
362                        );
363                    }
364
365                    let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
366                    let err_cost = partial
367                        .as_ref()
368                        .and_then(|p| p.cost_usd)
369                        .unwrap_or(Decimal::ZERO);
370                    self.step_results.push(StepResult::from_failure(
371                        *step_trace,
372                        step_name,
373                        &err_msg,
374                        err_duration,
375                        err_cost,
376                    ));
377
378                    if step_config.allow_failure() {
379                        self.has_allowed_failure = true;
380                        info!(
381                            run_id = %self.run_id,
382                            step = %step_name,
383                            error = %err_msg,
384                            "parallel step failed but allow_failure is set, continuing"
385                        );
386                        indexed_results[idx] = Some(Ok(allowed_failure_output(
387                            &err_msg,
388                            raw_response_output,
389                            partial.as_ref(),
390                        )));
391                    } else {
392                        indexed_results[idx] = Some(Err(err_msg.clone()));
393
394                        if first_error.is_none() {
395                            first_error = Some(err);
396                        }
397
398                        if fail_fast {
399                            join_set.abort_all();
400                        }
401                    }
402                }
403            }
404        }
405
406        if let Some(err) = first_error {
407            return Err(err);
408        }
409
410        self.persist_progress().await;
411
412        self.last_step_ids = step_records.iter().map(|(id, _, _, _)| *id).collect();
413
414        // Build results in original order.
415        let results: Vec<ParallelStepResult> = step_records
416            .iter()
417            .enumerate()
418            .map(|(idx, (step_id, _trace_id, name, _))| {
419                let output = match indexed_results[idx].take() {
420                    Some(Ok(o)) => o,
421                    _ => unreachable!("all steps succeeded if no error returned"),
422                };
423                ParallelStepResult {
424                    name: name.clone(),
425                    output,
426                    step_id: *step_id,
427                }
428            })
429            .collect();
430
431        Ok(results)
432    }
433}