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 environment_id: result.environment_id().map(String::from),
174 })
175 }
176}
177
178#[cfg(test)]
179mod tests {
180 use std::sync::Arc;
181 use std::time::Duration;
182
183 use ironflow_core::operations::agent::PermissionMode;
184 use ironflow_core::provider::{AgentConfig, AgentOutput, AgentProvider, InvokeFuture};
185 use serde_json::json;
186 use tokio::time::timeout;
187 use uuid::Uuid;
188
189 use super::{AgentExecutor, StepExecutor};
190
191 struct FixedUsageProvider {
193 output: AgentOutput,
194 }
195
196 impl AgentProvider for FixedUsageProvider {
197 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
198 Box::pin(async move { Ok(self.output.clone()) })
199 }
200 }
201
202 fn budget_config() -> AgentConfig {
203 let mut config = AgentConfig::new("hi");
204 config.max_budget_usd = Some(0.10);
205 config
206 }
207
208 #[tokio::test]
209 async fn agent_executor_propagates_cache_tokens() {
210 timeout(Duration::from_secs(10), async {
211 let mut output = AgentOutput::new(json!("ok"));
212 output.input_tokens = Some(100);
213 output.cache_read_input_tokens = Some(5000);
214 output.cache_creation_input_tokens = Some(200);
215 output.output_tokens = Some(50);
216 output.cost_usd = Some(0.02);
217 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
218
219 let config = budget_config();
220 let step = AgentExecutor::new(&config)
221 .execute(&provider)
222 .await
223 .expect("agent step succeeds");
224
225 assert_eq!(step.input_tokens, Some(100));
226 assert_eq!(step.cache_read_input_tokens, Some(5000));
227 assert_eq!(step.cache_creation_input_tokens, Some(200));
228 assert_eq!(step.output_tokens, Some(50));
229 assert_eq!(step.total_tokens(), 5350);
230 })
231 .await
232 .expect("test timed out");
233 }
234
235 #[tokio::test]
236 async fn agent_executor_propagates_account_id() {
237 timeout(Duration::from_secs(10), async {
238 let account_id = Uuid::now_v7();
239 let mut output = AgentOutput::new(json!("ok"));
240 output.cost_usd = Some(0.02);
241 output.account_id = Some(account_id.to_string());
242 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
243
244 let step = AgentExecutor::new(&budget_config())
245 .execute(&provider)
246 .await
247 .expect("agent step succeeds");
248 assert_eq!(step.account_id, Some(account_id));
249
250 let mut invalid = AgentOutput::new(json!("ok"));
251 invalid.cost_usd = Some(0.02);
252 invalid.account_id = Some("not-a-uuid".to_string());
253 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output: invalid });
254 let step = AgentExecutor::new(&budget_config())
255 .execute(&provider)
256 .await
257 .expect("agent step succeeds");
258 assert_eq!(step.account_id, None);
259 })
260 .await
261 .expect("test timed out");
262 }
263
264 #[tokio::test]
265 async fn agent_executor_propagates_environment_id() {
266 timeout(Duration::from_secs(10), async {
267 let mut output = AgentOutput::new(json!("ok"));
268 output.cost_usd = Some(0.02);
269 output.environment_id = Some("ironflow-env-0a1b2c".to_string());
270 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
271 let step = AgentExecutor::new(&budget_config())
272 .execute(&provider)
273 .await
274 .expect("agent step succeeds");
275 assert_eq!(step.environment_id.as_deref(), Some("ironflow-env-0a1b2c"));
276
277 let mut output = AgentOutput::new(json!("ok"));
278 output.cost_usd = Some(0.02);
279 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
280 let step = AgentExecutor::new(&budget_config())
281 .execute(&provider)
282 .await
283 .expect("agent step succeeds");
284 assert_eq!(step.environment_id, None);
285 })
286 .await
287 .expect("test timed out");
288 }
289
290 #[tokio::test]
291 async fn agent_executor_without_cache_tokens_yields_none() {
292 timeout(Duration::from_secs(10), async {
293 let mut output = AgentOutput::new(json!("ok"));
294 output.input_tokens = Some(100);
295 output.output_tokens = Some(50);
296 output.cost_usd = Some(0.02);
297 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
298
299 let config = budget_config();
300 let step = AgentExecutor::new(&config)
301 .execute(&provider)
302 .await
303 .expect("agent step succeeds");
304
305 assert_eq!(step.input_tokens, Some(100));
306 assert!(step.cache_read_input_tokens.is_none());
307 assert!(step.cache_creation_input_tokens.is_none());
308 assert_eq!(step.total_tokens(), 150);
309 })
310 .await
311 .expect("test timed out");
312 }
313
314 #[test]
315 fn parse_permission_mode_via_serde() {
316 let json = r#""auto""#;
317 let mode: PermissionMode = serde_json::from_str(json).unwrap();
318 assert!(matches!(mode, PermissionMode::Auto));
319 }
320
321 #[test]
322 fn parse_permission_mode_dont_ask() {
323 let json = r#""dont_ask""#;
324 let mode: PermissionMode = serde_json::from_str(json).unwrap();
325 assert!(matches!(mode, PermissionMode::DontAsk));
326 }
327
328 #[test]
329 fn parse_permission_mode_bypass() {
330 let json = r#""bypass""#;
331 let mode: PermissionMode = serde_json::from_str(json).unwrap();
332 assert!(matches!(mode, PermissionMode::BypassPermissions));
333 }
334
335 #[test]
336 fn parse_permission_mode_case_insensitive() {
337 let json = r#""AUTO""#;
338 let mode: PermissionMode = serde_json::from_str(json).unwrap();
339 assert!(matches!(mode, PermissionMode::Auto));
340
341 let json = r#""DONT_ASK""#;
342 let mode: PermissionMode = serde_json::from_str(json).unwrap();
343 assert!(matches!(mode, PermissionMode::DontAsk));
344 }
345
346 #[test]
347 fn parse_permission_mode_unknown_defaults() {
348 let json = r#""unknown""#;
349 let mode: PermissionMode = serde_json::from_str(json).unwrap();
350 assert!(matches!(mode, PermissionMode::Default));
351 }
352}