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