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