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