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            self.assign_agent_session(&mut config_with_trace, step.id, wave_position, name)
202                .await?;
203            step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
204            record_slots.push(slot);
205        }
206
207        let mut join_set = JoinSet::new();
208        let mut task_index: HashMap<Id, usize> = HashMap::new();
209        let parallel_timeout = self.guard_remaining_timeout();
210        for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
211            let provider = self.provider.clone();
212            // Each task owns its own handle: `intercept` is synchronous, so no
213            // borrow of the context is held across an await point.
214            let interceptor = self.interceptor.clone();
215            let config = config.clone();
216            let step_log_sender = self
217                .log_sender
218                .as_ref()
219                .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
220            let handle = join_set.spawn(async move {
221                let result = match parallel_timeout {
222                    Some(dur) => {
223                        match timeout(
224                            dur,
225                            execute_step_config_intercepted(
226                                &config,
227                                &provider,
228                                step_log_sender,
229                                interceptor.as_ref(),
230                            ),
231                        )
232                        .await
233                        {
234                            Ok(r) => r,
235                            Err(_elapsed) => {
236                                Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
237                                    elapsed_secs: 0,
238                                    max: 0,
239                                }))
240                            }
241                        }
242                    }
243                    None => {
244                        execute_step_config_intercepted(
245                            &config,
246                            &provider,
247                            step_log_sender,
248                            interceptor.as_ref(),
249                        )
250                        .await
251                    }
252                };
253                (idx, result)
254            });
255            task_index.insert(handle.id(), idx);
256        }
257
258        // JoinSet returns in completion order; indexed_results restores input order.
259        let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
260            vec![None; step_records.len()];
261        let mut first_error: Option<EngineError> = None;
262        // Earliest capacity wait of the wave: the run sleeps once every other
263        // step settled, and only the parked steps run again when it wakes.
264        let mut capacity_wait: Option<(Uuid, String, DateTime<Utc>)> = None;
265
266        while let Some(join_result) = join_set.join_next().await {
267            let (idx, step_result) = match join_result {
268                Ok(r) => r,
269                Err(e) => {
270                    let error_msg = format!("join error: {e}");
271                    if let Some(&idx) = task_index.get(&e.id()) {
272                        let (step_id, _, step_name, _) = &step_records[idx];
273                        let completed_at = Utc::now();
274                        error!(
275                            run_id = %self.run_id,
276                            step = %step_name,
277                            error = %error_msg,
278                            "parallel step panicked or was cancelled"
279                        );
280                        if let Err(store_err) = self
281                            .store
282                            .update_step(
283                                *step_id,
284                                StepUpdate {
285                                    status: Some(StepStatus::Failed),
286                                    error: Some(error_msg.clone()),
287                                    completed_at: Some(completed_at),
288                                    ..StepUpdate::default()
289                                },
290                            )
291                            .await
292                        {
293                            error!(
294                                run_id = %self.run_id,
295                                step_id = %step_id,
296                                error = %store_err,
297                                "failed to persist JoinError for step"
298                            );
299                        }
300                        indexed_results[idx] = Some(Err(error_msg.clone()));
301                    }
302                    if first_error.is_none() {
303                        first_error = Some(EngineError::StepConfig(error_msg));
304                    }
305                    if fail_fast {
306                        join_set.abort_all();
307                    }
308                    continue;
309                }
310            };
311
312            let (step_id, step_trace, step_name, step_config) = &step_records[idx];
313            let completed_at = Utc::now();
314
315            if let Err(err) = self
316                .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
317                .await
318            {
319                self.fail_step(*step_id, &err).await;
320                indexed_results[idx] = Some(Err(err.to_string()));
321                if first_error.is_none() {
322                    first_error = Some(err);
323                }
324                if fail_fast {
325                    join_set.abort_all();
326                }
327                continue;
328            }
329
330            match step_result {
331                Ok(output) => {
332                    self.total_cost_usd += output.cost_usd;
333                    self.total_duration_ms += output.duration_ms;
334
335                    // Record token usage in the guard for agent steps.
336                    if matches!(step_config, StepConfig::Agent(_)) {
337                        let tokens = output.total_tokens();
338                        if tokens > 0
339                            && let Err(guard_err) = self.guard_record_tokens(tokens)
340                        {
341                            if first_error.is_none() {
342                                first_error = Some(guard_err);
343                            }
344                            if fail_fast {
345                                join_set.abort_all();
346                            }
347                        }
348                    }
349
350                    let debug_messages_json = output.debug_messages_json();
351
352                    self.store
353                        .update_step(
354                            *step_id,
355                            StepUpdate {
356                                status: Some(StepStatus::Completed),
357                                output: Some(output.output.clone()),
358                                duration_ms: Some(output.duration_ms),
359                                cost_usd: Some(output.cost_usd),
360                                input_tokens: output.input_tokens,
361                                cache_read_input_tokens: output.cache_read_input_tokens,
362                                cache_creation_input_tokens: output.cache_creation_input_tokens,
363                                output_tokens: output.output_tokens,
364                                completed_at: Some(completed_at),
365                                debug_messages: debug_messages_json,
366                                account_id: output.account_id,
367                                environment_id: output.environment_id.clone(),
368                                ..StepUpdate::default()
369                            },
370                        )
371                        .await?;
372
373                    self.step_results.push(StepResult::from_success(
374                        *step_trace,
375                        step_name,
376                        &output,
377                    ));
378
379                    if let Some(ref bus) = self.event_bus
380                        && matches!(step_config, StepConfig::Agent(_))
381                    {
382                        let tokens = output.total_tokens();
383                        bus.publish(
384                            self.run_id,
385                            WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
386                                step_name: step_name.clone(),
387                                tokens,
388                                cost_usd: output.cost_usd,
389                            }),
390                        );
391                    }
392
393                    info!(
394                        run_id = %self.run_id,
395                        step = %step_name,
396                        trace_id = %step_trace,
397                        duration_ms = output.duration_ms,
398                        "parallel step completed"
399                    );
400
401                    indexed_results[idx] = Some(Ok(output));
402                }
403                Err(EngineError::Operation(OperationError::Agent(AgentError::CapacityWait {
404                    kind,
405                    wake_at,
406                }))) => {
407                    // Not a failure, even with `allow_failure`: see `run_step`.
408                    self.park_capacity_step(*step_id).await?;
409                    info!(
410                        run_id = %self.run_id,
411                        step = %step_name,
412                        kind = %kind,
413                        wake_at = %wake_at,
414                        "no provider capacity, parallel step parked until the run wakes"
415                    );
416                    if capacity_wait
417                        .as_ref()
418                        .is_none_or(|(_, _, earliest)| wake_at < *earliest)
419                    {
420                        capacity_wait = Some((*step_id, kind, wake_at));
421                    }
422                }
423                Err(err) => {
424                    let err_msg = err.to_string();
425                    let debug_messages_json = extract_debug_messages_from_error(&err);
426                    let partial = extract_partial_usage_from_error(&err);
427                    let raw_response_output = extract_raw_response_from_error(&err);
428
429                    if let Some(ref usage) = partial {
430                        if let Some(cost) = usage.cost_usd {
431                            self.total_cost_usd += cost;
432                        }
433                        if let Some(dur) = usage.duration_ms {
434                            self.total_duration_ms += dur;
435                        }
436                    }
437
438                    if let Err(store_err) = self
439                        .store
440                        .update_step(
441                            *step_id,
442                            StepUpdate {
443                                status: Some(StepStatus::Failed),
444                                error: Some(err_msg.clone()),
445                                output: raw_response_output.clone(),
446                                completed_at: Some(completed_at),
447                                debug_messages: debug_messages_json,
448                                duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
449                                cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
450                                input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
451                                cache_read_input_tokens: partial
452                                    .as_ref()
453                                    .and_then(|p| p.cache_read_input_tokens),
454                                cache_creation_input_tokens: partial
455                                    .as_ref()
456                                    .and_then(|p| p.cache_creation_input_tokens),
457                                output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
458                                ..StepUpdate::default()
459                            },
460                        )
461                        .await
462                    {
463                        error!(
464                            step_id = %step_id,
465                            error = %store_err,
466                            "failed to persist parallel step failure"
467                        );
468                    }
469
470                    let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
471                    let err_cost = partial
472                        .as_ref()
473                        .and_then(|p| p.cost_usd)
474                        .unwrap_or(Decimal::ZERO);
475                    self.step_results.push(StepResult::from_failure(
476                        *step_trace,
477                        step_name,
478                        &err_msg,
479                        err_duration,
480                        err_cost,
481                    ));
482
483                    if step_config.allow_failure() {
484                        self.has_allowed_failure = true;
485                        info!(
486                            run_id = %self.run_id,
487                            step = %step_name,
488                            error = %err_msg,
489                            "parallel step failed but allow_failure is set, continuing"
490                        );
491                        indexed_results[idx] = Some(Ok(allowed_failure_output(
492                            &err_msg,
493                            raw_response_output,
494                            partial.as_ref(),
495                        )));
496                    } else {
497                        indexed_results[idx] = Some(Err(err_msg.clone()));
498
499                        if first_error.is_none() {
500                            first_error = Some(err);
501                        }
502
503                        if fail_fast {
504                            join_set.abort_all();
505                        }
506                    }
507                }
508            }
509        }
510
511        if let Some(err) = first_error {
512            return Err(err);
513        }
514
515        if let Some((step_id, kind, wake_at)) = capacity_wait {
516            return Err(EngineError::CapacitySleeping {
517                run_id: self.run_id,
518                step_id,
519                kind,
520                wake_at,
521            });
522        }
523
524        self.persist_progress().await;
525
526        // Build results in original order, replayed and launched steps alike.
527        for (idx, (step_id, _trace_id, name, config)) in step_records.iter().enumerate() {
528            let mut output = match indexed_results[idx].take() {
529                Some(Ok(o)) => o,
530                _ => unreachable!("all steps succeeded if no error returned"),
531            };
532            output.artifacts = StepArtifacts::new(name, Some(*step_id), config.declared_outputs());
533            slots[record_slots[idx]] = Some(ParallelStepResult {
534                name: name.clone(),
535                output,
536                step_id: *step_id,
537            });
538        }
539        let results: Vec<ParallelStepResult> = slots.into_iter().flatten().collect();
540        self.last_step_ids = results.iter().map(|r| r.step_id).collect();
541
542        Ok(results)
543    }
544
545    /// Replay the steps of a parallel wave that completed in a previous
546    /// execution of the current attempt.
547    ///
548    /// Returns one slot per entry of `steps`, in order: `Some` holds the
549    /// replayed result of a step that completed at `position`, `None` marks a
550    /// step that must run. `last_step_ids` is left untouched, so the steps
551    /// launched next still depend on the steps before the wave.
552    fn replay_wave(
553        &mut self,
554        position: u32,
555        steps: &[(&str, StepConfig)],
556    ) -> Result<Vec<Option<ParallelStepResult>>, EngineError> {
557        let mut slots = Vec::with_capacity(steps.len());
558        for (name, config) in steps {
559            let Some(step) = self.replay_wave_steps.get(&(position, (*name).to_string())) else {
560                slots.push(None);
561                continue;
562            };
563            check_replay_identity(step, position, name, &config.kind())?;
564            if step.status.state != StepStatus::Completed {
565                slots.push(None);
566                continue;
567            }
568
569            let mut output = StepOutput::from(step);
570            let step_id = step.id;
571            // Cost is not added: `carry_over_run_totals` seeded
572            // `total_cost_usd` from the run totals persisted before the
573            // suspension, which already include this step.
574            self.total_duration_ms += output.duration_ms;
575            output.artifacts = StepArtifacts::new(name, Some(step_id), config.declared_outputs());
576
577            info!(
578                run_id = %self.run_id,
579                step = %name,
580                position,
581                "step replayed from previous execution"
582            );
583
584            slots.push(Some(ParallelStepResult {
585                name: (*name).to_string(),
586                output,
587                step_id,
588            }));
589        }
590        Ok(slots)
591    }
592}
593
594/// Reject a wave in which two steps share a name.
595///
596/// The name identifies a step of the wave in the run timeline, in its
597/// artifact handles and on resume, where two steps with one name would replay
598/// the same stored result.
599fn reject_duplicate_names(steps: &[(&str, StepConfig)]) -> Result<(), EngineError> {
600    let mut seen = HashSet::with_capacity(steps.len());
601    for (name, _) in steps {
602        if !seen.insert(*name) {
603            return Err(EngineError::StepConfig(format!(
604                "parallel wave has two steps named {name:?}; each step of a wave needs its own name"
605            )));
606        }
607    }
608    Ok(())
609}