use std::future::Future;
use std::sync::Arc;
use std::sync::mpsc::Receiver;
use std::time::Duration;
use crate::claude_agent::{ClaudeAgentRunnerConfig, stream_agent};
use crate::claude_headless::{ClaudeHeadlessRunner, ClaudeHeadlessRunnerConfig, stream_headless};
use crate::claude_terminal::{new_sdk_with_defaults, stream_via_thread};
use crate::sdk::config_loader::load_config;
use crate::sdk::errors::is_non_retryable_error;
use crate::sdk::pi_rpc::load_hardcoded_skills_appendix;
use crate::sdk::{
CancelToken, EffortLevel, LimitProbe, OmpRpcRunner, OmpRpcRunnerOptions, PiRpcRunner,
PiRpcRunnerOptions, PiRunner, PiRunnerOptions, ResolveError, ResolveOptions, ResolvedAgent,
RetryConfig, RunError, SeherTool, StreamChunk, resolve_agent, sdk_supports_tools,
split_model_ref, split_thinking_suffix,
};
#[derive(Default, Clone)]
#[expect(
clippy::type_complexity,
reason = "callback type is intentionally simple"
)]
pub struct RunAgentOptions {
pub working_dir: Option<std::path::PathBuf>,
pub resume: Option<String>,
pub tools: Vec<SeherTool>,
pub api_key: Option<String>,
pub timeout_ms: Option<u64>,
pub system_prompt: Option<String>,
pub cancel: CancelToken,
pub on_retry: Option<Arc<dyn Fn(u32, &str) + Send + Sync>>,
pub effort: Option<EffortLevel>,
}
#[derive(Debug)]
pub struct RunOutput {
pub text: String,
pub session_id: Option<String>,
}
#[derive(Debug, thiserror::Error)]
pub enum ProviderFallbackError {
#[error(transparent)]
Resolve(#[from] ResolveError),
#[error(transparent)]
Run(#[from] RunError),
}
#[derive(Debug)]
pub(crate) enum BackendChoice {
Pi {
provider: String,
model: String,
thinking: Option<String>,
},
Omp {
provider: String,
model: String,
thinking: Option<String>,
},
PiRust {
provider: String,
model: String,
thinking: Option<String>,
},
ClaudeAgent {
model: Option<String>,
effort: Option<EffortLevel>,
},
ClaudeHeadless {
model: Option<String>,
effort: Option<EffortLevel>,
},
ClaudeTerminal {
model: Option<String>,
effort: Option<EffortLevel>,
},
Unsupported { message: String },
}
fn effort_to_ts_pi_thinking(effort: EffortLevel) -> &'static str {
match effort {
EffortLevel::Low => "low",
EffortLevel::Medium => "medium",
EffortLevel::High => "high",
EffortLevel::XHigh => "xhigh",
EffortLevel::Max => "max",
}
}
fn normalize_ts_pi_suffix(thinking: Option<String>) -> Option<String> {
thinking.map(|value| {
match value.trim().to_ascii_lowercase().as_str() {
"off" | "none" | "0" => "off",
"minimal" | "min" => "minimal",
"low" | "1" => "low",
"medium" | "med" | "2" => "medium",
"high" | "3" => "high",
"xhigh" | "4" => "xhigh",
"max" | "5" => "max",
_ => value.as_str(),
}
.to_string()
})
}
fn normalize_pi_rust_suffix(thinking: Option<String>) -> Option<String> {
thinking.map(|value| {
match value.trim().to_ascii_lowercase().as_str() {
"max" | "5" => "xhigh",
_ => value.as_str(),
}
.to_string()
})
}
fn effort_to_pi_rust_thinking(effort: EffortLevel) -> &'static str {
match effort {
EffortLevel::Low => "low",
EffortLevel::Medium => "medium",
EffortLevel::High => "high",
EffortLevel::XHigh | EffortLevel::Max => "xhigh",
}
}
fn effort_from_suffix(suffix: &str) -> Option<EffortLevel> {
match suffix.trim().to_lowercase().as_str() {
"minimal" | "min" | "low" | "1" => Some(EffortLevel::Low),
"medium" | "med" | "2" => Some(EffortLevel::Medium),
"high" | "3" => Some(EffortLevel::High),
"xhigh" | "4" => Some(EffortLevel::XHigh),
"max" => Some(EffortLevel::Max),
_ => None,
}
}
fn claude_family_model_and_effort(
resolved: &ResolvedAgent,
effort: Option<EffortLevel>,
) -> (Option<String>, Option<EffortLevel>) {
let (model_name, suffix_thinking) = split_thinking_suffix(&resolved.model_id);
let effort = effort.or_else(|| suffix_thinking.and_then(effort_from_suffix));
let model = if model_name.is_empty() {
None
} else {
Some(model_name.to_string())
};
(model, effort)
}
fn ambient_api_key_for_provider(
provider: &str,
overrides: Option<&indexmap::IndexMap<String, String>>,
) -> Option<String> {
let provider = provider
.rsplit('/')
.next()
.unwrap_or(provider)
.to_ascii_lowercase();
let vars: &[&str] = match provider.as_str() {
"anthropic" | "claude" => &["ANTHROPIC_API_KEY"],
"openai" | "codex" | "openai-codex" => &["OPENAI_API_KEY"],
"google" | "gemini" | "google-gemini" => &["GOOGLE_API_KEY", "GEMINI_API_KEY"],
"mistral" => &["MISTRAL_API_KEY"],
"cohere" => &["COHERE_API_KEY"],
"groq" => &["GROQ_API_KEY"],
"xai" | "grok" => &["XAI_API_KEY"],
"deepseek" => &["DEEPSEEK_API_KEY"],
"together" => &["TOGETHER_API_KEY"],
"perplexity" => &["PERPLEXITY_API_KEY"],
"cerebras" => &["CEREBRAS_API_KEY"],
"fireworks" => &["FIREWORKS_API_KEY"],
"openrouter" => &["OPENROUTER_API_KEY"],
"huggingface" | "hf" => &["HUGGINGFACE_API_KEY", "HF_TOKEN"],
"ai21" => &["AI21_API_KEY"],
"nvidia" => &["NVIDIA_API_KEY"],
"moonshot" | "kimi" => &["MOONSHOT_API_KEY"],
"minimax" => &["MINIMAX_API_KEY"],
"dashscope" | "qwen" => &["DASHSCOPE_API_KEY"],
_ => return None,
};
vars.iter().find_map(|name| {
overrides
.and_then(|env| env.get(*name).cloned().filter(|value| !value.is_empty()))
.or_else(|| std::env::var(name).ok().filter(|value| !value.is_empty()))
})
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum DispatchError {
ToolsNotSupported { sdk: String },
}
pub(crate) fn choose_backend(
resolved: &ResolvedAgent,
opts: &RunAgentOptions,
) -> Result<BackendChoice, DispatchError> {
let sdk = resolved.sdk.as_str();
if !opts.tools.is_empty() && !sdk_supports_tools(sdk) {
return Err(DispatchError::ToolsNotSupported {
sdk: sdk.to_string(),
});
}
let effort = opts.effort.or(resolved.effort);
Ok(match sdk {
"pi" | "omp" => {
let (provider, model, suffix_thinking) =
split_model_ref(&resolved.provider, &resolved.model_id);
let thinking = effort
.map(|e| effort_to_ts_pi_thinking(e).to_string())
.or_else(|| normalize_ts_pi_suffix(suffix_thinking));
if sdk == "pi" {
BackendChoice::Pi {
provider,
model,
thinking,
}
} else {
BackendChoice::Omp {
provider,
model,
thinking,
}
}
}
"pi-rust" => {
let (provider, model, suffix_thinking) =
split_model_ref(&resolved.provider, &resolved.model_id);
let thinking = effort
.map(|e| effort_to_pi_rust_thinking(e).to_string())
.or_else(|| normalize_pi_rust_suffix(suffix_thinking));
BackendChoice::PiRust {
provider,
model,
thinking,
}
}
"claude" => {
let (model, effort) = claude_family_model_and_effort(resolved, effort);
BackendChoice::ClaudeAgent { model, effort }
}
"claude-headless" => {
let (model, effort) = claude_family_model_and_effort(resolved, effort);
BackendChoice::ClaudeHeadless { model, effort }
}
"claude-terminal" => {
let (model, effort) = claude_family_model_and_effort(resolved, effort);
BackendChoice::ClaudeTerminal { model, effort }
}
other => BackendChoice::Unsupported {
message: format!("unsupported sdk kind: {other}"),
},
})
}
pub(crate) fn fold_stream(rx: &Receiver<StreamChunk>) -> Result<RunOutput, RunError> {
let mut buffered = String::new();
let mut session_id: Option<String> = None;
loop {
match rx.recv() {
Ok(StreamChunk::Delta(d)) => buffered.push_str(&d),
Ok(StreamChunk::Session(id)) => session_id = Some(id),
Ok(StreamChunk::Done(text)) => {
return Ok(RunOutput {
text: if text.is_empty() { buffered } else { text },
session_id,
});
}
Ok(StreamChunk::Limit(error)) => {
return Err(RunError::Limit {
error,
partial: buffered,
});
}
Ok(StreamChunk::Error(msg)) => {
return Err(RunError::Other {
message: msg,
partial: buffered,
});
}
Err(_) => {
return Err(RunError::Other {
message: "seher dispatch channel closed".to_string(),
partial: buffered,
});
}
}
}
}
fn stream_pi_backend(
resolved: &ResolvedAgent,
prompt: String,
opts: RunAgentOptions,
provider: String,
model: String,
thinking: Option<String>,
configured_api_key: Option<String>,
) -> Receiver<StreamChunk> {
let api_key =
configured_api_key.or_else(|| ambient_api_key_for_provider(&provider, Some(&resolved.env)));
let pi_opts = PiRpcRunnerOptions {
provider: Some(provider),
model: Some(model),
thinking,
api_key,
system_prompt: opts.system_prompt,
append_system_prompt: load_hardcoded_skills_appendix(opts.working_dir.as_deref()),
working_directory: opts.working_dir,
env: resolved.env.clone(),
tools: opts.tools,
cancel: opts.cancel,
..Default::default()
};
PiRpcRunner::new(pi_opts).stream(prompt, opts.resume)
}
fn stream_omp_backend(
resolved: &ResolvedAgent,
prompt: String,
opts: RunAgentOptions,
provider: String,
model: String,
thinking: Option<String>,
configured_api_key: Option<String>,
) -> Receiver<StreamChunk> {
let api_key =
configured_api_key.or_else(|| ambient_api_key_for_provider(&provider, Some(&resolved.env)));
let omp_opts = OmpRpcRunnerOptions {
provider: Some(provider),
model: Some(model),
thinking,
api_key,
system_prompt: opts.system_prompt,
append_system_prompt: load_hardcoded_skills_appendix(opts.working_dir.as_deref()),
working_directory: opts.working_dir,
env: resolved.env.clone(),
tools: opts.tools,
cancel: opts.cancel,
..Default::default()
};
OmpRpcRunner::new(omp_opts).stream(prompt, opts.resume)
}
#[must_use]
#[expect(
clippy::too_many_lines,
reason = "backend routing keeps all SDK branches together"
)]
pub fn stream_for_resolved(
resolved: &ResolvedAgent,
prompt: String,
opts: RunAgentOptions,
) -> Receiver<StreamChunk> {
let configured_api_key = opts
.api_key
.clone()
.or_else(|| resolved.api.as_ref().and_then(|a| a.key.clone()));
match choose_backend(resolved, &opts) {
Ok(BackendChoice::Pi {
provider,
model,
thinking,
}) => stream_pi_backend(
resolved,
prompt,
opts,
provider,
model,
thinking,
configured_api_key,
),
Ok(BackendChoice::Omp {
provider,
model,
thinking,
}) => stream_omp_backend(
resolved,
prompt,
opts,
provider,
model,
thinking,
configured_api_key,
),
Ok(BackendChoice::PiRust {
provider,
model,
thinking,
}) => {
let pi_opts = PiRunnerOptions {
provider: Some(provider),
model: Some(model),
thinking,
api_key: configured_api_key,
system_prompt: opts.system_prompt,
working_directory: opts.working_dir,
env: resolved.env.clone(),
tools: opts.tools,
};
PiRunner::new(pi_opts).stream(prompt, opts.resume)
}
Ok(BackendChoice::ClaudeAgent { model, effort }) => {
let config = ClaudeAgentRunnerConfig {
model,
effort,
system_prompt: opts.system_prompt,
cwd: opts
.working_dir
.as_ref()
.map(|p| p.to_string_lossy().into_owned()),
resume_session_id: opts.resume,
tools: opts.tools,
env: resolved
.env
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
..Default::default()
};
stream_agent(config, prompt, resolved.provider.clone())
}
Ok(BackendChoice::ClaudeHeadless { model, effort }) => {
let config = ClaudeHeadlessRunnerConfig {
model,
effort,
system_prompt: opts.system_prompt,
cwd: opts
.working_dir
.as_ref()
.map(|p| p.to_string_lossy().into_owned()),
resume_session_id: opts.resume,
timeout_ms: opts.timeout_ms,
cancel: opts.cancel.clone(),
env: resolved.env.clone(),
..Default::default()
};
stream_headless(
ClaudeHeadlessRunner::new(config),
prompt,
resolved.provider.clone(),
)
}
Ok(BackendChoice::ClaudeTerminal { model, effort }) => {
let sdk = new_sdk_with_defaults(
None,
None,
model,
opts.system_prompt,
effort,
opts.timeout_ms,
opts.working_dir
.as_ref()
.map(|p| p.to_string_lossy().into_owned()),
resolved
.env
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect::<std::collections::HashMap<_, _>>(),
);
stream_via_thread(sdk, prompt, resolved.provider.clone(), opts.resume)
}
Ok(BackendChoice::Unsupported { message }) => {
let (tx, rx) = std::sync::mpsc::channel();
let _ = tx.send(StreamChunk::Error(message));
rx
}
Err(DispatchError::ToolsNotSupported { sdk }) => {
let (tx, rx) = std::sync::mpsc::channel();
let _ = tx.send(StreamChunk::Error(format!(
"sdk '{sdk}' does not support custom tools"
)));
rx
}
}
}
#[expect(
clippy::needless_pass_by_value,
clippy::type_complexity,
reason = "mirrors the public run_for_resolved signature; callback type is intentionally simple"
)]
pub(crate) fn run_with_retry_inner<F, S>(
mut run: F,
prompt: String,
opts: RunAgentOptions,
retry: &RetryConfig,
on_retry: Option<&dyn Fn(u32, &str)>,
mut sleep_fn: S,
) -> Result<RunOutput, RunError>
where
F: FnMut(String, RunAgentOptions) -> Result<RunOutput, RunError>,
S: FnMut(Duration),
{
let mut attempt = 1;
loop {
let result = run(prompt.clone(), opts.clone());
match result {
Ok(output) => return Ok(output),
Err(err) => {
if attempt >= retry.effective_max_attempts() || !retry.enabled {
return Err(err);
}
match &err {
RunError::Other { message, .. } => {
if is_non_retryable_error(message) || !retry.is_retryable_message(message) {
return Err(err);
}
}
RunError::Limit { .. } => {}
RunError::Timeout { .. } => return Err(err),
}
let retry_message = match &err {
RunError::Other { message, .. } => message.clone(),
RunError::Limit { error, .. } => error.to_string(),
RunError::Timeout { error, .. } => error.to_string(),
};
if let Some(cb) = on_retry {
cb(attempt, &retry_message);
}
let delay = retry.delay_for_attempt(attempt);
sleep_fn(delay);
attempt += 1;
}
}
}
}
pub fn run_for_resolved(
resolved: &ResolvedAgent,
prompt: String,
opts: RunAgentOptions,
) -> Result<RunOutput, RunError> {
let on_retry_holder = opts.on_retry.clone();
#[expect(
clippy::type_complexity,
reason = "callback type is intentionally simple"
)]
let on_retry: Option<&dyn Fn(u32, &str)> =
on_retry_holder.as_deref().map(|f| f as &dyn Fn(u32, &str));
run_with_retry_inner(
|p, o| {
let rx = stream_for_resolved(resolved, p, o);
fold_stream(&rx)
},
prompt,
opts,
&resolved.retry,
on_retry,
std::thread::sleep,
)
}
pub async fn run_with_provider_fallback(
resolve_options: ResolveOptions,
probe: &mut dyn LimitProbe,
prompt: String,
opts: RunAgentOptions,
) -> Result<RunOutput, ProviderFallbackError> {
run_with_provider_fallback_inner(
resolve_options,
probe,
prompt,
opts,
|resolved, prompt, opts| {
let resolved = resolved.clone();
async move {
tokio::task::spawn_blocking(move || run_for_resolved(&resolved, prompt, opts))
.await
.map_err(|error| RunError::Other {
message: format!("provider run task failed: {error}"),
partial: String::new(),
})?
}
},
)
.await
}
async fn run_with_provider_fallback_inner<F, Fut>(
resolve_options: ResolveOptions,
probe: &mut dyn LimitProbe,
prompt: String,
opts: RunAgentOptions,
mut run: F,
) -> Result<RunOutput, ProviderFallbackError>
where
F: FnMut(&ResolvedAgent, String, RunAgentOptions) -> Fut,
Fut: Future<Output = Result<RunOutput, RunError>>,
{
let config = match resolve_options.config.clone() {
Some(config) => config,
None => load_config(resolve_options.config_path.as_deref()).map_err(ResolveError::from)?,
};
let mut excluded = resolve_options.exclude_providers.clone();
let mut last_network_error = None;
loop {
let mut options = resolve_options.clone();
options.config = Some(config.clone());
options.exclude_providers = excluded.clone();
let resolved = match resolve_agent(options, probe).await {
Ok(resolved) => resolved,
Err(error @ ResolveError::NoMatching(_)) => {
return Err(last_network_error.map_or_else(
|| ProviderFallbackError::Resolve(error),
ProviderFallbackError::Run,
));
}
Err(error) => return Err(error.into()),
};
match run(&resolved, prompt.clone(), opts.clone()).await {
Ok(output) => return Ok(output),
Err(error) if error.is_network_error() && opts.resume.is_none() => {
excluded.push(resolved.provider.clone());
last_network_error = Some(error);
}
Err(error) => return Err(error.into()),
}
}
}
#[cfg(test)]
#[expect(
clippy::expect_used,
clippy::unwrap_used,
reason = "tests may panic on unexpected fixtures"
)]
mod tests {
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::mpsc::channel;
use super::*;
use crate::sdk::config::{
Config, ModelEntry, ProviderEntry, ResolvedSkillsConfig, RetryConfig,
};
use crate::sdk::errors::{LimitError, NETWORK_ERROR_REASON, TimeoutError};
use crate::sdk::{CancelToken, LimitProbe, ProbeFuture};
fn make_resolved(sdk: &str, provider: &str, model_id: &str) -> ResolvedAgent {
ResolvedAgent {
provider: provider.to_string(),
model_id: model_id.to_string(),
mode_key: "build".to_string(),
sdk: sdk.to_string(),
api: None,
skills: ResolvedSkillsConfig::default(),
retry: RetryConfig::default(),
env: indexmap::IndexMap::new(),
effort: None,
}
}
struct AvailableProbe;
impl LimitProbe for AvailableProbe {
fn probe<'a>(
&'a mut self,
_entry: &'a ProviderEntry,
_resolved: &'a ResolvedAgent,
) -> ProbeFuture<'a> {
Box::pin(async {
Ok::<_, Box<dyn std::error::Error>>(crate::codexbar::AgentLimit::NotLimited)
})
}
}
struct FailingProbe {
provider: String,
}
impl LimitProbe for FailingProbe {
fn probe<'a>(
&'a mut self,
_entry: &'a ProviderEntry,
resolved: &'a ResolvedAgent,
) -> ProbeFuture<'a> {
let should_fail = resolved.provider == self.provider;
Box::pin(async move {
if should_fail {
Err(std::io::Error::other("probe failed").into())
} else {
Ok(crate::codexbar::AgentLimit::NotLimited)
}
})
}
}
struct LimitedProbe;
impl LimitProbe for LimitedProbe {
fn probe<'a>(
&'a mut self,
_entry: &'a ProviderEntry,
_resolved: &'a ResolvedAgent,
) -> ProbeFuture<'a> {
Box::pin(async {
Ok::<_, Box<dyn std::error::Error>>(crate::codexbar::AgentLimit::Limited {
reset_time: None,
})
})
}
}
fn fallback_config(providers: &[(&str, &str)]) -> Config {
Config {
providers: providers
.iter()
.enumerate()
.map(|(order, (provider, sdk))| ProviderEntry {
key: (*provider).to_string(),
order,
provider: (*provider).to_string(),
sdk: (*sdk).to_string(),
priority: None,
api: None,
skills: None,
retry: None,
env: None,
effort: None,
models: indexmap::IndexMap::from([(
"build".to_string(),
ModelEntry {
model: format!("{provider}-model"),
priority: None,
effort: None,
},
)]),
})
.collect(),
..Config::default()
}
}
fn network_error(_message: &str) -> RunError {
RunError::Other {
message: NETWORK_ERROR_REASON.to_string(),
partial: "partial".to_string(),
}
}
fn dummy_tool() -> SeherTool {
SeherTool::new(
"dummy",
"dummy tool",
serde_json::json!({"type": "object", "properties": {}}),
Arc::new(|_| Ok(String::new())),
)
}
fn no_tools_opts() -> RunAgentOptions {
RunAgentOptions::default()
}
fn tools_opts() -> RunAgentOptions {
RunAgentOptions {
tools: vec![dummy_tool()],
..Default::default()
}
}
#[test]
fn choose_backend_pi_extracts_provider_model_and_thinking() {
let resolved = make_resolved("pi", "codex", "openai-codex/gpt-5.5:xhigh");
let choice =
choose_backend(&resolved, &no_tools_opts()).expect("pi backend is always valid");
match choice {
BackendChoice::Pi {
provider,
model,
thinking,
} => {
assert_eq!(provider, "openai-codex");
assert_eq!(model, "gpt-5.5");
assert_eq!(thinking, Some("xhigh".to_string()));
}
other => panic!("expected Pi, got {other:?}"),
}
}
#[test]
fn choose_backend_omp_extracts_thinking_and_explicit_effort_wins() {
let resolved = make_resolved("omp", "codex", "openai-codex/gpt-5.6:low");
let suffix_choice =
choose_backend(&resolved, &no_tools_opts()).expect("omp backend is valid");
match suffix_choice {
BackendChoice::Omp {
provider,
model,
thinking,
} => {
assert_eq!(provider, "openai-codex");
assert_eq!(model, "gpt-5.6");
assert_eq!(thinking, Some("low".to_string()));
}
other => panic!("expected Omp, got {other:?}"),
}
let choice = choose_backend(
&resolved,
&RunAgentOptions {
effort: Some(EffortLevel::Max),
..Default::default()
},
)
.expect("omp backend is valid");
match choice {
BackendChoice::Omp { thinking, .. } => {
assert_eq!(thinking, Some("max".to_string()));
}
other => panic!("expected Omp, got {other:?}"),
}
}
#[test]
fn pi_rpc_options_preserve_user_prompt_and_skills_appendix() {
let options = PiRpcRunnerOptions {
provider: Some("openai-codex".into()),
model: Some("gpt-5.5".into()),
system_prompt: Some("user prompt".into()),
append_system_prompt: Some("formatted skills".into()),
api_key: Some("configured-key".into()),
..Default::default()
};
assert_eq!(options.system_prompt.as_deref(), Some("user prompt"));
assert_eq!(
options.append_system_prompt.as_deref(),
Some("formatted skills")
);
assert_eq!(options.api_key.as_deref(), Some("configured-key"));
}
#[test]
fn choose_backend_pi_rust_routes_to_in_process_backend() {
let resolved = make_resolved("pi-rust", "anthropic", "claude-sonnet-4-5:high");
let choice = choose_backend(&resolved, &no_tools_opts()).expect("pi-rust backend is valid");
match choice {
BackendChoice::PiRust {
provider,
model,
thinking,
} => {
assert_eq!(provider, "anthropic");
assert_eq!(model, "claude-sonnet-4-5");
assert_eq!(thinking, Some("high".to_string()));
}
other => panic!("expected PiRust, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_sets_bare_model_name() {
let resolved = make_resolved("claude", "claude", "sonnet");
let choice =
choose_backend(&resolved, &no_tools_opts()).expect("claude backend is always valid");
match choice {
BackendChoice::ClaudeAgent { model, effort: _ } => {
assert_eq!(model, Some("sonnet".to_string()));
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_strips_thinking_suffix_from_model() {
let resolved = make_resolved("claude", "claude", "sonnet:high");
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("claude with thinking suffix is valid");
match choice {
BackendChoice::ClaudeAgent { model, effort } => {
assert_eq!(model, Some("sonnet".to_string()));
assert_eq!(effort, Some(EffortLevel::High));
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_explicit_effort_overrides_suffix_thinking() {
let resolved = make_resolved("claude", "claude", "sonnet:low");
let opts = RunAgentOptions {
effort: Some(EffortLevel::Max),
..Default::default()
};
let choice = choose_backend(&resolved, &opts).expect("claude with effort is valid");
match choice {
BackendChoice::ClaudeAgent { effort, .. } => {
assert_eq!(effort, Some(EffortLevel::Max));
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_off_suffix_has_no_effort_equivalent() {
let resolved = make_resolved("claude", "claude", "sonnet:off");
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("claude with off suffix is still valid");
match choice {
BackendChoice::ClaudeAgent { model, effort } => {
assert_eq!(model, Some("sonnet".to_string()));
assert_eq!(effort, None);
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_med_alias_maps_to_effort_medium() {
let resolved = make_resolved("claude", "claude", "sonnet:med");
let choice =
choose_backend(&resolved, &no_tools_opts()).expect("claude with med alias is valid");
match choice {
BackendChoice::ClaudeAgent { model, effort } => {
assert_eq!(model, Some("sonnet".to_string()));
assert_eq!(effort, Some(EffortLevel::Medium));
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_max_suffix_maps_to_effort_max() {
let resolved = make_resolved("claude", "claude", "sonnet:max");
let choice =
choose_backend(&resolved, &no_tools_opts()).expect("claude with max suffix is valid");
match choice {
BackendChoice::ClaudeAgent { model, effort } => {
assert_eq!(model, Some("sonnet".to_string()));
assert_eq!(effort, Some(EffortLevel::Max));
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_headless_uses_suffix_thinking_as_effort() {
let resolved = make_resolved("claude-headless", "claude", "sonnet:xhigh");
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("claude-headless with thinking suffix is valid");
match choice {
BackendChoice::ClaudeHeadless { effort, .. } => {
assert_eq!(effort, Some(EffortLevel::XHigh));
}
other => panic!("expected ClaudeHeadless, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_terminal_uses_suffix_thinking_as_effort() {
let resolved = make_resolved("claude-terminal", "claude", "sonnet:medium");
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("claude-terminal with thinking suffix is valid");
match choice {
BackendChoice::ClaudeTerminal { effort, .. } => {
assert_eq!(effort, Some(EffortLevel::Medium));
}
other => panic!("expected ClaudeTerminal, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_agent_carries_resolved_effort() {
let mut resolved = make_resolved("claude", "claude", "sonnet");
resolved.effort = Some(EffortLevel::High);
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("claude with resolved effort is valid");
match choice {
BackendChoice::ClaudeAgent { effort, .. } => {
assert_eq!(effort, Some(EffortLevel::High));
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_explicit_effort_overrides_resolved_effort() {
let mut resolved = make_resolved("claude", "claude", "sonnet");
resolved.effort = Some(EffortLevel::Low);
let opts = RunAgentOptions {
effort: Some(EffortLevel::Max),
..Default::default()
};
let choice = choose_backend(&resolved, &opts).expect("claude with effort is valid");
match choice {
BackendChoice::ClaudeAgent { effort, .. } => {
assert_eq!(effort, Some(EffortLevel::Max));
}
other => panic!("expected ClaudeAgent, got {other:?}"),
}
}
#[test]
fn choose_backend_pi_explicit_effort_overrides_suffix_thinking() {
let resolved = make_resolved("pi", "codex", "openai-codex/gpt-5.5:low");
let opts = RunAgentOptions {
effort: Some(EffortLevel::Max),
..Default::default()
};
let choice = choose_backend(&resolved, &opts).expect("pi backend is always valid");
match choice {
BackendChoice::Pi { thinking, .. } => {
assert_eq!(thinking, Some("max".to_string()));
}
other => panic!("expected Pi, got {other:?}"),
}
}
#[test]
fn choose_backend_pi_rust_explicit_max_retains_xhigh_compatibility() {
let resolved = make_resolved("pi-rust", "codex", "openai-codex/gpt-5.5:low");
let opts = RunAgentOptions {
effort: Some(EffortLevel::Max),
..Default::default()
};
let choice = choose_backend(&resolved, &opts).expect("pi-rust backend is valid");
match choice {
BackendChoice::PiRust { thinking, .. } => {
assert_eq!(thinking, Some("xhigh".to_string()));
}
other => panic!("expected PiRust, got {other:?}"),
}
}
#[test]
fn choose_backend_pi_preserves_max_suffix() {
let resolved = make_resolved("pi", "codex", "openai-codex/gpt-5.5:max");
let choice =
choose_backend(&resolved, &no_tools_opts()).expect("pi backend is always valid");
match choice {
BackendChoice::Pi {
provider,
model,
thinking,
} => {
assert_eq!(provider, "openai-codex");
assert_eq!(model, "gpt-5.5");
assert_eq!(thinking, Some("max".to_string()));
}
other => panic!("expected Pi, got {other:?}"),
}
}
#[test]
fn choose_backend_claude_headless_routes_when_no_tools() {
let resolved = make_resolved("claude-headless", "claude", "sonnet");
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("claude-headless with no tools is valid");
assert!(matches!(choice, BackendChoice::ClaudeHeadless { .. }));
}
#[test]
fn choose_backend_claude_headless_errors_when_tools_non_empty() {
let resolved = make_resolved("claude-headless", "claude", "sonnet");
let err = choose_backend(&resolved, &tools_opts())
.expect_err("claude-headless does not support tools");
assert_eq!(
err,
DispatchError::ToolsNotSupported {
sdk: "claude-headless".to_string()
}
);
}
#[test]
fn choose_backend_claude_terminal_routes_when_no_tools() {
let resolved = make_resolved("claude-terminal", "claude", "sonnet");
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("claude-terminal with no tools is valid");
assert!(matches!(choice, BackendChoice::ClaudeTerminal { .. }));
}
#[test]
fn choose_backend_claude_terminal_errors_when_tools_non_empty() {
let resolved = make_resolved("claude-terminal", "claude", "sonnet");
let err = choose_backend(&resolved, &tools_opts())
.expect_err("claude-terminal does not support tools");
assert_eq!(
err,
DispatchError::ToolsNotSupported {
sdk: "claude-terminal".to_string()
}
);
}
#[test]
fn choose_backend_unknown_sdk_returns_unsupported_variant() {
let resolved = make_resolved("unknown-foo", "unknown-foo", "some-model");
let choice = choose_backend(&resolved, &no_tools_opts())
.expect("unknown sdk returns Unsupported variant, not Err");
assert!(matches!(choice, BackendChoice::Unsupported { .. }));
}
#[test]
fn stream_for_resolved_unknown_sdk_sends_error_chunk_then_closes() {
let resolved = make_resolved("unknown-foo", "unknown-foo", "some-model");
let rx = stream_for_resolved(&resolved, "prompt".to_string(), no_tools_opts());
let chunk = rx.recv().expect("at least one chunk on unknown sdk");
assert!(
matches!(chunk, StreamChunk::Error(_)),
"expected Error chunk, got {chunk:?}"
);
assert!(
rx.recv().is_err(),
"channel must close after the error chunk"
);
}
#[test]
fn stream_for_resolved_tools_on_headless_sdk_sends_error_chunk() {
let resolved = make_resolved("claude-headless", "claude", "sonnet");
let rx = stream_for_resolved(&resolved, "prompt".to_string(), tools_opts());
let chunk = rx.recv().expect("error chunk for tools-on-headless");
assert!(
matches!(chunk, StreamChunk::Error(_)),
"expected Error chunk, got {chunk:?}"
);
assert!(
rx.recv().is_err(),
"channel must close after the error chunk"
);
}
#[test]
fn fold_stream_accumulates_deltas_and_returns_on_done() {
let (tx, rx) = channel();
tx.send(StreamChunk::Session("s1".to_string()))
.expect("send Session");
tx.send(StreamChunk::Delta("hi ".to_string()))
.expect("send Delta 1");
tx.send(StreamChunk::Delta("there".to_string()))
.expect("send Delta 2");
tx.send(StreamChunk::Done(String::new()))
.expect("send Done");
drop(tx);
let out = fold_stream(&rx).expect("fold succeeds");
assert_eq!(out.text, "hi there");
assert_eq!(out.session_id, Some("s1".to_string()));
}
#[test]
fn fold_stream_done_with_non_empty_text_takes_precedence_over_buffered() {
let (tx, rx) = channel();
tx.send(StreamChunk::Delta("ignored delta".to_string()))
.expect("send Delta");
tx.send(StreamChunk::Done("full text".to_string()))
.expect("send Done");
drop(tx);
let out = fold_stream(&rx).expect("fold succeeds");
assert_eq!(out.text, "full text");
}
#[test]
fn fold_stream_limit_returns_err_with_partial_text() {
let (tx, rx) = channel();
tx.send(StreamChunk::Delta("partial".to_string()))
.expect("send Delta");
tx.send(StreamChunk::Limit(LimitError {
provider: "claude".to_string(),
reset_at: None,
}))
.expect("send Limit");
drop(tx);
let err = fold_stream(&rx).expect_err("fold returns Err on Limit");
assert!(
matches!(err, RunError::Limit { ref partial, .. } if partial == "partial"),
"expected Limit with partial='partial', got {err:?}"
);
}
#[test]
fn fold_stream_channel_close_without_done_returns_err() {
let (tx, rx) = channel();
tx.send(StreamChunk::Delta("x".to_string()))
.expect("send Delta");
drop(tx);
let err = fold_stream(&rx).expect_err("fold returns Err on channel close");
match &err {
RunError::Other { message, partial } => {
assert_eq!(partial, "x");
assert!(
message.contains("channel") || message.contains("closed"),
"message should mention channel closure, got: {message}"
);
}
other => panic!("expected RunError::Other, got {other:?}"),
}
}
#[test]
fn fold_stream_maps_only_exact_network_marker_to_network_error() {
let (tx, rx) = channel();
tx.send(StreamChunk::Delta("partial".to_string()))
.expect("send Delta");
tx.send(StreamChunk::Error(NETWORK_ERROR_REASON.to_string()))
.expect("send network error");
drop(tx);
let err = fold_stream(&rx).expect_err("network marker should fail");
assert!(matches!(
err,
RunError::Other { ref message, ref partial }
if message == NETWORK_ERROR_REASON && partial == "partial"
));
let (tx, rx) = channel();
tx.send(StreamChunk::Error("network_error: timeout".to_string()))
.expect("send ordinary error");
drop(tx);
let err = fold_stream(&rx).expect_err("ordinary error should fail");
assert!(
matches!(err, RunError::Other { message, .. } if message == "network_error: timeout")
);
}
#[test]
fn provider_fallback_network_error_moves_from_a_to_b() {
let config = fallback_config(&[("a", "pi"), ("b", "pi")]);
let resolve_options = ResolveOptions {
config: Some(config),
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let mut calls = Vec::new();
let result = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|resolved, _, _| {
calls.push(resolved.provider.clone());
let result = if resolved.provider == "a" {
Err(network_error("a unavailable"))
} else {
Ok(RunOutput {
text: "ok".to_string(),
session_id: None,
})
};
async move { result }
},
))
.expect("fallback should succeed");
assert_eq!(result.text, "ok");
assert_eq!(calls, vec!["a".to_string(), "b".to_string()]);
}
#[tokio::test(flavor = "current_thread")]
async fn public_provider_fallback_uses_the_real_runner_wrapper() {
let resolve_options = ResolveOptions {
config: Some(fallback_config(&[("a", "claude-headless")])),
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let cancel = CancelToken::new();
cancel.cancel();
let result = run_with_provider_fallback(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions {
cancel,
..Default::default()
},
)
.await;
assert!(matches!(result, Err(ProviderFallbackError::Run(_))));
}
#[test]
fn provider_fallback_honors_initial_exclusions() {
let config = fallback_config(&[("a", "pi"), ("b", "pi")]);
let resolve_options = ResolveOptions {
config: Some(config),
exclude_providers: vec!["a".to_string()],
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let mut calls = Vec::new();
let result = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|resolved, _, _| {
calls.push(resolved.provider.clone());
let result = Ok(RunOutput {
text: "ok".to_string(),
session_id: None,
});
async move { result }
},
))
.expect("eligible provider should run");
assert_eq!(result.text, "ok");
assert_eq!(calls, vec!["b".to_string()]);
}
#[test]
fn provider_fallback_returns_last_network_error_after_finite_exhaustion() {
let config = fallback_config(&[("a", "pi"), ("b", "pi")]);
let resolve_options = ResolveOptions {
config: Some(config),
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let mut calls = Vec::new();
let error = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|resolved, _, _| {
calls.push(resolved.provider.clone());
let result = Err(network_error(&format!("{} unavailable", resolved.provider)));
async move { result }
},
))
.expect_err("all network failures should be returned");
assert_eq!(calls, vec!["a".to_string(), "b".to_string()]);
match error {
ProviderFallbackError::Run(RunError::Other { message, .. }) => {
assert_eq!(message, NETWORK_ERROR_REASON);
}
other => panic!("expected last network error, got {other:?}"),
}
}
#[test]
fn provider_fallback_keeps_resumed_session_pinned() {
let resolve_options = ResolveOptions {
config: Some(fallback_config(&[("a", "pi"), ("b", "pi")])),
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let mut calls = Vec::new();
let opts = RunAgentOptions {
resume: Some("session-a".to_string()),
..Default::default()
};
let error = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
opts,
|resolved, prompt, opts| {
assert_eq!(prompt, "prompt");
assert_eq!(opts.resume.as_deref(), Some("session-a"));
calls.push(resolved.provider.clone());
let result = Err(network_error("network_error"));
async move { result }
},
))
.expect_err("resumed network failure should stay terminal");
assert_eq!(calls, vec!["a".to_string()]);
assert!(matches!(
error,
ProviderFallbackError::Run(RunError::Other { message, .. })
if message == NETWORK_ERROR_REASON
));
}
#[test]
fn provider_fallback_keeps_timeout_and_cancellation_terminal() {
for failure in [
RunError::Timeout {
error: TimeoutError {
ms: 1,
label: "test",
},
partial: "partial".to_string(),
},
RunError::Other {
message: "pi session cancelled before worker startup".to_string(),
partial: "partial".to_string(),
},
] {
let resolve_options = ResolveOptions {
config: Some(fallback_config(&[("a", "pi"), ("b", "pi")])),
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let mut calls = Vec::new();
let mut failure = Some(failure);
let error = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|resolved, _, _| {
calls.push(resolved.provider.clone());
let result = Err(failure.take().expect("runner is called once"));
async move { result }
},
))
.expect_err("terminal errors must stay on the selected provider");
assert_eq!(calls, vec!["a".to_string()]);
match error {
ProviderFallbackError::Run(RunError::Timeout { .. } | RunError::Other { .. }) => {}
other => panic!("expected terminal run error, got {other:?}"),
}
}
}
#[test]
fn provider_fallback_surfaces_resolver_rate_limits() {
let resolve_options = ResolveOptions {
config: Some(fallback_config(&[("a", "pi"), ("b", "pi")])),
no_wait: true,
..Default::default()
};
let mut probe = LimitedProbe;
let error = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|_, _, _| async {
Err(RunError::Other {
message: "runner must not be called".to_string(),
partial: String::new(),
})
},
))
.expect_err("resolver rate limits must remain resolution errors");
assert!(matches!(
error,
ProviderFallbackError::Resolve(ResolveError::AllLimited(_))
));
}
#[test]
fn provider_fallback_surfaces_config_errors_as_resolution_errors() {
let resolve_options = ResolveOptions {
config_path: Some(PathBuf::from(
"/definitely-missing-seher-config/config.yaml",
)),
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let error = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|_, _, _| async {
Err(RunError::Other {
message: "runner must not be called".to_string(),
partial: String::new(),
})
},
))
.expect_err("config errors must remain resolution errors");
assert!(matches!(
error,
ProviderFallbackError::Resolve(ResolveError::Config(_))
));
}
#[test]
fn provider_fallback_returns_original_error_when_next_resolution_has_no_match() {
let resolve_options = ResolveOptions {
config: Some(fallback_config(&[("a", "pi"), ("b", "pi")])),
no_wait: true,
..Default::default()
};
let mut probe = FailingProbe {
provider: "b".to_string(),
};
let mut calls = Vec::new();
let error = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|resolved, _, _| {
calls.push(resolved.provider.clone());
let result = Err(network_error("a network failure"));
async move { result }
},
))
.expect_err("the original network error should survive re-resolution failure");
assert_eq!(calls, vec!["a".to_string()]);
match error {
ProviderFallbackError::Run(RunError::Other { message, .. }) => {
assert_eq!(message, NETWORK_ERROR_REASON);
}
other => panic!("expected original network error, got {other:?}"),
}
}
#[test]
fn provider_fallback_preserves_non_network_and_rate_limit_errors() {
for failure in [
RunError::Other {
message: "HTTP 400".to_string(),
partial: "partial".to_string(),
},
RunError::Limit {
error: LimitError {
provider: "a".to_string(),
reset_at: None,
},
partial: "partial".to_string(),
},
] {
let config = fallback_config(&[("a", "pi"), ("b", "pi")]);
let resolve_options = ResolveOptions {
config: Some(config),
no_wait: true,
..Default::default()
};
let mut probe = AvailableProbe;
let mut calls = Vec::new();
let mut failure = Some(failure);
let error = futures::executor::block_on(run_with_provider_fallback_inner(
resolve_options,
&mut probe,
"prompt".to_string(),
RunAgentOptions::default(),
|resolved, _, _| {
calls.push(resolved.provider.clone());
let result = Err(failure.take().expect("runner is called once"));
async move { result }
},
))
.expect_err("non-network errors must stay on the selected provider");
assert_eq!(calls, vec!["a".to_string()]);
assert!(matches!(error, ProviderFallbackError::Run(_)));
}
}
fn retry_config_with_client_errors() -> RetryConfig {
RetryConfig {
retry_client_errors: true,
..RetryConfig::default()
}
}
fn other_error(message: &str) -> RunError {
RunError::Other {
message: message.to_string(),
partial: String::new(),
}
}
fn limit_error() -> RunError {
RunError::Limit {
error: LimitError {
provider: "claude".to_string(),
reset_at: None,
},
partial: String::new(),
}
}
fn timeout_error() -> RunError {
RunError::Timeout {
error: TimeoutError {
ms: 1000,
label: "test",
},
partial: String::new(),
}
}
#[test]
fn retry_succeeds_on_first_attempt() {
let run = |_prompt: String, _opts: RunAgentOptions| {
Ok(RunOutput {
text: "hello".to_string(),
session_id: None,
})
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert_eq!(result.unwrap().text, "hello");
}
#[test]
fn retry_limit_error_then_success() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
if calls == 1 {
Err(limit_error())
} else {
Ok(RunOutput {
text: "ok".to_string(),
session_id: None,
})
}
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert_eq!(result.unwrap().text, "ok");
assert_eq!(calls, 2);
}
#[test]
fn retry_401_with_client_errors_enabled_then_success() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
if calls == 1 {
Err(other_error("Anthropic API error (HTTP 401): auth_error"))
} else {
Ok(RunOutput {
text: "ok".to_string(),
session_id: None,
})
}
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&retry_config_with_client_errors(),
None,
|_| {},
);
assert_eq!(result.unwrap().text, "ok");
assert_eq!(calls, 2);
}
#[test]
fn pi_process_crash_marker_never_retries_retryable_stderr() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(other_error(
"Pi RPC non-retryable: Pi RPC process exited while prompting: HTTP 500",
))
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert!(result.is_err());
assert_eq!(calls, 1);
}
#[test]
fn retry_401_without_client_errors_fails_immediately() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(other_error("Anthropic API error (HTTP 401): auth_error"))
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert!(result.is_err());
assert_eq!(calls, 1);
}
#[test]
fn retry_404_with_client_errors_enabled_then_success() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
if calls == 1 {
Err(other_error("Anthropic API error (HTTP 404): not found"))
} else {
Ok(RunOutput {
text: "ok".to_string(),
session_id: None,
})
}
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&retry_config_with_client_errors(),
None,
|_| {},
);
assert_eq!(result.unwrap().text, "ok");
assert_eq!(calls, 2);
}
#[test]
fn retry_404_without_client_errors_fails_immediately() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(other_error("Anthropic API error (HTTP 404): not found"))
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert!(result.is_err());
assert_eq!(calls, 1);
}
#[test]
fn retry_500_always_retries_regardless_of_client_error_flag() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
if calls == 1 {
Err(other_error("Anthropic API error (HTTP 500): internal"))
} else {
Ok(RunOutput {
text: "ok".to_string(),
session_id: None,
})
}
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert_eq!(result.unwrap().text, "ok");
assert_eq!(calls, 2);
}
#[test]
fn retry_403_fails_immediately() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(other_error("Anthropic API error (HTTP 403): forbidden"))
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&retry_config_with_client_errors(),
None,
|_| {},
);
assert!(result.is_err());
assert_eq!(calls, 1);
}
#[test]
fn retry_disabled_fails_immediately_on_transient_error() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(other_error("Anthropic API error (HTTP 429): rate limited"))
};
let cfg = RetryConfig {
enabled: false,
..RetryConfig::default()
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&cfg,
None,
|_| {},
);
assert!(result.is_err());
assert_eq!(calls, 1);
}
#[test]
fn retry_exhausted_returns_last_error() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(other_error(&format!("HTTP 500 attempt {calls}")))
};
let cfg = RetryConfig {
max_attempts: 3,
..RetryConfig::default()
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&cfg,
None,
|_| {},
);
match result {
Err(RunError::Other { message, .. }) => {
assert!(message.contains("attempt 3"), "got: {message}");
}
other => panic!("expected attempt 3 error, got {other:?}"),
}
assert_eq!(calls, 3);
}
#[test]
fn retry_timeout_fails_immediately() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(timeout_error())
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert!(result.is_err());
assert_eq!(calls, 1);
}
#[test]
fn retry_network_error_fails_immediately() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
Err(network_error("connection reset"))
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
None,
|_| {},
);
assert!(matches!(
result,
Err(RunError::Other { message, .. }) if message == NETWORK_ERROR_REASON
));
assert_eq!(calls, 1);
}
#[test]
fn retry_delay_doubles_each_attempt_until_max() {
let run = |_prompt: String, _opts: RunAgentOptions| Err(other_error("HTTP 500"));
let mut sleeps: Vec<u64> = Vec::new();
let cfg = RetryConfig {
max_attempts: 6,
..RetryConfig::default()
};
let _ = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&cfg,
None,
|d| sleeps.push(d.as_secs()),
);
assert_eq!(sleeps, vec![2, 4, 8, 16, 32]);
}
#[test]
fn retry_delay_clamps_at_max_delay_secs() {
let run = |_prompt: String, _opts: RunAgentOptions| Err(other_error("HTTP 500"));
let mut sleeps: Vec<u64> = Vec::new();
let cfg = RetryConfig {
max_attempts: 5,
initial_delay_secs: 10,
max_delay_secs: 15,
multiplier: 2.0,
..RetryConfig::default()
};
let _ = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&cfg,
None,
|d| sleeps.push(d.as_secs()),
);
assert_eq!(sleeps, vec![10, 15, 15, 15]);
}
#[test]
fn retry_invokes_on_callback_once_per_retry() {
let mut calls = 0;
let run = |_prompt: String, _opts: RunAgentOptions| {
calls += 1;
if calls == 1 {
Err(other_error("HTTP 500"))
} else {
Ok(RunOutput {
text: "ok".to_string(),
session_id: None,
})
}
};
let callback_attempts = std::cell::RefCell::new(Vec::new());
let cb = |attempt: u32, _msg: &str| {
callback_attempts.borrow_mut().push(attempt);
};
let result = run_with_retry_inner(
run,
"prompt".to_string(),
RunAgentOptions::default(),
&RetryConfig::default(),
Some(&cb),
|_| {},
);
assert_eq!(result.unwrap().text, "ok");
assert_eq!(*callback_attempts.borrow(), vec![1]);
}
}