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;
43
44pub struct Model;
75
76impl Model {
77 pub const SONNET: &str = "sonnet";
81 pub const OPUS: &str = "opus";
83 pub const HAIKU: &str = "haiku";
85
86 pub const HAIKU_45: &str = "claude-haiku-4-5-20251001";
90
91 pub const SONNET_46: &str = "claude-sonnet-4-6";
95 pub const OPUS_46: &str = "claude-opus-4-6";
97
98 pub const SONNET_46_1M: &str = "claude-sonnet-4-6[1m]";
102 pub const OPUS_46_1M: &str = "claude-opus-4-6[1m]";
104
105 pub const OPUS_47: &str = "claude-opus-4-7";
109 pub const OPUS_47_1M: &str = "claude-opus-4-7[1m]";
111
112 pub const OPUS_48: &str = "claude-opus-4-8";
116 pub const OPUS_48_1M: &str = "claude-opus-4-8[1m]";
118
119 pub const FABLE_5: &str = "claude-fable-5";
123 pub const MYTHOS_5: &str = "claude-mythos-5";
125 pub const OPUS_5: &str = "claude-opus-5";
127 pub const OPUS_5_1M: &str = "claude-opus-5[1m]";
129 pub const SONNET_5: &str = "claude-sonnet-5";
131 pub const SONNET_5_1M: &str = "claude-sonnet-5[1m]";
133}
134
135#[derive(Debug, Default, Clone, Copy, Serialize)]
140pub enum PermissionMode {
141 #[default]
143 Default,
144 Auto,
146 DontAsk,
148 BypassPermissions,
152}
153
154impl<'de> Deserialize<'de> for PermissionMode {
155 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
156 where
157 D: serde::Deserializer<'de>,
158 {
159 let s = String::deserialize(deserializer)?;
160 Ok(match s.to_lowercase().replace('_', "").as_str() {
161 "auto" => Self::Auto,
162 "dontask" => Self::DontAsk,
163 "bypass" | "bypasspermissions" => Self::BypassPermissions,
164 _ => Self::Default,
165 })
166 }
167}
168
169#[must_use = "an Agent does nothing until .run() is awaited"]
199pub struct Agent {
200 config: AgentConfig,
201 dry_run: Option<bool>,
202 retry_policy: Option<RetryPolicy>,
203 log_sink: Option<Arc<dyn LogSink>>,
204}
205
206impl Agent {
207 pub fn new() -> Self {
212 Self {
213 config: AgentConfig::new(""),
214 dry_run: None,
215 retry_policy: None,
216 log_sink: None,
217 }
218 }
219
220 pub fn from_config(config: impl Into<AgentConfig>) -> Self {
239 Self {
240 config: config.into(),
241 dry_run: None,
242 retry_policy: None,
243 log_sink: None,
244 }
245 }
246
247 pub fn system_prompt(mut self, prompt: &str) -> Self {
249 self.config.system_prompt = Some(prompt.to_string());
250 self
251 }
252
253 pub fn prompt(mut self, prompt: &str) -> Self {
255 self.config.prompt = prompt.to_string();
256 self
257 }
258
259 pub fn model(mut self, model: impl Into<String>) -> Self {
266 self.config.model = model.into();
267 self
268 }
269
270 pub fn allowed_tools(mut self, tools: &[&str]) -> Self {
275 self.config.allowed_tools = tools.iter().map(|s| s.to_string()).collect();
276 self
277 }
278
279 pub fn max_turns(mut self, turns: u32) -> Self {
285 assert!(turns > 0, "max_turns must be greater than 0");
286 self.config.max_turns = Some(turns);
287 self
288 }
289
290 pub fn max_budget_usd(mut self, budget: f64) -> Self {
296 assert!(
297 budget.is_finite() && budget > 0.0,
298 "budget must be a positive finite number, got {budget}"
299 );
300 self.config.max_budget_usd = Some(budget);
301 self
302 }
303
304 pub fn working_dir(mut self, dir: &str) -> Self {
306 self.config.working_dir = Some(dir.to_string());
307 self
308 }
309
310 pub fn mcp_config(mut self, config: &str) -> Self {
312 self.config.mcp_config = Some(config.to_string());
313 self
314 }
315
316 pub fn permission_mode(mut self, mode: PermissionMode) -> Self {
320 self.config.permission_mode = mode;
321 self
322 }
323
324 pub fn output<T: JsonSchema>(mut self) -> Self {
355 let schema = schema_for!(T);
356 self.config.json_schema = match to_string(&schema) {
357 Ok(s) => Some(s),
358 Err(e) => {
359 warn!(error = %e, type_name = any::type_name::<T>(), "failed to serialize JSON schema, structured output disabled");
360 None
361 }
362 };
363 self
364 }
365
366 pub fn output_schema_raw(mut self, schema: &str) -> Self {
390 self.config.json_schema = Some(schema.to_string());
391 self
392 }
393
394 pub fn retry(mut self, max_retries: u32) -> Self {
422 self.retry_policy = Some(RetryPolicy::new(max_retries));
423 self
424 }
425
426 pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
453 self.retry_policy = Some(policy);
454 self
455 }
456
457 pub fn dry_run(mut self, enabled: bool) -> Self {
466 self.dry_run = Some(enabled);
467 self
468 }
469
470 pub fn log_sink(mut self, sink: Arc<dyn LogSink>) -> Self {
497 self.log_sink = Some(sink);
498 self
499 }
500
501 pub fn verbose(mut self) -> Self {
531 self.config.verbose = true;
532 self
533 }
534
535 pub fn resume(mut self, session_id: &str) -> Self {
571 assert!(!session_id.is_empty(), "session_id must not be empty");
572 assert!(
573 session_id
574 .chars()
575 .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'),
576 "session_id must only contain alphanumeric characters, hyphens, or underscores, got: {session_id}"
577 );
578 self.config.resume_session_id = Some(session_id.to_string());
579 self
580 }
581
582 #[tracing::instrument(name = "agent", skip_all, fields(model = %self.config.model, prompt_len = self.config.prompt.len()))]
603 pub async fn run(self, provider: &dyn AgentProvider) -> Result<AgentResult, OperationError> {
604 assert!(
605 !self.config.prompt.trim().is_empty(),
606 "prompt must not be empty - call .prompt(\"...\") before .run()"
607 );
608
609 if crate::dry_run::effective_dry_run(self.dry_run) {
610 info!(
611 prompt_len = self.config.prompt.len(),
612 "[dry-run] agent call skipped"
613 );
614 let mut output =
615 AgentOutput::new(Value::String("[dry-run] agent call skipped".to_string()));
616 output.cost_usd = Some(0.0);
617 output.input_tokens = Some(0);
618 output.output_tokens = Some(0);
619 return Ok(AgentResult { output });
620 }
621
622 let result = self.invoke_once(provider).await;
623
624 let default_schema_retry = RetryPolicy::new(2);
625 let policy = match &self.retry_policy {
626 Some(p) => p,
627 None if self.config.json_schema.is_some() => &default_schema_retry,
628 None => return result,
629 };
630
631 if let Err(ref err) = result {
633 if !crate::retry::is_retryable(err) {
634 return result;
635 }
636 } else {
637 return result;
638 }
639
640 let mut last_result = result;
641
642 for attempt in 0..policy.max_retries {
643 let delay = policy.delay_for_attempt(attempt);
644 let retry_reason = if matches!(
645 &last_result,
646 Err(OperationError::Agent(
647 crate::error::AgentError::SchemaValidation { .. }
648 ))
649 ) {
650 "structured_output was null (CLI non-determinism)"
651 } else {
652 "transient failure"
653 };
654 warn!(
655 attempt = attempt + 1,
656 max_retries = policy.max_retries,
657 delay_ms = delay.as_millis() as u64,
658 reason = retry_reason,
659 "retrying agent invocation"
660 );
661 time::sleep(delay).await;
662
663 last_result = self.invoke_once(provider).await;
664
665 match &last_result {
666 Ok(_) => return last_result,
667 Err(err) if !crate::retry::is_retryable(err) => return last_result,
668 _ => {}
669 }
670 }
671
672 last_result
673 }
674
675 async fn invoke_once(
677 &self,
678 provider: &dyn AgentProvider,
679 ) -> Result<AgentResult, OperationError> {
680 #[cfg(feature = "prometheus")]
681 let model_label = self.config.model.to_string();
682
683 let invoke_result = match self.log_sink {
684 Some(ref sink) => provider.invoke_with_logs(&self.config, sink.clone()).await,
685 None => provider.invoke(&self.config).await,
686 };
687 let output = match invoke_result {
688 Ok(output) => output,
689 Err(e) => {
690 #[cfg(feature = "prometheus")]
691 {
692 metrics::counter!(metric_names::AGENT_TOTAL, "model" => model_label.clone(), "status" => metric_names::STATUS_ERROR).increment(1);
693 }
694 return Err(OperationError::Agent(e));
695 }
696 };
697
698 info!(
699 duration_ms = output.duration_ms,
700 cost_usd = output.cost_usd,
701 input_tokens = output.input_tokens,
702 output_tokens = output.output_tokens,
703 model = output.model,
704 "agent completed"
705 );
706
707 #[cfg(feature = "prometheus")]
708 {
709 metrics::counter!(metric_names::AGENT_TOTAL, "model" => model_label.clone(), "status" => metric_names::STATUS_SUCCESS).increment(1);
710 metrics::histogram!(metric_names::AGENT_DURATION_SECONDS, "model" => model_label.clone())
711 .record(output.duration_ms as f64 / 1000.0);
712 if let Some(cost) = output.cost_usd {
713 metrics::gauge!(metric_names::AGENT_COST_USD_TOTAL, "model" => model_label.clone())
714 .increment(cost);
715 }
716 if let Some(tokens) = output.input_tokens {
717 metrics::counter!(metric_names::AGENT_TOKENS_INPUT_TOTAL, "model" => model_label.clone()).increment(tokens);
718 }
719 if let Some(tokens) = output.output_tokens {
720 metrics::counter!(metric_names::AGENT_TOKENS_OUTPUT_TOTAL, "model" => model_label)
721 .increment(tokens);
722 }
723 }
724
725 Ok(AgentResult { output })
726 }
727}
728
729impl Default for Agent {
730 fn default() -> Self {
731 Self::new()
732 }
733}
734
735#[derive(Debug)]
740pub struct AgentResult {
741 output: AgentOutput,
742}
743
744impl AgentResult {
745 pub fn text(&self) -> &str {
750 match self.output.value.as_str() {
751 Some(s) => s,
752 None => {
753 warn!(
754 value_type = self.output.value.to_string(),
755 "agent output is not a string, returning empty"
756 );
757 ""
758 }
759 }
760 }
761
762 pub fn value(&self) -> &Value {
764 &self.output.value
765 }
766
767 pub fn json<T: DeserializeOwned>(&self) -> Result<T, OperationError> {
777 from_value(self.output.value.clone()).map_err(OperationError::deserialize::<T>)
778 }
779
780 pub fn into_json<T: DeserializeOwned>(self) -> Result<T, OperationError> {
786 from_value(self.output.value).map_err(OperationError::deserialize::<T>)
787 }
788
789 #[cfg(test)]
794 pub(crate) fn from_output(output: AgentOutput) -> Self {
795 Self { output }
796 }
797
798 pub fn session_id(&self) -> Option<&str> {
800 self.output.session_id.as_deref()
801 }
802
803 pub fn cost_usd(&self) -> Option<f64> {
805 self.output.cost_usd
806 }
807
808 pub fn input_tokens(&self) -> Option<u64> {
810 self.output.input_tokens
811 }
812
813 pub fn output_tokens(&self) -> Option<u64> {
815 self.output.output_tokens
816 }
817
818 pub fn duration_ms(&self) -> u64 {
820 self.output.duration_ms
821 }
822
823 pub fn model(&self) -> Option<&str> {
825 self.output.model.as_deref()
826 }
827
828 pub fn debug_messages(&self) -> Option<&[DebugMessage]> {
834 self.output.debug_messages.as_deref()
835 }
836}
837
838#[cfg(test)]
839mod tests {
840 use super::*;
841 use crate::error::AgentError;
842 use crate::provider::InvokeFuture;
843 use serde_json::json;
844
845 struct TestProvider {
846 output: AgentOutput,
847 }
848
849 impl AgentProvider for TestProvider {
850 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
851 Box::pin(async move {
852 Ok(AgentOutput {
853 value: self.output.value.clone(),
854 session_id: self.output.session_id.clone(),
855 cost_usd: self.output.cost_usd,
856 input_tokens: self.output.input_tokens,
857 output_tokens: self.output.output_tokens,
858 model: self.output.model.clone(),
859 duration_ms: self.output.duration_ms,
860 debug_messages: None,
861 })
862 })
863 }
864 }
865
866 struct ConfigCapture {
867 output: AgentOutput,
868 }
869
870 impl AgentProvider for ConfigCapture {
871 fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
872 let config_json = serde_json::to_value(config).unwrap();
873 Box::pin(async move {
874 Ok(AgentOutput {
875 value: config_json,
876 session_id: self.output.session_id.clone(),
877 cost_usd: self.output.cost_usd,
878 input_tokens: self.output.input_tokens,
879 output_tokens: self.output.output_tokens,
880 model: self.output.model.clone(),
881 duration_ms: self.output.duration_ms,
882 debug_messages: None,
883 })
884 })
885 }
886 }
887
888 fn default_output() -> AgentOutput {
889 AgentOutput {
890 value: json!("test output"),
891 session_id: Some("sess-123".to_string()),
892 cost_usd: Some(0.05),
893 input_tokens: Some(100),
894 output_tokens: Some(50),
895 model: Some("sonnet".to_string()),
896 duration_ms: 1500,
897 debug_messages: None,
898 }
899 }
900
901 #[test]
904 fn model_constants_have_expected_values() {
905 assert_eq!(Model::SONNET, "sonnet");
906 assert_eq!(Model::OPUS, "opus");
907 assert_eq!(Model::HAIKU, "haiku");
908 assert_eq!(Model::HAIKU_45, "claude-haiku-4-5-20251001");
909 assert_eq!(Model::SONNET_46, "claude-sonnet-4-6");
910 assert_eq!(Model::OPUS_46, "claude-opus-4-6");
911 assert_eq!(Model::SONNET_46_1M, "claude-sonnet-4-6[1m]");
912 assert_eq!(Model::OPUS_46_1M, "claude-opus-4-6[1m]");
913 assert_eq!(Model::OPUS_47, "claude-opus-4-7");
914 assert_eq!(Model::OPUS_47_1M, "claude-opus-4-7[1m]");
915 assert_eq!(Model::OPUS_48, "claude-opus-4-8");
916 assert_eq!(Model::OPUS_48_1M, "claude-opus-4-8[1m]");
917 assert_eq!(Model::FABLE_5, "claude-fable-5");
918 assert_eq!(Model::MYTHOS_5, "claude-mythos-5");
919 assert_eq!(Model::OPUS_5, "claude-opus-5");
920 assert_eq!(Model::OPUS_5_1M, "claude-opus-5[1m]");
921 assert_eq!(Model::SONNET_5, "claude-sonnet-5");
922 assert_eq!(Model::SONNET_5_1M, "claude-sonnet-5[1m]");
923 }
924
925 #[tokio::test]
928 async fn agent_new_default_values() {
929 let provider = ConfigCapture {
930 output: default_output(),
931 };
932 let result = Agent::new().prompt("hi").run(&provider).await.unwrap();
933
934 let config = result.value();
935 assert_eq!(config["system_prompt"], json!(null));
936 assert_eq!(config["prompt"], json!("hi"));
937 assert_eq!(config["model"], json!("sonnet"));
938 assert_eq!(config["allowed_tools"], json!([]));
939 assert_eq!(config["max_turns"], json!(null));
940 assert_eq!(config["max_budget_usd"], json!(null));
941 assert_eq!(config["working_dir"], json!(null));
942 assert_eq!(config["mcp_config"], json!(null));
943 assert_eq!(config["permission_mode"], json!("Default"));
944 assert_eq!(config["json_schema"], json!(null));
945 }
946
947 #[tokio::test]
948 async fn agent_default_matches_new() {
949 let provider = ConfigCapture {
950 output: default_output(),
951 };
952 let result_new = Agent::new().prompt("x").run(&provider).await.unwrap();
953 let result_default = Agent::default().prompt("x").run(&provider).await.unwrap();
954
955 assert_eq!(result_new.value(), result_default.value());
956 }
957
958 #[tokio::test]
961 async fn builder_methods_store_values_correctly() {
962 let provider = ConfigCapture {
963 output: default_output(),
964 };
965 let result = Agent::new()
966 .system_prompt("you are a bot")
967 .prompt("do something")
968 .model(Model::OPUS)
969 .allowed_tools(&["Read", "Write"])
970 .max_turns(5)
971 .max_budget_usd(1.5)
972 .working_dir("/tmp")
973 .mcp_config("{}")
974 .permission_mode(PermissionMode::Auto)
975 .run(&provider)
976 .await
977 .unwrap();
978
979 let config = result.value();
980 assert_eq!(config["system_prompt"], json!("you are a bot"));
981 assert_eq!(config["prompt"], json!("do something"));
982 assert_eq!(config["model"], json!("opus"));
983 assert_eq!(config["allowed_tools"], json!(["Read", "Write"]));
984 assert_eq!(config["max_turns"], json!(5));
985 assert_eq!(config["max_budget_usd"], json!(1.5));
986 assert_eq!(config["working_dir"], json!("/tmp"));
987 assert_eq!(config["mcp_config"], json!("{}"));
988 assert_eq!(config["permission_mode"], json!("Auto"));
989 }
990
991 #[test]
994 #[should_panic(expected = "max_turns must be greater than 0")]
995 fn max_turns_zero_panics() {
996 let _ = Agent::new().max_turns(0);
997 }
998
999 #[test]
1000 #[should_panic(expected = "budget must be a positive finite number")]
1001 fn max_budget_negative_panics() {
1002 let _ = Agent::new().max_budget_usd(-1.0);
1003 }
1004
1005 #[test]
1006 #[should_panic(expected = "budget must be a positive finite number")]
1007 fn max_budget_nan_panics() {
1008 let _ = Agent::new().max_budget_usd(f64::NAN);
1009 }
1010
1011 #[test]
1012 #[should_panic(expected = "budget must be a positive finite number")]
1013 fn max_budget_infinity_panics() {
1014 let _ = Agent::new().max_budget_usd(f64::INFINITY);
1015 }
1016
1017 #[tokio::test]
1020 async fn agent_result_text_with_string_value() {
1021 let provider = TestProvider {
1022 output: AgentOutput {
1023 value: json!("hello world"),
1024 ..default_output()
1025 },
1026 };
1027 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1028 assert_eq!(result.text(), "hello world");
1029 }
1030
1031 #[tokio::test]
1032 async fn agent_result_text_with_non_string_value() {
1033 let provider = TestProvider {
1034 output: AgentOutput {
1035 value: json!(42),
1036 ..default_output()
1037 },
1038 };
1039 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1040 assert_eq!(result.text(), "");
1041 }
1042
1043 #[tokio::test]
1044 async fn agent_result_text_with_null_value() {
1045 let provider = TestProvider {
1046 output: AgentOutput {
1047 value: json!(null),
1048 ..default_output()
1049 },
1050 };
1051 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1052 assert_eq!(result.text(), "");
1053 }
1054
1055 #[tokio::test]
1056 async fn agent_result_json_successful_deserialize() {
1057 #[derive(Deserialize, PartialEq, Debug)]
1058 struct MyOutput {
1059 name: String,
1060 count: u32,
1061 }
1062 let provider = TestProvider {
1063 output: AgentOutput {
1064 value: json!({"name": "test", "count": 7}),
1065 ..default_output()
1066 },
1067 };
1068 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1069 let parsed: MyOutput = result.json().unwrap();
1070 assert_eq!(parsed.name, "test");
1071 assert_eq!(parsed.count, 7);
1072 }
1073
1074 #[tokio::test]
1075 async fn agent_result_json_failed_deserialize() {
1076 #[derive(Debug, Deserialize)]
1077 #[allow(dead_code)]
1078 struct MyOutput {
1079 name: String,
1080 }
1081 let provider = TestProvider {
1082 output: AgentOutput {
1083 value: json!(42),
1084 ..default_output()
1085 },
1086 };
1087 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1088 let err = result.json::<MyOutput>().unwrap_err();
1089 assert!(matches!(err, OperationError::Deserialize { .. }));
1090 }
1091
1092 #[tokio::test]
1093 async fn agent_result_accessors() {
1094 let provider = TestProvider {
1095 output: AgentOutput {
1096 value: json!("v"),
1097 session_id: Some("s-1".to_string()),
1098 cost_usd: Some(0.123),
1099 input_tokens: Some(999),
1100 output_tokens: Some(456),
1101 model: Some("opus".to_string()),
1102 duration_ms: 2000,
1103 debug_messages: None,
1104 },
1105 };
1106 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1107 assert_eq!(result.session_id(), Some("s-1"));
1108 assert_eq!(result.cost_usd(), Some(0.123));
1109 assert_eq!(result.input_tokens(), Some(999));
1110 assert_eq!(result.output_tokens(), Some(456));
1111 assert_eq!(result.duration_ms(), 2000);
1112 assert_eq!(result.model(), Some("opus"));
1113 }
1114
1115 #[tokio::test]
1118 async fn resume_passes_session_id_in_config() {
1119 let provider = ConfigCapture {
1120 output: default_output(),
1121 };
1122 let result = Agent::new()
1123 .prompt("followup")
1124 .resume("sess-abc")
1125 .run(&provider)
1126 .await
1127 .unwrap();
1128
1129 let config = result.value();
1130 assert_eq!(config["resume_session_id"], json!("sess-abc"));
1131 }
1132
1133 #[tokio::test]
1134 async fn no_resume_has_null_session_id() {
1135 let provider = ConfigCapture {
1136 output: default_output(),
1137 };
1138 let result = Agent::new()
1139 .prompt("first call")
1140 .run(&provider)
1141 .await
1142 .unwrap();
1143
1144 let config = result.value();
1145 assert_eq!(config["resume_session_id"], json!(null));
1146 }
1147
1148 #[test]
1149 #[should_panic(expected = "session_id must not be empty")]
1150 fn resume_empty_session_id_panics() {
1151 let _ = Agent::new().resume("");
1152 }
1153
1154 #[test]
1155 #[should_panic(expected = "session_id must only contain")]
1156 fn resume_invalid_chars_panics() {
1157 let _ = Agent::new().resume("sess;rm -rf /");
1158 }
1159
1160 #[test]
1161 fn resume_valid_formats_accepted() {
1162 let _ = Agent::new().resume("sess-abc123");
1163 let _ = Agent::new().resume("a1b2c3d4_session");
1164 let _ = Agent::new().resume("abc-DEF-123_456");
1165 }
1166
1167 #[tokio::test]
1168 #[should_panic(expected = "prompt must not be empty")]
1169 async fn run_without_prompt_panics() {
1170 let provider = TestProvider {
1171 output: default_output(),
1172 };
1173 let _ = Agent::new().run(&provider).await;
1174 }
1175
1176 #[tokio::test]
1177 #[should_panic(expected = "prompt must not be empty")]
1178 async fn run_with_whitespace_only_prompt_panics() {
1179 let provider = TestProvider {
1180 output: default_output(),
1181 };
1182 let _ = Agent::new().prompt(" ").run(&provider).await;
1183 }
1184
1185 #[tokio::test]
1188 async fn model_accepts_custom_string() {
1189 let provider = ConfigCapture {
1190 output: default_output(),
1191 };
1192 let result = Agent::new()
1193 .prompt("hi")
1194 .model("mistral-large-latest")
1195 .run(&provider)
1196 .await
1197 .unwrap();
1198 assert_eq!(result.value()["model"], json!("mistral-large-latest"));
1199 }
1200
1201 #[tokio::test]
1202 async fn verbose_sets_config_flag() {
1203 let provider = ConfigCapture {
1204 output: default_output(),
1205 };
1206 let result = Agent::new()
1207 .prompt("hi")
1208 .verbose()
1209 .run(&provider)
1210 .await
1211 .unwrap();
1212 assert_eq!(result.value()["verbose"], json!(true));
1213 }
1214
1215 #[tokio::test]
1216 async fn verbose_not_set_by_default() {
1217 let provider = ConfigCapture {
1218 output: default_output(),
1219 };
1220 let result = Agent::new().prompt("hi").run(&provider).await.unwrap();
1221 assert_eq!(result.value()["verbose"], json!(false));
1222 }
1223
1224 #[tokio::test]
1225 async fn debug_messages_none_without_verbose() {
1226 let provider = TestProvider {
1227 output: default_output(),
1228 };
1229 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1230 assert!(result.debug_messages().is_none());
1231 }
1232
1233 #[tokio::test]
1234 async fn model_accepts_owned_string() {
1235 let provider = ConfigCapture {
1236 output: default_output(),
1237 };
1238 let model_name = String::from("gpt-4o");
1239 let result = Agent::new()
1240 .prompt("hi")
1241 .model(model_name)
1242 .run(&provider)
1243 .await
1244 .unwrap();
1245 assert_eq!(result.value()["model"], json!("gpt-4o"));
1246 }
1247
1248 #[tokio::test]
1249 async fn into_json_success() {
1250 #[derive(Deserialize, PartialEq, Debug)]
1251 struct Out {
1252 name: String,
1253 }
1254 let provider = TestProvider {
1255 output: AgentOutput {
1256 value: json!({"name": "test"}),
1257 ..default_output()
1258 },
1259 };
1260 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1261 let parsed: Out = result.into_json().unwrap();
1262 assert_eq!(parsed.name, "test");
1263 }
1264
1265 #[tokio::test]
1266 async fn into_json_failure() {
1267 #[derive(Debug, Deserialize)]
1268 #[allow(dead_code)]
1269 struct Out {
1270 name: String,
1271 }
1272 let provider = TestProvider {
1273 output: AgentOutput {
1274 value: json!(42),
1275 ..default_output()
1276 },
1277 };
1278 let result = Agent::new().prompt("test").run(&provider).await.unwrap();
1279 let err = result.into_json::<Out>().unwrap_err();
1280 assert!(matches!(err, OperationError::Deserialize { .. }));
1281 }
1282
1283 #[test]
1284 fn from_output_creates_result() {
1285 let output = AgentOutput {
1286 value: json!("hello"),
1287 ..default_output()
1288 };
1289 let result = AgentResult::from_output(output);
1290 assert_eq!(result.text(), "hello");
1291 assert_eq!(result.cost_usd(), Some(0.05));
1292 }
1293
1294 #[test]
1295 #[should_panic(expected = "budget must be a positive finite number")]
1296 fn max_budget_zero_panics() {
1297 let _ = Agent::new().max_budget_usd(0.0);
1298 }
1299
1300 #[test]
1301 fn model_constant_equality() {
1302 assert_eq!(Model::SONNET, "sonnet");
1303 assert_ne!(Model::SONNET, Model::OPUS);
1304 }
1305
1306 #[test]
1307 fn permission_mode_serialize_deserialize_roundtrip() {
1308 for mode in [
1309 PermissionMode::Default,
1310 PermissionMode::Auto,
1311 PermissionMode::DontAsk,
1312 PermissionMode::BypassPermissions,
1313 ] {
1314 let json = to_string(&mode).unwrap();
1315 let back: PermissionMode = serde_json::from_str(&json).unwrap();
1316 assert_eq!(format!("{:?}", mode), format!("{:?}", back));
1317 }
1318 }
1319
1320 #[test]
1323 fn retry_builder_stores_policy() {
1324 let agent = Agent::new().retry(3);
1325 assert!(agent.retry_policy.is_some());
1326 assert_eq!(agent.retry_policy.unwrap().max_retries(), 3);
1327 }
1328
1329 #[test]
1330 fn retry_policy_builder_stores_custom_policy() {
1331 use crate::retry::RetryPolicy;
1332 let policy = RetryPolicy::new(5).backoff(Duration::from_secs(1));
1333 let agent = Agent::new().retry_policy(policy);
1334 let p = agent.retry_policy.unwrap();
1335 assert_eq!(p.max_retries(), 5);
1336 }
1337
1338 #[test]
1339 fn no_retry_by_default() {
1340 let agent = Agent::new();
1341 assert!(agent.retry_policy.is_none());
1342 }
1343
1344 use std::sync::Arc;
1347 use std::sync::atomic::{AtomicU32, Ordering};
1348 use std::time::Duration;
1349
1350 struct FailNTimesProvider {
1351 fail_count: AtomicU32,
1352 failures_before_success: u32,
1353 output: AgentOutput,
1354 }
1355
1356 impl AgentProvider for FailNTimesProvider {
1357 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1358 Box::pin(async move {
1359 let current = self.fail_count.fetch_add(1, Ordering::SeqCst);
1360 if current < self.failures_before_success {
1361 Err(AgentError::ProcessFailed {
1362 exit_code: 1,
1363 stderr: format!("transient failure #{}", current + 1),
1364 })
1365 } else {
1366 Ok(AgentOutput {
1367 value: self.output.value.clone(),
1368 session_id: self.output.session_id.clone(),
1369 cost_usd: self.output.cost_usd,
1370 input_tokens: self.output.input_tokens,
1371 output_tokens: self.output.output_tokens,
1372 model: self.output.model.clone(),
1373 duration_ms: self.output.duration_ms,
1374 debug_messages: None,
1375 })
1376 }
1377 })
1378 }
1379 }
1380
1381 #[tokio::test]
1382 async fn retry_succeeds_after_transient_failures() {
1383 let provider = FailNTimesProvider {
1384 fail_count: AtomicU32::new(0),
1385 failures_before_success: 2,
1386 output: default_output(),
1387 };
1388 let result = Agent::new()
1389 .prompt("test")
1390 .retry_policy(crate::retry::RetryPolicy::new(3).backoff(Duration::from_millis(1)))
1391 .run(&provider)
1392 .await;
1393
1394 assert!(result.is_ok());
1395 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 3); }
1397
1398 #[tokio::test]
1399 async fn retry_exhausted_returns_last_error() {
1400 let provider = FailNTimesProvider {
1401 fail_count: AtomicU32::new(0),
1402 failures_before_success: 10, output: default_output(),
1404 };
1405 let result = Agent::new()
1406 .prompt("test")
1407 .retry_policy(crate::retry::RetryPolicy::new(2).backoff(Duration::from_millis(1)))
1408 .run(&provider)
1409 .await;
1410
1411 assert!(result.is_err());
1412 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 3);
1414 }
1415
1416 #[tokio::test]
1417 async fn retry_does_not_retry_prompt_too_large() {
1418 let call_count = Arc::new(AtomicU32::new(0));
1419 let count = call_count.clone();
1420
1421 struct CountingNonRetryable {
1422 count: Arc<AtomicU32>,
1423 }
1424 impl AgentProvider for CountingNonRetryable {
1425 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1426 self.count.fetch_add(1, Ordering::SeqCst);
1427 Box::pin(async move {
1428 Err(AgentError::PromptTooLarge {
1429 chars: 1_000_000,
1430 estimated_tokens: 250_000,
1431 model_limit: 200_000,
1432 })
1433 })
1434 }
1435 }
1436
1437 let provider = CountingNonRetryable { count };
1438 let result = Agent::new()
1439 .prompt("test")
1440 .retry_policy(crate::retry::RetryPolicy::new(3).backoff(Duration::from_millis(1)))
1441 .run(&provider)
1442 .await;
1443
1444 assert!(result.is_err());
1445 assert_eq!(call_count.load(Ordering::SeqCst), 1);
1446 }
1447
1448 #[tokio::test]
1449 async fn retry_retries_schema_validation_errors() {
1450 let call_count = Arc::new(AtomicU32::new(0));
1451 let count = call_count.clone();
1452
1453 struct SchemaFailProvider {
1454 count: Arc<AtomicU32>,
1455 }
1456 impl AgentProvider for SchemaFailProvider {
1457 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1458 self.count.fetch_add(1, Ordering::SeqCst);
1459 Box::pin(async move {
1460 Err(AgentError::SchemaValidation {
1461 expected: "object".to_string(),
1462 got: "null".to_string(),
1463 debug_messages: Vec::new(),
1464 partial_usage: Box::default(),
1465 raw_response: None,
1466 })
1467 })
1468 }
1469 }
1470
1471 let provider = SchemaFailProvider { count };
1472 let result = Agent::new()
1473 .prompt("test")
1474 .retry_policy(crate::retry::RetryPolicy::new(2).backoff(Duration::from_millis(1)))
1475 .run(&provider)
1476 .await;
1477
1478 assert!(result.is_err());
1479 assert_eq!(call_count.load(Ordering::SeqCst), 3);
1481 }
1482
1483 #[tokio::test]
1484 async fn schema_validation_succeeds_on_retry() {
1485 let call_count = Arc::new(AtomicU32::new(0));
1486 let count = call_count.clone();
1487
1488 struct SchemaFailThenSucceed {
1489 count: Arc<AtomicU32>,
1490 output: AgentOutput,
1491 }
1492 impl AgentProvider for SchemaFailThenSucceed {
1493 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1494 let current = self.count.fetch_add(1, Ordering::SeqCst);
1495 let output = self.output.clone();
1496 Box::pin(async move {
1497 if current == 0 {
1498 Err(AgentError::SchemaValidation {
1499 expected: "structured_output field".to_string(),
1500 got: "null".to_string(),
1501 debug_messages: Vec::new(),
1502 partial_usage: Box::default(),
1503 raw_response: None,
1504 })
1505 } else {
1506 Ok(output)
1507 }
1508 })
1509 }
1510 }
1511
1512 let provider = SchemaFailThenSucceed {
1513 count,
1514 output: default_output(),
1515 };
1516 let result = Agent::new()
1517 .prompt("test")
1518 .retry_policy(crate::retry::RetryPolicy::new(1).backoff(Duration::from_millis(1)))
1519 .run(&provider)
1520 .await;
1521
1522 assert!(result.is_ok());
1523 assert_eq!(call_count.load(Ordering::SeqCst), 2);
1524 }
1525
1526 #[tokio::test]
1527 async fn auto_retry_applied_when_json_schema_set() {
1528 let call_count = Arc::new(AtomicU32::new(0));
1529 let count = call_count.clone();
1530
1531 struct AlwaysSchemaFail {
1532 count: Arc<AtomicU32>,
1533 }
1534 impl AgentProvider for AlwaysSchemaFail {
1535 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1536 self.count.fetch_add(1, Ordering::SeqCst);
1537 Box::pin(async move {
1538 Err(AgentError::SchemaValidation {
1539 expected: "object".to_string(),
1540 got: "null".to_string(),
1541 debug_messages: Vec::new(),
1542 partial_usage: Box::default(),
1543 raw_response: None,
1544 })
1545 })
1546 }
1547 }
1548
1549 let provider = AlwaysSchemaFail { count };
1550 let result = Agent::new()
1551 .prompt("test")
1552 .output_schema_raw(r#"{"type":"object"}"#)
1553 .run(&provider)
1554 .await;
1555
1556 assert!(result.is_err());
1557 assert_eq!(call_count.load(Ordering::SeqCst), 3);
1559 }
1560
1561 #[tokio::test]
1562 async fn no_retry_without_policy() {
1563 let provider = FailNTimesProvider {
1564 fail_count: AtomicU32::new(0),
1565 failures_before_success: 1,
1566 output: default_output(),
1567 };
1568 let result = Agent::new().prompt("test").run(&provider).await;
1569
1570 assert!(result.is_err());
1571 assert_eq!(provider.fail_count.load(Ordering::SeqCst), 1);
1572 }
1573
1574 use crate::test_support::VecSink;
1577
1578 struct SinkCapture {
1579 output: AgentOutput,
1580 saw_logs: Arc<AtomicU32>,
1581 }
1582
1583 impl AgentProvider for SinkCapture {
1584 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
1585 Box::pin(async {
1586 Ok(AgentOutput {
1587 value: self.output.value.clone(),
1588 session_id: self.output.session_id.clone(),
1589 cost_usd: self.output.cost_usd,
1590 input_tokens: self.output.input_tokens,
1591 output_tokens: self.output.output_tokens,
1592 model: self.output.model.clone(),
1593 duration_ms: self.output.duration_ms,
1594 debug_messages: None,
1595 })
1596 })
1597 }
1598
1599 fn invoke_with_logs<'a>(
1600 &'a self,
1601 config: &'a AgentConfig,
1602 log_sink: Arc<dyn LogSink>,
1603 ) -> InvokeFuture<'a> {
1604 self.saw_logs.fetch_add(1, Ordering::SeqCst);
1605 log_sink.log("stdout", "streaming line");
1606 self.invoke(config)
1607 }
1608 }
1609
1610 #[tokio::test]
1611 async fn log_sink_routes_to_invoke_with_logs() {
1612 let saw_logs = Arc::new(AtomicU32::new(0));
1613 let provider = SinkCapture {
1614 output: default_output(),
1615 saw_logs: saw_logs.clone(),
1616 };
1617 let sink: Arc<dyn LogSink> = VecSink::new();
1618
1619 let result = Agent::new()
1620 .prompt("test")
1621 .log_sink(sink)
1622 .run(&provider)
1623 .await;
1624
1625 assert!(result.is_ok());
1626 assert_eq!(saw_logs.load(Ordering::SeqCst), 1);
1627 }
1628
1629 #[tokio::test]
1630 async fn no_log_sink_routes_to_invoke() {
1631 let saw_logs = Arc::new(AtomicU32::new(0));
1632 let provider = SinkCapture {
1633 output: default_output(),
1634 saw_logs: saw_logs.clone(),
1635 };
1636
1637 let result = Agent::new().prompt("test").run(&provider).await;
1638
1639 assert!(result.is_ok());
1640 assert_eq!(saw_logs.load(Ordering::SeqCst), 0);
1641 }
1642
1643 #[tokio::test]
1644 async fn log_sink_receives_provider_lines() {
1645 let saw_logs = Arc::new(AtomicU32::new(0));
1646 let provider = SinkCapture {
1647 output: default_output(),
1648 saw_logs: saw_logs.clone(),
1649 };
1650 let sink = VecSink::new();
1651
1652 let _ = Agent::new()
1653 .prompt("test")
1654 .log_sink(sink.clone() as Arc<dyn LogSink>)
1655 .run(&provider)
1656 .await;
1657
1658 let lines = sink.0.lock().unwrap();
1659 assert_eq!(lines.len(), 1);
1660 assert_eq!(lines[0].0, "stdout");
1661 assert_eq!(lines[0].1, "streaming line");
1662 }
1663}