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::{
CancelToken, PiRunner, PiRunnerOptions, ResolvedAgent, RetryConfig, RunError, SeherTool,
StreamChunk, 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>>,
}
#[derive(Debug)]
pub struct RunOutput {
pub text: String,
pub session_id: Option<String>,
}
#[derive(Debug)]
pub(crate) enum BackendChoice {
Pi {
provider: String,
model: String,
thinking: Option<String>,
},
ClaudeAgent {
model: Option<String>,
},
ClaudeHeadless {
model: Option<String>,
},
ClaudeTerminal {
model: Option<String>,
},
Unsupported {
message: String,
},
}
#[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(),
});
}
Ok(match sdk {
"pi" => {
let (provider, model, thinking) =
split_model_ref(&resolved.provider, &resolved.model_id);
BackendChoice::Pi {
provider,
model,
thinking,
}
}
"claude" => {
let (model_name, _) = split_thinking_suffix(&resolved.model_id);
BackendChoice::ClaudeAgent {
model: if model_name.is_empty() {
None
} else {
Some(model_name.to_string())
},
}
}
"claude-headless" => {
let (model_name, _) = split_thinking_suffix(&resolved.model_id);
BackendChoice::ClaudeHeadless {
model: if model_name.is_empty() {
None
} else {
Some(model_name.to_string())
},
}
}
"claude-terminal" => {
let (model_name, _) = split_thinking_suffix(&resolved.model_id);
BackendChoice::ClaudeTerminal {
model: if model_name.is_empty() {
None
} else {
Some(model_name.to_string())
},
}
}
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,
});
}
}
}
}
#[must_use]
pub fn stream_for_resolved(
resolved: &ResolvedAgent,
prompt: String,
opts: RunAgentOptions,
) -> Receiver<StreamChunk> {
let api_key = opts
.api_key
.clone()
.or_else(|| resolved.api.as_ref().and_then(|a| a.key.clone()));
match choose_backend(resolved, &opts) {
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
}
Ok(BackendChoice::Pi {
provider,
model,
thinking,
}) => {
let pi_opts = PiRunnerOptions {
provider: Some(provider),
model: Some(model),
thinking,
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 }) => {
let config = ClaudeAgentRunnerConfig {
model,
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 }) => {
let config = ClaudeHeadlessRunnerConfig {
model,
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 }) => {
let sdk = new_sdk_with_defaults(
None,
None,
model,
opts.system_prompt,
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
}
}
}
#[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::Timeout { .. } => return Err(err),
RunError::Limit { .. } => {}
RunError::Other { message, .. } => {
if !retry.is_retryable_message(message) {
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,
)
}
#[cfg(test)]
#[expect(
clippy::expect_used,
clippy::unwrap_used,
reason = "tests may panic on unexpected fixtures"
)]
mod tests {
use std::sync::Arc;
use std::sync::mpsc::channel;
use super::*;
use crate::sdk::config::{ResolvedSkillsConfig, RetryConfig};
use crate::sdk::errors::{LimitError, TimeoutError};
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(),
}
}
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_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 } => {
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 } => {
assert_eq!(model, Some("sonnet".to_string()));
}
other => panic!("expected ClaudeAgent, 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:?}"),
}
}
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 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_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]);
}
}