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