ironflow_engine/context/steps/
decision.rs1use std::collections::BTreeMap;
12
13use chrono::Utc;
14use serde_json::{Value, from_value, to_value};
15use tracing::info;
16use uuid::Uuid;
17
18use ironflow_core::decision::{DecisionOutput, DecisionUsage};
19use ironflow_store::models::{NewStep, StepKind, StepStatus, StepUpdate, step_trace_id};
20
21use crate::config::DecisionConfig;
22use crate::context::WorkflowContext;
23use crate::context::lifecycle::check_replay_identity;
24use crate::decision::DecisionAnswers;
25use crate::error::EngineError;
26use crate::executor::{DecisionExecution, StepArtifacts, StepOutput, StepResult, execute_decision};
27use crate::notify::{
28 WorkflowApprovalRequiredEvent, WorkflowEvent, WorkflowStepCompletedEvent,
29 WorkflowStepStartedEvent,
30};
31use crate::plan::lock_plan;
32
33impl WorkflowContext {
34 pub async fn decision<T: DecisionAnswers>(
90 &mut self,
91 name: &str,
92 config: DecisionConfig<T>,
93 ) -> Result<T, EngineError> {
94 let output = self.decision_output(name, config.erase()).await?;
95 Ok(T::from_output(&output)?)
96 }
97
98 async fn decision_output(
100 &mut self,
101 name: &str,
102 config: DecisionConfig,
103 ) -> Result<DecisionOutput, EngineError> {
104 if let Some(plan) = self.plan().cloned() {
107 self.position += 1;
108 let mut recorder = lock_plan(&plan);
109 if recorder.record(name, StepKind::Decision, &self.workflow_name, None) {
110 recorder.set_last(vec![name.to_string()]);
111 }
112 return Ok(DecisionOutput {
113 model: None,
114 answers: BTreeMap::new(),
115 usage: DecisionUsage::default(),
116 });
117 }
118
119 if let Some(output) = self.decision_replay(name, &config).await? {
120 return Ok(output);
121 }
122 self.decision_execute(name, config).await
123 }
124
125 async fn decision_replay(
132 &mut self,
133 name: &str,
134 _config: &DecisionConfig,
135 ) -> Result<Option<DecisionOutput>, EngineError> {
136 let position = self.position;
137
138 let Some(existing) = self.replay_steps.get(&position).cloned() else {
139 return Ok(None);
140 };
141 check_replay_identity(&existing, position, name, &StepKind::Decision)?;
142
143 self.position += 1;
144
145 let stored: DecisionOutput = existing
146 .output
147 .clone()
148 .ok_or_else(|| {
149 EngineError::StepConfig(format!(
150 "decision step '{name}' has no stored output to replay"
151 ))
152 })
153 .and_then(|v| from_value(v).map_err(EngineError::from))?;
154
155 if existing.status.state == StepStatus::AwaitingApproval {
158 self.store
159 .update_step(
160 existing.id,
161 StepUpdate {
162 status: Some(StepStatus::Completed),
163 completed_at: Some(Utc::now()),
164 ..StepUpdate::default()
165 },
166 )
167 .await?;
168 info!(
169 run_id = %self.run_id,
170 step = %name,
171 position,
172 "decision step replayed (approved after escalation)"
173 );
174 } else {
175 info!(
176 run_id = %self.run_id,
177 step = %name,
178 position,
179 "decision step replayed from previous execution"
180 );
181 }
182
183 self.last_step_ids = vec![existing.id];
189 Ok(Some(stored))
190 }
191
192 async fn decision_execute(
195 &mut self,
196 name: &str,
197 config: DecisionConfig,
198 ) -> Result<DecisionOutput, EngineError> {
199 self.check_guard_timeout()?;
200
201 let position = self.position;
202 self.position += 1;
203
204 let provider =
205 self.decision_provider
206 .clone()
207 .ok_or_else(|| EngineError::NoDecisionProvider {
208 step: name.to_string(),
209 })?;
210
211 let trace_id = step_trace_id(self.run_id, name, position);
212 let step = self
213 .store
214 .create_step(NewStep {
215 run_id: self.run_id,
216 trace_id,
217 name: name.to_string(),
218 kind: StepKind::Decision,
219 position,
220 input: Some(to_value(&config)?),
221 is_error_handler: false,
222 })
223 .await?;
224 self.start_step(step.id, Utc::now()).await?;
225
226 if let Some(ref bus) = self.event_bus {
227 bus.publish(
228 self.run_id,
229 WorkflowEvent::StepStarted(WorkflowStepStartedEvent {
230 step_name: name.to_string(),
231 step_index: position,
232 timestamp: Utc::now(),
233 }),
234 );
235 }
236
237 let execution = match execute_decision(&provider, &config).await {
238 Ok(execution) => execution,
239 Err(err) => {
240 self.fail_step(step.id, &err).await;
241 return Err(err);
242 }
243 };
244
245 self.total_cost_usd += execution.cost_usd;
247 self.total_duration_ms += execution.duration_ms;
248
249 let output_value = to_value(&execution.output)?;
250
251 let escalated = config
252 .escalate_below
253 .zip(execution.output.min_confidence())
254 .map(|(threshold, min)| min < threshold)
255 .unwrap_or(false);
256
257 if escalated {
258 return self
259 .decision_escalate(name, position, step.id, &config, &execution, output_value)
260 .await;
261 }
262
263 let step_output = StepOutput {
264 output: output_value.clone(),
265 duration_ms: execution.duration_ms,
266 cost_usd: execution.cost_usd,
267 input_tokens: Some(execution.input_tokens),
268 cache_read_input_tokens: None,
269 cache_creation_input_tokens: None,
270 output_tokens: Some(execution.output_tokens),
271 model: execution.output.model.as_ref().map(ToString::to_string),
272 debug_messages: None,
273 artifacts: StepArtifacts::default(),
274 account_id: None,
275 };
276
277 let completed_at = Utc::now();
278 self.store
279 .update_step(
280 step.id,
281 StepUpdate {
282 status: Some(StepStatus::Completed),
283 output: Some(output_value),
284 duration_ms: Some(execution.duration_ms),
285 cost_usd: Some(execution.cost_usd),
286 input_tokens: Some(execution.input_tokens),
287 output_tokens: Some(execution.output_tokens),
288 completed_at: Some(completed_at),
289 ..StepUpdate::default()
290 },
291 )
292 .await?;
293
294 self.step_results
295 .push(StepResult::from_success(trace_id, name, &step_output));
296 self.persist_progress().await;
297 self.last_step_ids = vec![step.id];
298
299 info!(
300 run_id = %self.run_id,
301 step = %name,
302 trace_id = %trace_id,
303 cost_usd = %execution.cost_usd,
304 "decision step completed"
305 );
306
307 if let Some(ref bus) = self.event_bus {
308 bus.publish(
309 self.run_id,
310 WorkflowEvent::StepCompleted(WorkflowStepCompletedEvent {
311 step_name: name.to_string(),
312 step_index: position,
313 duration_ms: execution.duration_ms,
314 output_summary: None,
315 }),
316 );
317 }
318
319 Ok(execution.output)
320 }
321
322 async fn decision_escalate(
325 &mut self,
326 name: &str,
327 position: u32,
328 step_id: Uuid,
329 config: &DecisionConfig,
330 execution: &DecisionExecution,
331 output_value: Value,
332 ) -> Result<DecisionOutput, EngineError> {
333 let threshold = config.escalate_below.unwrap_or_default();
334 let min = execution.output.min_confidence().unwrap_or_default();
335
336 self.store
337 .update_step(
338 step_id,
339 StepUpdate {
340 status: Some(StepStatus::AwaitingApproval),
341 output: Some(output_value),
342 duration_ms: Some(execution.duration_ms),
343 cost_usd: Some(execution.cost_usd),
344 input_tokens: Some(execution.input_tokens),
345 output_tokens: Some(execution.output_tokens),
346 ..StepUpdate::default()
347 },
348 )
349 .await?;
350 self.last_step_ids = vec![step_id];
351
352 info!(
353 run_id = %self.run_id,
354 step = %name,
355 position,
356 confidence = min,
357 threshold,
358 "decision escalated to human approval"
359 );
360
361 if let Some(ref bus) = self.event_bus {
362 bus.publish(
363 self.run_id,
364 WorkflowEvent::ApprovalRequired(WorkflowApprovalRequiredEvent {
365 step_name: name.to_string(),
366 step_index: position,
367 approval_id: step_id,
368 }),
369 );
370 }
371
372 Err(EngineError::ApprovalRequired {
373 run_id: self.run_id,
374 step_id,
375 message: format!(
376 "decision '{name}' escalated: confidence {min:.3} below threshold {threshold:.3}"
377 ),
378 })
379 }
380}