1use 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::context::lifecycle::check_replay_identity;
23use crate::error::EngineError;
24use crate::executor::{
25 ParallelStepResult, StepArtifacts, StepOutput, StepResult, execute_step_config_intercepted,
26};
27use crate::guard::WorkflowRejection;
28use crate::log_sender::StepLogSender;
29use crate::notify::{WorkflowAgentStepTokensUsedEvent, WorkflowEvent};
30use crate::plan::{lock_plan, planned_output};
31
32impl WorkflowContext {
33 pub async fn parallel(
82 &mut self,
83 steps: Vec<(&str, StepConfig)>,
84 fail_fast: bool,
85 ) -> Result<Vec<ParallelStepResult>, EngineError> {
86 if steps.is_empty() {
87 return Ok(Vec::new());
88 }
89 reject_duplicate_names(&steps)?;
90
91 if let Some(plan) = self.plan().cloned() {
94 self.position += 1;
95 let mut results = Vec::with_capacity(steps.len());
96 let mut names = Vec::with_capacity(steps.len());
97 {
98 let mut recorder = lock_plan(&plan);
99 let group = recorder.next_group();
100 for (name, config) in &steps {
101 let wave = Some(group.clone());
102 if !recorder.record(name, config.kind(), &self.workflow_name, wave) {
103 break;
104 }
105 names.push((*name).to_string());
106 let estimate = recorder.estimate_for(name);
107 let mut output = planned_output(config, estimate);
108 output.artifacts = StepArtifacts::new(name, None, config.declared_outputs());
109 results.push(ParallelStepResult {
110 name: (*name).to_string(),
111 output,
112 step_id: Uuid::now_v7(),
113 });
114 }
115 recorder.set_last(names);
116 }
117 return Ok(results);
118 }
119
120 self.check_guard_timeout()?;
122
123 let wave_position = self.position;
124 self.position += 1;
125
126 let mut slots = self.replay_wave(wave_position, &steps)?;
132 if slots.iter().all(Option::is_some) {
133 let results: Vec<ParallelStepResult> = slots.into_iter().flatten().collect();
134 self.last_step_ids = results.iter().map(|r| r.step_id).collect();
135 return Ok(results);
136 }
137
138 let wave_budget: Decimal = steps
141 .iter()
142 .zip(&slots)
143 .filter(|(_, slot)| slot.is_none())
144 .filter_map(|((_, config), _)| match config {
145 StepConfig::Agent(agent_config) => Some(agent_config.max_budget_usd),
146 _ => None,
147 })
148 .map(step_budget_usd)
149 .sum();
150 self.check_run_budget(wave_budget)?;
151
152 let now = Utc::now();
153 let mut step_records: Vec<(Uuid, Uuid, String, StepConfig)> =
154 Vec::with_capacity(steps.len());
155 let mut record_slots: Vec<usize> = Vec::with_capacity(steps.len());
157
158 for (slot, (name, config)) in steps.iter().enumerate() {
159 if slots[slot].is_some() {
160 continue;
161 }
162 let kind = config.kind();
163 let trace_id = step_trace_id(self.run_id, name, wave_position);
164 let step = self
165 .store
166 .create_step(NewStep {
167 run_id: self.run_id,
168 trace_id,
169 name: name.to_string(),
170 kind,
171 position: wave_position,
172 input: Some(to_value(config)?),
173 is_error_handler: false,
174 })
175 .await?;
176
177 self.start_step(step.id, now).await?;
178
179 if let Err(err) = self.prepare_step_inputs(config, wave_position).await {
182 self.fail_step(step.id, &err).await;
183 if !config.allow_failure() {
184 return Err(err);
185 }
186 self.has_allowed_failure = true;
187 info!(
188 run_id = %self.run_id,
189 step = %name,
190 error = %err,
191 "parallel step input preparation failed but allow_failure is set, skipping"
192 );
193 continue;
194 }
195
196 let mut config_with_trace = config.clone();
197 self.scope_step_config(&mut config_with_trace, name);
198 step_records.push((step.id, trace_id, name.to_string(), config_with_trace));
199 record_slots.push(slot);
200 }
201
202 let mut join_set = JoinSet::new();
203 let mut task_index: HashMap<Id, usize> = HashMap::new();
204 let parallel_timeout = self.guard_remaining_timeout();
205 for (idx, (step_id, _trace_id, step_name, config)) in step_records.iter().enumerate() {
206 let provider = self.provider.clone();
207 let interceptor = self.interceptor.clone();
210 let config = config.clone();
211 let step_log_sender = self
212 .log_sender
213 .as_ref()
214 .map(|s| StepLogSender::new(s.clone(), self.run_id, *step_id, step_name.clone()));
215 let handle = join_set.spawn(async move {
216 let result = match parallel_timeout {
217 Some(dur) => {
218 match timeout(
219 dur,
220 execute_step_config_intercepted(
221 &config,
222 &provider,
223 step_log_sender,
224 interceptor.as_ref(),
225 ),
226 )
227 .await
228 {
229 Ok(r) => r,
230 Err(_elapsed) => {
231 Err(EngineError::from(WorkflowRejection::WorkflowTimeout {
232 elapsed_secs: 0,
233 max: 0,
234 }))
235 }
236 }
237 }
238 None => {
239 execute_step_config_intercepted(
240 &config,
241 &provider,
242 step_log_sender,
243 interceptor.as_ref(),
244 )
245 .await
246 }
247 };
248 (idx, result)
249 });
250 task_index.insert(handle.id(), idx);
251 }
252
253 let mut indexed_results: Vec<Option<Result<StepOutput, String>>> =
255 vec![None; step_records.len()];
256 let mut first_error: Option<EngineError> = None;
257
258 while let Some(join_result) = join_set.join_next().await {
259 let (idx, step_result) = match join_result {
260 Ok(r) => r,
261 Err(e) => {
262 let error_msg = format!("join error: {e}");
263 if let Some(&idx) = task_index.get(&e.id()) {
264 let (step_id, _, step_name, _) = &step_records[idx];
265 let completed_at = Utc::now();
266 error!(
267 run_id = %self.run_id,
268 step = %step_name,
269 error = %error_msg,
270 "parallel step panicked or was cancelled"
271 );
272 if let Err(store_err) = self
273 .store
274 .update_step(
275 *step_id,
276 StepUpdate {
277 status: Some(StepStatus::Failed),
278 error: Some(error_msg.clone()),
279 completed_at: Some(completed_at),
280 ..StepUpdate::default()
281 },
282 )
283 .await
284 {
285 error!(
286 run_id = %self.run_id,
287 step_id = %step_id,
288 error = %store_err,
289 "failed to persist JoinError for step"
290 );
291 }
292 indexed_results[idx] = Some(Err(error_msg.clone()));
293 }
294 if first_error.is_none() {
295 first_error = Some(EngineError::StepConfig(error_msg));
296 }
297 if fail_fast {
298 join_set.abort_all();
299 }
300 continue;
301 }
302 };
303
304 let (step_id, step_trace, step_name, step_config) = &step_records[idx];
305 let completed_at = Utc::now();
306
307 if let Err(err) = self
308 .store_step_outputs(step_config, *step_id, step_name, step_result.is_ok())
309 .await
310 {
311 self.fail_step(*step_id, &err).await;
312 indexed_results[idx] = Some(Err(err.to_string()));
313 if first_error.is_none() {
314 first_error = Some(err);
315 }
316 if fail_fast {
317 join_set.abort_all();
318 }
319 continue;
320 }
321
322 match step_result {
323 Ok(output) => {
324 self.total_cost_usd += output.cost_usd;
325 self.total_duration_ms += output.duration_ms;
326
327 if matches!(step_config, StepConfig::Agent(_)) {
329 let tokens = output.total_tokens();
330 if tokens > 0
331 && let Err(guard_err) = self.guard_record_tokens(tokens)
332 {
333 if first_error.is_none() {
334 first_error = Some(guard_err);
335 }
336 if fail_fast {
337 join_set.abort_all();
338 }
339 }
340 }
341
342 let debug_messages_json = output.debug_messages_json();
343
344 self.store
345 .update_step(
346 *step_id,
347 StepUpdate {
348 status: Some(StepStatus::Completed),
349 output: Some(output.output.clone()),
350 duration_ms: Some(output.duration_ms),
351 cost_usd: Some(output.cost_usd),
352 input_tokens: output.input_tokens,
353 cache_read_input_tokens: output.cache_read_input_tokens,
354 cache_creation_input_tokens: output.cache_creation_input_tokens,
355 output_tokens: output.output_tokens,
356 completed_at: Some(completed_at),
357 debug_messages: debug_messages_json,
358 account_id: output.account_id,
359 ..StepUpdate::default()
360 },
361 )
362 .await?;
363
364 self.step_results.push(StepResult::from_success(
365 *step_trace,
366 step_name,
367 &output,
368 ));
369
370 if let Some(ref bus) = self.event_bus
371 && matches!(step_config, StepConfig::Agent(_))
372 {
373 let tokens = output.total_tokens();
374 bus.publish(
375 self.run_id,
376 WorkflowEvent::AgentStepTokensUsed(WorkflowAgentStepTokensUsedEvent {
377 step_name: step_name.clone(),
378 tokens,
379 cost_usd: output.cost_usd,
380 }),
381 );
382 }
383
384 info!(
385 run_id = %self.run_id,
386 step = %step_name,
387 trace_id = %step_trace,
388 duration_ms = output.duration_ms,
389 "parallel step completed"
390 );
391
392 indexed_results[idx] = Some(Ok(output));
393 }
394 Err(err) => {
395 let err_msg = err.to_string();
396 let debug_messages_json = extract_debug_messages_from_error(&err);
397 let partial = extract_partial_usage_from_error(&err);
398 let raw_response_output = extract_raw_response_from_error(&err);
399
400 if let Some(ref usage) = partial {
401 if let Some(cost) = usage.cost_usd {
402 self.total_cost_usd += cost;
403 }
404 if let Some(dur) = usage.duration_ms {
405 self.total_duration_ms += dur;
406 }
407 }
408
409 if let Err(store_err) = self
410 .store
411 .update_step(
412 *step_id,
413 StepUpdate {
414 status: Some(StepStatus::Failed),
415 error: Some(err_msg.clone()),
416 output: raw_response_output.clone(),
417 completed_at: Some(completed_at),
418 debug_messages: debug_messages_json,
419 duration_ms: partial.as_ref().and_then(|p| p.duration_ms),
420 cost_usd: partial.as_ref().and_then(|p| p.cost_usd),
421 input_tokens: partial.as_ref().and_then(|p| p.input_tokens),
422 cache_read_input_tokens: partial
423 .as_ref()
424 .and_then(|p| p.cache_read_input_tokens),
425 cache_creation_input_tokens: partial
426 .as_ref()
427 .and_then(|p| p.cache_creation_input_tokens),
428 output_tokens: partial.as_ref().and_then(|p| p.output_tokens),
429 ..StepUpdate::default()
430 },
431 )
432 .await
433 {
434 error!(
435 step_id = %step_id,
436 error = %store_err,
437 "failed to persist parallel step failure"
438 );
439 }
440
441 let err_duration = partial.as_ref().and_then(|p| p.duration_ms).unwrap_or(0);
442 let err_cost = partial
443 .as_ref()
444 .and_then(|p| p.cost_usd)
445 .unwrap_or(Decimal::ZERO);
446 self.step_results.push(StepResult::from_failure(
447 *step_trace,
448 step_name,
449 &err_msg,
450 err_duration,
451 err_cost,
452 ));
453
454 if step_config.allow_failure() {
455 self.has_allowed_failure = true;
456 info!(
457 run_id = %self.run_id,
458 step = %step_name,
459 error = %err_msg,
460 "parallel step failed but allow_failure is set, continuing"
461 );
462 indexed_results[idx] = Some(Ok(allowed_failure_output(
463 &err_msg,
464 raw_response_output,
465 partial.as_ref(),
466 )));
467 } else {
468 indexed_results[idx] = Some(Err(err_msg.clone()));
469
470 if first_error.is_none() {
471 first_error = Some(err);
472 }
473
474 if fail_fast {
475 join_set.abort_all();
476 }
477 }
478 }
479 }
480 }
481
482 if let Some(err) = first_error {
483 return Err(err);
484 }
485
486 self.persist_progress().await;
487
488 for (idx, (step_id, _trace_id, name, config)) in step_records.iter().enumerate() {
490 let mut output = match indexed_results[idx].take() {
491 Some(Ok(o)) => o,
492 _ => unreachable!("all steps succeeded if no error returned"),
493 };
494 output.artifacts = StepArtifacts::new(name, Some(*step_id), config.declared_outputs());
495 slots[record_slots[idx]] = Some(ParallelStepResult {
496 name: name.clone(),
497 output,
498 step_id: *step_id,
499 });
500 }
501 let results: Vec<ParallelStepResult> = slots.into_iter().flatten().collect();
502 self.last_step_ids = results.iter().map(|r| r.step_id).collect();
503
504 Ok(results)
505 }
506
507 fn replay_wave(
515 &mut self,
516 position: u32,
517 steps: &[(&str, StepConfig)],
518 ) -> Result<Vec<Option<ParallelStepResult>>, EngineError> {
519 let mut slots = Vec::with_capacity(steps.len());
520 for (name, config) in steps {
521 let Some(step) = self.replay_wave_steps.get(&(position, (*name).to_string())) else {
522 slots.push(None);
523 continue;
524 };
525 check_replay_identity(step, position, name, &config.kind())?;
526 if step.status.state != StepStatus::Completed {
527 slots.push(None);
528 continue;
529 }
530
531 let mut output = StepOutput::from(step);
532 let step_id = step.id;
533 self.total_duration_ms += output.duration_ms;
537 output.artifacts = StepArtifacts::new(name, Some(step_id), config.declared_outputs());
538
539 info!(
540 run_id = %self.run_id,
541 step = %name,
542 position,
543 "step replayed from previous execution"
544 );
545
546 slots.push(Some(ParallelStepResult {
547 name: (*name).to_string(),
548 output,
549 step_id,
550 }));
551 }
552 Ok(slots)
553 }
554}
555
556fn reject_duplicate_names(steps: &[(&str, StepConfig)]) -> Result<(), EngineError> {
562 let mut seen = HashSet::with_capacity(steps.len());
563 for (name, _) in steps {
564 if !seen.insert(*name) {
565 return Err(EngineError::StepConfig(format!(
566 "parallel wave has two steps named {name:?}; each step of a wave needs its own name"
567 )));
568 }
569 }
570 Ok(())
571}