Skip to main content

ironflow_engine/context/steps/
parallel.rs

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