ironflow_engine/executor/
agent.rs1use std::sync::Arc;
4use std::time::Instant;
5
6use rust_decimal::Decimal;
7use tracing::{info, warn};
8use uuid::Uuid;
9
10use ironflow_core::operations::agent::Agent;
11use ironflow_core::pricing::{CostBreakdown, StaticPricing, spawn_log};
12use ironflow_core::provider::{AgentConfig, AgentProvider, LogSink};
13use ironflow_store::entities::StepKind;
14
15use crate::error::EngineError;
16use crate::log_sender::StepLogSender;
17use crate::notify::LogStream;
18
19use super::{StepArtifacts, StepExecutor, StepOutput};
20
21pub struct AgentExecutor<'a> {
27 config: &'a AgentConfig,
28 log_sender: Option<StepLogSender>,
29}
30
31impl<'a> AgentExecutor<'a> {
32 pub fn new(config: &'a AgentConfig) -> Self {
34 Self {
35 config,
36 log_sender: None,
37 }
38 }
39
40 pub fn with_log_sender(mut self, sender: StepLogSender) -> Self {
42 self.log_sender = Some(sender);
43 self
44 }
45}
46
47impl StepExecutor for AgentExecutor<'_> {
48 fn kind(&self) -> StepKind {
49 StepKind::Agent
50 }
51
52 async fn execute(&self, provider: &Arc<dyn AgentProvider>) -> Result<StepOutput, EngineError> {
53 let start = Instant::now();
54
55 if let Some(ref sender) = self.log_sender {
56 sender.emit(
57 LogStream::System,
58 &format!("agent step started (model={})", self.config.model),
59 );
60 }
61
62 if self.config.json_schema.is_some() && self.config.max_turns == Some(1) {
63 warn!(
64 "structured output (json_schema) requires max_turns >= 2; \
65 max_turns is set to 1, the agent will likely fail with error_max_turns"
66 );
67 }
68
69 let mut agent = Agent::from_config(self.config.clone());
70 if let Some(ref sender) = self.log_sender {
71 agent = agent.log_sink(Arc::new(sender.clone()) as Arc<dyn LogSink>);
72 }
73 let result = agent.run(provider.as_ref()).await?;
74
75 let duration_ms = start.elapsed().as_millis() as u64;
76 let cost = Decimal::try_from(result.cost_usd().unwrap_or(0.0)).unwrap_or(Decimal::ZERO);
77 let input_tokens = result.input_tokens();
78 let cache_read_tokens = result.cache_read_input_tokens();
79 let cache_creation_tokens = result.cache_creation_input_tokens();
80 let output_tokens = result.output_tokens();
81
82 info!(
83 step_kind = "agent",
84 model = %self.config.model,
85 cost_usd = %cost,
86 input_tokens = ?input_tokens,
87 cache_read_input_tokens = ?cache_read_tokens,
88 cache_creation_input_tokens = ?cache_creation_tokens,
89 output_tokens = ?output_tokens,
90 duration_ms,
91 "agent step completed"
92 );
93
94 let pricing = StaticPricing::new();
95 let breakdown = CostBreakdown::compute_with_cache(
96 &pricing,
97 &self.config.model,
98 input_tokens.unwrap_or(0),
99 cache_read_tokens.unwrap_or(0),
100 cache_creation_tokens.unwrap_or(0),
101 output_tokens.unwrap_or(0),
102 );
103 spawn_log("agent", &self.config.model, breakdown);
104
105 #[cfg(feature = "prometheus")]
106 {
107 use ironflow_core::metric_names::{
108 AGENT_COST_USD_TOTAL, AGENT_DURATION_SECONDS, AGENT_TOKENS_CACHE_READ_TOTAL,
109 AGENT_TOKENS_CACHE_WRITE_TOTAL, AGENT_TOKENS_INPUT_TOTAL,
110 AGENT_TOKENS_OUTPUT_TOTAL, AGENT_TOTAL, STATUS_SUCCESS,
111 };
112 use metrics::{counter, gauge, histogram};
113 let model_label = self.config.model.clone();
114 counter!(AGENT_TOTAL, "model" => model_label.clone(), "status" => STATUS_SUCCESS)
115 .increment(1);
116 histogram!(AGENT_DURATION_SECONDS, "model" => model_label.clone())
117 .record(duration_ms as f64 / 1000.0);
118 gauge!(AGENT_COST_USD_TOTAL, "model" => model_label.clone())
119 .increment(cost.to_string().parse::<f64>().unwrap_or(0.0));
120 if let Some(inp) = input_tokens {
121 counter!(AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(inp);
122 }
123 if let Some(t) = cache_read_tokens {
124 counter!(AGENT_TOKENS_CACHE_READ_TOTAL, "model" => model_label.clone())
125 .increment(t);
126 }
127 if let Some(t) = cache_creation_tokens {
128 counter!(AGENT_TOKENS_CACHE_WRITE_TOTAL, "model" => model_label.clone())
129 .increment(t);
130 }
131 if let Some(out) = output_tokens {
132 counter!(AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label).increment(out);
133 }
134 }
135
136 if let Some(ref sender) = self.log_sender {
137 sender.emit(
138 LogStream::System,
139 &format!(
140 "agent step completed (cost=${cost}, tokens_in={}, cache_read={}, cache_write={}, tokens_out={})",
141 input_tokens.unwrap_or(0),
142 cache_read_tokens.unwrap_or(0),
143 cache_creation_tokens.unwrap_or(0),
144 output_tokens.unwrap_or(0),
145 ),
146 );
147 }
148
149 let debug_messages = result.debug_messages().map(|msgs| msgs.to_vec());
150 let account_id = match result.account_id() {
151 Some(raw) => match Uuid::parse_str(raw) {
152 Ok(id) => Some(id),
153 Err(e) => {
154 warn!(account_id = raw, error = %e, "agent output carries an invalid account id");
155 None
156 }
157 },
158 None => None,
159 };
160
161 Ok(StepOutput {
162 output: result.value().clone(),
163 duration_ms,
164 cost_usd: cost,
165 input_tokens,
166 cache_read_input_tokens: cache_read_tokens,
167 cache_creation_input_tokens: cache_creation_tokens,
168 output_tokens,
169 model: result.model().map(String::from),
170 debug_messages,
171 artifacts: StepArtifacts::default(),
172 account_id,
173 })
174 }
175}
176
177#[cfg(test)]
178mod tests {
179 use std::sync::Arc;
180 use std::time::Duration;
181
182 use ironflow_core::operations::agent::PermissionMode;
183 use ironflow_core::provider::{AgentConfig, AgentOutput, AgentProvider, InvokeFuture};
184 use serde_json::json;
185 use tokio::time::timeout;
186 use uuid::Uuid;
187
188 use super::{AgentExecutor, StepExecutor};
189
190 struct FixedUsageProvider {
192 output: AgentOutput,
193 }
194
195 impl AgentProvider for FixedUsageProvider {
196 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
197 Box::pin(async move { Ok(self.output.clone()) })
198 }
199 }
200
201 fn budget_config() -> AgentConfig {
202 let mut config = AgentConfig::new("hi");
203 config.max_budget_usd = Some(0.10);
204 config
205 }
206
207 #[tokio::test]
208 async fn agent_executor_propagates_cache_tokens() {
209 timeout(Duration::from_secs(10), async {
210 let mut output = AgentOutput::new(json!("ok"));
211 output.input_tokens = Some(100);
212 output.cache_read_input_tokens = Some(5000);
213 output.cache_creation_input_tokens = Some(200);
214 output.output_tokens = Some(50);
215 output.cost_usd = Some(0.02);
216 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
217
218 let config = budget_config();
219 let step = AgentExecutor::new(&config)
220 .execute(&provider)
221 .await
222 .expect("agent step succeeds");
223
224 assert_eq!(step.input_tokens, Some(100));
225 assert_eq!(step.cache_read_input_tokens, Some(5000));
226 assert_eq!(step.cache_creation_input_tokens, Some(200));
227 assert_eq!(step.output_tokens, Some(50));
228 assert_eq!(step.total_tokens(), 5350);
229 })
230 .await
231 .expect("test timed out");
232 }
233
234 #[tokio::test]
235 async fn agent_executor_propagates_account_id() {
236 timeout(Duration::from_secs(10), async {
237 let account_id = Uuid::now_v7();
238 let mut output = AgentOutput::new(json!("ok"));
239 output.cost_usd = Some(0.02);
240 output.account_id = Some(account_id.to_string());
241 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
242
243 let step = AgentExecutor::new(&budget_config())
244 .execute(&provider)
245 .await
246 .expect("agent step succeeds");
247 assert_eq!(step.account_id, Some(account_id));
248
249 let mut invalid = AgentOutput::new(json!("ok"));
250 invalid.cost_usd = Some(0.02);
251 invalid.account_id = Some("not-a-uuid".to_string());
252 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output: invalid });
253 let step = AgentExecutor::new(&budget_config())
254 .execute(&provider)
255 .await
256 .expect("agent step succeeds");
257 assert_eq!(step.account_id, None);
258 })
259 .await
260 .expect("test timed out");
261 }
262
263 #[tokio::test]
264 async fn agent_executor_without_cache_tokens_yields_none() {
265 timeout(Duration::from_secs(10), async {
266 let mut output = AgentOutput::new(json!("ok"));
267 output.input_tokens = Some(100);
268 output.output_tokens = Some(50);
269 output.cost_usd = Some(0.02);
270 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
271
272 let config = budget_config();
273 let step = AgentExecutor::new(&config)
274 .execute(&provider)
275 .await
276 .expect("agent step succeeds");
277
278 assert_eq!(step.input_tokens, Some(100));
279 assert!(step.cache_read_input_tokens.is_none());
280 assert!(step.cache_creation_input_tokens.is_none());
281 assert_eq!(step.total_tokens(), 150);
282 })
283 .await
284 .expect("test timed out");
285 }
286
287 #[test]
288 fn parse_permission_mode_via_serde() {
289 let json = r#""auto""#;
290 let mode: PermissionMode = serde_json::from_str(json).unwrap();
291 assert!(matches!(mode, PermissionMode::Auto));
292 }
293
294 #[test]
295 fn parse_permission_mode_dont_ask() {
296 let json = r#""dont_ask""#;
297 let mode: PermissionMode = serde_json::from_str(json).unwrap();
298 assert!(matches!(mode, PermissionMode::DontAsk));
299 }
300
301 #[test]
302 fn parse_permission_mode_bypass() {
303 let json = r#""bypass""#;
304 let mode: PermissionMode = serde_json::from_str(json).unwrap();
305 assert!(matches!(mode, PermissionMode::BypassPermissions));
306 }
307
308 #[test]
309 fn parse_permission_mode_case_insensitive() {
310 let json = r#""AUTO""#;
311 let mode: PermissionMode = serde_json::from_str(json).unwrap();
312 assert!(matches!(mode, PermissionMode::Auto));
313
314 let json = r#""DONT_ASK""#;
315 let mode: PermissionMode = serde_json::from_str(json).unwrap();
316 assert!(matches!(mode, PermissionMode::DontAsk));
317 }
318
319 #[test]
320 fn parse_permission_mode_unknown_defaults() {
321 let json = r#""unknown""#;
322 let mode: PermissionMode = serde_json::from_str(json).unwrap();
323 assert!(matches!(mode, PermissionMode::Default));
324 }
325}