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