1use std::any;
29use std::sync::Arc;
30
31use schemars::{JsonSchema, schema_for};
32use serde::de::DeserializeOwned;
33use serde::{Deserialize, Serialize};
34use serde_json::{Value, from_value, to_string};
35use tokio::time;
36use tracing::{info, warn};
37
38use crate::error::OperationError;
39#[cfg(feature = "prometheus")]
40use crate::metric_names;
41use crate::provider::{
42 AgentConfig, AgentOutput, AgentProvider, DebugMessage, LogSink, assert_environment_id_valid,
43};
44use crate::retry::RetryPolicy;
45use crate::trace_context::WorkflowTraceContext;
46
47pub struct Model;
78
79impl Model {
80 pub const SONNET: &str = "sonnet";
84 pub const OPUS: &str = "opus";
86 pub const HAIKU: &str = "haiku";
88
89 pub const HAIKU_45: &str = "claude-haiku-4-5-20251001";
93
94 pub const SONNET_46: &str = "claude-sonnet-4-6";
98 pub const OPUS_46: &str = "claude-opus-4-6";
100
101 pub const SONNET_46_1M: &str = "claude-sonnet-4-6[1m]";
105 pub const OPUS_46_1M: &str = "claude-opus-4-6[1m]";
107
108 pub const OPUS_47: &str = "claude-opus-4-7";
112 pub const OPUS_47_1M: &str = "claude-opus-4-7[1m]";
114
115 pub const OPUS_48: &str = "claude-opus-4-8";
119 pub const OPUS_48_1M: &str = "claude-opus-4-8[1m]";
121
122 pub const FABLE_51: &str = "claude-fable-5-1";
126 pub const FABLE_5: &str = "claude-fable-5";
128 pub const MYTHOS_51: &str = "claude-mythos-5-1";
130 pub const MYTHOS_5: &str = "claude-mythos-5";
132 pub const OPUS_55: &str = "claude-opus-5-5";
134 pub const OPUS_5: &str = "claude-opus-5";
136 pub const OPUS_5_1M: &str = "claude-opus-5[1m]";
138 pub const SONNET_55: &str = "claude-sonnet-5-5";
140 pub const SONNET_5: &str = "claude-sonnet-5";
142 pub const SONNET_5_1M: &str = "claude-sonnet-5[1m]";
144}
145
146#[derive(Debug, Default, Clone, Copy, Serialize)]
151pub enum PermissionMode {
152 #[default]
154 Default,
155 Auto,
157 DontAsk,
159 BypassPermissions,
163}
164
165impl<'de> Deserialize<'de> for PermissionMode {
166 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
167 where
168 D: serde::Deserializer<'de>,
169 {
170 let s = String::deserialize(deserializer)?;
171 Ok(match s.to_lowercase().replace('_', "").as_str() {
172 "auto" => Self::Auto,
173 "dontask" => Self::DontAsk,
174 "bypass" | "bypasspermissions" => Self::BypassPermissions,
175 _ => Self::Default,
176 })
177 }
178}
179
180#[must_use = "an Agent does nothing until .run() is awaited"]
210pub struct Agent {
211 config: AgentConfig,
212 dry_run: Option<bool>,
213 retry_policy: Option<RetryPolicy>,
214 log_sink: Option<Arc<dyn LogSink>>,
215}
216
217impl Agent {
218 pub fn new() -> Self {
223 Self {
224 config: AgentConfig::new(""),
225 dry_run: None,
226 retry_policy: None,
227 log_sink: None,
228 }
229 }
230
231 pub fn from_config(config: impl Into<AgentConfig>) -> Self {
250 Self {
251 config: config.into(),
252 dry_run: None,
253 retry_policy: None,
254 log_sink: None,
255 }
256 }
257
258 pub fn system_prompt(mut self, prompt: &str) -> Self {
260 self.config.system_prompt = Some(prompt.to_string());
261 self
262 }
263
264 pub fn prompt(mut self, prompt: &str) -> Self {
266 self.config.prompt = prompt.to_string();
267 self
268 }
269
270 pub fn model(mut self, model: impl Into<String>) -> Self {
277 self.config.model = model.into();
278 self
279 }
280
281 pub fn allowed_tools(mut self, tools: &[&str]) -> Self {
286 self.config.allowed_tools = tools.iter().map(|s| s.to_string()).collect();
287 self
288 }
289
290 pub fn max_turns(mut self, turns: u32) -> Self {
296 assert!(turns > 0, "max_turns must be greater than 0");
297 self.config.max_turns = Some(turns);
298 self
299 }
300
301 pub fn max_budget_usd(mut self, budget: f64) -> Self {
307 assert!(
308 budget.is_finite() && budget > 0.0,
309 "budget must be a positive finite number, got {budget}"
310 );
311 self.config.max_budget_usd = Some(budget);
312 self
313 }
314
315 pub fn working_dir(mut self, dir: &str) -> Self {
317 self.config.working_dir = Some(dir.to_string());
318 self
319 }
320
321 pub fn mcp_config(mut self, config: &str) -> Self {
323 self.config.mcp_config = Some(config.to_string());
324 self
325 }
326
327 pub fn permission_mode(mut self, mode: PermissionMode) -> Self {
331 self.config.permission_mode = mode;
332 self
333 }
334
335 pub fn output<T: JsonSchema>(mut self) -> Self {
366 let schema = schema_for!(T);
367 self.config.json_schema = match to_string(&schema) {
368 Ok(s) => Some(s),
369 Err(e) => {
370 warn!(error = %e, type_name = any::type_name::<T>(), "failed to serialize JSON schema, structured output disabled");
371 None
372 }
373 };
374 self
375 }
376
377 pub fn output_schema_raw(mut self, schema: &str) -> Self {
401 self.config.json_schema = Some(schema.to_string());
402 self
403 }
404
405 pub fn retry(mut self, max_retries: u32) -> Self {
433 self.retry_policy = Some(RetryPolicy::new(max_retries));
434 self
435 }
436
437 pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
464 self.retry_policy = Some(policy);
465 self
466 }
467
468 pub fn dry_run(mut self, enabled: bool) -> Self {
477 self.dry_run = Some(enabled);
478 self
479 }
480
481 pub fn log_sink(mut self, sink: Arc<dyn LogSink>) -> Self {
508 self.log_sink = Some(sink);
509 self
510 }
511
512 pub fn trace_context(mut self, ctx: WorkflowTraceContext) -> Self {
537 self.config.trace_context = Some(ctx);
538 self
539 }
540
541 pub fn verbose(mut self) -> Self {
571 self.config.verbose = true;
572 self
573 }
574
575 pub fn resume(mut self, session_id: &str) -> Self {
611 assert!(!session_id.is_empty(), "session_id must not be empty");
612 assert!(
613 session_id
614 .chars()
615 .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'),
616 "session_id must only contain alphanumeric characters, hyphens, or underscores, got: {session_id}"
617 );
618 self.config.resume_session_id = Some(session_id.to_string());
619 self
620 }
621
622 pub fn resume_environment(mut self, environment_id: &str) -> Self {
661 assert_environment_id_valid(environment_id);
662 self.config.resume_environment_id = Some(environment_id.to_string());
663 self
664 }
665
666 #[tracing::instrument(name = "agent", skip_all, fields(model = %self.config.model, prompt_len = self.config.prompt.len()))]
687 pub async fn run(self, provider: &dyn AgentProvider) -> Result<AgentResult, OperationError> {
688 assert!(
689 !self.config.prompt.trim().is_empty(),
690 "prompt must not be empty - call .prompt(\"...\") before .run()"
691 );
692
693 if crate::dry_run::effective_dry_run(self.dry_run) {
694 info!(
695 prompt_len = self.config.prompt.len(),
696 "[dry-run] agent call skipped"
697 );
698 let mut output =
699 AgentOutput::new(Value::String("[dry-run] agent call skipped".to_string()));
700 output.cost_usd = Some(0.0);
701 output.input_tokens = Some(0);
702 output.output_tokens = Some(0);
703 return Ok(AgentResult { output });
704 }
705
706 let result = self.invoke_once(provider).await;
707
708 let default_schema_retry = RetryPolicy::new(2);
709 let policy = match &self.retry_policy {
710 Some(p) => p,
711 None if self.config.json_schema.is_some() => &default_schema_retry,
712 None => return result,
713 };
714
715 if let Err(ref err) = result {
717 if !crate::retry::is_retryable(err) {
718 return result;
719 }
720 } else {
721 return result;
722 }
723
724 let mut last_result = result;
725
726 for attempt in 0..policy.max_retries {
727 let delay = policy.delay_for_attempt(attempt);
728 let retry_reason = if matches!(
729 &last_result,
730 Err(OperationError::Agent(
731 crate::error::AgentError::SchemaValidation { .. }
732 ))
733 ) {
734 "structured_output was null (CLI non-determinism)"
735 } else {
736 "transient failure"
737 };
738 warn!(
739 attempt = attempt + 1,
740 max_retries = policy.max_retries,
741 delay_ms = delay.as_millis() as u64,
742 reason = retry_reason,
743 "retrying agent invocation"
744 );
745 time::sleep(delay).await;
746
747 last_result = self.invoke_once(provider).await;
748
749 match &last_result {
750 Ok(_) => return last_result,
751 Err(err) if !crate::retry::is_retryable(err) => return last_result,
752 _ => {}
753 }
754 }
755
756 last_result
757 }
758
759 async fn invoke_once(
761 &self,
762 provider: &dyn AgentProvider,
763 ) -> Result<AgentResult, OperationError> {
764 #[cfg(feature = "prometheus")]
765 let model_label = self.config.model.to_string();
766
767 let invoke_result = match self.log_sink {
768 Some(ref sink) => provider.invoke_with_logs(&self.config, sink.clone()).await,
769 None => provider.invoke(&self.config).await,
770 };
771 let output = match invoke_result {
772 Ok(output) => output,
773 Err(e) => {
774 #[cfg(feature = "prometheus")]
775 {
776 metrics::counter!(metric_names::AGENT_TOTAL, "model" => model_label.clone(), "status" => metric_names::STATUS_ERROR).increment(1);
777 }
778 return Err(OperationError::Agent(e));
779 }
780 };
781
782 info!(
783 duration_ms = output.duration_ms,
784 cost_usd = output.cost_usd,
785 input_tokens = output.input_tokens,
786 cache_read_input_tokens = output.cache_read_input_tokens,
787 cache_creation_input_tokens = output.cache_creation_input_tokens,
788 output_tokens = output.output_tokens,
789 model = output.model,
790 "agent completed"
791 );
792
793 #[cfg(feature = "prometheus")]
794 {
795 metrics::counter!(metric_names::AGENT_TOTAL, "model" => model_label.clone(), "status" => metric_names::STATUS_SUCCESS).increment(1);
796 metrics::histogram!(metric_names::AGENT_DURATION_SECONDS, "model" => model_label.clone())
797 .record(output.duration_ms as f64 / 1000.0);
798 if let Some(cost) = output.cost_usd {
799 metrics::gauge!(metric_names::AGENT_COST_USD_TOTAL, "model" => model_label.clone())
800 .increment(cost);
801 }
802 if let Some(tokens) = output.input_tokens {
803 metrics::counter!(metric_names::AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(tokens);
804 }
805 if let Some(t) = output.cache_read_input_tokens {
806 metrics::counter!(metric_names::AGENT_TOKENS_CACHE_READ_TOTAL, "model" => model_label.clone()).increment(t);
807 }
808 if let Some(t) = output.cache_creation_input_tokens {
809 metrics::counter!(metric_names::AGENT_TOKENS_CACHE_WRITE_TOTAL, "model" => model_label.clone()).increment(t);
810 }
811 if let Some(tokens) = output.output_tokens {
812 metrics::counter!(metric_names::AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label)
813 .increment(tokens);
814 }
815 }
816
817 Ok(AgentResult { output })
818 }
819}
820
821impl Default for Agent {
822 fn default() -> Self {
823 Self::new()
824 }
825}
826
827#[derive(Debug)]
832pub struct AgentResult {
833 output: AgentOutput,
834}
835
836impl AgentResult {
837 pub fn text(&self) -> &str {
842 match self.output.value.as_str() {
843 Some(s) => s,
844 None => {
845 warn!(
846 value_type = self.output.value.to_string(),
847 "agent output is not a string, returning empty"
848 );
849 ""
850 }
851 }
852 }
853
854 pub fn value(&self) -> &Value {
856 &self.output.value
857 }
858
859 pub fn json<T: DeserializeOwned>(&self) -> Result<T, OperationError> {
869 from_value(self.output.value.clone()).map_err(OperationError::deserialize::<T>)
870 }
871
872 pub fn into_json<T: DeserializeOwned>(self) -> Result<T, OperationError> {
878 from_value(self.output.value).map_err(OperationError::deserialize::<T>)
879 }
880
881 #[cfg(test)]
886 pub(crate) fn from_output(output: AgentOutput) -> Self {
887 Self { output }
888 }
889
890 pub fn session_id(&self) -> Option<&str> {
892 self.output.session_id.as_deref()
893 }
894
895 pub fn cost_usd(&self) -> Option<f64> {
897 self.output.cost_usd
898 }
899
900 pub fn input_tokens(&self) -> Option<u64> {
906 self.output.input_tokens
907 }
908
909 pub fn cache_read_input_tokens(&self) -> Option<u64> {
924 self.output.cache_read_input_tokens
925 }
926
927 pub fn cache_creation_input_tokens(&self) -> Option<u64> {
942 self.output.cache_creation_input_tokens
943 }
944
945 pub fn output_tokens(&self) -> Option<u64> {
947 self.output.output_tokens
948 }
949
950 pub fn duration_ms(&self) -> u64 {
952 self.output.duration_ms
953 }
954
955 pub fn model(&self) -> Option<&str> {
957 self.output.model.as_deref()
958 }
959
960 pub fn debug_messages(&self) -> Option<&[DebugMessage]> {
966 self.output.debug_messages.as_deref()
967 }
968
969 pub fn account_id(&self) -> Option<&str> {
986 self.output.account_id.as_deref()
987 }
988
989 pub fn environment_id(&self) -> Option<&str> {
1007 self.output.environment_id.as_deref()
1008 }
1009}
1010
1011#[cfg(test)]
1012mod tests {
1013 use super::*;
1014 use crate::error::AgentError;
1015 use crate::provider::InvokeFuture;
1016 use serde_json::json;
1017
1018 struct TestProvider {
1019 output: AgentOutput,
1020 }
1021
1022 impl AgentProvider for TestProvider {
1023 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1024 Box::pin(async move {
1025 Ok(AgentOutput {
1026 value: self.output.value.clone(),
1027 session_id: self.output.session_id.clone(),
1028 cost_usd: self.output.cost_usd,
1029 input_tokens: self.output.input_tokens,
1030 cache_read_input_tokens: None,
1031 cache_creation_input_tokens: None,
1032 output_tokens: self.output.output_tokens,
1033 model: self.output.model.clone(),
1034 duration_ms: self.output.duration_ms,
1035 debug_messages: None,
1036 account_id: None,
1037 environment_id: self.output.environment_id.clone(),
1038 })
1039 })
1040 }
1041 }
1042
1043 struct ConfigCapture {
1044 output: AgentOutput,
1045 }
1046
1047 impl AgentProvider for ConfigCapture {
1048 fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
1049 let config_json = serde_json::to_value(config).unwrap();
1050 Box::pin(async move {
1051 Ok(AgentOutput {
1052 value: config_json,
1053 session_id: self.output.session_id.clone(),
1054 cost_usd: self.output.cost_usd,
1055 input_tokens: self.output.input_tokens,
1056 cache_read_input_tokens: None,
1057 cache_creation_input_tokens: None,
1058 output_tokens: self.output.output_tokens,
1059 model: self.output.model.clone(),
1060 duration_ms: self.output.duration_ms,
1061 debug_messages: None,
1062 account_id: None,
1063 environment_id: None,
1064 })
1065 })
1066 }
1067 }
1068
1069 fn default_output() -> AgentOutput {
1070 AgentOutput {
1071 value: json!("test output"),
1072 session_id: Some("sess-123".to_string()),
1073 cost_usd: Some(0.05),
1074 input_tokens: Some(100),
1075 cache_read_input_tokens: None,
1076 cache_creation_input_tokens: None,
1077 output_tokens: Some(50),
1078 model: Some("sonnet".to_string()),
1079 duration_ms: 1500,
1080 debug_messages: None,
1081 account_id: None,
1082 environment_id: None,
1083 }
1084 }
1085
1086 #[test]
1089 fn model_constants_have_expected_values() {
1090 assert_eq!(Model::SONNET, "sonnet");
1091 assert_eq!(Model::OPUS, "opus");
1092 assert_eq!(Model::HAIKU, "haiku");
1093 assert_eq!(Model::HAIKU_45, "claude-haiku-4-5-20251001");
1094 assert_eq!(Model::SONNET_46, "claude-sonnet-4-6");
1095 assert_eq!(Model::OPUS_46, "claude-opus-4-6");
1096 assert_eq!(Model::SONNET_46_1M, "claude-sonnet-4-6[1m]");
1097 assert_eq!(Model::OPUS_46_1M, "claude-opus-4-6[1m]");
1098 assert_eq!(Model::OPUS_47, "claude-opus-4-7");
1099 assert_eq!(Model::OPUS_47_1M, "claude-opus-4-7[1m]");
1100 assert_eq!(Model::OPUS_48, "claude-opus-4-8");
1101 assert_eq!(Model::OPUS_48_1M, "claude-opus-4-8[1m]");
1102 assert_eq!(Model::FABLE_5, "claude-fable-5");
1103 assert_eq!(Model::FABLE_51, "claude-fable-5-1");
1104 assert_eq!(Model::MYTHOS_5, "claude-mythos-5");
1105 assert_eq!(Model::MYTHOS_51, "claude-mythos-5-1");
1106 assert_eq!(Model::OPUS_55, "claude-opus-5-5");
1107 assert_eq!(Model::SONNET_55, "claude-sonnet-5-5");
1108 assert_eq!(Model::OPUS_5, "claude-opus-5");
1109 assert_eq!(Model::OPUS_5_1M, "claude-opus-5[1m]");
1110 assert_eq!(Model::SONNET_5, "claude-sonnet-5");
1111 assert_eq!(Model::SONNET_5_1M, "claude-sonnet-5[1m]");
1112 }
1113
1114 #[tokio::test]
1117 async fn agent_new_default_values() {
1118 let provider = ConfigCapture {
1119 output: default_output(),
1120 };
1121 let result = Agent::new().prompt("hi").run(&provider).await.unwrap();
1122
1123 let config = result.value();
1124 assert_eq!(config["system_prompt"], json!(null));
1125 assert_eq!(config["prompt"], json!("hi"));
1126 assert_eq!(config["model"], json!("sonnet"));
1127 assert_eq!(config["allowed_tools"], json!([]));
1128 assert_eq!(config["max_turns"], json!(null));
1129 assert_eq!(config["max_budget_usd"], json!(null));
1130 assert_eq!(config["working_dir"], json!(null));
1131 assert_eq!(config["mcp_config"], json!(null));
1132 assert_eq!(config["permission_mode"], json!("Default"));
1133 assert_eq!(config["json_schema"], json!(null));
1134 }
1135
1136 #[tokio::test]
1137 async fn agent_default_matches_new() {
1138 let provider = ConfigCapture {
1139 output: default_output(),
1140 };
1141 let result_new = Agent::new().prompt("x").run(&provider).await.unwrap();
1142 let result_default = Agent::default().prompt("x").run(&provider).await.unwrap();
1143
1144 assert_eq!(result_new.value(), result_default.value());
1145 }
1146
1147 #[tokio::test]
1150 async fn builder_methods_store_values_correctly() {
1151 let provider = ConfigCapture {
1152 output: default_output(),
1153 };
1154 let result = Agent::new()
1155 .system_prompt("you are a bot")
1156 .prompt("do something")
1157 .model(Model::OPUS)
1158 .allowed_tools(&["Read", "Write"])
1159 .max_turns(5)
1160 .max_budget_usd(1.5)
1161 .working_dir("/tmp")
1162 .mcp_config("{}")
1163 .permission_mode(PermissionMode::Auto)
1164 .run(&provider)
1165 .await
1166 .unwrap();
1167
1168 let config = result.value();
1169 assert_eq!(config["system_prompt"], json!("you are a bot"));
1170 assert_eq!(config["prompt"], json!("do something"));
1171 assert_eq!(config["model"], json!("opus"));
1172 assert_eq!(config["allowed_tools"], json!(["Read", "Write"]));
1173 assert_eq!(config["max_turns"], json!(5));
1174 assert_eq!(config["max_budget_usd"], json!(1.5));
1175 assert_eq!(config["working_dir"], json!("/tmp"));
1176 assert_eq!(config["mcp_config"], json!("{}"));
1177 assert_eq!(config["permission_mode"], json!("Auto"));
1178 }
1179
1180 #[test]
1183 #[should_panic(expected = "max_turns must be greater than 0")]
1184 fn max_turns_zero_panics() {
1185 let _ = Agent::new().max_turns(0);
1186 }
1187
1188 #[test]
1189 #[should_panic(expected = "budget must be a positive finite number")]
1190 fn max_budget_negative_panics() {
1191 let _ = Agent::new().max_budget_usd(-1.0);
1192 }
1193
1194 #[test]
1195 #[should_panic(expected = "budget must be a positive finite number")]
1196 fn max_budget_nan_panics() {
1197 let _ = Agent::new().max_budget_usd(f64::NAN);
1198 }
1199
1200 #[test]
1201 #[should_panic(expected = "budget must be a positive finite number")]
1202 fn max_budget_infinity_panics() {
1203 let _ = Agent::new().max_budget_usd(f64::INFINITY);
1204 }
1205
1206 #[tokio::test]
1209 async fn agent_result_text_with_string_value() {
1210 let provider = TestProvider {
1211 output: AgentOutput {
1212 value: json!("hello world"),
1213 ..default_output()
1214 },
1215 };
1216 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1217 assert_eq!(result.text(), "hello world");
1218 }
1219
1220 #[tokio::test]
1221 async fn agent_result_text_with_non_string_value() {
1222 let provider = TestProvider {
1223 output: AgentOutput {
1224 value: json!(42),
1225 ..default_output()
1226 },
1227 };
1228 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1229 assert_eq!(result.text(), "");
1230 }
1231
1232 #[tokio::test]
1233 async fn agent_result_text_with_null_value() {
1234 let provider = TestProvider {
1235 output: AgentOutput {
1236 value: json!(null),
1237 ..default_output()
1238 },
1239 };
1240 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1241 assert_eq!(result.text(), "");
1242 }
1243
1244 #[tokio::test]
1245 async fn agent_result_json_successful_deserialize() {
1246 #[derive(Deserialize, PartialEq, Debug)]
1247 struct MyOutput {
1248 name: String,
1249 count: u32,
1250 }
1251 let provider = TestProvider {
1252 output: AgentOutput {
1253 value: json!({"name": "test", "count": 7}),
1254 ..default_output()
1255 },
1256 };
1257 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1258 let parsed: MyOutput = result.json().unwrap();
1259 assert_eq!(parsed.name, "test");
1260 assert_eq!(parsed.count, 7);
1261 }
1262
1263 #[tokio::test]
1264 async fn agent_result_json_failed_deserialize() {
1265 #[derive(Debug, Deserialize)]
1266 #[allow(dead_code)]
1267 struct MyOutput {
1268 name: String,
1269 }
1270 let provider = TestProvider {
1271 output: AgentOutput {
1272 value: json!(42),
1273 ..default_output()
1274 },
1275 };
1276 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1277 let err = result.json::<MyOutput>().unwrap_err();
1278 assert!(matches!(err, OperationError::Deserialize { .. }));
1279 }
1280
1281 #[tokio::test]
1282 async fn agent_result_accessors() {
1283 let provider = TestProvider {
1284 output: AgentOutput {
1285 value: json!("v"),
1286 session_id: Some("s-1".to_string()),
1287 cost_usd: Some(0.123),
1288 input_tokens: Some(999),
1289 cache_read_input_tokens: None,
1290 cache_creation_input_tokens: None,
1291 output_tokens: Some(456),
1292 model: Some("opus".to_string()),
1293 duration_ms: 2000,
1294 debug_messages: None,
1295 account_id: None,
1296 environment_id: None,
1297 },
1298 };
1299 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1300 assert_eq!(result.session_id(), Some("s-1"));
1301 assert_eq!(result.cost_usd(), Some(0.123));
1302 assert_eq!(result.input_tokens(), Some(999));
1303 assert_eq!(result.output_tokens(), Some(456));
1304 assert_eq!(result.duration_ms(), 2000);
1305 assert_eq!(result.model(), Some("opus"));
1306 }
1307
1308 #[tokio::test]
1311 async fn resume_passes_session_id_in_config() {
1312 let provider = ConfigCapture {
1313 output: default_output(),
1314 };
1315 let result = Agent::new()
1316 .prompt("followup")
1317 .resume("sess-abc")
1318 .run(&provider)
1319 .await
1320 .unwrap();
1321
1322 let config = result.value();
1323 assert_eq!(config["resume_session_id"], json!("sess-abc"));
1324 }
1325
1326 #[tokio::test]
1327 async fn no_resume_has_null_session_id() {
1328 let provider = ConfigCapture {
1329 output: default_output(),
1330 };
1331 let result = Agent::new()
1332 .prompt("first call")
1333 .run(&provider)
1334 .await
1335 .unwrap();
1336
1337 let config = result.value();
1338 assert_eq!(config["resume_session_id"], json!(null));
1339 }
1340
1341 #[test]
1342 #[should_panic(expected = "session_id must not be empty")]
1343 fn resume_empty_session_id_panics() {
1344 let _ = Agent::new().resume("");
1345 }
1346
1347 #[test]
1348 #[should_panic(expected = "session_id must only contain")]
1349 fn resume_invalid_chars_panics() {
1350 let _ = Agent::new().resume("sess;rm -rf /");
1351 }
1352
1353 #[test]
1354 fn resume_valid_formats_accepted() {
1355 let _ = Agent::new().resume("sess-abc123");
1356 let _ = Agent::new().resume("a1b2c3d4_session");
1357 let _ = Agent::new().resume("abc-DEF-123_456");
1358 }
1359
1360 #[tokio::test]
1363 async fn resume_environment_passes_id_in_config() {
1364 let provider = ConfigCapture {
1365 output: default_output(),
1366 };
1367 let result = Agent::new()
1368 .prompt("followup")
1369 .resume("sess-abc")
1370 .resume_environment("ironflow-env-0192f0c1")
1371 .run(&provider)
1372 .await
1373 .unwrap();
1374
1375 let config = result.value();
1376 assert_eq!(
1377 config["resume_environment_id"],
1378 json!("ironflow-env-0192f0c1")
1379 );
1380 assert_eq!(config["resume_session_id"], json!("sess-abc"));
1381 }
1382
1383 #[tokio::test]
1384 async fn no_resume_environment_has_null_id() {
1385 let provider = ConfigCapture {
1386 output: default_output(),
1387 };
1388 let result = Agent::new()
1389 .prompt("first call")
1390 .run(&provider)
1391 .await
1392 .unwrap();
1393
1394 let config = result.value();
1395 assert_eq!(config["resume_environment_id"], json!(null));
1396 }
1397
1398 #[test]
1399 #[should_panic(expected = "environment_id must not be empty")]
1400 fn resume_environment_empty_panics() {
1401 let _ = Agent::new().resume_environment("");
1402 }
1403
1404 #[test]
1405 #[should_panic(expected = "environment_id must only contain")]
1406 fn resume_environment_invalid_chars_panics() {
1407 let _ = Agent::new().resume_environment("../other_ns/claim");
1408 }
1409
1410 #[tokio::test]
1411 async fn agent_result_exposes_environment_id() {
1412 let mut output = default_output();
1413 output.environment_id = Some("ironflow-env-42".to_string());
1414 let provider = TestProvider { output };
1415 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1416 assert_eq!(result.environment_id(), Some("ironflow-env-42"));
1417
1418 let provider = TestProvider {
1419 output: default_output(),
1420 };
1421 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1422 assert_eq!(result.environment_id(), None);
1423 }
1424
1425 #[tokio::test]
1426 #[should_panic(expected = "prompt must not be empty")]
1427 async fn run_without_prompt_panics() {
1428 let provider = TestProvider {
1429 output: default_output(),
1430 };
1431 let _ = Agent::new().run(&provider).await;
1432 }
1433
1434 #[tokio::test]
1435 #[should_panic(expected = "prompt must not be empty")]
1436 async fn run_with_whitespace_only_prompt_panics() {
1437 let provider = TestProvider {
1438 output: default_output(),
1439 };
1440 let _ = Agent::new().prompt(" ").run(&provider).await;
1441 }
1442
1443 #[tokio::test]
1446 async fn model_accepts_custom_string() {
1447 let provider = ConfigCapture {
1448 output: default_output(),
1449 };
1450 let result = Agent::new()
1451 .prompt("hi")
1452 .model("mistral-large-latest")
1453 .run(&provider)
1454 .await
1455 .unwrap();
1456 assert_eq!(result.value()["model"], json!("mistral-large-latest"));
1457 }
1458
1459 #[tokio::test]
1460 async fn verbose_sets_config_flag() {
1461 let provider = ConfigCapture {
1462 output: default_output(),
1463 };
1464 let result = Agent::new()
1465 .prompt("hi")
1466 .verbose()
1467 .run(&provider)
1468 .await
1469 .unwrap();
1470 assert_eq!(result.value()["verbose"], json!(true));
1471 }
1472
1473 #[tokio::test]
1474 async fn verbose_not_set_by_default() {
1475 let provider = ConfigCapture {
1476 output: default_output(),
1477 };
1478 let result = Agent::new().prompt("hi").run(&provider).await.unwrap();
1479 assert_eq!(result.value()["verbose"], json!(false));
1480 }
1481
1482 #[tokio::test]
1483 async fn debug_messages_none_without_verbose() {
1484 let provider = TestProvider {
1485 output: default_output(),
1486 };
1487 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1488 assert!(result.debug_messages().is_none());
1489 }
1490
1491 #[tokio::test]
1492 async fn model_accepts_owned_string() {
1493 let provider = ConfigCapture {
1494 output: default_output(),
1495 };
1496 let model_name = String::from("gpt-4o");
1497 let result = Agent::new()
1498 .prompt("hi")
1499 .model(model_name)
1500 .run(&provider)
1501 .await
1502 .unwrap();
1503 assert_eq!(result.value()["model"], json!("gpt-4o"));
1504 }
1505
1506 #[tokio::test]
1507 async fn into_json_success() {
1508 #[derive(Deserialize, PartialEq, Debug)]
1509 struct Out {
1510 name: String,
1511 }
1512 let provider = TestProvider {
1513 output: AgentOutput {
1514 value: json!({"name": "test"}),
1515 ..default_output()
1516 },
1517 };
1518 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1519 let parsed: Out = result.into_json().unwrap();
1520 assert_eq!(parsed.name, "test");
1521 }
1522
1523 #[tokio::test]
1524 async fn into_json_failure() {
1525 #[derive(Debug, Deserialize)]
1526 #[allow(dead_code)]
1527 struct Out {
1528 name: String,
1529 }
1530 let provider = TestProvider {
1531 output: AgentOutput {
1532 value: json!(42),
1533 ..default_output()
1534 },
1535 };
1536 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1537 let err = result.into_json::<Out>().unwrap_err();
1538 assert!(matches!(err, OperationError::Deserialize { .. }));
1539 }
1540
1541 #[test]
1542 fn from_output_creates_result() {
1543 let output = AgentOutput {
1544 value: json!("hello"),
1545 ..default_output()
1546 };
1547 let result = AgentResult::from_output(output);
1548 assert_eq!(result.text(), "hello");
1549 assert_eq!(result.cost_usd(), Some(0.05));
1550 }
1551
1552 #[test]
1553 #[should_panic(expected = "budget must be a positive finite number")]
1554 fn max_budget_zero_panics() {
1555 let _ = Agent::new().max_budget_usd(0.0);
1556 }
1557
1558 #[test]
1559 fn model_constant_equality() {
1560 assert_eq!(Model::SONNET, "sonnet");
1561 assert_ne!(Model::SONNET, Model::OPUS);
1562 }
1563
1564 #[test]
1565 fn permission_mode_serialize_deserialize_roundtrip() {
1566 for mode in [
1567 PermissionMode::Default,
1568 PermissionMode::Auto,
1569 PermissionMode::DontAsk,
1570 PermissionMode::BypassPermissions,
1571 ] {
1572 let json = to_string(&mode).unwrap();
1573 let back: PermissionMode = serde_json::from_str(&json).unwrap();
1574 assert_eq!(format!("{:?}", mode), format!("{:?}", back));
1575 }
1576 }
1577
1578 #[test]
1581 fn retry_builder_stores_policy() {
1582 let agent = Agent::new().retry(3);
1583 assert!(agent.retry_policy.is_some());
1584 assert_eq!(agent.retry_policy.unwrap().max_retries(), 3);
1585 }
1586
1587 #[test]
1588 fn retry_policy_builder_stores_custom_policy() {
1589 use crate::retry::RetryPolicy;
1590 let policy = RetryPolicy::new(5).backoff(Duration::from_secs(1));
1591 let agent = Agent::new().retry_policy(policy);
1592 let p = agent.retry_policy.unwrap();
1593 assert_eq!(p.max_retries(), 5);
1594 }
1595
1596 #[test]
1597 fn no_retry_by_default() {
1598 let agent = Agent::new();
1599 assert!(agent.retry_policy.is_none());
1600 }
1601
1602 use std::sync::Arc;
1605 use std::sync::atomic::{AtomicU32, Ordering};
1606 use std::time::Duration;
1607
1608 struct FailNTimesProvider {
1609 fail_count: AtomicU32,
1610 failures_before_success: u32,
1611 output: AgentOutput,
1612 }
1613
1614 impl AgentProvider for FailNTimesProvider {
1615 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1616 Box::pin(async move {
1617 let current = self.fail_count.fetch_add(1, Ordering::SeqCst);
1618 if current < self.failures_before_success {
1619 Err(AgentError::ProcessFailed {
1620 exit_code: 1,
1621 stderr: format!("transient failure #{}", current + 1),
1622 })
1623 } else {
1624 Ok(AgentOutput {
1625 value: self.output.value.clone(),
1626 session_id: self.output.session_id.clone(),
1627 cost_usd: self.output.cost_usd,
1628 input_tokens: self.output.input_tokens,
1629 cache_read_input_tokens: None,
1630 cache_creation_input_tokens: None,
1631 output_tokens: self.output.output_tokens,
1632 model: self.output.model.clone(),
1633 duration_ms: self.output.duration_ms,
1634 debug_messages: None,
1635 account_id: None,
1636 environment_id: None,
1637 })
1638 }
1639 })
1640 }
1641 }
1642
1643 #[tokio::test]
1644 async fn retry_succeeds_after_transient_failures() {
1645 let provider = FailNTimesProvider {
1646 fail_count: AtomicU32::new(0),
1647 failures_before_success: 2,
1648 output: default_output(),
1649 };
1650 let result = Agent::new()
1651 .prompt("test")
1652 .retry_policy(crate::retry::RetryPolicy::new(3).backoff(Duration::from_millis(1)))
1653 .run(&provider)
1654 .await;
1655
1656 assert!(result.is_ok());
1657 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 3); }
1659
1660 #[tokio::test]
1661 async fn retry_exhausted_returns_last_error() {
1662 let provider = FailNTimesProvider {
1663 fail_count: AtomicU32::new(0),
1664 failures_before_success: 10, output: default_output(),
1666 };
1667 let result = Agent::new()
1668 .prompt("test")
1669 .retry_policy(crate::retry::RetryPolicy::new(2).backoff(Duration::from_millis(1)))
1670 .run(&provider)
1671 .await;
1672
1673 assert!(result.is_err());
1674 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 3);
1676 }
1677
1678 #[tokio::test]
1679 async fn retry_does_not_retry_prompt_too_large() {
1680 let call_count = Arc::new(AtomicU32::new(0));
1681 let count = call_count.clone();
1682
1683 struct CountingNonRetryable {
1684 count: Arc<AtomicU32>,
1685 }
1686 impl AgentProvider for CountingNonRetryable {
1687 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1688 self.count.fetch_add(1, Ordering::SeqCst);
1689 Box::pin(async move {
1690 Err(AgentError::PromptTooLarge {
1691 chars: 1_000_000,
1692 estimated_tokens: 250_000,
1693 model_limit: 200_000,
1694 })
1695 })
1696 }
1697 }
1698
1699 let provider = CountingNonRetryable { count };
1700 let result = Agent::new()
1701 .prompt("test")
1702 .retry_policy(crate::retry::RetryPolicy::new(3).backoff(Duration::from_millis(1)))
1703 .run(&provider)
1704 .await;
1705
1706 assert!(result.is_err());
1707 assert_eq!(call_count.load(Ordering::SeqCst), 1);
1708 }
1709
1710 #[tokio::test]
1711 async fn retry_retries_schema_validation_errors() {
1712 let call_count = Arc::new(AtomicU32::new(0));
1713 let count = call_count.clone();
1714
1715 struct SchemaFailProvider {
1716 count: Arc<AtomicU32>,
1717 }
1718 impl AgentProvider for SchemaFailProvider {
1719 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1720 self.count.fetch_add(1, Ordering::SeqCst);
1721 Box::pin(async move {
1722 Err(AgentError::SchemaValidation {
1723 expected: "object".to_string(),
1724 got: "null".to_string(),
1725 debug_messages: Vec::new(),
1726 partial_usage: Box::default(),
1727 raw_response: None,
1728 })
1729 })
1730 }
1731 }
1732
1733 let provider = SchemaFailProvider { count };
1734 let result = Agent::new()
1735 .prompt("test")
1736 .retry_policy(crate::retry::RetryPolicy::new(2).backoff(Duration::from_millis(1)))
1737 .run(&provider)
1738 .await;
1739
1740 assert!(result.is_err());
1741 assert_eq!(call_count.load(Ordering::SeqCst), 3);
1743 }
1744
1745 #[tokio::test]
1746 async fn schema_validation_succeeds_on_retry() {
1747 let call_count = Arc::new(AtomicU32::new(0));
1748 let count = call_count.clone();
1749
1750 struct SchemaFailThenSucceed {
1751 count: Arc<AtomicU32>,
1752 output: AgentOutput,
1753 }
1754 impl AgentProvider for SchemaFailThenSucceed {
1755 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1756 let current = self.count.fetch_add(1, Ordering::SeqCst);
1757 let output = self.output.clone();
1758 Box::pin(async move {
1759 if current == 0 {
1760 Err(AgentError::SchemaValidation {
1761 expected: "structured_output field".to_string(),
1762 got: "null".to_string(),
1763 debug_messages: Vec::new(),
1764 partial_usage: Box::default(),
1765 raw_response: None,
1766 })
1767 } else {
1768 Ok(output)
1769 }
1770 })
1771 }
1772 }
1773
1774 let provider = SchemaFailThenSucceed {
1775 count,
1776 output: default_output(),
1777 };
1778 let result = Agent::new()
1779 .prompt("test")
1780 .retry_policy(crate::retry::RetryPolicy::new(1).backoff(Duration::from_millis(1)))
1781 .run(&provider)
1782 .await;
1783
1784 assert!(result.is_ok());
1785 assert_eq!(call_count.load(Ordering::SeqCst), 2);
1786 }
1787
1788 #[tokio::test]
1789 async fn auto_retry_applied_when_json_schema_set() {
1790 let call_count = Arc::new(AtomicU32::new(0));
1791 let count = call_count.clone();
1792
1793 struct AlwaysSchemaFail {
1794 count: Arc<AtomicU32>,
1795 }
1796 impl AgentProvider for AlwaysSchemaFail {
1797 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1798 self.count.fetch_add(1, Ordering::SeqCst);
1799 Box::pin(async move {
1800 Err(AgentError::SchemaValidation {
1801 expected: "object".to_string(),
1802 got: "null".to_string(),
1803 debug_messages: Vec::new(),
1804 partial_usage: Box::default(),
1805 raw_response: None,
1806 })
1807 })
1808 }
1809 }
1810
1811 let provider = AlwaysSchemaFail { count };
1812 let result = Agent::new()
1813 .prompt("test")
1814 .output_schema_raw(r#"{"type":"object"}"#)
1815 .run(&provider)
1816 .await;
1817
1818 assert!(result.is_err());
1819 assert_eq!(call_count.load(Ordering::SeqCst), 3);
1821 }
1822
1823 #[tokio::test]
1824 async fn no_retry_without_policy() {
1825 let provider = FailNTimesProvider {
1826 fail_count: AtomicU32::new(0),
1827 failures_before_success: 1,
1828 output: default_output(),
1829 };
1830 let result = Agent::new().prompt("test").run(&provider).await;
1831
1832 assert!(result.is_err());
1833 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 1);
1834 }
1835
1836 use crate::test_support::VecSink;
1839
1840 struct SinkCapture {
1841 output: AgentOutput,
1842 saw_logs: Arc<AtomicU32>,
1843 }
1844
1845 impl AgentProvider for SinkCapture {
1846 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1847 Box::pin(async {
1848 Ok(AgentOutput {
1849 value: self.output.value.clone(),
1850 session_id: self.output.session_id.clone(),
1851 cost_usd: self.output.cost_usd,
1852 input_tokens: self.output.input_tokens,
1853 cache_read_input_tokens: None,
1854 cache_creation_input_tokens: None,
1855 output_tokens: self.output.output_tokens,
1856 model: self.output.model.clone(),
1857 duration_ms: self.output.duration_ms,
1858 debug_messages: None,
1859 account_id: None,
1860 environment_id: None,
1861 })
1862 })
1863 }
1864
1865 fn invoke_with_logs<'a>(
1866 &'a self,
1867 config: &'a AgentConfig,
1868 log_sink: Arc<dyn LogSink>,
1869 ) -> InvokeFuture<'a> {
1870 self.saw_logs.fetch_add(1, Ordering::SeqCst);
1871 log_sink.log("stdout", "streaming line");
1872 self.invoke(config)
1873 }
1874 }
1875
1876 #[tokio::test]
1877 async fn log_sink_routes_to_invoke_with_logs() {
1878 let saw_logs = Arc::new(AtomicU32::new(0));
1879 let provider = SinkCapture {
1880 output: default_output(),
1881 saw_logs: saw_logs.clone(),
1882 };
1883 let sink: Arc<dyn LogSink> = VecSink::new();
1884
1885 let result = Agent::new()
1886 .prompt("test")
1887 .log_sink(sink)
1888 .run(&provider)
1889 .await;
1890
1891 assert!(result.is_ok());
1892 assert_eq!(saw_logs.load(Ordering::SeqCst), 1);
1893 }
1894
1895 #[tokio::test]
1896 async fn no_log_sink_routes_to_invoke() {
1897 let saw_logs = Arc::new(AtomicU32::new(0));
1898 let provider = SinkCapture {
1899 output: default_output(),
1900 saw_logs: saw_logs.clone(),
1901 };
1902
1903 let result = Agent::new().prompt("test").run(&provider).await;
1904
1905 assert!(result.is_ok());
1906 assert_eq!(saw_logs.load(Ordering::SeqCst), 0);
1907 }
1908
1909 #[tokio::test]
1910 async fn log_sink_receives_provider_lines() {
1911 let saw_logs = Arc::new(AtomicU32::new(0));
1912 let provider = SinkCapture {
1913 output: default_output(),
1914 saw_logs: saw_logs.clone(),
1915 };
1916 let sink = VecSink::new();
1917
1918 let _ = Agent::new()
1919 .prompt("test")
1920 .log_sink(sink.clone() as Arc<dyn LogSink>)
1921 .run(&provider)
1922 .await;
1923
1924 let lines = sink.0.lock().unwrap();
1925 assert_eq!(lines.len(), 1);
1926 assert_eq!(lines[0].0, "stdout");
1927 assert_eq!(lines[0].1, "streaming line");
1928 }
1929}