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