ironflow_engine/executor/
agent.rs1use std::sync::Arc;
4use std::time::Instant;
5
6use rust_decimal::Decimal;
7use tracing::{info, warn};
8
9use ironflow_core::operations::agent::Agent;
10use ironflow_core::pricing::{CostBreakdown, StaticPricing, spawn_log};
11use ironflow_core::provider::{AgentConfig, AgentProvider, LogSink};
12use ironflow_store::entities::StepKind;
13
14use crate::error::EngineError;
15use crate::log_sender::StepLogSender;
16use crate::notify::LogStream;
17
18use super::{StepExecutor, StepOutput};
19
20pub struct AgentExecutor<'a> {
26 config: &'a AgentConfig,
27 log_sender: Option<StepLogSender>,
28}
29
30impl<'a> AgentExecutor<'a> {
31 pub fn new(config: &'a AgentConfig) -> Self {
33 Self {
34 config,
35 log_sender: None,
36 }
37 }
38
39 pub fn with_log_sender(mut self, sender: StepLogSender) -> Self {
41 self.log_sender = Some(sender);
42 self
43 }
44}
45
46impl StepExecutor for AgentExecutor<'_> {
47 fn kind(&self) -> StepKind {
48 StepKind::Agent
49 }
50
51 async fn execute(&self, provider: &Arc<dyn AgentProvider>) -> Result<StepOutput, EngineError> {
52 let start = Instant::now();
53
54 if let Some(ref sender) = self.log_sender {
55 sender.emit(
56 LogStream::System,
57 &format!("agent step started (model={})", self.config.model),
58 );
59 }
60
61 if self.config.json_schema.is_some() && self.config.max_turns == Some(1) {
62 warn!(
63 "structured output (json_schema) requires max_turns >= 2; \
64 max_turns is set to 1, the agent will likely fail with error_max_turns"
65 );
66 }
67
68 let mut agent = Agent::from_config(self.config.clone());
69 if let Some(ref sender) = self.log_sender {
70 agent = agent.log_sink(Arc::new(sender.clone()) as Arc<dyn LogSink>);
71 }
72 let result = agent.run(provider.as_ref()).await?;
73
74 let duration_ms = start.elapsed().as_millis() as u64;
75 let cost = Decimal::try_from(result.cost_usd().unwrap_or(0.0)).unwrap_or(Decimal::ZERO);
76 let input_tokens = result.input_tokens();
77 let output_tokens = result.output_tokens();
78
79 info!(
80 step_kind = "agent",
81 model = %self.config.model,
82 cost_usd = %cost,
83 input_tokens = ?input_tokens,
84 output_tokens = ?output_tokens,
85 duration_ms,
86 "agent step completed"
87 );
88
89 let pricing = StaticPricing::new();
90 let breakdown = CostBreakdown::compute(
91 &pricing,
92 &self.config.model,
93 input_tokens.unwrap_or(0),
94 output_tokens.unwrap_or(0),
95 );
96 spawn_log("agent", &self.config.model, breakdown);
97
98 #[cfg(feature = "prometheus")]
99 {
100 use ironflow_core::metric_names::{
101 AGENT_COST_USD_TOTAL, AGENT_DURATION_SECONDS, AGENT_TOKENS_INPUT_TOTAL,
102 AGENT_TOKENS_OUTPUT_TOTAL, AGENT_TOTAL, STATUS_SUCCESS,
103 };
104 use metrics::{counter, gauge, histogram};
105 let model_label = self.config.model.clone();
106 counter!(AGENT_TOTAL, "model" => model_label.clone(), "status" => STATUS_SUCCESS)
107 .increment(1);
108 histogram!(AGENT_DURATION_SECONDS, "model" => model_label.clone())
109 .record(duration_ms as f64 / 1000.0);
110 gauge!(AGENT_COST_USD_TOTAL, "model" => model_label.clone())
111 .increment(cost.to_string().parse::<f64>().unwrap_or(0.0));
112 if let Some(inp) = input_tokens {
113 counter!(AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(inp);
114 }
115 if let Some(out) = output_tokens {
116 counter!(AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label).increment(out);
117 }
118 }
119
120 if let Some(ref sender) = self.log_sender {
121 sender.emit(
122 LogStream::System,
123 &format!(
124 "agent step completed (cost=${cost}, tokens_in={}, tokens_out={})",
125 input_tokens.unwrap_or(0),
126 output_tokens.unwrap_or(0),
127 ),
128 );
129 }
130
131 let debug_messages = result.debug_messages().map(|msgs| msgs.to_vec());
132
133 Ok(StepOutput {
134 output: result.value().clone(),
135 duration_ms,
136 cost_usd: cost,
137 input_tokens,
138 output_tokens,
139 model: result.model().map(String::from),
140 debug_messages,
141 })
142 }
143}
144
145#[cfg(test)]
146mod tests {
147 use ironflow_core::operations::agent::PermissionMode;
148
149 #[test]
150 fn parse_permission_mode_via_serde() {
151 let json = r#""auto""#;
152 let mode: PermissionMode = serde_json::from_str(json).unwrap();
153 assert!(matches!(mode, PermissionMode::Auto));
154 }
155
156 #[test]
157 fn parse_permission_mode_dont_ask() {
158 let json = r#""dont_ask""#;
159 let mode: PermissionMode = serde_json::from_str(json).unwrap();
160 assert!(matches!(mode, PermissionMode::DontAsk));
161 }
162
163 #[test]
164 fn parse_permission_mode_bypass() {
165 let json = r#""bypass""#;
166 let mode: PermissionMode = serde_json::from_str(json).unwrap();
167 assert!(matches!(mode, PermissionMode::BypassPermissions));
168 }
169
170 #[test]
171 fn parse_permission_mode_case_insensitive() {
172 let json = r#""AUTO""#;
173 let mode: PermissionMode = serde_json::from_str(json).unwrap();
174 assert!(matches!(mode, PermissionMode::Auto));
175
176 let json = r#""DONT_ASK""#;
177 let mode: PermissionMode = serde_json::from_str(json).unwrap();
178 assert!(matches!(mode, PermissionMode::DontAsk));
179 }
180
181 #[test]
182 fn parse_permission_mode_unknown_defaults() {
183 let json = r#""unknown""#;
184 let mode: PermissionMode = serde_json::from_str(json).unwrap();
185 assert!(matches!(mode, PermissionMode::Default));
186 }
187}