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