ironflow_engine/context/steps/
parallel.rs1use 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, 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 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 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 results.push(ParallelStepResult {
94 name: (*name).to_string(),
95 output: planned_output(config, estimate),
96 step_id: Uuid::now_v7(),
97 });
98 }
99 recorder.set_last(names);
100 }
101 return Ok(results);
102 }
103
104 self.check_guard_timeout()?;
106
107 let wave_budget: Decimal = steps
110 .iter()
111 .filter_map(|(_, config)| match config {
112 StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
113 _ => None,
114 })
115 .map(step_budget_usd)
116 .sum();
117 self.check_run_budget(wave_budget)?;
118
119 let wave_position = self.position;
120 self.position += 1;
121
122 let now = Utc::now();
123 let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
124 Vec::with_capacity(steps.len());
125
126 for (name, config) in &steps {
127 let kind = config.kind();
128 let trace_id = step_trace_id(self.run_id, name, wave_position);
129 let step = self
130 .store
131 .create_step(NewStep {
132 run_id: self.run_id,
133 trace_id,
134 name: name.to_string(),
135 kind,
136 position: wave_position,
137 input: Some(to_value(config)?),
138 is_error_handler: false,
139 })
140 .await?;
141
142 self.start_step(step.id, now).await?;
143
144 if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
147 self.fail_step(step.id, &err).await;
148 if !config.allow_failure() {
149 return Err(err);
150 }
151 self.has_allowed_failure = true;
152 info!(
153 run_id = %self.run_id,
154 step = %name,
155 error = %err,
156 "parallel step input preparation failed but allow_failure is set, skipping"
157 );
158 continue;
159 }
160
161 let mut config_with_trace = config.clone();
162 let step_trace = self.trace_context.child();
163 match config_with_trace {
164 StepConfig::Agent(ref mut agent_config) => {
165 agent_config.trace_context = Some(step_trace);
166 }
167 StepConfig::Http(ref mut http_config) => {
168 http_config.trace_context = Some(step_trace);
169 }
170 _ => {}
171 }
172 step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
173 }
174
175 let mut join_set = JoinSet::new();
176 let mut task_index: HashMap<Id, usize> = HashMap::new();
177 let parallel_timeout = self.guard_remaining_timeout();
178 for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
179 let provider = self.provider.clone();
180 let interceptor = self.interceptor.clone();
183 let config = config.clone();
184 let step_log_sender = self
185 .log_sender
186 .as_ref()
187 .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
188 let handle = join_set.spawn(async move {
189 let result = match parallel_timeout {
190 Some(dur) => {
191 match timeout(
192 dur,
193 execute_step_config_intercepted(
194 &config,
195 &provider,
196 step_log_sender,
197 interceptor.as_ref(),
198 ),
199 )
200 .await
201 {
202 Ok(r) => r,
203 Err(_elapsed) => {
204 Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
205 elapsed_secs: 0,
206 max: 0,
207 }))
208 }
209 }
210 }
211 None => {
212 execute_step_config_intercepted(
213 &config,
214 &provider,
215 step_log_sender,
216 interceptor.as_ref(),
217 )
218 .await
219 }
220 };
221 (idx, result)
222 });
223 task_index.insert(handle.id(), idx);
224 }
225
226 let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
228 vec![None; step_records.len()];
229 let mut first_error: Option<EngineError> = None;
230
231 while let Some(join_result) = join_set.join_next().await {
232 let (idx, step_result) = match join_result {
233 Ok(r) => r,
234 Err(e) => {
235 let error_msg = format!("join error: {e}");
236 if let Some(&idx) = task_index.get(&e.id()) {
237 let (step_id, _, step_name, _) = &step_records[idx];
238 let completed_at = Utc::now();
239 error!(
240 run_id = %self.run_id,
241 step = %step_name,
242 error = %error_msg,
243 "parallel step panicked or was cancelled"
244 );
245 if let Err(store_err) = self
246 .store
247 .update_step(
248 *step_id,
249 StepUpdate {
250 status: Some(StepStatus::Failed),
251 error: Some(error_msg.clone()),
252 completed_at: Some(completed_at),
253 ..StepUpdate::default()
254 },
255 )
256 .await
257 {
258 error!(
259 run_id = %self.run_id,
260 step_id = %step_id,
261 error = %store_err,
262 "failed to persist JoinError for step"
263 );
264 }
265 indexed_results[idx] = Some(Err(error_msg.clone()));
266 }
267 if first_error.is_none() {
268 first_error = Some(EngineError::StepConfig(error_msg));
269 }
270 if fail_fast {
271 join_set.abort_all();
272 }
273 continue;
274 }
275 };
276
277 let (step_id, step_trace, step_name, step_config) = &step_records[idx];
278 let completed_at = Utc::now();
279
280 if let Err(err) = self
281 .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
282 .await
283 {
284 self.fail_step(*step_id, &err).await;
285 indexed_results[idx] = Some(Err(err.to_string()));
286 if first_error.is_none() {
287 first_error = Some(err);
288 }
289 if fail_fast {
290 join_set.abort_all();
291 }
292 continue;
293 }
294
295 match step_result {
296 Ok(output) => {
297 self.total_cost_usd += output.cost_usd;
298 self.total_duration_ms += output.duration_ms;
299
300 if matches!(step_config, StepConfig::Agent(_)) {
302 let tokens = output
303 .input_tokens
304 .unwrap_or(0)
305 .saturating_add(output.output_tokens.unwrap_or(0));
306 if tokens > 0
307 && let Err(guard_err) = self.guard_record_tokens(tokens)
308 {
309 if first_error.is_none() {
310 first_error = Some(guard_err);
311 }
312 if fail_fast {
313 join_set.abort_all();
314 }
315 }
316 }
317
318 let debug_messages_json = output.debug_messages_json();
319
320 self.store
321 .update_step(
322 *step_id,
323 StepUpdate {
324 status: Some(StepStatus::Completed),
325 output: Some(output.output.clone()),
326 duration_ms: Some(output.duration_ms),
327 cost_usd: Some(output.cost_usd),
328 input_tokens: output.input_tokens,
329 output_tokens: output.output_tokens,
330 completed_at: Some(completed_at),
331 debug_messages: debug_messages_json,
332 ..StepUpdate::default()
333 },
334 )
335 .await?;
336
337 self.step_results.push(StepResult::from_success(
338 *step_trace,
339 step_name,
340 &output,
341 ));
342
343 if let Some(ref bus) = self.event_bus
344 && matches!(step_config, StepConfig::Agent(_))
345 {
346 let tokens = output
347 .input_tokens
348 .unwrap_or(0)
349 .saturating_add(output.output_tokens.unwrap_or(0));
350 bus.publish(
351 self.run_id,
352 WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
353 step_name: step_name.clone(),
354 tokens,
355 cost_usd: output.cost_usd,
356 }),
357 );
358 }
359
360 info!(
361 run_id = %self.run_id,
362 step = %step_name,
363 trace_id = %step_trace,
364 duration_ms = output.duration_ms,
365 "parallel step completed"
366 );
367
368 indexed_results[idx] = Some(Ok(output));
369 }
370 Err(err) => {
371 let err_msg = err.to_string();
372 let debug_messages_json = extract_debug_messages_from_error(&err);
373 let partial = extract_partial_usage_from_error(&err);
374 let raw_response_output = extract_raw_response_from_error(&err);
375
376 if let Some(ref usage) = partial {
377 if let Some(cost) = usage.cost_usd {
378 self.total_cost_usd += cost;
379 }
380 if let Some(dur) = usage.duration_ms {
381 self.total_duration_ms += dur;
382 }
383 }
384
385 if let Err(store_err) = self
386 .store
387 .update_step(
388 *step_id,
389 StepUpdate {
390 status: Some(StepStatus::Failed),
391 error: Some(err_msg.clone()),
392 output: raw_response_output.clone(),
393 completed_at: Some(completed_at),
394 debug_messages: debug_messages_json,
395 duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
396 cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
397 input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
398 output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
399 ..StepUpdate::default()
400 },
401 )
402 .await
403 {
404 error!(
405 step_id = %step_id,
406 error = %store_err,
407 "failed to persist parallel step failure"
408 );
409 }
410
411 let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
412 let err_cost = partial
413 .as_ref()
414 .and_then(|p| p.cost_usd)
415 .unwrap_or(Decimal::ZERO);
416 self.step_results.push(StepResult::from_failure(
417 *step_trace,
418 step_name,
419 &err_msg,
420 err_duration,
421 err_cost,
422 ));
423
424 if step_config.allow_failure() {
425 self.has_allowed_failure = true;
426 info!(
427 run_id = %self.run_id,
428 step = %step_name,
429 error = %err_msg,
430 "parallel step failed but allow_failure is set, continuing"
431 );
432 indexed_results[idx] = Some(Ok(allowed_failure_output(
433 &err_msg,
434 raw_response_output,
435 partial.as_ref(),
436 )));
437 } else {
438 indexed_results[idx] = Some(Err(err_msg.clone()));
439
440 if first_error.is_none() {
441 first_error = Some(err);
442 }
443
444 if fail_fast {
445 join_set.abort_all();
446 }
447 }
448 }
449 }
450 }
451
452 if let Some(err) = first_error {
453 return Err(err);
454 }
455
456 self.persist_progress().await;
457
458 self.last_step_ids = step_records.iter().map(|(id, _, _, _)| *id).collect();
459
460 let results: Vec<ParallelStepResult> = step_records
462 .iter()
463 .enumerate()
464 .map(|(idx, (step_id, _trace_id, name, _))| {
465 let output = match indexed_results[idx].take() {
466 Some(Ok(o)) => o,
467 _ => unreachable!("all steps succeeded if no error returned"),
468 };
469 ParallelStepResult {
470 name: name.clone(),
471 output,
472 step_id: *step_id,
473 }
474 })
475 .collect();
476
477 Ok(results)
478 }
479}