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