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