1use std::collections::HashMap;
33use std::mem::replace;
34use std::sync::{Arc, Mutex, MutexGuard};
35use std::time::Duration;
36
37use rust_decimal::Decimal;
38use serde::{Deserialize, Serialize};
39use serde_json::{Value, json};
40
41use ironflow_store::entities::{RunFilter, RunStatus, StepKind, StepStatus};
42use ironflow_store::store::Store;
43
44use crate::config::StepConfig;
45use crate::error::EngineError;
46use crate::executor::{StepArtifacts, StepOutput};
47
48pub const DEFAULT_PLAN_MAX_DEPTH: u32 = 3;
50
51pub const MAX_PLANNED_STEPS: usize = 1000;
56
57pub const DEFAULT_ESTIMATE_SAMPLE_RUNS: u32 = 20;
59
60#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
79#[serde(tag = "state", rename_all = "snake_case")]
80pub enum ConditionResult {
81 Evaluated {
83 expression: String,
86 value: bool,
88 },
89 Skipped {
92 reason: String,
94 },
95 Unevaluable {
97 expression: String,
99 reason: String,
101 },
102}
103
104#[derive(Debug, Clone, Serialize, Deserialize)]
125pub struct PlannedStep {
126 pub name: String,
128 pub kind: StepKind,
130 pub workflow: String,
132 pub depth: u32,
134 pub depends_on: Vec<String>,
140 pub condition: Option<ConditionResult>,
146 pub parallel_group: Option<String>,
149 #[serde(with = "opt_duration_ms")]
151 pub estimated_duration: Option<Duration>,
152}
153
154#[derive(Debug, Clone, Serialize, Deserialize)]
172pub struct ExecutionPlan {
173 pub workflow: String,
175 pub steps: Vec<PlannedStep>,
177 #[serde(with = "opt_duration_ms")]
179 pub estimated_duration: Option<Duration>,
180 pub max_depth: u32,
182 pub truncated: bool,
184 pub incomplete_reason: Option<String>,
186}
187
188#[derive(Debug, Clone)]
199pub struct PlanOptions {
200 pub max_depth: u32,
202 pub estimate_durations: bool,
204 pub sample_runs: u32,
206}
207
208impl Default for PlanOptions {
209 fn default() -> Self {
210 Self {
211 max_depth: DEFAULT_PLAN_MAX_DEPTH,
212 estimate_durations: true,
213 sample_runs: DEFAULT_ESTIMATE_SAMPLE_RUNS,
214 }
215 }
216}
217
218mod opt_duration_ms {
220 use std::time::Duration;
221
222 use serde::{Deserialize, Deserializer, Serialize, Serializer};
223
224 pub(super) fn serialize<S: Serializer>(
226 value: &Option<Duration>,
227 serializer: S,
228 ) -> Result<S::Ok, S::Error> {
229 let millis = value.map(|d| d.as_millis() as u64);
230 millis.serialize(serializer)
231 }
232
233 pub(super) fn deserialize<'de, D: Deserializer<'de>>(
235 deserializer: D,
236 ) -> Result<Option<Duration>, D::Error> {
237 let millis = Option::<u64>::deserialize(deserializer)?;
238 Ok(millis.map(Duration::from_millis))
239 }
240}
241
242pub(crate) type SharedPlanRecorder = Arc<Mutex<PlanRecorder>>;
244
245pub(crate) struct PlanRecorder {
247 payload: Value,
248 workflow: String,
249 steps: Vec<PlannedStep>,
250 last_names: Vec<String>,
251 pending_condition: Option<ConditionResult>,
252 depth: u32,
253 max_depth: u32,
254 parallel_groups: u32,
255 estimates: HashMap<String, Duration>,
256 truncated: bool,
257 incomplete_reason: Option<String>,
258}
259
260impl PlanRecorder {
261 pub(crate) fn new(
263 workflow: String,
264 payload: Value,
265 max_depth: u32,
266 estimates: HashMap<String, Duration>,
267 ) -> Self {
268 Self {
269 payload,
270 workflow,
271 steps: Vec::new(),
272 last_names: Vec::new(),
273 pending_condition: None,
274 depth: 0,
275 max_depth,
276 parallel_groups: 0,
277 estimates,
278 truncated: false,
279 incomplete_reason: None,
280 }
281 }
282
283 pub(crate) fn payload(&self) -> Value {
285 self.payload.clone()
286 }
287
288 pub(crate) fn swap_payload(&mut self, next: Value) -> Value {
293 replace(&mut self.payload, next)
294 }
295
296 pub(crate) fn set_condition(&mut self, condition: ConditionResult) {
298 self.pending_condition = Some(condition);
299 }
300
301 pub(crate) fn estimate_for(&self, name: &str) -> Option<Duration> {
303 self.estimates.get(name).copied()
304 }
305
306 pub(crate) fn seed_estimate(&mut self, name: &str, duration: Duration) {
311 self.estimates.entry(name.to_string()).or_insert(duration);
312 }
313
314 pub(crate) fn record(
319 &mut self,
320 name: &str,
321 kind: StepKind,
322 workflow: &str,
323 parallel_group: Option<String>,
324 ) -> bool {
325 if self.steps.len() >= MAX_PLANNED_STEPS {
326 self.truncated = true;
327 if self.incomplete_reason.is_none() {
328 self.incomplete_reason = Some(format!("step cap of {MAX_PLANNED_STEPS} reached"));
329 }
330 return false;
331 }
332
333 let estimated_duration = self.estimates.get(name).copied();
334 self.steps.push(PlannedStep {
335 name: name.to_string(),
336 kind,
337 workflow: workflow.to_string(),
338 depth: self.depth,
339 depends_on: self.last_names.clone(),
340 condition: self.pending_condition.take(),
341 parallel_group,
342 estimated_duration,
343 });
344 true
345 }
346
347 pub(crate) fn set_last(&mut self, names: Vec<String>) {
349 self.last_names = names;
350 }
351
352 pub(crate) fn next_group(&mut self) -> String {
354 self.parallel_groups += 1;
355 format!("parallel-{}", self.parallel_groups)
356 }
357
358 pub(crate) fn enter_workflow(&mut self) -> bool {
360 if self.depth + 1 > self.max_depth {
361 self.truncated = true;
362 if self.incomplete_reason.is_none() {
363 self.incomplete_reason = Some(format!(
364 "sub-workflow expansion stopped at depth {}",
365 self.max_depth
366 ));
367 }
368 return false;
369 }
370 self.depth += 1;
371 true
372 }
373
374 pub(crate) fn leave_workflow(&mut self) {
376 self.depth = self.depth.saturating_sub(1);
377 }
378
379 pub(crate) fn fail(&mut self, reason: String) {
381 self.truncated = true;
382 if self.incomplete_reason.is_none() {
383 self.incomplete_reason = Some(reason);
384 }
385 }
386
387 pub(crate) fn snapshot(&self) -> ExecutionPlan {
389 ExecutionPlan {
390 workflow: self.workflow.clone(),
391 estimated_duration: total_estimate(&self.steps),
392 steps: self.steps.clone(),
393 max_depth: self.max_depth,
394 truncated: self.truncated,
395 incomplete_reason: self.incomplete_reason.clone(),
396 }
397 }
398
399 pub(crate) fn into_plan(self) -> ExecutionPlan {
401 ExecutionPlan {
402 workflow: self.workflow,
403 estimated_duration: total_estimate(&self.steps),
404 steps: self.steps,
405 max_depth: self.max_depth,
406 truncated: self.truncated,
407 incomplete_reason: self.incomplete_reason,
408 }
409 }
410}
411
412pub(crate) fn lock_plan(plan: &SharedPlanRecorder) -> MutexGuard<'_, PlanRecorder> {
417 plan.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
418}
419
420fn total_estimate(steps: &[PlannedStep]) -> Option<Duration> {
426 let mut total = Duration::ZERO;
427 let mut seen_any = false;
428 let mut current_group: Option<&str> = None;
429 let mut group_max = Duration::ZERO;
430
431 for step in steps {
432 match step.parallel_group.as_deref() {
433 Some(group) if current_group == Some(group) => {
434 if let Some(estimate) = step.estimated_duration {
435 seen_any = true;
436 group_max = group_max.max(estimate);
437 }
438 }
439 Some(group) => {
440 if current_group.is_some() {
441 total += group_max;
442 }
443 current_group = Some(group);
444 group_max = step.estimated_duration.unwrap_or(Duration::ZERO);
445 if step.estimated_duration.is_some() {
446 seen_any = true;
447 }
448 }
449 None => {
450 if current_group.is_some() {
451 total += group_max;
452 current_group = None;
453 group_max = Duration::ZERO;
454 }
455 if let Some(estimate) = step.estimated_duration {
456 seen_any = true;
457 total += estimate;
458 }
459 }
460 }
461 }
462
463 if current_group.is_some() {
464 total += group_max;
465 }
466
467 seen_any.then_some(total)
468}
469
470pub(crate) fn planned_output(config: &StepConfig, estimate: Option<Duration>) -> StepOutput {
475 let output = match config {
476 StepConfig::Shell(_) => json!({"stdout": "", "stderr": "", "exit_code": 0}),
477 StepConfig::Http(_) => json!({"status": 200, "headers": {}, "body": ""}),
478 StepConfig::Agent(_) => json!({}),
479 StepConfig::Workflow(c) => json!({
480 "run_id": Value::Null,
481 "workflow_name": c.workflow_name,
482 "status": "completed",
483 "cost_usd": 0,
484 "duration_ms": 0,
485 }),
486 StepConfig::Approval(_) | StepConfig::Decision(_) | StepConfig::Delay(_) => Value::Null,
487 };
488
489 StepOutput {
490 output,
491 duration_ms: estimate.map(|d| d.as_millis() as u64).unwrap_or(0),
492 cost_usd: Decimal::ZERO,
493 input_tokens: None,
494 cache_read_input_tokens: None,
495 cache_creation_input_tokens: None,
496 output_tokens: None,
497 model: None,
498 debug_messages: None,
499 artifacts: StepArtifacts::default(),
500 account_id: None,
501 environment_id: None,
502 }
503}
504
505pub(crate) fn planned_custom_output(estimate: Option<Duration>) -> StepOutput {
507 StepOutput {
508 output: json!({}),
509 duration_ms: estimate.map(|d| d.as_millis() as u64).unwrap_or(0),
510 cost_usd: Decimal::ZERO,
511 input_tokens: None,
512 cache_read_input_tokens: None,
513 cache_creation_input_tokens: None,
514 output_tokens: None,
515 model: None,
516 debug_messages: None,
517 artifacts: StepArtifacts::default(),
518 account_id: None,
519 environment_id: None,
520 }
521}
522
523pub async fn estimate_durations(
549 store: &Arc<dyn Store>,
550 workflow_name: &str,
551 sample_runs: u32,
552) -> Result<HashMap<String, Duration>, EngineError> {
553 let page = store
554 .list_runs(
555 RunFilter {
556 workflow_name: Some(workflow_name.to_string()),
557 status: Some(RunStatus::Completed),
558 has_steps: Some(true),
559 ..RunFilter::default()
560 },
561 1,
562 sample_runs.clamp(1, 100),
563 )
564 .await?;
565
566 let mut totals: HashMap<String, (u64, u64)> = HashMap::new();
567 for run in &page.items {
568 for step in store.list_steps(run.id).await? {
569 if step.status.state != StepStatus::Completed {
570 continue;
571 }
572 let entry = totals.entry(step.name).or_insert((0, 0));
573 entry.0 += step.duration_ms;
574 entry.1 += 1;
575 }
576 }
577
578 Ok(totals
579 .into_iter()
580 .filter(|(_, (_, count))| *count > 0)
581 .map(|(name, (sum, count))| (name, Duration::from_millis(sum / count)))
582 .collect())
583}
584
585impl ExecutionPlan {
586 pub fn estimated_duration_ms(&self) -> Option<u64> {
606 self.estimated_duration.map(|d| d.as_millis() as u64)
607 }
608}
609
610impl PlannedStep {
611 pub fn estimated_duration_ms(&self) -> Option<u64> {
634 self.estimated_duration.map(|d| d.as_millis() as u64)
635 }
636}
637
638#[cfg(test)]
639mod tests {
640 use std::collections::HashMap;
641 use std::thread::spawn;
642 use std::time::Duration;
643
644 use serde_json::{from_value, json, to_value};
645
646 use crate::config::{HttpConfig, ShellConfig};
647
648 use super::*;
649
650 fn step(name: &str, group: Option<&str>, estimate: Option<u64>) -> PlannedStep {
651 PlannedStep {
652 name: name.to_string(),
653 kind: StepKind::Shell,
654 workflow: "wf".to_string(),
655 depth: 0,
656 depends_on: Vec::new(),
657 condition: None,
658 parallel_group: group.map(str::to_string),
659 estimated_duration: estimate.map(Duration::from_millis),
660 }
661 }
662
663 #[test]
664 fn total_estimate_sums_sequential_steps() {
665 let steps = vec![
666 step("a", None, Some(100)),
667 step("b", None, Some(250)),
668 step("c", None, Some(50)),
669 ];
670 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(400)));
671 }
672
673 #[test]
674 fn total_estimate_counts_a_parallel_group_once_at_its_slowest() {
675 let steps = vec![
676 step("build", None, Some(100)),
677 step("t1", Some("parallel-1"), Some(300)),
678 step("t2", Some("parallel-1"), Some(700)),
679 step("t3", Some("parallel-1"), Some(200)),
680 step("deploy", None, Some(100)),
681 ];
682 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(900)));
683 }
684
685 #[test]
686 fn total_estimate_handles_a_trailing_parallel_group() {
687 let steps = vec![
688 step("build", None, Some(100)),
689 step("t1", Some("parallel-1"), Some(300)),
690 step("t2", Some("parallel-1"), Some(700)),
691 ];
692 assert_eq!(total_estimate(&steps), Some(Duration::from_millis(800)));
693 }
694
695 #[test]
696 fn total_estimate_is_none_without_any_estimate() {
697 let steps = vec![step("a", None, None), step("b", None, None)];
698 assert_eq!(total_estimate(&steps), None);
699 }
700
701 #[test]
702 fn planned_step_duration_round_trips_as_milliseconds() {
703 let original = step("a", None, Some(1234));
704 let value = to_value(&original).expect("serialize");
705 assert_eq!(value["estimated_duration"], 1234);
706
707 let back: PlannedStep = from_value(value).expect("deserialize");
708 assert_eq!(back.estimated_duration, Some(Duration::from_millis(1234)));
709 }
710
711 #[test]
712 fn planned_step_duration_round_trips_when_absent() {
713 let original = step("a", None, None);
714 let value = to_value(&original).expect("serialize");
715 assert!(value["estimated_duration"].is_null());
716
717 let back: PlannedStep = from_value(value).expect("deserialize");
718 assert_eq!(back.estimated_duration, None);
719 }
720
721 #[test]
722 fn planned_shell_and_http_outputs_look_successful() {
723 let shell = planned_output(&StepConfig::Shell(ShellConfig::new("echo hi")), None);
724 assert!(shell.is_success());
725
726 let http = planned_output(
727 &StepConfig::Http(HttpConfig::get("https://example.com")),
728 None,
729 );
730 assert!(http.is_success());
731 }
732
733 #[test]
734 fn planned_output_carries_the_estimate_as_its_duration() {
735 let output = planned_output(
736 &StepConfig::Shell(ShellConfig::new("echo hi")),
737 Some(Duration::from_millis(900)),
738 );
739 assert_eq!(output.duration_ms, 900);
740 assert_eq!(output.cost_usd, Decimal::ZERO);
741 }
742
743 #[test]
744 fn planned_custom_output_is_an_empty_object() {
745 let output = planned_custom_output(None);
746 assert_eq!(output.output, json!({}));
747 assert_eq!(output.duration_ms, 0);
748 }
749
750 #[test]
751 fn condition_result_serializes_its_state_tag() {
752 let evaluated = to_value(ConditionResult::Evaluated {
753 expression: "env == prod".to_string(),
754 value: false,
755 })
756 .expect("serialize");
757 assert_eq!(evaluated["state"], "evaluated");
758 assert_eq!(evaluated["value"], false);
759
760 let skipped = to_value(ConditionResult::Skipped {
761 reason: "not prod".to_string(),
762 })
763 .expect("serialize");
764 assert_eq!(skipped["state"], "skipped");
765
766 let unevaluable = to_value(ConditionResult::Unevaluable {
767 expression: "build succeeded".to_string(),
768 reason: "depends on a step output".to_string(),
769 })
770 .expect("serialize");
771 assert_eq!(unevaluable["state"], "unevaluable");
772 }
773
774 #[test]
775 fn recorder_records_dependencies_and_conditions() {
776 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
777 assert!(recorder.record("a", StepKind::Shell, "wf", None));
778 recorder.set_last(vec!["a".to_string()]);
779 recorder.set_condition(ConditionResult::Skipped {
780 reason: "nope".to_string(),
781 });
782 assert!(recorder.record("b", StepKind::Shell, "wf", None));
783
784 let plan = recorder.into_plan();
785 assert_eq!(plan.steps.len(), 2);
786 assert_eq!(plan.steps[1].depends_on, vec!["a".to_string()]);
787 assert!(matches!(
788 plan.steps[1].condition,
789 Some(ConditionResult::Skipped { .. })
790 ));
791 assert!(plan.steps[0].condition.is_none());
792 }
793
794 #[test]
795 fn recorder_stops_at_the_step_cap() {
796 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
797 for index in 0..MAX_PLANNED_STEPS {
798 assert!(recorder.record(&format!("s{index}"), StepKind::Shell, "wf", None));
799 }
800 assert!(!recorder.record("overflow", StepKind::Shell, "wf", None));
801
802 let plan = recorder.into_plan();
803 assert_eq!(plan.steps.len(), MAX_PLANNED_STEPS);
804 assert!(plan.truncated);
805 assert!(
806 plan.incomplete_reason
807 .expect("a reason")
808 .contains("step cap")
809 );
810 }
811
812 #[test]
813 fn recorder_refuses_to_expand_past_the_depth_limit() {
814 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 1, HashMap::new());
815 assert!(recorder.enter_workflow());
816 assert!(!recorder.enter_workflow());
817
818 let plan = recorder.snapshot();
819 assert!(plan.truncated);
820 assert!(
821 plan.incomplete_reason
822 .expect("a reason")
823 .contains("depth 1")
824 );
825 }
826
827 #[test]
828 fn recorder_allocates_successive_parallel_group_names() {
829 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
830 assert_eq!(recorder.next_group(), "parallel-1");
831 assert_eq!(recorder.next_group(), "parallel-2");
832 }
833
834 #[test]
835 fn recorder_swaps_and_restores_the_payload() {
836 let mut recorder =
837 PlanRecorder::new("wf".to_string(), json!({"env": "prod"}), 3, HashMap::new());
838 let previous = recorder.swap_payload(json!({"env": "dev"}));
839 assert_eq!(previous, json!({"env": "prod"}));
840 assert_eq!(recorder.payload(), json!({"env": "dev"}));
841 recorder.swap_payload(previous);
842 assert_eq!(recorder.payload(), json!({"env": "prod"}));
843 }
844
845 #[test]
846 fn recorder_keeps_the_first_failure_reason() {
847 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, HashMap::new());
848 recorder.fail("first".to_string());
849 recorder.fail("second".to_string());
850 let plan = recorder.into_plan();
851 assert_eq!(plan.incomplete_reason.as_deref(), Some("first"));
852 assert!(plan.truncated);
853 }
854
855 #[test]
856 fn recorder_uses_the_history_estimate_for_a_known_step() {
857 let estimates = HashMap::from([("build".to_string(), Duration::from_millis(400))]);
858 let mut recorder = PlanRecorder::new("wf".to_string(), json!({}), 3, estimates);
859 assert_eq!(
860 recorder.estimate_for("build"),
861 Some(Duration::from_millis(400))
862 );
863 assert_eq!(recorder.estimate_for("unknown"), None);
864 recorder.record("build", StepKind::Shell, "wf", None);
865 let plan = recorder.into_plan();
866 assert_eq!(plan.estimated_duration, Some(Duration::from_millis(400)));
867 }
868
869 #[test]
870 fn lock_plan_recovers_from_poisoning() {
871 let shared: SharedPlanRecorder = Arc::new(Mutex::new(PlanRecorder::new(
872 "wf".to_string(),
873 json!({}),
874 3,
875 HashMap::new(),
876 )));
877 let poisoner = Arc::clone(&shared);
878 let _ = spawn(move || {
879 let _guard = poisoner.lock().expect("lock");
880 panic!("poison the mutex");
881 })
882 .join();
883
884 assert_eq!(lock_plan(&shared).payload(), json!({}));
885 }
886}