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::error::OperationError;
11use ironflow_core::operations::agent::{Agent, AgentResult};
12use ironflow_core::pricing::{CostBreakdown, StaticPricing, spawn_log};
13use ironflow_core::provider::{AgentConfig, AgentProvider, LogSink};
14use ironflow_core::providers::claude::is_session_not_found;
15use ironflow_store::entities::StepKind;
16
17use crate::error::EngineError;
18use crate::log_sender::StepLogSender;
19use crate::notify::LogStream;
20
21use super::{StepArtifacts, StepExecutor, StepOutput};
22
23pub struct AgentExecutor<'a> {
29 config: &'a AgentConfig,
30 log_sender: Option<StepLogSender>,
31}
32
33impl<'a> AgentExecutor<'a> {
34 pub fn new(config: &'a AgentConfig) -> Self {
36 Self {
37 config,
38 log_sender: None,
39 }
40 }
41
42 pub fn with_log_sender(mut self, sender: StepLogSender) -> Self {
44 self.log_sender = Some(sender);
45 self
46 }
47
48 fn emit_system(&self, line: &str) {
49 if let Some(ref sender) = self.log_sender {
50 sender.emit(LogStream::System, line);
51 }
52 }
53
54 async fn run_agent(
55 &self,
56 config: AgentConfig,
57 provider: &Arc<dyn AgentProvider>,
58 ) -> Result<AgentResult, OperationError> {
59 let mut agent = Agent::from_config(config);
60 if let Some(ref sender) = self.log_sender {
61 agent = agent.log_sink(Arc::new(sender.clone()) as Arc<dyn LogSink>);
62 }
63 agent.run(provider.as_ref()).await
64 }
65
66 async fn resume(
71 &self,
72 provider: &Arc<dyn AgentProvider>,
73 session_id: &str,
74 resume_prompt: &str,
75 ) -> Result<AgentResult, OperationError> {
76 info!(session_id, "agent step resumed from session");
77 self.emit_system(&format!("agent step resumed from session {session_id}"));
78
79 let mut resumed = self.config.clone();
80 resumed.prompt = resume_prompt.to_string();
81 match self.run_agent(resumed, provider).await {
82 Err(OperationError::Agent(ref err)) if is_session_not_found(err) => {
83 warn!(
84 session_id,
85 error = %err,
86 "session not found, restarting the agent from scratch"
87 );
88 self.emit_system(&format!(
89 "session {session_id} not found, restarting the agent from scratch"
90 ));
91 let mut fresh = self.config.clone();
92 fresh.resume_session_id = None;
93 fresh.session_id = Some(session_id.to_string());
94 self.run_agent(fresh, provider).await
95 }
96 other => other,
97 }
98 }
99}
100
101impl StepExecutor for AgentExecutor<'_> {
102 fn kind(&self) -> StepKind {
103 StepKind::Agent
104 }
105
106 async fn execute(&self, provider: &Arc<dyn AgentProvider>) -> Result<StepOutput, EngineError> {
107 let start = Instant::now();
108
109 if let Some(ref sender) = self.log_sender {
110 sender.emit(
111 LogStream::System,
112 &format!("agent step started (model={})", self.config.model),
113 );
114 }
115
116 if self.config.json_schema.is_some() && self.config.max_turns == Some(1) {
117 warn!(
118 "structured output (json_schema) requires max_turns >= 2; \
119 max_turns is set to 1, the agent will likely fail with error_max_turns"
120 );
121 }
122
123 let result = match (&self.config.resume_session_id, &self.config.resume_prompt) {
124 (Some(session_id), Some(resume_prompt)) => {
125 self.resume(provider, session_id, resume_prompt).await?
126 }
127 _ => self.run_agent(self.config.clone(), provider).await?,
128 };
129
130 let duration_ms = start.elapsed().as_millis() as u64;
131 let cost = Decimal::try_from(result.cost_usd().unwrap_or(0.0)).unwrap_or(Decimal::ZERO);
132 let input_tokens = result.input_tokens();
133 let cache_read_tokens = result.cache_read_input_tokens();
134 let cache_creation_tokens = result.cache_creation_input_tokens();
135 let output_tokens = result.output_tokens();
136
137 info!(
138 step_kind = "agent",
139 model = %self.config.model,
140 cost_usd = %cost,
141 input_tokens = ?input_tokens,
142 cache_read_input_tokens = ?cache_read_tokens,
143 cache_creation_input_tokens = ?cache_creation_tokens,
144 output_tokens = ?output_tokens,
145 duration_ms,
146 "agent step completed"
147 );
148
149 let pricing = StaticPricing::new();
150 let breakdown = CostBreakdown::compute_with_cache(
151 &pricing,
152 &self.config.model,
153 input_tokens.unwrap_or(0),
154 cache_read_tokens.unwrap_or(0),
155 cache_creation_tokens.unwrap_or(0),
156 output_tokens.unwrap_or(0),
157 );
158 spawn_log("agent", &self.config.model, breakdown);
159
160 #[cfg(feature = "prometheus")]
161 {
162 use ironflow_core::metric_names::{
163 AGENT_COST_USD_TOTAL, AGENT_DURATION_SECONDS, AGENT_TOKENS_CACHE_READ_TOTAL,
164 AGENT_TOKENS_CACHE_WRITE_TOTAL, AGENT_TOKENS_INPUT_TOTAL,
165 AGENT_TOKENS_OUTPUT_TOTAL, AGENT_TOTAL, STATUS_SUCCESS,
166 };
167 use metrics::{counter, gauge, histogram};
168 let model_label = self.config.model.clone();
169 counter!(AGENT_TOTAL, "model" => model_label.clone(), "status" => STATUS_SUCCESS)
170 .increment(1);
171 histogram!(AGENT_DURATION_SECONDS, "model" => model_label.clone())
172 .record(duration_ms as f64 / 1000.0);
173 gauge!(AGENT_COST_USD_TOTAL, "model" => model_label.clone())
174 .increment(cost.to_string().parse::<f64>().unwrap_or(0.0));
175 if let Some(inp) = input_tokens {
176 counter!(AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(inp);
177 }
178 if let Some(t) = cache_read_tokens {
179 counter!(AGENT_TOKENS_CACHE_READ_TOTAL, "model" => model_label.clone())
180 .increment(t);
181 }
182 if let Some(t) = cache_creation_tokens {
183 counter!(AGENT_TOKENS_CACHE_WRITE_TOTAL, "model" => model_label.clone())
184 .increment(t);
185 }
186 if let Some(out) = output_tokens {
187 counter!(AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label).increment(out);
188 }
189 }
190
191 if let Some(ref sender) = self.log_sender {
192 sender.emit(
193 LogStream::System,
194 &format!(
195 "agent step completed (cost=${cost}, tokens_in={}, cache_read={}, cache_write={}, tokens_out={})",
196 input_tokens.unwrap_or(0),
197 cache_read_tokens.unwrap_or(0),
198 cache_creation_tokens.unwrap_or(0),
199 output_tokens.unwrap_or(0),
200 ),
201 );
202 }
203
204 let debug_messages = result.debug_messages().map(|msgs| msgs.to_vec());
205 let account_id = match result.account_id() {
206 Some(raw) => match Uuid::parse_str(raw) {
207 Ok(id) => Some(id),
208 Err(e) => {
209 warn!(account_id = raw, error = %e, "agent output carries an invalid account id");
210 None
211 }
212 },
213 None => None,
214 };
215
216 Ok(StepOutput {
217 output: result.value().clone(),
218 duration_ms,
219 cost_usd: cost,
220 input_tokens,
221 cache_read_input_tokens: cache_read_tokens,
222 cache_creation_input_tokens: cache_creation_tokens,
223 output_tokens,
224 model: result.model().map(String::from),
225 debug_messages,
226 artifacts: StepArtifacts::default(),
227 account_id,
228 environment_id: result.environment_id().map(String::from),
229 })
230 }
231}
232
233#[cfg(test)]
234mod tests {
235 use std::sync::{Arc, Mutex};
236 use std::time::Duration;
237
238 use ironflow_core::error::AgentError;
239 use ironflow_core::operations::agent::PermissionMode;
240 use ironflow_core::provider::{AgentConfig, AgentOutput, AgentProvider, InvokeFuture};
241 use serde_json::json;
242 use tokio::time::timeout;
243 use uuid::Uuid;
244
245 use super::{AgentExecutor, StepExecutor};
246
247 struct FixedUsageProvider {
249 output: AgentOutput,
250 }
251
252 impl AgentProvider for FixedUsageProvider {
253 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
254 Box::pin(async move { Ok(self.output.clone()) })
255 }
256 }
257
258 struct SessionProvider {
261 seen: Mutex<Vec<AgentConfig>>,
262 resume_error: Option<String>,
263 }
264
265 impl SessionProvider {
266 fn new(resume_error: Option<&str>) -> Arc<Self> {
267 Arc::new(Self {
268 seen: Mutex::new(Vec::new()),
269 resume_error: resume_error.map(String::from),
270 })
271 }
272
273 fn seen(&self) -> Vec<AgentConfig> {
274 self.seen.lock().unwrap().clone()
275 }
276 }
277
278 impl AgentProvider for SessionProvider {
279 fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
280 Box::pin(async move {
281 self.seen.lock().unwrap().push(config.clone());
282 match (&config.resume_session_id, &self.resume_error) {
283 (Some(_), Some(stderr)) => Err(AgentError::ProcessFailed {
284 exit_code: 1,
285 stderr: stderr.clone(),
286 }),
287 _ => {
288 let mut output = AgentOutput::new(json!("ok"));
289 output.cost_usd = Some(0.02);
290 Ok(output)
291 }
292 }
293 })
294 }
295 }
296
297 const SID: &str = "0192f0c1-7d2e-7a4b-9c3d-1e2f3a4b5c6d";
298
299 #[tokio::test]
300 async fn agent_resume_executor_sends_resume_prompt() {
301 timeout(Duration::from_secs(10), async {
302 let recorder = SessionProvider::new(None);
303 let provider: Arc<dyn AgentProvider> = recorder.clone();
304 let config = budget_config().resume(SID).resume_prompt("go on");
305
306 AgentExecutor::new(&config)
307 .execute(&provider)
308 .await
309 .expect("resumed step succeeds");
310
311 let seen = recorder.seen();
312 assert_eq!(seen.len(), 1);
313 assert_eq!(seen[0].prompt, "go on");
314 assert_eq!(seen[0].resume_session_id.as_deref(), Some(SID));
315 })
316 .await
317 .expect("test timed out");
318 }
319
320 #[tokio::test]
321 async fn agent_resume_executor_falls_back_when_session_is_missing() {
322 timeout(Duration::from_secs(10), async {
323 let recorder =
324 SessionProvider::new(Some("No conversation found with session ID: 0192f0c1"));
325 let provider: Arc<dyn AgentProvider> = recorder.clone();
326 let config = budget_config().resume(SID).resume_prompt("go on");
327
328 AgentExecutor::new(&config)
329 .execute(&provider)
330 .await
331 .expect("a missing session never fails the step");
332
333 let seen = recorder.seen();
334 assert_eq!(seen.len(), 2);
335 assert_eq!(seen[1].prompt, "hi");
336 assert_eq!(seen[1].resume_session_id, None);
337 assert_eq!(seen[1].session_id.as_deref(), Some(SID));
338 })
339 .await
340 .expect("test timed out");
341 }
342
343 #[tokio::test]
344 async fn agent_resume_executor_propagates_other_errors() {
345 timeout(Duration::from_secs(10), async {
346 let recorder = SessionProvider::new(Some("permission denied"));
347 let provider: Arc<dyn AgentProvider> = recorder.clone();
348 let config = budget_config().resume(SID).resume_prompt("go on");
349
350 let err = AgentExecutor::new(&config)
351 .execute(&provider)
352 .await
353 .expect_err("other errors are not masked");
354
355 assert!(err.to_string().contains("permission denied"), "{err}");
356 assert_eq!(recorder.seen().len(), 1);
357 })
358 .await
359 .expect("test timed out");
360 }
361
362 #[tokio::test]
363 async fn agent_resume_executor_leaves_user_resume_alone() {
364 timeout(Duration::from_secs(10), async {
365 let recorder = SessionProvider::new(None);
366 let provider: Arc<dyn AgentProvider> = recorder.clone();
367 let config = budget_config().resume(SID);
369
370 AgentExecutor::new(&config)
371 .execute(&provider)
372 .await
373 .expect("step succeeds");
374
375 let seen = recorder.seen();
376 assert_eq!(seen.len(), 1);
377 assert_eq!(seen[0].prompt, "hi");
378 assert_eq!(seen[0].resume_session_id.as_deref(), Some(SID));
379 })
380 .await
381 .expect("test timed out");
382 }
383
384 fn budget_config() -> AgentConfig {
385 let mut config = AgentConfig::new("hi");
386 config.max_budget_usd = Some(0.10);
387 config
388 }
389
390 #[tokio::test]
391 async fn agent_executor_propagates_cache_tokens() {
392 timeout(Duration::from_secs(10), async {
393 let mut output = AgentOutput::new(json!("ok"));
394 output.input_tokens = Some(100);
395 output.cache_read_input_tokens = Some(5000);
396 output.cache_creation_input_tokens = Some(200);
397 output.output_tokens = Some(50);
398 output.cost_usd = Some(0.02);
399 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
400
401 let config = budget_config();
402 let step = AgentExecutor::new(&config)
403 .execute(&provider)
404 .await
405 .expect("agent step succeeds");
406
407 assert_eq!(step.input_tokens, Some(100));
408 assert_eq!(step.cache_read_input_tokens, Some(5000));
409 assert_eq!(step.cache_creation_input_tokens, Some(200));
410 assert_eq!(step.output_tokens, Some(50));
411 assert_eq!(step.total_tokens(), 5350);
412 })
413 .await
414 .expect("test timed out");
415 }
416
417 #[tokio::test]
418 async fn agent_executor_propagates_account_id() {
419 timeout(Duration::from_secs(10), async {
420 let account_id = Uuid::now_v7();
421 let mut output = AgentOutput::new(json!("ok"));
422 output.cost_usd = Some(0.02);
423 output.account_id = Some(account_id.to_string());
424 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
425
426 let step = AgentExecutor::new(&budget_config())
427 .execute(&provider)
428 .await
429 .expect("agent step succeeds");
430 assert_eq!(step.account_id, Some(account_id));
431
432 let mut invalid = AgentOutput::new(json!("ok"));
433 invalid.cost_usd = Some(0.02);
434 invalid.account_id = Some("not-a-uuid".to_string());
435 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output: invalid });
436 let step = AgentExecutor::new(&budget_config())
437 .execute(&provider)
438 .await
439 .expect("agent step succeeds");
440 assert_eq!(step.account_id, None);
441 })
442 .await
443 .expect("test timed out");
444 }
445
446 #[tokio::test]
447 async fn agent_executor_propagates_environment_id() {
448 timeout(Duration::from_secs(10), async {
449 let mut output = AgentOutput::new(json!("ok"));
450 output.cost_usd = Some(0.02);
451 output.environment_id = Some("ironflow-env-0a1b2c".to_string());
452 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
453 let step = AgentExecutor::new(&budget_config())
454 .execute(&provider)
455 .await
456 .expect("agent step succeeds");
457 assert_eq!(step.environment_id.as_deref(), Some("ironflow-env-0a1b2c"));
458
459 let mut output = AgentOutput::new(json!("ok"));
460 output.cost_usd = Some(0.02);
461 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
462 let step = AgentExecutor::new(&budget_config())
463 .execute(&provider)
464 .await
465 .expect("agent step succeeds");
466 assert_eq!(step.environment_id, None);
467 })
468 .await
469 .expect("test timed out");
470 }
471
472 #[tokio::test]
473 async fn agent_executor_without_cache_tokens_yields_none() {
474 timeout(Duration::from_secs(10), async {
475 let mut output = AgentOutput::new(json!("ok"));
476 output.input_tokens = Some(100);
477 output.output_tokens = Some(50);
478 output.cost_usd = Some(0.02);
479 let provider: Arc<dyn AgentProvider> = Arc::new(FixedUsageProvider { output });
480
481 let config = budget_config();
482 let step = AgentExecutor::new(&config)
483 .execute(&provider)
484 .await
485 .expect("agent step succeeds");
486
487 assert_eq!(step.input_tokens, Some(100));
488 assert!(step.cache_read_input_tokens.is_none());
489 assert!(step.cache_creation_input_tokens.is_none());
490 assert_eq!(step.total_tokens(), 150);
491 })
492 .await
493 .expect("test timed out");
494 }
495
496 #[test]
497 fn parse_permission_mode_via_serde() {
498 let json = r#""auto""#;
499 let mode: PermissionMode = serde_json::from_str(json).unwrap();
500 assert!(matches!(mode, PermissionMode::Auto));
501 }
502
503 #[test]
504 fn parse_permission_mode_dont_ask() {
505 let json = r#""dont_ask""#;
506 let mode: PermissionMode = serde_json::from_str(json).unwrap();
507 assert!(matches!(mode, PermissionMode::DontAsk));
508 }
509
510 #[test]
511 fn parse_permission_mode_bypass() {
512 let json = r#""bypass""#;
513 let mode: PermissionMode = serde_json::from_str(json).unwrap();
514 assert!(matches!(mode, PermissionMode::BypassPermissions));
515 }
516
517 #[test]
518 fn parse_permission_mode_case_insensitive() {
519 let json = r#""AUTO""#;
520 let mode: PermissionMode = serde_json::from_str(json).unwrap();
521 assert!(matches!(mode, PermissionMode::Auto));
522
523 let json = r#""DONT_ASK""#;
524 let mode: PermissionMode = serde_json::from_str(json).unwrap();
525 assert!(matches!(mode, PermissionMode::DontAsk));
526 }
527
528 #[test]
529 fn parse_permission_mode_unknown_defaults() {
530 let json = r#""unknown""#;
531 let mode: PermissionMode = serde_json::from_str(json).unwrap();
532 assert!(matches!(mode, PermissionMode::Default));
533 }
534}