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::{AgentConfig, AgentOutput, AgentProvider, DebugMessage, LogSink};
42use crate::retry::RetryPolicy;
43use crate::trace_context::WorkflowTraceContext;
44
45pub struct Model;
76
77impl Model {
78 pub const SONNET: &str = "sonnet";
82 pub const OPUS: &str = "opus";
84 pub const HAIKU: &str = "haiku";
86
87 pub const HAIKU_45: &str = "claude-haiku-4-5-20251001";
91
92 pub const SONNET_46: &str = "claude-sonnet-4-6";
96 pub const OPUS_46: &str = "claude-opus-4-6";
98
99 pub const SONNET_46_1M: &str = "claude-sonnet-4-6[1m]";
103 pub const OPUS_46_1M: &str = "claude-opus-4-6[1m]";
105
106 pub const OPUS_47: &str = "claude-opus-4-7";
110 pub const OPUS_47_1M: &str = "claude-opus-4-7[1m]";
112
113 pub const OPUS_48: &str = "claude-opus-4-8";
117 pub const OPUS_48_1M: &str = "claude-opus-4-8[1m]";
119
120 pub const FABLE_51: &str = "claude-fable-5-1";
124 pub const FABLE_5: &str = "claude-fable-5";
126 pub const MYTHOS_51: &str = "claude-mythos-5-1";
128 pub const MYTHOS_5: &str = "claude-mythos-5";
130 pub const OPUS_55: &str = "claude-opus-5-5";
132 pub const OPUS_5: &str = "claude-opus-5";
134 pub const OPUS_5_1M: &str = "claude-opus-5[1m]";
136 pub const SONNET_55: &str = "claude-sonnet-5-5";
138 pub const SONNET_5: &str = "claude-sonnet-5";
140 pub const SONNET_5_1M: &str = "claude-sonnet-5[1m]";
142}
143
144#[derive(Debug, Default, Clone, Copy, Serialize)]
149pub enum PermissionMode {
150 #[default]
152 Default,
153 Auto,
155 DontAsk,
157 BypassPermissions,
161}
162
163impl<'de> Deserialize<'de> for PermissionMode {
164 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
165 where
166 D: serde::Deserializer<'de>,
167 {
168 let s = String::deserialize(deserializer)?;
169 Ok(match s.to_lowercase().replace('_', "").as_str() {
170 "auto" => Self::Auto,
171 "dontask" => Self::DontAsk,
172 "bypass" | "bypasspermissions" => Self::BypassPermissions,
173 _ => Self::Default,
174 })
175 }
176}
177
178#[must_use = "an Agent does nothing until .run() is awaited"]
208pub struct Agent {
209 config: AgentConfig,
210 dry_run: Option<bool>,
211 retry_policy: Option<RetryPolicy>,
212 log_sink: Option<Arc<dyn LogSink>>,
213}
214
215impl Agent {
216 pub fn new() -> Self {
221 Self {
222 config: AgentConfig::new(""),
223 dry_run: None,
224 retry_policy: None,
225 log_sink: None,
226 }
227 }
228
229 pub fn from_config(config: impl Into<AgentConfig>) -> Self {
248 Self {
249 config: config.into(),
250 dry_run: None,
251 retry_policy: None,
252 log_sink: None,
253 }
254 }
255
256 pub fn system_prompt(mut self, prompt: &str) -> Self {
258 self.config.system_prompt = Some(prompt.to_string());
259 self
260 }
261
262 pub fn prompt(mut self, prompt: &str) -> Self {
264 self.config.prompt = prompt.to_string();
265 self
266 }
267
268 pub fn model(mut self, model: impl Into<String>) -> Self {
275 self.config.model = model.into();
276 self
277 }
278
279 pub fn allowed_tools(mut self, tools: &[&str]) -> Self {
284 self.config.allowed_tools = tools.iter().map(|s| s.to_string()).collect();
285 self
286 }
287
288 pub fn max_turns(mut self, turns: u32) -> Self {
294 assert!(turns > 0, "max_turns must be greater than 0");
295 self.config.max_turns = Some(turns);
296 self
297 }
298
299 pub fn max_budget_usd(mut self, budget: f64) -> Self {
305 assert!(
306 budget.is_finite() && budget > 0.0,
307 "budget must be a positive finite number, got {budget}"
308 );
309 self.config.max_budget_usd = Some(budget);
310 self
311 }
312
313 pub fn working_dir(mut self, dir: &str) -> Self {
315 self.config.working_dir = Some(dir.to_string());
316 self
317 }
318
319 pub fn mcp_config(mut self, config: &str) -> Self {
321 self.config.mcp_config = Some(config.to_string());
322 self
323 }
324
325 pub fn permission_mode(mut self, mode: PermissionMode) -> Self {
329 self.config.permission_mode = mode;
330 self
331 }
332
333 pub fn output<T: JsonSchema>(mut self) -> Self {
364 let schema = schema_for!(T);
365 self.config.json_schema = match to_string(&schema) {
366 Ok(s) => Some(s),
367 Err(e) => {
368 warn!(error = %e, type_name = any::type_name::<T>(), "failed to serialize JSON schema, structured output disabled");
369 None
370 }
371 };
372 self
373 }
374
375 pub fn output_schema_raw(mut self, schema: &str) -> Self {
399 self.config.json_schema = Some(schema.to_string());
400 self
401 }
402
403 pub fn retry(mut self, max_retries: u32) -> Self {
431 self.retry_policy = Some(RetryPolicy::new(max_retries));
432 self
433 }
434
435 pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
462 self.retry_policy = Some(policy);
463 self
464 }
465
466 pub fn dry_run(mut self, enabled: bool) -> Self {
475 self.dry_run = Some(enabled);
476 self
477 }
478
479 pub fn log_sink(mut self, sink: Arc<dyn LogSink>) -> Self {
506 self.log_sink = Some(sink);
507 self
508 }
509
510 pub fn trace_context(mut self, ctx: WorkflowTraceContext) -> Self {
535 self.config.trace_context = Some(ctx);
536 self
537 }
538
539 pub fn verbose(mut self) -> Self {
569 self.config.verbose = true;
570 self
571 }
572
573 pub fn resume(mut self, session_id: &str) -> Self {
609 assert!(!session_id.is_empty(), "session_id must not be empty");
610 assert!(
611 session_id
612 .chars()
613 .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'),
614 "session_id must only contain alphanumeric characters, hyphens, or underscores, got: {session_id}"
615 );
616 self.config.resume_session_id = Some(session_id.to_string());
617 self
618 }
619
620 #[tracing::instrument(name = "agent", skip_all, fields(model = %self.config.model, prompt_len = self.config.prompt.len()))]
641 pub async fn run(self, provider: &dyn AgentProvider) -> Result<AgentResult, OperationError> {
642 assert!(
643 !self.config.prompt.trim().is_empty(),
644 "prompt must not be empty - call .prompt(\"...\") before .run()"
645 );
646
647 if crate::dry_run::effective_dry_run(self.dry_run) {
648 info!(
649 prompt_len = self.config.prompt.len(),
650 "[dry-run] agent call skipped"
651 );
652 let mut output =
653 AgentOutput::new(Value::String("[dry-run] agent call skipped".to_string()));
654 output.cost_usd = Some(0.0);
655 output.input_tokens = Some(0);
656 output.output_tokens = Some(0);
657 return Ok(AgentResult { output });
658 }
659
660 let result = self.invoke_once(provider).await;
661
662 let default_schema_retry = RetryPolicy::new(2);
663 let policy = match &self.retry_policy {
664 Some(p) => p,
665 None if self.config.json_schema.is_some() => &default_schema_retry,
666 None => return result,
667 };
668
669 if let Err(ref err) = result {
671 if !crate::retry::is_retryable(err) {
672 return result;
673 }
674 } else {
675 return result;
676 }
677
678 let mut last_result = result;
679
680 for attempt in 0..policy.max_retries {
681 let delay = policy.delay_for_attempt(attempt);
682 let retry_reason = if matches!(
683 &last_result,
684 Err(OperationError::Agent(
685 crate::error::AgentError::SchemaValidation { .. }
686 ))
687 ) {
688 "structured_output was null (CLI non-determinism)"
689 } else {
690 "transient failure"
691 };
692 warn!(
693 attempt = attempt + 1,
694 max_retries = policy.max_retries,
695 delay_ms = delay.as_millis() as u64,
696 reason = retry_reason,
697 "retrying agent invocation"
698 );
699 time::sleep(delay).await;
700
701 last_result = self.invoke_once(provider).await;
702
703 match &last_result {
704 Ok(_) => return last_result,
705 Err(err) if !crate::retry::is_retryable(err) => return last_result,
706 _ => {}
707 }
708 }
709
710 last_result
711 }
712
713 async fn invoke_once(
715 &self,
716 provider: &dyn AgentProvider,
717 ) -> Result<AgentResult, OperationError> {
718 #[cfg(feature = "prometheus")]
719 let model_label = self.config.model.to_string();
720
721 let invoke_result = match self.log_sink {
722 Some(ref sink) => provider.invoke_with_logs(&self.config, sink.clone()).await,
723 None => provider.invoke(&self.config).await,
724 };
725 let output = match invoke_result {
726 Ok(output) => output,
727 Err(e) => {
728 #[cfg(feature = "prometheus")]
729 {
730 metrics::counter!(metric_names::AGENT_TOTAL, "model" => model_label.clone(), "status" => metric_names::STATUS_ERROR).increment(1);
731 }
732 return Err(OperationError::Agent(e));
733 }
734 };
735
736 info!(
737 duration_ms = output.duration_ms,
738 cost_usd = output.cost_usd,
739 input_tokens = output.input_tokens,
740 cache_read_input_tokens = output.cache_read_input_tokens,
741 cache_creation_input_tokens = output.cache_creation_input_tokens,
742 output_tokens = output.output_tokens,
743 model = output.model,
744 "agent completed"
745 );
746
747 #[cfg(feature = "prometheus")]
748 {
749 metrics::counter!(metric_names::AGENT_TOTAL, "model" => model_label.clone(), "status" => metric_names::STATUS_SUCCESS).increment(1);
750 metrics::histogram!(metric_names::AGENT_DURATION_SECONDS, "model" => model_label.clone())
751 .record(output.duration_ms as f64 / 1000.0);
752 if let Some(cost) = output.cost_usd {
753 metrics::gauge!(metric_names::AGENT_COST_USD_TOTAL, "model" => model_label.clone())
754 .increment(cost);
755 }
756 if let Some(tokens) = output.input_tokens {
757 metrics::counter!(metric_names::AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(tokens);
758 }
759 if let Some(t) = output.cache_read_input_tokens {
760 metrics::counter!(metric_names::AGENT_TOKENS_CACHE_READ_TOTAL, "model" => model_label.clone()).increment(t);
761 }
762 if let Some(t) = output.cache_creation_input_tokens {
763 metrics::counter!(metric_names::AGENT_TOKENS_CACHE_WRITE_TOTAL, "model" => model_label.clone()).increment(t);
764 }
765 if let Some(tokens) = output.output_tokens {
766 metrics::counter!(metric_names::AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label)
767 .increment(tokens);
768 }
769 }
770
771 Ok(AgentResult { output })
772 }
773}
774
775impl Default for Agent {
776 fn default() -> Self {
777 Self::new()
778 }
779}
780
781#[derive(Debug)]
786pub struct AgentResult {
787 output: AgentOutput,
788}
789
790impl AgentResult {
791 pub fn text(&self) -> &str {
796 match self.output.value.as_str() {
797 Some(s) => s,
798 None => {
799 warn!(
800 value_type = self.output.value.to_string(),
801 "agent output is not a string, returning empty"
802 );
803 ""
804 }
805 }
806 }
807
808 pub fn value(&self) -> &Value {
810 &self.output.value
811 }
812
813 pub fn json<T: DeserializeOwned>(&self) -> Result<T, OperationError> {
823 from_value(self.output.value.clone()).map_err(OperationError::deserialize::<T>)
824 }
825
826 pub fn into_json<T: DeserializeOwned>(self) -> Result<T, OperationError> {
832 from_value(self.output.value).map_err(OperationError::deserialize::<T>)
833 }
834
835 #[cfg(test)]
840 pub(crate) fn from_output(output: AgentOutput) -> Self {
841 Self { output }
842 }
843
844 pub fn session_id(&self) -> Option<&str> {
846 self.output.session_id.as_deref()
847 }
848
849 pub fn cost_usd(&self) -> Option<f64> {
851 self.output.cost_usd
852 }
853
854 pub fn input_tokens(&self) -> Option<u64> {
860 self.output.input_tokens
861 }
862
863 pub fn cache_read_input_tokens(&self) -> Option<u64> {
878 self.output.cache_read_input_tokens
879 }
880
881 pub fn cache_creation_input_tokens(&self) -> Option<u64> {
896 self.output.cache_creation_input_tokens
897 }
898
899 pub fn output_tokens(&self) -> Option<u64> {
901 self.output.output_tokens
902 }
903
904 pub fn duration_ms(&self) -> u64 {
906 self.output.duration_ms
907 }
908
909 pub fn model(&self) -> Option<&str> {
911 self.output.model.as_deref()
912 }
913
914 pub fn debug_messages(&self) -> Option<&[DebugMessage]> {
920 self.output.debug_messages.as_deref()
921 }
922}
923
924#[cfg(test)]
925mod tests {
926 use super::*;
927 use crate::error::AgentError;
928 use crate::provider::InvokeFuture;
929 use serde_json::json;
930
931 struct TestProvider {
932 output: AgentOutput,
933 }
934
935 impl AgentProvider for TestProvider {
936 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
937 Box::pin(async move {
938 Ok(AgentOutput {
939 value: self.output.value.clone(),
940 session_id: self.output.session_id.clone(),
941 cost_usd: self.output.cost_usd,
942 input_tokens: self.output.input_tokens,
943 cache_read_input_tokens: None,
944 cache_creation_input_tokens: None,
945 output_tokens: self.output.output_tokens,
946 model: self.output.model.clone(),
947 duration_ms: self.output.duration_ms,
948 debug_messages: None,
949 })
950 })
951 }
952 }
953
954 struct ConfigCapture {
955 output: AgentOutput,
956 }
957
958 impl AgentProvider for ConfigCapture {
959 fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
960 let config_json = serde_json::to_value(config).unwrap();
961 Box::pin(async move {
962 Ok(AgentOutput {
963 value: config_json,
964 session_id: self.output.session_id.clone(),
965 cost_usd: self.output.cost_usd,
966 input_tokens: self.output.input_tokens,
967 cache_read_input_tokens: None,
968 cache_creation_input_tokens: None,
969 output_tokens: self.output.output_tokens,
970 model: self.output.model.clone(),
971 duration_ms: self.output.duration_ms,
972 debug_messages: None,
973 })
974 })
975 }
976 }
977
978 fn default_output() -> AgentOutput {
979 AgentOutput {
980 value: json!("test output"),
981 session_id: Some("sess-123".to_string()),
982 cost_usd: Some(0.05),
983 input_tokens: Some(100),
984 cache_read_input_tokens: None,
985 cache_creation_input_tokens: None,
986 output_tokens: Some(50),
987 model: Some("sonnet".to_string()),
988 duration_ms: 1500,
989 debug_messages: None,
990 }
991 }
992
993 #[test]
996 fn model_constants_have_expected_values() {
997 assert_eq!(Model::SONNET, "sonnet");
998 assert_eq!(Model::OPUS, "opus");
999 assert_eq!(Model::HAIKU, "haiku");
1000 assert_eq!(Model::HAIKU_45, "claude-haiku-4-5-20251001");
1001 assert_eq!(Model::SONNET_46, "claude-sonnet-4-6");
1002 assert_eq!(Model::OPUS_46, "claude-opus-4-6");
1003 assert_eq!(Model::SONNET_46_1M, "claude-sonnet-4-6[1m]");
1004 assert_eq!(Model::OPUS_46_1M, "claude-opus-4-6[1m]");
1005 assert_eq!(Model::OPUS_47, "claude-opus-4-7");
1006 assert_eq!(Model::OPUS_47_1M, "claude-opus-4-7[1m]");
1007 assert_eq!(Model::OPUS_48, "claude-opus-4-8");
1008 assert_eq!(Model::OPUS_48_1M, "claude-opus-4-8[1m]");
1009 assert_eq!(Model::FABLE_5, "claude-fable-5");
1010 assert_eq!(Model::FABLE_51, "claude-fable-5-1");
1011 assert_eq!(Model::MYTHOS_5, "claude-mythos-5");
1012 assert_eq!(Model::MYTHOS_51, "claude-mythos-5-1");
1013 assert_eq!(Model::OPUS_55, "claude-opus-5-5");
1014 assert_eq!(Model::SONNET_55, "claude-sonnet-5-5");
1015 assert_eq!(Model::OPUS_5, "claude-opus-5");
1016 assert_eq!(Model::OPUS_5_1M, "claude-opus-5[1m]");
1017 assert_eq!(Model::SONNET_5, "claude-sonnet-5");
1018 assert_eq!(Model::SONNET_5_1M, "claude-sonnet-5[1m]");
1019 }
1020
1021 #[tokio::test]
1024 async fn agent_new_default_values() {
1025 let provider = ConfigCapture {
1026 output: default_output(),
1027 };
1028 let result = Agent::new().prompt("hi").run(&provider).await.unwrap();
1029
1030 let config = result.value();
1031 assert_eq!(config["system_prompt"], json!(null));
1032 assert_eq!(config["prompt"], json!("hi"));
1033 assert_eq!(config["model"], json!("sonnet"));
1034 assert_eq!(config["allowed_tools"], json!([]));
1035 assert_eq!(config["max_turns"], json!(null));
1036 assert_eq!(config["max_budget_usd"], json!(null));
1037 assert_eq!(config["working_dir"], json!(null));
1038 assert_eq!(config["mcp_config"], json!(null));
1039 assert_eq!(config["permission_mode"], json!("Default"));
1040 assert_eq!(config["json_schema"], json!(null));
1041 }
1042
1043 #[tokio::test]
1044 async fn agent_default_matches_new() {
1045 let provider = ConfigCapture {
1046 output: default_output(),
1047 };
1048 let result_new = Agent::new().prompt("x").run(&provider).await.unwrap();
1049 let result_default = Agent::default().prompt("x").run(&provider).await.unwrap();
1050
1051 assert_eq!(result_new.value(), result_default.value());
1052 }
1053
1054 #[tokio::test]
1057 async fn builder_methods_store_values_correctly() {
1058 let provider = ConfigCapture {
1059 output: default_output(),
1060 };
1061 let result = Agent::new()
1062 .system_prompt("you are a bot")
1063 .prompt("do something")
1064 .model(Model::OPUS)
1065 .allowed_tools(&["Read", "Write"])
1066 .max_turns(5)
1067 .max_budget_usd(1.5)
1068 .working_dir("/tmp")
1069 .mcp_config("{}")
1070 .permission_mode(PermissionMode::Auto)
1071 .run(&provider)
1072 .await
1073 .unwrap();
1074
1075 let config = result.value();
1076 assert_eq!(config["system_prompt"], json!("you are a bot"));
1077 assert_eq!(config["prompt"], json!("do something"));
1078 assert_eq!(config["model"], json!("opus"));
1079 assert_eq!(config["allowed_tools"], json!(["Read", "Write"]));
1080 assert_eq!(config["max_turns"], json!(5));
1081 assert_eq!(config["max_budget_usd"], json!(1.5));
1082 assert_eq!(config["working_dir"], json!("/tmp"));
1083 assert_eq!(config["mcp_config"], json!("{}"));
1084 assert_eq!(config["permission_mode"], json!("Auto"));
1085 }
1086
1087 #[test]
1090 #[should_panic(expected = "max_turns must be greater than 0")]
1091 fn max_turns_zero_panics() {
1092 let _ = Agent::new().max_turns(0);
1093 }
1094
1095 #[test]
1096 #[should_panic(expected = "budget must be a positive finite number")]
1097 fn max_budget_negative_panics() {
1098 let _ = Agent::new().max_budget_usd(-1.0);
1099 }
1100
1101 #[test]
1102 #[should_panic(expected = "budget must be a positive finite number")]
1103 fn max_budget_nan_panics() {
1104 let _ = Agent::new().max_budget_usd(f64::NAN);
1105 }
1106
1107 #[test]
1108 #[should_panic(expected = "budget must be a positive finite number")]
1109 fn max_budget_infinity_panics() {
1110 let _ = Agent::new().max_budget_usd(f64::INFINITY);
1111 }
1112
1113 #[tokio::test]
1116 async fn agent_result_text_with_string_value() {
1117 let provider = TestProvider {
1118 output: AgentOutput {
1119 value: json!("hello world"),
1120 ..default_output()
1121 },
1122 };
1123 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1124 assert_eq!(result.text(), "hello world");
1125 }
1126
1127 #[tokio::test]
1128 async fn agent_result_text_with_non_string_value() {
1129 let provider = TestProvider {
1130 output: AgentOutput {
1131 value: json!(42),
1132 ..default_output()
1133 },
1134 };
1135 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1136 assert_eq!(result.text(), "");
1137 }
1138
1139 #[tokio::test]
1140 async fn agent_result_text_with_null_value() {
1141 let provider = TestProvider {
1142 output: AgentOutput {
1143 value: json!(null),
1144 ..default_output()
1145 },
1146 };
1147 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1148 assert_eq!(result.text(), "");
1149 }
1150
1151 #[tokio::test]
1152 async fn agent_result_json_successful_deserialize() {
1153 #[derive(Deserialize, PartialEq, Debug)]
1154 struct MyOutput {
1155 name: String,
1156 count: u32,
1157 }
1158 let provider = TestProvider {
1159 output: AgentOutput {
1160 value: json!({"name": "test", "count": 7}),
1161 ..default_output()
1162 },
1163 };
1164 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1165 let parsed: MyOutput = result.json().unwrap();
1166 assert_eq!(parsed.name, "test");
1167 assert_eq!(parsed.count, 7);
1168 }
1169
1170 #[tokio::test]
1171 async fn agent_result_json_failed_deserialize() {
1172 #[derive(Debug, Deserialize)]
1173 #[allow(dead_code)]
1174 struct MyOutput {
1175 name: String,
1176 }
1177 let provider = TestProvider {
1178 output: AgentOutput {
1179 value: json!(42),
1180 ..default_output()
1181 },
1182 };
1183 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1184 let err = result.json::<MyOutput>().unwrap_err();
1185 assert!(matches!(err, OperationError::Deserialize { .. }));
1186 }
1187
1188 #[tokio::test]
1189 async fn agent_result_accessors() {
1190 let provider = TestProvider {
1191 output: AgentOutput {
1192 value: json!("v"),
1193 session_id: Some("s-1".to_string()),
1194 cost_usd: Some(0.123),
1195 input_tokens: Some(999),
1196 cache_read_input_tokens: None,
1197 cache_creation_input_tokens: None,
1198 output_tokens: Some(456),
1199 model: Some("opus".to_string()),
1200 duration_ms: 2000,
1201 debug_messages: None,
1202 },
1203 };
1204 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1205 assert_eq!(result.session_id(), Some("s-1"));
1206 assert_eq!(result.cost_usd(), Some(0.123));
1207 assert_eq!(result.input_tokens(), Some(999));
1208 assert_eq!(result.output_tokens(), Some(456));
1209 assert_eq!(result.duration_ms(), 2000);
1210 assert_eq!(result.model(), Some("opus"));
1211 }
1212
1213 #[tokio::test]
1216 async fn resume_passes_session_id_in_config() {
1217 let provider = ConfigCapture {
1218 output: default_output(),
1219 };
1220 let result = Agent::new()
1221 .prompt("followup")
1222 .resume("sess-abc")
1223 .run(&provider)
1224 .await
1225 .unwrap();
1226
1227 let config = result.value();
1228 assert_eq!(config["resume_session_id"], json!("sess-abc"));
1229 }
1230
1231 #[tokio::test]
1232 async fn no_resume_has_null_session_id() {
1233 let provider = ConfigCapture {
1234 output: default_output(),
1235 };
1236 let result = Agent::new()
1237 .prompt("first call")
1238 .run(&provider)
1239 .await
1240 .unwrap();
1241
1242 let config = result.value();
1243 assert_eq!(config["resume_session_id"], json!(null));
1244 }
1245
1246 #[test]
1247 #[should_panic(expected = "session_id must not be empty")]
1248 fn resume_empty_session_id_panics() {
1249 let _ = Agent::new().resume("");
1250 }
1251
1252 #[test]
1253 #[should_panic(expected = "session_id must only contain")]
1254 fn resume_invalid_chars_panics() {
1255 let _ = Agent::new().resume("sess;rm -rf /");
1256 }
1257
1258 #[test]
1259 fn resume_valid_formats_accepted() {
1260 let _ = Agent::new().resume("sess-abc123");
1261 let _ = Agent::new().resume("a1b2c3d4_session");
1262 let _ = Agent::new().resume("abc-DEF-123_456");
1263 }
1264
1265 #[tokio::test]
1266 #[should_panic(expected = "prompt must not be empty")]
1267 async fn run_without_prompt_panics() {
1268 let provider = TestProvider {
1269 output: default_output(),
1270 };
1271 let _ = Agent::new().run(&provider).await;
1272 }
1273
1274 #[tokio::test]
1275 #[should_panic(expected = "prompt must not be empty")]
1276 async fn run_with_whitespace_only_prompt_panics() {
1277 let provider = TestProvider {
1278 output: default_output(),
1279 };
1280 let _ = Agent::new().prompt(" ").run(&provider).await;
1281 }
1282
1283 #[tokio::test]
1286 async fn model_accepts_custom_string() {
1287 let provider = ConfigCapture {
1288 output: default_output(),
1289 };
1290 let result = Agent::new()
1291 .prompt("hi")
1292 .model("mistral-large-latest")
1293 .run(&provider)
1294 .await
1295 .unwrap();
1296 assert_eq!(result.value()["model"], json!("mistral-large-latest"));
1297 }
1298
1299 #[tokio::test]
1300 async fn verbose_sets_config_flag() {
1301 let provider = ConfigCapture {
1302 output: default_output(),
1303 };
1304 let result = Agent::new()
1305 .prompt("hi")
1306 .verbose()
1307 .run(&provider)
1308 .await
1309 .unwrap();
1310 assert_eq!(result.value()["verbose"], json!(true));
1311 }
1312
1313 #[tokio::test]
1314 async fn verbose_not_set_by_default() {
1315 let provider = ConfigCapture {
1316 output: default_output(),
1317 };
1318 let result = Agent::new().prompt("hi").run(&provider).await.unwrap();
1319 assert_eq!(result.value()["verbose"], json!(false));
1320 }
1321
1322 #[tokio::test]
1323 async fn debug_messages_none_without_verbose() {
1324 let provider = TestProvider {
1325 output: default_output(),
1326 };
1327 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1328 assert!(result.debug_messages().is_none());
1329 }
1330
1331 #[tokio::test]
1332 async fn model_accepts_owned_string() {
1333 let provider = ConfigCapture {
1334 output: default_output(),
1335 };
1336 let model_name = String::from("gpt-4o");
1337 let result = Agent::new()
1338 .prompt("hi")
1339 .model(model_name)
1340 .run(&provider)
1341 .await
1342 .unwrap();
1343 assert_eq!(result.value()["model"], json!("gpt-4o"));
1344 }
1345
1346 #[tokio::test]
1347 async fn into_json_success() {
1348 #[derive(Deserialize, PartialEq, Debug)]
1349 struct Out {
1350 name: String,
1351 }
1352 let provider = TestProvider {
1353 output: AgentOutput {
1354 value: json!({"name": "test"}),
1355 ..default_output()
1356 },
1357 };
1358 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1359 let parsed: Out = result.into_json().unwrap();
1360 assert_eq!(parsed.name, "test");
1361 }
1362
1363 #[tokio::test]
1364 async fn into_json_failure() {
1365 #[derive(Debug, Deserialize)]
1366 #[allow(dead_code)]
1367 struct Out {
1368 name: String,
1369 }
1370 let provider = TestProvider {
1371 output: AgentOutput {
1372 value: json!(42),
1373 ..default_output()
1374 },
1375 };
1376 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1377 let err = result.into_json::<Out>().unwrap_err();
1378 assert!(matches!(err, OperationError::Deserialize { .. }));
1379 }
1380
1381 #[test]
1382 fn from_output_creates_result() {
1383 let output = AgentOutput {
1384 value: json!("hello"),
1385 ..default_output()
1386 };
1387 let result = AgentResult::from_output(output);
1388 assert_eq!(result.text(), "hello");
1389 assert_eq!(result.cost_usd(), Some(0.05));
1390 }
1391
1392 #[test]
1393 #[should_panic(expected = "budget must be a positive finite number")]
1394 fn max_budget_zero_panics() {
1395 let _ = Agent::new().max_budget_usd(0.0);
1396 }
1397
1398 #[test]
1399 fn model_constant_equality() {
1400 assert_eq!(Model::SONNET, "sonnet");
1401 assert_ne!(Model::SONNET, Model::OPUS);
1402 }
1403
1404 #[test]
1405 fn permission_mode_serialize_deserialize_roundtrip() {
1406 for mode in [
1407 PermissionMode::Default,
1408 PermissionMode::Auto,
1409 PermissionMode::DontAsk,
1410 PermissionMode::BypassPermissions,
1411 ] {
1412 let json = to_string(&mode).unwrap();
1413 let back: PermissionMode = serde_json::from_str(&json).unwrap();
1414 assert_eq!(format!("{:?}", mode), format!("{:?}", back));
1415 }
1416 }
1417
1418 #[test]
1421 fn retry_builder_stores_policy() {
1422 let agent = Agent::new().retry(3);
1423 assert!(agent.retry_policy.is_some());
1424 assert_eq!(agent.retry_policy.unwrap().max_retries(), 3);
1425 }
1426
1427 #[test]
1428 fn retry_policy_builder_stores_custom_policy() {
1429 use crate::retry::RetryPolicy;
1430 let policy = RetryPolicy::new(5).backoff(Duration::from_secs(1));
1431 let agent = Agent::new().retry_policy(policy);
1432 let p = agent.retry_policy.unwrap();
1433 assert_eq!(p.max_retries(), 5);
1434 }
1435
1436 #[test]
1437 fn no_retry_by_default() {
1438 let agent = Agent::new();
1439 assert!(agent.retry_policy.is_none());
1440 }
1441
1442 use std::sync::Arc;
1445 use std::sync::atomic::{AtomicU32, Ordering};
1446 use std::time::Duration;
1447
1448 struct FailNTimesProvider {
1449 fail_count: AtomicU32,
1450 failures_before_success: u32,
1451 output: AgentOutput,
1452 }
1453
1454 impl AgentProvider for FailNTimesProvider {
1455 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1456 Box::pin(async move {
1457 let current = self.fail_count.fetch_add(1, Ordering::SeqCst);
1458 if current < self.failures_before_success {
1459 Err(AgentError::ProcessFailed {
1460 exit_code: 1,
1461 stderr: format!("transient failure #{}", current + 1),
1462 })
1463 } else {
1464 Ok(AgentOutput {
1465 value: self.output.value.clone(),
1466 session_id: self.output.session_id.clone(),
1467 cost_usd: self.output.cost_usd,
1468 input_tokens: self.output.input_tokens,
1469 cache_read_input_tokens: None,
1470 cache_creation_input_tokens: None,
1471 output_tokens: self.output.output_tokens,
1472 model: self.output.model.clone(),
1473 duration_ms: self.output.duration_ms,
1474 debug_messages: None,
1475 })
1476 }
1477 })
1478 }
1479 }
1480
1481 #[tokio::test]
1482 async fn retry_succeeds_after_transient_failures() {
1483 let provider = FailNTimesProvider {
1484 fail_count: AtomicU32::new(0),
1485 failures_before_success: 2,
1486 output: default_output(),
1487 };
1488 let result = Agent::new()
1489 .prompt("test")
1490 .retry_policy(crate::retry::RetryPolicy::new(3).backoff(Duration::from_millis(1)))
1491 .run(&provider)
1492 .await;
1493
1494 assert!(result.is_ok());
1495 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 3); }
1497
1498 #[tokio::test]
1499 async fn retry_exhausted_returns_last_error() {
1500 let provider = FailNTimesProvider {
1501 fail_count: AtomicU32::new(0),
1502 failures_before_success: 10, output: default_output(),
1504 };
1505 let result = Agent::new()
1506 .prompt("test")
1507 .retry_policy(crate::retry::RetryPolicy::new(2).backoff(Duration::from_millis(1)))
1508 .run(&provider)
1509 .await;
1510
1511 assert!(result.is_err());
1512 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 3);
1514 }
1515
1516 #[tokio::test]
1517 async fn retry_does_not_retry_prompt_too_large() {
1518 let call_count = Arc::new(AtomicU32::new(0));
1519 let count = call_count.clone();
1520
1521 struct CountingNonRetryable {
1522 count: Arc<AtomicU32>,
1523 }
1524 impl AgentProvider for CountingNonRetryable {
1525 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1526 self.count.fetch_add(1, Ordering::SeqCst);
1527 Box::pin(async move {
1528 Err(AgentError::PromptTooLarge {
1529 chars: 1_000_000,
1530 estimated_tokens: 250_000,
1531 model_limit: 200_000,
1532 })
1533 })
1534 }
1535 }
1536
1537 let provider = CountingNonRetryable { count };
1538 let result = Agent::new()
1539 .prompt("test")
1540 .retry_policy(crate::retry::RetryPolicy::new(3).backoff(Duration::from_millis(1)))
1541 .run(&provider)
1542 .await;
1543
1544 assert!(result.is_err());
1545 assert_eq!(call_count.load(Ordering::SeqCst), 1);
1546 }
1547
1548 #[tokio::test]
1549 async fn retry_retries_schema_validation_errors() {
1550 let call_count = Arc::new(AtomicU32::new(0));
1551 let count = call_count.clone();
1552
1553 struct SchemaFailProvider {
1554 count: Arc<AtomicU32>,
1555 }
1556 impl AgentProvider for SchemaFailProvider {
1557 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1558 self.count.fetch_add(1, Ordering::SeqCst);
1559 Box::pin(async move {
1560 Err(AgentError::SchemaValidation {
1561 expected: "object".to_string(),
1562 got: "null".to_string(),
1563 debug_messages: Vec::new(),
1564 partial_usage: Box::default(),
1565 raw_response: None,
1566 })
1567 })
1568 }
1569 }
1570
1571 let provider = SchemaFailProvider { count };
1572 let result = Agent::new()
1573 .prompt("test")
1574 .retry_policy(crate::retry::RetryPolicy::new(2).backoff(Duration::from_millis(1)))
1575 .run(&provider)
1576 .await;
1577
1578 assert!(result.is_err());
1579 assert_eq!(call_count.load(Ordering::SeqCst), 3);
1581 }
1582
1583 #[tokio::test]
1584 async fn schema_validation_succeeds_on_retry() {
1585 let call_count = Arc::new(AtomicU32::new(0));
1586 let count = call_count.clone();
1587
1588 struct SchemaFailThenSucceed {
1589 count: Arc<AtomicU32>,
1590 output: AgentOutput,
1591 }
1592 impl AgentProvider for SchemaFailThenSucceed {
1593 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1594 let current = self.count.fetch_add(1, Ordering::SeqCst);
1595 let output = self.output.clone();
1596 Box::pin(async move {
1597 if current == 0 {
1598 Err(AgentError::SchemaValidation {
1599 expected: "structured_output field".to_string(),
1600 got: "null".to_string(),
1601 debug_messages: Vec::new(),
1602 partial_usage: Box::default(),
1603 raw_response: None,
1604 })
1605 } else {
1606 Ok(output)
1607 }
1608 })
1609 }
1610 }
1611
1612 let provider = SchemaFailThenSucceed {
1613 count,
1614 output: default_output(),
1615 };
1616 let result = Agent::new()
1617 .prompt("test")
1618 .retry_policy(crate::retry::RetryPolicy::new(1).backoff(Duration::from_millis(1)))
1619 .run(&provider)
1620 .await;
1621
1622 assert!(result.is_ok());
1623 assert_eq!(call_count.load(Ordering::SeqCst), 2);
1624 }
1625
1626 #[tokio::test]
1627 async fn auto_retry_applied_when_json_schema_set() {
1628 let call_count = Arc::new(AtomicU32::new(0));
1629 let count = call_count.clone();
1630
1631 struct AlwaysSchemaFail {
1632 count: Arc<AtomicU32>,
1633 }
1634 impl AgentProvider for AlwaysSchemaFail {
1635 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1636 self.count.fetch_add(1, Ordering::SeqCst);
1637 Box::pin(async move {
1638 Err(AgentError::SchemaValidation {
1639 expected: "object".to_string(),
1640 got: "null".to_string(),
1641 debug_messages: Vec::new(),
1642 partial_usage: Box::default(),
1643 raw_response: None,
1644 })
1645 })
1646 }
1647 }
1648
1649 let provider = AlwaysSchemaFail { count };
1650 let result = Agent::new()
1651 .prompt("test")
1652 .output_schema_raw(r#"{"type":"object"}"#)
1653 .run(&provider)
1654 .await;
1655
1656 assert!(result.is_err());
1657 assert_eq!(call_count.load(Ordering::SeqCst), 3);
1659 }
1660
1661 #[tokio::test]
1662 async fn no_retry_without_policy() {
1663 let provider = FailNTimesProvider {
1664 fail_count: AtomicU32::new(0),
1665 failures_before_success: 1,
1666 output: default_output(),
1667 };
1668 let result = Agent::new().prompt("test").run(&provider).await;
1669
1670 assert!(result.is_err());
1671 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 1);
1672 }
1673
1674 use crate::test_support::VecSink;
1677
1678 struct SinkCapture {
1679 output: AgentOutput,
1680 saw_logs: Arc<AtomicU32>,
1681 }
1682
1683 impl AgentProvider for SinkCapture {
1684 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1685 Box::pin(async {
1686 Ok(AgentOutput {
1687 value: self.output.value.clone(),
1688 session_id: self.output.session_id.clone(),
1689 cost_usd: self.output.cost_usd,
1690 input_tokens: self.output.input_tokens,
1691 cache_read_input_tokens: None,
1692 cache_creation_input_tokens: None,
1693 output_tokens: self.output.output_tokens,
1694 model: self.output.model.clone(),
1695 duration_ms: self.output.duration_ms,
1696 debug_messages: None,
1697 })
1698 })
1699 }
1700
1701 fn invoke_with_logs<'a>(
1702 &'a self,
1703 config: &'a AgentConfig,
1704 log_sink: Arc<dyn LogSink>,
1705 ) -> InvokeFuture<'a> {
1706 self.saw_logs.fetch_add(1, Ordering::SeqCst);
1707 log_sink.log("stdout", "streaming line");
1708 self.invoke(config)
1709 }
1710 }
1711
1712 #[tokio::test]
1713 async fn log_sink_routes_to_invoke_with_logs() {
1714 let saw_logs = Arc::new(AtomicU32::new(0));
1715 let provider = SinkCapture {
1716 output: default_output(),
1717 saw_logs: saw_logs.clone(),
1718 };
1719 let sink: Arc<dyn LogSink> = VecSink::new();
1720
1721 let result = Agent::new()
1722 .prompt("test")
1723 .log_sink(sink)
1724 .run(&provider)
1725 .await;
1726
1727 assert!(result.is_ok());
1728 assert_eq!(saw_logs.load(Ordering::SeqCst), 1);
1729 }
1730
1731 #[tokio::test]
1732 async fn no_log_sink_routes_to_invoke() {
1733 let saw_logs = Arc::new(AtomicU32::new(0));
1734 let provider = SinkCapture {
1735 output: default_output(),
1736 saw_logs: saw_logs.clone(),
1737 };
1738
1739 let result = Agent::new().prompt("test").run(&provider).await;
1740
1741 assert!(result.is_ok());
1742 assert_eq!(saw_logs.load(Ordering::SeqCst), 0);
1743 }
1744
1745 #[tokio::test]
1746 async fn log_sink_receives_provider_lines() {
1747 let saw_logs = Arc::new(AtomicU32::new(0));
1748 let provider = SinkCapture {
1749 output: default_output(),
1750 saw_logs: saw_logs.clone(),
1751 };
1752 let sink = VecSink::new();
1753
1754 let _ = Agent::new()
1755 .prompt("test")
1756 .log_sink(sink.clone() as Arc<dyn LogSink>)
1757 .run(&provider)
1758 .await;
1759
1760 let lines = sink.0.lock().unwrap();
1761 assert_eq!(lines.len(), 1);
1762 assert_eq!(lines[0].0, "stdout");
1763 assert_eq!(lines[0].1, "streaming line");
1764 }
1765}