Skip to main content

ironflow_engine/context/steps/
parallel.rs

1//! Parallel step wave for [`WorkflowContext`].
2
3use std::collections::{HashMap, HashSet};
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::context::lifecycle::check_replay_identity;
23use crate::error::EngineError;
24use crate::executor::{
25    ParallelStepResult, StepArtifacts, StepOutput, StepResult, execute_step_config_intercepted,
26};
27use crate::guard::WorkflowRejection;
28use crate::log_sender::StepLogSender;
29use crate::notify::{WorkflowAgentStepTokensUsedEvent, WorkflowEvent};
30use crate::plan::{lock_plan, planned_output};
31
32impl WorkflowContext {
33    /// Execute multiple steps concurrently (wait-all model).
34    ///
35    /// All steps in the batch execute in parallel via `tokio::JoinSet`.
36    /// Each step is recorded with the same `position` (execution wave).
37    /// Dependencies on previous steps are recorded automatically.
38    ///
39    /// When `fail_fast` is true, remaining steps are aborted on the first
40    /// failure. When false, all steps run to completion and the first
41    /// error is returned.
42    ///
43    /// Every step of a wave must have its own name: the name identifies the
44    /// step in the run timeline, in its artifact handles, on resume and in the
45    /// `ironflow.io/step` pod label. Two steps sharing that label would let
46    /// the K8s ephemeral provider delete one step's pod when starting the other.
47    ///
48    /// On resume, each step of the wave that already completed in a prior
49    /// execution of the current attempt is replayed from the store; only the
50    /// other steps of the wave are launched again.
51    ///
52    /// # Errors
53    ///
54    /// Returns [`EngineError::StepConfig`] if two steps of the wave share a
55    /// name, before anything runs. Returns [`EngineError::ReplayDivergence`]
56    /// when the step recorded at a wave position has a different kind.
57    /// Returns [`EngineError`] if any step fails.
58    ///
59    /// # Examples
60    ///
61    /// ```no_run
62    /// use ironflow_engine::context::WorkflowContext;
63    /// use ironflow_engine::config::{StepConfig, ShellConfig};
64    /// use ironflow_engine::error::EngineError;
65    ///
66    /// # async fn example(ctx: &mut WorkflowContext) -> Result<(), EngineError> {
67    /// let results = ctx.parallel(
68    ///     vec![
69    ///         ("test-unit", StepConfig::Shell(ShellConfig::new("cargo test --lib"))),
70    ///         ("lint", StepConfig::Shell(ShellConfig::new("cargo clippy"))),
71    ///     ],
72    ///     true,
73    /// ).await?;
74    ///
75    /// for r in &results {
76    ///     println!("{}: {:?}", r.name, r.output.output);
77    /// }
78    /// # Ok(())
79    /// # }
80    /// ```
81    pub async fn parallel(
82        &mut self,
83        steps: Vec<(&str, StepConfig)>,
84        fail_fast: bool,
85    ) -> Result<Vec<ParallelStepResult>, EngineError> {
86        if steps.is_empty() {
87            return Ok(Vec::new());
88        }
89        reject_duplicate_names(&steps)?;
90
91        // Plan mode: record the whole wave under one parallel group and return
92        // synthetic outputs. No step record is created and nothing runs.
93        if let Some(plan) = self.plan().cloned() {
94            self.position += 1;
95            let mut results = Vec::with_capacity(steps.len());
96            let mut names = Vec::with_capacity(steps.len());
97            {
98                let mut recorder = lock_plan(&plan);
99                let group = recorder.next_group();
100                for (name, config) in &steps {
101                    let wave = Some(group.clone());
102                    if !recorder.record(name, config.kind(), &self.workflow_name, wave) {
103                        break;
104                    }
105                    names.push((*name).to_string());
106                    let estimate = recorder.estimate_for(name);
107                    let mut output = planned_output(config, estimate);
108                    output.artifacts = StepArtifacts::new(name, None, config.declared_outputs());
109                    results.push(ParallelStepResult {
110                        name: (*name).to_string(),
111                        output,
112                        step_id: Uuid::now_v7(),
113                    });
114                }
115                recorder.set_last(names);
116            }
117            return Ok(results);
118        }
119
120        // Guard timeout: checked before launching the wave.
121        self.check_guard_timeout()?;
122
123        let wave_position = self.position;
124        self.position += 1;
125
126        // Replay: a step of this wave that already completed in a prior
127        // execution of the current attempt returns its cached output; only the
128        // other steps are launched. When the whole wave completed, nothing is
129        // created or launched. Mirrors the replay-before-budget-check ordering
130        // of `execute_step`.
131        let mut slots = self.replay_wave(wave_position, &steps)?;
132        if slots.iter().all(Option::is_some) {
133            let results: Vec<ParallelStepResult> = slots.into_iter().flatten().collect();
134            self.last_step_ids = results.iter().map(|r| r.step_id).collect();
135            return Ok(results);
136        }
137
138        // Cost cap: the steps left to run are charged at once. Refused before
139        // any step record is created, so nothing in the wave starts.
140        let wave_budget: Decimal = steps
141            .iter()
142            .zip(&slots)
143            .filter(|(_, slot)| slot.is_none())
144            .filter_map(|((_, config), _)| match config {
145                StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
146                _ => None,
147            })
148            .map(step_budget_usd)
149            .sum();
150        self.check_run_budget(wave_budget)?;
151
152        let now = Utc::now();
153        let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
154            Vec::with_capacity(steps.len());
155        // Index in `steps` of each entry of `step_records`.
156        let mut record_slots: Vec<usize> = Vec::with_capacity(steps.len());
157
158        for (slot, (name, config)) in steps.iter().enumerate() {
159            if slots[slot].is_some() {
160                continue;
161            }
162            let kind = config.kind();
163            let trace_id = step_trace_id(self.run_id, name, wave_position);
164            let step = self
165                .store
166                .create_step(NewStep {
167                    run_id: self.run_id,
168                    trace_id,
169                    name: name.to_string(),
170                    kind,
171                    position: wave_position,
172                    input: Some(to_value(config)?),
173                    is_error_handler: false,
174                })
175                .await?;
176
177            self.start_step(step.id, now).await?;
178
179            // Inputs are materialized before any step in the wave starts, so a
180            // missing one fails the wave rather than a half-run command.
181            if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
182                self.fail_step(step.id, &err).await;
183                if !config.allow_failure() {
184                    return Err(err);
185                }
186                self.has_allowed_failure = true;
187                info!(
188                    run_id = %self.run_id,
189                    step = %name,
190                    error = %err,
191                    "parallel step input preparation failed but allow_failure is set, skipping"
192                );
193                continue;
194            }
195
196            let mut config_with_trace = config.clone();
197            self.scope_step_config(&mut config_with_trace, name);
198            step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
199            record_slots.push(slot);
200        }
201
202        let mut join_set = JoinSet::new();
203        let mut task_index: HashMap<Id, usize> = HashMap::new();
204        let parallel_timeout = self.guard_remaining_timeout();
205        for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
206            let provider = self.provider.clone();
207            // Each task owns its own handle: `intercept` is synchronous, so no
208            // borrow of the context is held across an await point.
209            let interceptor = self.interceptor.clone();
210            let config = config.clone();
211            let step_log_sender = self
212                .log_sender
213                .as_ref()
214                .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
215            let handle = join_set.spawn(async move {
216                let result = match parallel_timeout {
217                    Some(dur) => {
218                        match timeout(
219                            dur,
220                            execute_step_config_intercepted(
221                                &config,
222                                &provider,
223                                step_log_sender,
224                                interceptor.as_ref(),
225                            ),
226                        )
227                        .await
228                        {
229                            Ok(r) => r,
230                            Err(_elapsed) => {
231                                Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
232                                    elapsed_secs: 0,
233                                    max: 0,
234                                }))
235                            }
236                        }
237                    }
238                    None => {
239                        execute_step_config_intercepted(
240                            &config,
241                            &provider,
242                            step_log_sender,
243                            interceptor.as_ref(),
244                        )
245                        .await
246                    }
247                };
248                (idx, result)
249            });
250            task_index.insert(handle.id(), idx);
251        }
252
253        // JoinSet returns in completion order; indexed_results restores input order.
254        let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
255            vec![None; step_records.len()];
256        let mut first_error: Option<EngineError> = None;
257
258        while let Some(join_result) = join_set.join_next().await {
259            let (idx, step_result) = match join_result {
260                Ok(r) => r,
261                Err(e) => {
262                    let error_msg = format!("join error: {e}");
263                    if let Some(&idx) = task_index.get(&e.id()) {
264                        let (step_id, _, step_name, _) = &step_records[idx];
265                        let completed_at = Utc::now();
266                        error!(
267                            run_id = %self.run_id,
268                            step = %step_name,
269                            error = %error_msg,
270                            "parallel step panicked or was cancelled"
271                        );
272                        if let Err(store_err) = self
273                            .store
274                            .update_step(
275                                *step_id,
276                                StepUpdate {
277                                    status: Some(StepStatus::Failed),
278                                    error: Some(error_msg.clone()),
279                                    completed_at: Some(completed_at),
280                                    ..StepUpdate::default()
281                                },
282                            )
283                            .await
284                        {
285                            error!(
286                                run_id = %self.run_id,
287                                step_id = %step_id,
288                                error = %store_err,
289                                "failed to persist JoinError for step"
290                            );
291                        }
292                        indexed_results[idx] = Some(Err(error_msg.clone()));
293                    }
294                    if first_error.is_none() {
295                        first_error = Some(EngineError::StepConfig(error_msg));
296                    }
297                    if fail_fast {
298                        join_set.abort_all();
299                    }
300                    continue;
301                }
302            };
303
304            let (step_id, step_trace, step_name, step_config) = &step_records[idx];
305            let completed_at = Utc::now();
306
307            if let Err(err) = self
308                .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
309                .await
310            {
311                self.fail_step(*step_id, &err).await;
312                indexed_results[idx] = Some(Err(err.to_string()));
313                if first_error.is_none() {
314                    first_error = Some(err);
315                }
316                if fail_fast {
317                    join_set.abort_all();
318                }
319                continue;
320            }
321
322            match step_result {
323                Ok(output) => {
324                    self.total_cost_usd += output.cost_usd;
325                    self.total_duration_ms += output.duration_ms;
326
327                    // Record token usage in the guard for agent steps.
328                    if matches!(step_config, StepConfig::Agent(_)) {
329                        let tokens = output.total_tokens();
330                        if tokens > 0
331                            && let Err(guard_err) = self.guard_record_tokens(tokens)
332                        {
333                            if first_error.is_none() {
334                                first_error = Some(guard_err);
335                            }
336                            if fail_fast {
337                                join_set.abort_all();
338                            }
339                        }
340                    }
341
342                    let debug_messages_json = output.debug_messages_json();
343
344                    self.store
345                        .update_step(
346                            *step_id,
347                            StepUpdate {
348                                status: Some(StepStatus::Completed),
349                                output: Some(output.output.clone()),
350                                duration_ms: Some(output.duration_ms),
351                                cost_usd: Some(output.cost_usd),
352                                input_tokens: output.input_tokens,
353                                cache_read_input_tokens: output.cache_read_input_tokens,
354                                cache_creation_input_tokens: output.cache_creation_input_tokens,
355                                output_tokens: output.output_tokens,
356                                completed_at: Some(completed_at),
357                                debug_messages: debug_messages_json,
358                                ..StepUpdate::default()
359                            },
360                        )
361                        .await?;
362
363                    self.step_results.push(StepResult::from_success(
364                        *step_trace,
365                        step_name,
366                        &output,
367                    ));
368
369                    if let Some(ref bus) = self.event_bus
370                        && matches!(step_config, StepConfig::Agent(_))
371                    {
372                        let tokens = output.total_tokens();
373                        bus.publish(
374                            self.run_id,
375                            WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
376                                step_name: step_name.clone(),
377                                tokens,
378                                cost_usd: output.cost_usd,
379                            }),
380                        );
381                    }
382
383                    info!(
384                        run_id = %self.run_id,
385                        step = %step_name,
386                        trace_id = %step_trace,
387                        duration_ms = output.duration_ms,
388                        "parallel step completed"
389                    );
390
391                    indexed_results[idx] = Some(Ok(output));
392                }
393                Err(err) => {
394                    let err_msg = err.to_string();
395                    let debug_messages_json = extract_debug_messages_from_error(&err);
396                    let partial = extract_partial_usage_from_error(&err);
397                    let raw_response_output = extract_raw_response_from_error(&err);
398
399                    if let Some(ref usage) = partial {
400                        if let Some(cost) = usage.cost_usd {
401                            self.total_cost_usd += cost;
402                        }
403                        if let Some(dur) = usage.duration_ms {
404                            self.total_duration_ms += dur;
405                        }
406                    }
407
408                    if let Err(store_err) = self
409                        .store
410                        .update_step(
411                            *step_id,
412                            StepUpdate {
413                                status: Some(StepStatus::Failed),
414                                error: Some(err_msg.clone()),
415                                output: raw_response_output.clone(),
416                                completed_at: Some(completed_at),
417                                debug_messages: debug_messages_json,
418                                duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
419                                cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
420                                input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
421                                cache_read_input_tokens: partial
422                                    .as_ref()
423                                    .and_then(|p| p.cache_read_input_tokens),
424                                cache_creation_input_tokens: partial
425                                    .as_ref()
426                                    .and_then(|p| p.cache_creation_input_tokens),
427                                output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
428                                ..StepUpdate::default()
429                            },
430                        )
431                        .await
432                    {
433                        error!(
434                            step_id = %step_id,
435                            error = %store_err,
436                            "failed to persist parallel step failure"
437                        );
438                    }
439
440                    let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
441                    let err_cost = partial
442                        .as_ref()
443                        .and_then(|p| p.cost_usd)
444                        .unwrap_or(Decimal::ZERO);
445                    self.step_results.push(StepResult::from_failure(
446                        *step_trace,
447                        step_name,
448                        &err_msg,
449                        err_duration,
450                        err_cost,
451                    ));
452
453                    if step_config.allow_failure() {
454                        self.has_allowed_failure = true;
455                        info!(
456                            run_id = %self.run_id,
457                            step = %step_name,
458                            error = %err_msg,
459                            "parallel step failed but allow_failure is set, continuing"
460                        );
461                        indexed_results[idx] = Some(Ok(allowed_failure_output(
462                            &err_msg,
463                            raw_response_output,
464                            partial.as_ref(),
465                        )));
466                    } else {
467                        indexed_results[idx] = Some(Err(err_msg.clone()));
468
469                        if first_error.is_none() {
470                            first_error = Some(err);
471                        }
472
473                        if fail_fast {
474                            join_set.abort_all();
475                        }
476                    }
477                }
478            }
479        }
480
481        if let Some(err) = first_error {
482            return Err(err);
483        }
484
485        self.persist_progress().await;
486
487        // Build results in original order, replayed and launched steps alike.
488        for (idx, (step_id, _trace_id, name, config)) in step_records.iter().enumerate() {
489            let mut output = match indexed_results[idx].take() {
490                Some(Ok(o)) => o,
491                _ => unreachable!("all steps succeeded if no error returned"),
492            };
493            output.artifacts = StepArtifacts::new(name, Some(*step_id), config.declared_outputs());
494            slots[record_slots[idx]] = Some(ParallelStepResult {
495                name: name.clone(),
496                output,
497                step_id: *step_id,
498            });
499        }
500        let results: Vec<ParallelStepResult> = slots.into_iter().flatten().collect();
501        self.last_step_ids = results.iter().map(|r| r.step_id).collect();
502
503        Ok(results)
504    }
505
506    /// Replay the steps of a parallel wave that completed in a previous
507    /// execution of the current attempt.
508    ///
509    /// Returns one slot per entry of `steps`, in order: `Some` holds the
510    /// replayed result of a step that completed at `position`, `None` marks a
511    /// step that must run. `last_step_ids` is left untouched, so the steps
512    /// launched next still depend on the steps before the wave.
513    fn replay_wave(
514        &mut self,
515        position: u32,
516        steps: &[(&str, StepConfig)],
517    ) -> Result<Vec<Option<ParallelStepResult>>, EngineError> {
518        let mut slots = Vec::with_capacity(steps.len());
519        for (name, config) in steps {
520            let Some(step) = self.replay_wave_steps.get(&(position, (*name).to_string())) else {
521                slots.push(None);
522                continue;
523            };
524            check_replay_identity(step, position, name, &config.kind())?;
525            if step.status.state != StepStatus::Completed {
526                slots.push(None);
527                continue;
528            }
529
530            let mut output = StepOutput::from(step);
531            let step_id = step.id;
532            // Cost is not added: `carry_over_run_totals` seeded
533            // `total_cost_usd` from the run totals persisted before the
534            // suspension, which already include this step.
535            self.total_duration_ms += output.duration_ms;
536            output.artifacts = StepArtifacts::new(name, Some(step_id), config.declared_outputs());
537
538            info!(
539                run_id = %self.run_id,
540                step = %name,
541                position,
542                "step replayed from previous execution"
543            );
544
545            slots.push(Some(ParallelStepResult {
546                name: (*name).to_string(),
547                output,
548                step_id,
549            }));
550        }
551        Ok(slots)
552    }
553}
554
555/// Reject a wave in which two steps share a name.
556///
557/// The name identifies a step of the wave in the run timeline, in its
558/// artifact handles and on resume, where two steps with one name would replay
559/// the same stored result.
560fn reject_duplicate_names(steps: &[(&str, StepConfig)]) -> Result<(), EngineError> {
561    let mut seen = HashSet::with_capacity(steps.len());
562    for (name, _) in steps {
563        if !seen.insert(*name) {
564            return Err(EngineError::StepConfig(format!(
565                "parallel wave has two steps named {name:?}; each step of a wave needs its own name"
566            )));
567        }
568    }
569    Ok(())
570}