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