use anyhow::Result;
use serde::Serialize;
use std::collections::VecDeque;
use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use std::time::Instant;
use tokio::sync::{mpsc, oneshot};
use crate::api::types::ImageSource;
#[cfg(test)]
use crate::api::ProviderStream;
use crate::api::{ApiEvent, ApiFailure, ApiFailureKind, ContentBlock, Message, Provider};
use crate::checkpoint::{PendingCheckpoint, TurnCheckpoint};
use crate::compact::{self};
use crate::config::{HookTrigger, ModelBinding};
use crate::cost::CostTracker;
use crate::permissions::{PermissionChecker, PermissionResponse, PermissionResult};
use crate::plugin::PluginRegistry;
use crate::tools::ToolRegistry;
pub type SteeringQueue = Arc<Mutex<VecDeque<String>>>;
const MAX_PARALLEL_TOOLS: usize = 8;
pub struct Engine {
provider: Box<dyn Provider>,
tools: ToolRegistry,
permissions: PermissionChecker,
messages: Vec<Message>,
system_prompt: String,
model: String,
model_binding: Option<ModelBinding>,
max_tokens: u32,
context_window: usize,
auto_compact_threshold: f64,
steering: SteeringQueue,
pending_images: Vec<ImageSource>,
plugins: Option<Arc<PluginRegistry>>,
checkpoint_enabled: bool,
pending_checkpoint: Option<PendingCheckpoint>,
last_checkpoint: Option<TurnCheckpoint>,
tool_trace: Vec<ToolTraceEntry>,
model_trace: Vec<ModelTraceEntry>,
trace_started_at: Option<Instant>,
trace_duration_ms: Option<u64>,
transcript_checkpoint: Option<PathBuf>,
pub cost: CostTracker,
last_request_usage: Option<RequestUsageBaseline>,
last_compaction_notice: Option<String>,
last_failure: Option<FailureRecord>,
retry_backoff_base: std::time::Duration,
}
struct RequestUsageBaseline {
prompt_tokens: usize,
message_count: usize,
}
#[derive(Clone, Debug, Serialize)]
pub struct ToolTraceEntry {
pub id: String,
pub name: String,
pub input: serde_json::Value,
pub output: String,
pub is_error: bool,
pub read_only: bool,
pub started_after_ms: u64,
pub duration_ms: u64,
}
#[derive(Clone, Debug, Serialize)]
pub struct ModelTraceEntry {
pub index: usize,
pub started_after_ms: u64,
pub duration_ms: u64,
pub status: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub failure: Option<ApiFailureKind>,
#[serde(skip_serializing_if = "Option::is_none")]
pub usage: Option<ModelRoundUsage>,
}
#[derive(Clone, Debug, Serialize, PartialEq)]
pub struct FailureRecord {
pub kind: ApiFailureKind,
pub retryable: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub http_status: Option<u16>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retry_after_ms: Option<u64>,
pub attempts: u32,
}
impl FailureRecord {
pub fn cancelled(attempts: u32) -> Self {
Self {
kind: ApiFailureKind::Cancelled,
retryable: false,
http_status: None,
retry_after_ms: None,
attempts,
}
}
pub fn unclassified() -> Self {
Self {
kind: ApiFailureKind::Other,
retryable: false,
http_status: None,
retry_after_ms: None,
attempts: 1,
}
}
fn from_failure(failure: &ApiFailure, attempts: u32) -> Self {
Self {
kind: failure.kind,
retryable: failure.kind.retryable(),
http_status: failure.http_status,
retry_after_ms: failure
.retry_after
.map(|duration| duration.as_millis() as u64),
attempts,
}
}
}
#[derive(Clone, Debug, Serialize)]
pub struct ModelRoundUsage {
pub input_tokens: u32,
pub output_tokens: u32,
pub cache_read_tokens: u32,
pub cache_creation_tokens: u32,
pub cost_usd: Option<f64>,
}
#[derive(Clone, Debug, Serialize)]
pub struct ExecutionTiming {
pub total_duration_ms: u64,
pub model_rounds: Vec<ModelTraceEntry>,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
pub struct ContextUsageSnapshot {
pub estimated_tokens: usize,
pub context_window: usize,
pub compact_threshold_tokens: usize,
pub provider_anchored: bool,
}
impl ContextUsageSnapshot {
pub fn utilization_percent(&self) -> usize {
self.estimated_tokens
.saturating_mul(100)
.checked_div(self.context_window.max(1))
.unwrap_or(0)
}
pub fn compact_threshold_percent(&self) -> usize {
self.compact_threshold_tokens
.saturating_mul(100)
.checked_div(self.context_window.max(1))
.unwrap_or(0)
}
pub fn headroom_tokens(&self) -> usize {
self.context_window.saturating_sub(self.estimated_tokens)
}
pub fn compact_headroom_tokens(&self) -> usize {
self.compact_threshold_tokens
.saturating_sub(self.estimated_tokens)
}
pub fn short_status(&self) -> String {
format!(
"ctx {}{}/{} ({}%)",
if self.provider_anchored { "" } else { "~" },
format_token_count(self.estimated_tokens),
format_token_count(self.context_window),
self.utilization_percent()
)
}
}
fn format_token_count(tokens: usize) -> String {
if tokens >= 1_000_000 {
format!("{:.1}m", tokens as f64 / 1_000_000.0)
} else if tokens >= 1_000 {
format!("{:.0}k", tokens as f64 / 1_000.0)
} else {
tokens.to_string()
}
}
fn format_compaction_notice(
strategy: &str,
before_tokens: usize,
after_tokens: usize,
context_window: usize,
before_messages: usize,
after_messages: usize,
) -> String {
let before_percent = before_tokens.saturating_mul(100) / context_window.max(1);
let after_percent = after_tokens.saturating_mul(100) / context_window.max(1);
format!(
"Compacted via {strategy}: {} → {} messages; context ~{} → ~{} tokens \
({}% → {}% of {}; ~{} freed)",
before_messages,
after_messages,
before_tokens,
after_tokens,
before_percent,
after_percent,
context_window,
before_tokens.saturating_sub(after_tokens),
)
}
struct TimedToolOutput {
output: crate::tools::ToolOutput,
started_after_ms: u64,
duration_ms: u64,
}
pub enum StreamEvent {
ModelRequest,
Reasoning,
Text(String),
Retry(String),
Notice(String),
ContextUsage(ContextUsageSnapshot),
SteeringSent(String),
ToolStart {
name: String,
summary: String,
input: serde_json::Value,
},
ToolResult {
is_error: bool,
content: String,
},
ToolRunning {
index: usize,
},
ToolFinished {
index: usize,
is_error: bool,
content: String,
},
ToolOutput {
index: usize,
content: String,
},
PermissionRequest {
tool_name: String,
summary: String,
input: serde_json::Value,
respond: oneshot::Sender<PermissionResponse>,
},
PermissionRequestWithDiff {
tool_name: String,
summary: String,
diff: String,
input: serde_json::Value,
respond: oneshot::Sender<PermissionResponse>,
},
Interrupted,
Error(String),
Done,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Continuation {
AwaitUserTurn,
ResumeTask,
}
const MAX_TRANSIENT_RETRIES: u32 = 3;
const DEFAULT_RETRY_BACKOFF_BASE: std::time::Duration = std::time::Duration::from_secs(1);
const MAX_RETRY_BACKOFF: std::time::Duration = std::time::Duration::from_secs(30);
const MAX_RETRY_AFTER: std::time::Duration = std::time::Duration::from_secs(60);
impl Engine {
pub fn new(
provider: Box<dyn Provider>,
tools: ToolRegistry,
permissions: PermissionChecker,
model: &str,
) -> Self {
Self {
provider,
tools,
permissions,
messages: Vec::new(),
system_prompt: String::new(),
model: model.to_string(),
model_binding: None,
max_tokens: 16384,
context_window: crate::model::built_in_metadata(model).context_window,
auto_compact_threshold: 0.8,
steering: SteeringQueue::default(),
pending_images: Vec::new(),
plugins: None,
checkpoint_enabled: true,
pending_checkpoint: None,
last_checkpoint: None,
tool_trace: Vec::new(),
model_trace: Vec::new(),
trace_started_at: None,
trace_duration_ms: None,
transcript_checkpoint: None,
cost: CostTracker::new(model),
last_request_usage: None,
last_compaction_notice: None,
last_failure: None,
retry_backoff_base: DEFAULT_RETRY_BACKOFF_BASE,
}
}
#[cfg(test)]
pub(crate) fn for_tests(
provider: Box<dyn Provider>,
steering: SteeringQueue,
mode: crate::permissions::PermissionMode,
) -> Self {
Self {
provider,
tools: ToolRegistry::without_agent_for_tests(),
permissions: PermissionChecker::new(mode),
messages: vec![],
system_prompt: String::new(),
model: "test".to_string(),
model_binding: None,
max_tokens: 1000,
context_window: crate::model::built_in_metadata("test").context_window,
auto_compact_threshold: 0.8,
steering,
pending_images: Vec::new(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
tool_trace: Vec::new(),
model_trace: Vec::new(),
trace_started_at: None,
trace_duration_ms: None,
transcript_checkpoint: None,
cost: CostTracker::new("test"),
last_request_usage: None,
last_compaction_notice: None,
last_failure: None,
retry_backoff_base: DEFAULT_RETRY_BACKOFF_BASE,
}
}
pub fn set_plugins(&mut self, plugins: Arc<PluginRegistry>) {
self.plugins = Some(plugins);
}
async fn fire_hook(&self, trigger: &HookTrigger) {
if let Some(plugins) = &self.plugins {
if let Err(error) = plugins.execute_side_effects(trigger, None).await {
tracing::warn!("plugin hook {trigger:?} failed: {error}");
}
}
}
fn begin_checkpoint(&mut self) {
if !self.checkpoint_enabled {
return;
}
self.last_checkpoint = None;
self.pending_checkpoint = match PendingCheckpoint::capture() {
Ok(checkpoint) => Some(checkpoint),
Err(error) => {
tracing::debug!("turn checkpoint unavailable: {error}");
None
}
};
}
fn finish_checkpoint(&mut self) {
let Some(pending) = self.pending_checkpoint.take() else {
return;
};
match pending.finish() {
Ok(checkpoint) => self.last_checkpoint = Some(checkpoint),
Err(error) => tracing::warn!("could not finish turn checkpoint: {error}"),
}
}
pub fn last_turn_diff(&self) -> String {
self.last_checkpoint
.as_ref()
.map(TurnCheckpoint::diff)
.unwrap_or_else(|| {
"No turn checkpoint is available (checkpoints require a Git worktree).".to_string()
})
}
pub fn undo_last_turn(&mut self) -> Result<String> {
anyhow::ensure!(!self.jobs().snapshots().iter().any(|job| job.status.active()),
"Wait for or cancel background jobs before undoing a turn; they may still be changing files.");
let checkpoint = self.last_checkpoint.as_ref().ok_or_else(|| {
anyhow::anyhow!("No turn checkpoint is available (checkpoints require a Git worktree).")
})?;
let result = checkpoint.undo()?;
self.last_checkpoint = None;
self.provider.reset_session();
self.messages.push(Message::user(
"[Claux checkpoint] The user invoked /undo-turn. The previous turn's \
checkpointed filesystem changes were reverted. Re-read affected files \
before relying on the previous turn's results.",
));
Ok(result)
}
pub fn steering_queue(&self) -> SteeringQueue {
self.steering.clone()
}
pub fn queue_image(&mut self, image: ImageSource) -> usize {
self.pending_images.push(image);
self.pending_images.len()
}
fn take_user_message(&mut self, text: &str) -> Message {
if self.pending_images.is_empty() {
Message::user(text)
} else {
Message::user_with_images(text, std::mem::take(&mut self.pending_images))
}
}
pub fn inject_steering(&mut self) -> Vec<String> {
let drained: Vec<String> = {
let mut q = self.steering.lock().expect("steering queue poisoned");
q.drain(..).collect()
};
for text in &drained {
self.messages.push(Message::user(text));
}
drained
}
pub fn steering_pending(&self) -> bool {
!self
.steering
.lock()
.expect("steering queue poisoned")
.is_empty()
}
pub const SKIPPED_FOR_STEERING: &'static str =
"Skipped: superseded by a new user message before this tool ran.";
async fn execute_tool_steerable(
&self,
name: &str,
input: serde_json::Value,
cancel: &tokio_util::sync::CancellationToken,
progress: tokio::sync::watch::Sender<String>,
) -> crate::tools::ToolOutput {
let token = cancel.child_token();
let steering = self.steering.clone();
let watch_token = token.clone();
let watcher = tokio::spawn(async move {
loop {
if !steering.lock().expect("steering queue poisoned").is_empty() {
watch_token.cancel();
return;
}
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
}
});
let output = self
.tools
.execute_with_progress(name, input, token, Some(progress))
.await;
watcher.abort();
output
}
pub fn set_auto_compact_threshold(&mut self, threshold: f64) {
self.auto_compact_threshold = threshold.clamp(0.0, 1.0);
}
pub fn set_max_tokens(&mut self, max_tokens: u32) {
self.max_tokens = max_tokens.max(1);
}
pub fn last_failure(&self) -> Option<&FailureRecord> {
self.last_failure.as_ref()
}
#[cfg(test)]
pub(crate) fn set_retry_backoff_base(&mut self, base: std::time::Duration) {
self.retry_backoff_base = base;
}
fn retry_delay(
&self,
attempt: u32,
retry_after: Option<std::time::Duration>,
) -> std::time::Duration {
if let Some(retry_after) = retry_after {
return retry_after.min(MAX_RETRY_AFTER);
}
let exponent = attempt.saturating_sub(1).min(16);
let base = self
.retry_backoff_base
.saturating_mul(1u32 << exponent)
.min(MAX_RETRY_BACKOFF);
let jitter_permille = (uuid::Uuid::new_v4().as_u128() % 251) as u32;
base + base.mul_f64(jitter_permille as f64 / 1000.0)
}
async fn wait_before_retry(
delay: std::time::Duration,
cancel: &tokio_util::sync::CancellationToken,
) -> bool {
tokio::select! {
_ = tokio::time::sleep(delay) => true,
_ = cancel.cancelled() => false,
}
}
fn failure_of(error: &anyhow::Error) -> Option<ApiFailure> {
error.downcast_ref::<ApiFailure>().cloned()
}
fn record_failure(&mut self, failure: &ApiFailure, attempts: u32) {
self.last_failure = Some(FailureRecord::from_failure(failure, attempts));
}
pub fn set_model_metadata(&mut self, metadata: crate::model::ModelMetadata) {
self.context_window = metadata.context_window;
self.cost.set_pricing_override(metadata.pricing);
}
pub fn set_system_prompt(&mut self, prompt: String) {
self.system_prompt = prompt;
}
pub fn messages(&self) -> &[Message] {
&self.messages
}
pub fn jobs(&self) -> std::sync::Arc<crate::tools::jobs::JobManager> {
self.tools.jobs.clone()
}
pub fn tool_trace(&self) -> &[ToolTraceEntry] {
&self.tool_trace
}
pub fn execution_timing(&self) -> ExecutionTiming {
ExecutionTiming {
total_duration_ms: self.trace_duration_ms.unwrap_or_else(|| {
self.trace_started_at
.map(|started| started.elapsed().as_millis() as u64)
.unwrap_or_default()
}),
model_rounds: self.model_trace.clone(),
}
}
pub fn set_transcript_checkpoint(&mut self, path: PathBuf) {
self.transcript_checkpoint = Some(path);
}
fn checkpoint_transcript(&self) {
let Some(path) = self.transcript_checkpoint.as_deref() else {
return;
};
let transcript = crate::output::OneShotTranscript::running(
self.model(),
&self.cost,
self.messages(),
self.tool_trace(),
self.execution_timing(),
);
if let Err(error) = crate::output::write_transcript(path, &transcript) {
tracing::warn!(
"could not checkpoint transcript {}: {error}",
path.display()
);
}
}
fn start_recording(&mut self) {
self.tool_trace.clear();
self.model_trace.clear();
self.trace_duration_ms = None;
self.trace_started_at = Some(Instant::now());
}
fn finish_recording(&mut self) {
self.trace_duration_ms = self
.trace_started_at
.map(|started| started.elapsed().as_millis() as u64);
}
fn trace_offset_ms(&self) -> u64 {
self.trace_started_at
.map(|started| started.elapsed().as_millis() as u64)
.unwrap_or_default()
}
#[cfg(test)]
pub fn messages_mut(&mut self) -> &mut Vec<Message> {
&mut self.messages
}
pub fn set_messages(&mut self, messages: Vec<Message>) {
self.provider.reset_session();
self.permissions.reset_session();
self.tools.reset_session();
self.cost.reset_usage();
self.steering
.lock()
.expect("steering queue poisoned")
.clear();
self.messages = messages;
self.tool_trace.clear();
self.model_trace.clear();
self.trace_started_at = None;
self.trace_duration_ms = None;
self.pending_checkpoint = None;
self.last_checkpoint = None;
self.last_request_usage = None;
}
pub fn model(&self) -> &str {
&self.model
}
pub fn set_model_binding(&mut self, binding: ModelBinding) {
self.model_binding = Some(binding);
}
pub fn model_binding(&self) -> Option<&ModelBinding> {
self.model_binding.as_ref()
}
pub fn set_theme(&mut self, _theme: crate::theme::ThemeName) {
}
pub fn message_count(&self) -> usize {
self.messages.len()
}
fn estimated_context_tokens(&self) -> usize {
let Some(baseline) = &self.last_request_usage else {
return compact::estimate_tokens(&self.messages);
};
if baseline.message_count > self.messages.len() {
return compact::estimate_tokens(&self.messages);
}
baseline.prompt_tokens + compact::estimate_tokens(&self.messages[baseline.message_count..])
}
pub fn context_usage(&self) -> ContextUsageSnapshot {
let provider_anchored = self
.last_request_usage
.as_ref()
.is_some_and(|baseline| baseline.message_count <= self.messages.len());
ContextUsageSnapshot {
estimated_tokens: self.estimated_context_tokens(),
context_window: self.context_window,
compact_threshold_tokens: (self.context_window as f64 * self.auto_compact_threshold)
as usize,
provider_anchored,
}
}
pub fn context_status(&self) -> String {
self.context_usage().short_status()
}
pub fn context_report(&self) -> String {
let usage = self.context_usage();
format!(
"Context: {} / {} estimated tokens ({}%)\n\
Estimate source: {}\n\
Auto-compact: {} tokens ({}%); {} tokens until threshold\n\
Window headroom: {} tokens",
usage.estimated_tokens,
usage.context_window,
usage.utilization_percent(),
if usage.provider_anchored {
"provider usage plus estimated message delta"
} else {
"message estimate until the next provider usage report"
},
usage.compact_threshold_tokens,
usage.compact_threshold_percent(),
usage.compact_headroom_tokens(),
usage.headroom_tokens(),
)
}
fn record_request_usage(&mut self, usage: &crate::api::types::Usage, message_count: usize) {
let prompt_tokens = usage.input_tokens as usize
+ usage.cache_read_tokens as usize
+ usage.cache_creation_tokens as usize;
if prompt_tokens == 0 {
return; }
self.last_request_usage = Some(RequestUsageBaseline {
prompt_tokens,
message_count,
});
}
async fn maybe_auto_compact_with_cancel(
&mut self,
cancel: &tokio_util::sync::CancellationToken,
continuation: Continuation,
) -> Result<Option<String>> {
if self.auto_compact_threshold <= 0.0 {
return Ok(None);
}
let current_tokens = self.estimated_context_tokens();
let threshold_tokens = (self.context_window as f64 * self.auto_compact_threshold) as usize;
if current_tokens > threshold_tokens {
tracing::info!(
"Auto-compact triggered: {} tokens > {} (threshold: {:.0}% of {})",
current_tokens,
threshold_tokens,
self.auto_compact_threshold * 100.0,
self.context_window
);
self.compact_with_cancel(cancel, continuation).await?;
let notice = self
.last_compaction_notice
.clone()
.unwrap_or_else(|| "conversation auto-compacted to free context".to_string());
tracing::info!("Auto-compact completed: {}", notice);
Ok(Some(notice))
} else {
Ok(None)
}
}
pub async fn compact(&mut self) -> Result<String> {
self.compact_with_cancel(
&tokio_util::sync::CancellationToken::new(),
Continuation::AwaitUserTurn,
)
.await
}
async fn compact_with_cancel(
&mut self,
cancel: &tokio_util::sync::CancellationToken,
continuation: Continuation,
) -> Result<String> {
self.last_compaction_notice = None;
if self.messages.is_empty() {
return Ok("Nothing to compact.".to_string());
}
let before_context = self.estimated_context_tokens();
self.summarize_conversation(self.messages.clone(), before_context, cancel, continuation)
.await
}
async fn summarize_conversation(
&mut self,
messages: Vec<Message>,
before_context: usize,
cancel: &tokio_util::sync::CancellationToken,
continuation: Continuation,
) -> Result<String> {
let summary_prompt = compact::SUMMARY_PROMPT;
let old_count = messages.len();
let old_message_tokens = compact::estimate_tokens(&messages);
let original_request = compact::original_request(&messages);
let mut summary_messages = messages;
summary_messages.push(Message::user(summary_prompt));
let mut rx = self
.provider
.stream(
&summary_messages,
&self.system_prompt,
&[],
self.max_tokens,
cancel.child_token(),
)
.await?;
let mut summary = String::new();
let mut completed = false;
loop {
let event = tokio::select! {
_ = cancel.cancelled() => anyhow::bail!("Compaction cancelled by user"),
event = rx.recv() => event,
};
let Some(event) = event else { break };
match event {
ApiEvent::Text(t) => summary.push_str(&t),
ApiEvent::Usage(usage) => self.cost.add_usage(&usage),
ApiEvent::Done => {
completed = true;
break;
}
ApiEvent::Error(failure) => {
let message = format!("Compact error: {}", failure.message);
return Err(anyhow::Error::new(ApiFailure::new(failure.kind, message)));
}
_ => {}
}
}
if !completed {
if cancel.is_cancelled() {
anyhow::bail!("Compaction cancelled by user");
}
anyhow::bail!("Compact error: API stream ended without completion");
}
if cancel.is_cancelled() {
self.provider.reset_session();
anyhow::bail!("Compaction cancelled by user");
}
if summary.trim().is_empty() {
self.provider.reset_session();
anyhow::bail!(
"Compact error: provider returned an empty task handoff; history preserved"
);
}
let mut compacted = vec![
original_request
.unwrap_or_else(|| Message::user("Here is a summary of our conversation so far:")),
Message::assistant_text(&summary),
];
if continuation == Continuation::ResumeTask {
compacted.push(Message::user(
"Continue with the outstanding task described above.",
));
}
if compact::estimate_tokens(&compacted) >= old_message_tokens {
self.provider.reset_session();
anyhow::bail!("Compact error: task handoff did not reduce context; history preserved");
}
self.commit_compacted_messages(compacted);
let after_context = self.estimated_context_tokens();
self.last_compaction_notice = Some(format_compaction_notice(
"summary",
before_context,
after_context,
self.context_window,
old_count,
self.messages.len(),
));
Ok(format!(
"{}\n\n\x1b[2m{summary}\x1b[0m",
self.last_compaction_notice.as_deref().unwrap_or_default()
))
}
fn commit_compacted_messages(&mut self, messages: Vec<Message>) {
self.messages = messages;
self.provider.reset_session();
self.last_request_usage = None;
}
fn unique_tool_use_id(
&self,
id: String,
batch: &[(String, String, serde_json::Value)],
) -> String {
let mut taken: std::collections::HashSet<&str> =
batch.iter().map(|(id, _, _)| id.as_str()).collect();
let history_ids: Vec<String> = self
.messages
.iter()
.filter_map(|message| match &message.content {
crate::api::types::MessageContent::Blocks(blocks) => Some(blocks),
crate::api::types::MessageContent::Text(_) => None,
})
.flatten()
.filter_map(|block| match block {
ContentBlock::ToolUse { id, .. } => Some(id.clone()),
_ => None,
})
.collect();
taken.extend(history_ids.iter().map(String::as_str));
let base = if id.trim().is_empty() {
tracing::warn!("provider omitted a tool_use id; synthesizing one");
format!("call_{}", batch.len() + 1)
} else {
id
};
if !taken.contains(base.as_str()) {
return base;
}
let mut suffix = 2;
loop {
let candidate = format!("{base}#{suffix}");
if !taken.contains(candidate.as_str()) {
tracing::warn!(
"provider repeated tool_use id {base:?}; renamed duplicate to {candidate:?}"
);
return candidate;
}
suffix += 1;
}
}
fn failure_kind(error: &anyhow::Error) -> ApiFailureKind {
error
.downcast_ref::<ApiFailure>()
.map(|failure| failure.kind)
.unwrap_or(ApiFailureKind::Other)
}
fn malformed_tool_retry_prompt(err: &str) -> String {
let detail = crate::utils::truncate_str(err, 512);
format!(
"Your previous response was rejected before any tools executed because one or more \
tool calls contained invalid JSON arguments ({detail}). Reissue the entire intended \
tool-call batch with valid JSON arguments. Do not assume any tool from the rejected \
response ran."
)
}
pub const INTERRUPTED_BY_USER: &'static str = "Interrupted by user.";
pub async fn submit(
&mut self,
user_input: &str,
cancel: tokio_util::sync::CancellationToken,
) -> Result<String> {
let (tx, mut rx) = mpsc::channel::<StreamEvent>(256);
self.start_recording();
self.begin_checkpoint();
let collector = tokio::spawn(async move {
let mut text = String::new();
while let Some(event) = rx.recv().await {
match event {
StreamEvent::Text(t) => text.push_str(&t),
StreamEvent::Retry(_) => text.clear(),
_ => {}
}
}
text
});
let message = self.take_user_message(user_input);
let result = self.run_turn(message, tx, false, cancel).await;
self.finish_recording();
let text = collector.await.unwrap_or_default();
self.finish_checkpoint();
self.fire_hook(&HookTrigger::OnTurnEnd).await;
result?;
Ok(text)
}
pub async fn submit_message(
&mut self,
message: Message,
cancel: tokio_util::sync::CancellationToken,
) -> Result<String> {
let (tx, mut rx) = mpsc::channel::<StreamEvent>(256);
self.start_recording();
self.begin_checkpoint();
let collector = tokio::spawn(async move {
let mut text = String::new();
while let Some(event) = rx.recv().await {
match event {
StreamEvent::Text(t) => text.push_str(&t),
StreamEvent::Retry(_) => text.clear(),
_ => {}
}
}
text
});
let result = self.run_turn(message, tx, false, cancel).await;
self.finish_recording();
let text = collector.await.unwrap_or_default();
self.finish_checkpoint();
self.fire_hook(&HookTrigger::OnTurnEnd).await;
result?;
Ok(text)
}
pub async fn submit_streaming(
&mut self,
user_input: &str,
tx: mpsc::Sender<StreamEvent>,
cancel: tokio_util::sync::CancellationToken,
) -> Result<()> {
self.start_recording();
self.begin_checkpoint();
let message = self.take_user_message(user_input);
let result = self.run_turn(message, tx, true, cancel).await;
self.finish_recording();
self.finish_checkpoint();
self.fire_hook(&HookTrigger::OnTurnEnd).await;
result
}
async fn run_turn(
&mut self,
user_message: Message,
tx: mpsc::Sender<StreamEvent>,
interactive: bool,
cancel: tokio_util::sync::CancellationToken,
) -> Result<()> {
if cancel.is_cancelled() {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
let compact_notice = match self
.maybe_auto_compact_with_cancel(&cancel, Continuation::AwaitUserTurn)
.await
{
Ok(notice) => notice,
Err(_) if cancel.is_cancelled() => {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
Err(error) => return Err(error),
};
if let Some(notice) = compact_notice {
let _ = tx.send(StreamEvent::Notice(notice)).await;
let _ = tx
.send(StreamEvent::ContextUsage(self.context_usage()))
.await;
}
self.messages.push(user_message);
self.checkpoint_transcript();
self.last_failure = None;
let mut recovery_attempts = 0;
const MAX_RECOVERY: u32 = 3;
let mut transient_attempts: u32 = 0;
let mut malformed_tool_retries = 0;
const MAX_MALFORMED_TOOL_RETRIES: u32 = 1;
let mut retry_prompt: Option<String> = None;
let mut turn_had_meaningful_response = false;
loop {
for text in self.inject_steering() {
let _ = tx.send(StreamEvent::SteeringSent(text)).await;
}
if cancel.is_cancelled() {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
let tool_defs = self.tools.definitions();
let mut effective_system_prompt = retry_prompt
.as_ref()
.map(|prompt| format!("{}\n\n{prompt}", self.system_prompt));
let jobs = self.jobs().snapshots();
if !jobs.is_empty() {
let states = jobs
.iter()
.map(|job| format!("{}: {}", job.id, job.status.label()))
.collect::<Vec<_>>()
.join("\n");
let prompt =
effective_system_prompt.get_or_insert_with(|| self.system_prompt.clone());
prompt.push_str(&format!("\n\nSession background jobs:\n{states}\nUse Jobs to inspect output before reporting results. Do not start another job merely to check on an existing job."));
}
let model_started_after_ms = self.trace_offset_ms();
let model_started = Instant::now();
let _ = tx.send(StreamEvent::ModelRequest).await;
let stream_result = self
.provider
.stream(
&self.messages,
effective_system_prompt
.as_deref()
.unwrap_or(&self.system_prompt),
&tool_defs,
self.max_tokens,
cancel.clone(),
)
.await;
let sent_message_count = self.messages.len();
let mut rx = match stream_result {
Ok(rx) => rx,
Err(e) => {
let failure = Self::failure_of(&e);
let kind = failure.as_ref().map(|f| f.kind);
let transient = failure.as_ref().is_some_and(|f| f.kind.retryable())
&& transient_attempts < MAX_TRANSIENT_RETRIES
&& !cancel.is_cancelled();
self.model_trace.push(ModelTraceEntry {
index: self.model_trace.len() + 1,
started_after_ms: model_started_after_ms,
duration_ms: model_started.elapsed().as_millis() as u64,
status: if transient { "retry" } else { "error" }.to_string(),
failure: kind,
usage: None,
});
self.checkpoint_transcript();
if cancel.is_cancelled() {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
if transient {
let failure = failure.expect("transient implies a classified failure");
transient_attempts += 1;
let delay = self.retry_delay(transient_attempts, failure.retry_after);
let _ = tx
.send(StreamEvent::Retry(format!(
"provider {}; retrying in {:.0}s (attempt {transient_attempts}/{MAX_TRANSIENT_RETRIES})",
failure.kind.as_str().replace('_', " "),
delay.as_secs_f64()
)))
.await;
if !Self::wait_before_retry(delay, &cancel).await {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
continue;
}
let err_str = e.to_string();
match Self::failure_kind(&e) {
ApiFailureKind::MalformedToolArguments
if malformed_tool_retries < MAX_MALFORMED_TOOL_RETRIES =>
{
malformed_tool_retries += 1;
retry_prompt = Some(Self::malformed_tool_retry_prompt(&err_str));
let _ = tx
.send(StreamEvent::Retry(
"model returned malformed tool arguments; retrying once"
.to_string(),
))
.await;
continue;
}
ApiFailureKind::OutputLimitExceeded if self.max_tokens < 64_000 => {
self.max_tokens = (self.max_tokens * 2).min(64_000);
let _ = tx
.send(StreamEvent::Retry(
"provider hit the output limit; retrying with a larger budget"
.to_string(),
))
.await;
continue;
}
ApiFailureKind::ContextExceeded if recovery_attempts < MAX_RECOVERY => {
recovery_attempts += 1;
let _ = tx
.send(StreamEvent::Retry(
"provider rejected the context; compacting and retrying"
.to_string(),
))
.await;
let _ = tx
.send(StreamEvent::Notice(
"compacting conversation...".to_string(),
))
.await;
if let Err(error) = self
.compact_with_cancel(&cancel, Continuation::ResumeTask)
.await
{
if cancel.is_cancelled() {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
return Err(error);
}
if let Some(notice) = self.last_compaction_notice.clone() {
let _ = tx.send(StreamEvent::Notice(notice)).await;
}
let _ = tx
.send(StreamEvent::ContextUsage(self.context_usage()))
.await;
continue;
}
_ => {}
}
if let Some(failure) = &failure {
self.record_failure(failure, transient_attempts + 1);
}
let _ = tx.send(StreamEvent::Error(err_str.clone())).await;
return Err(e);
}
};
let mut text_buf = String::new();
let mut reasoning_text = String::new();
let mut reasoning_details = Vec::new();
let mut tool_uses: Vec<(String, String, serde_json::Value)> = Vec::new();
let mut had_error = false;
let mut stream_interrupted = false;
let mut committed = false;
let mut model_status = "completed";
let mut model_failure: Option<ApiFailureKind> = None;
let mut pending_retry_delay: Option<std::time::Duration> = None;
let mut model_usage = None;
loop {
let event = tokio::select! {
event = rx.recv() => match event {
Some(event) => event,
None => {
let failure =
ApiFailure::protocol_error("API stream ended without completion");
let _ = tx.send(StreamEvent::Error(failure.message.clone())).await;
self.model_trace.push(ModelTraceEntry {
index: self.model_trace.len() + 1,
started_after_ms: model_started_after_ms,
duration_ms: model_started.elapsed().as_millis() as u64,
status: "error".to_string(),
failure: Some(failure.kind),
usage: model_usage,
});
self.checkpoint_transcript();
self.record_failure(&failure, transient_attempts + 1);
return Err(anyhow::Error::new(failure));
}
},
_ = cancel.cancelled() => {
stream_interrupted = true;
break;
}
};
match event {
ApiEvent::Text(t) => {
let _ = tx.send(StreamEvent::Text(t.clone())).await;
text_buf.push_str(&t);
}
ApiEvent::Reasoning { text, details } => {
let _ = tx.send(StreamEvent::Reasoning).await;
if let Some(text) = text {
reasoning_text.push_str(&text);
}
reasoning_details.extend(details);
}
ApiEvent::ToolUse { id, name, input } => {
self.fire_hook(&HookTrigger::OnToolStart).await;
let summary = self.tools.summarize(&name, &input);
let _ = tx
.send(StreamEvent::ToolStart {
name: name.clone(),
summary,
input: input.clone(),
})
.await;
committed = true;
let id = self.unique_tool_use_id(id, &tool_uses);
tool_uses.push((id, name, input));
}
ApiEvent::Usage(usage) => {
model_usage = Some(ModelRoundUsage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
cache_read_tokens: usage.cache_read_tokens,
cache_creation_tokens: usage.cache_creation_tokens,
cost_usd: usage.provider_cost_usd,
});
self.record_request_usage(&usage, sent_message_count);
self.cost.add_usage(&usage);
let _ = tx
.send(StreamEvent::ContextUsage(self.context_usage()))
.await;
}
ApiEvent::Done => break,
ApiEvent::Error(failure) => {
if !committed {
match failure.kind {
kind if kind.retryable()
&& transient_attempts < MAX_TRANSIENT_RETRIES =>
{
transient_attempts += 1;
let delay =
self.retry_delay(transient_attempts, failure.retry_after);
let _ = tx
.send(StreamEvent::Retry(format!(
"provider {}; retrying in {:.0}s (attempt {transient_attempts}/{MAX_TRANSIENT_RETRIES})",
kind.as_str().replace('_', " "),
delay.as_secs_f64()
)))
.await;
pending_retry_delay = Some(delay);
had_error = true;
model_status = "retry";
model_failure = Some(kind);
break;
}
ApiFailureKind::MalformedToolArguments
if malformed_tool_retries < MAX_MALFORMED_TOOL_RETRIES =>
{
malformed_tool_retries += 1;
model_failure = Some(ApiFailureKind::MalformedToolArguments);
retry_prompt =
Some(Self::malformed_tool_retry_prompt(&failure.message));
let _ = tx
.send(StreamEvent::Retry(
"model returned malformed tool arguments; retrying once"
.to_string(),
))
.await;
had_error = true;
model_status = "retry";
break;
}
ApiFailureKind::OutputLimitExceeded if self.max_tokens < 64_000 => {
self.max_tokens = (self.max_tokens * 2).min(64_000);
model_failure = Some(ApiFailureKind::OutputLimitExceeded);
let _ = tx
.send(StreamEvent::Retry(
"provider hit the output limit; retrying with a larger budget"
.to_string(),
))
.await;
had_error = true;
model_status = "retry";
break;
}
ApiFailureKind::ContextExceeded
if recovery_attempts < MAX_RECOVERY =>
{
recovery_attempts += 1;
model_failure = Some(ApiFailureKind::ContextExceeded);
let _ = tx
.send(StreamEvent::Retry(
"provider rejected the context; compacting and retrying"
.to_string(),
))
.await;
let _ = tx
.send(StreamEvent::Notice(
"compacting conversation...".to_string(),
))
.await;
if let Err(error) = self
.compact_with_cancel(&cancel, Continuation::ResumeTask)
.await
{
if cancel.is_cancelled() {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
return Err(error);
}
if let Some(notice) = self.last_compaction_notice.clone() {
let _ = tx.send(StreamEvent::Notice(notice)).await;
}
let _ = tx
.send(StreamEvent::ContextUsage(self.context_usage()))
.await;
had_error = true;
model_status = "retry";
break;
}
_ => {}
}
}
let _ = tx.send(StreamEvent::Error(failure.message.clone())).await;
self.model_trace.push(ModelTraceEntry {
index: self.model_trace.len() + 1,
started_after_ms: model_started_after_ms,
duration_ms: model_started.elapsed().as_millis() as u64,
status: "error".to_string(),
failure: Some(failure.kind),
usage: model_usage,
});
self.checkpoint_transcript();
self.record_failure(&failure, transient_attempts + 1);
let message = format!("API error: {}", failure.message);
let mut terminal = ApiFailure::new(failure.kind, message);
terminal.http_status = failure.http_status;
terminal.retry_after = failure.retry_after;
return Err(anyhow::Error::new(terminal));
}
}
}
let empty_completion = !stream_interrupted
&& !had_error
&& !turn_had_meaningful_response
&& text_buf.trim().is_empty()
&& tool_uses.is_empty();
self.model_trace.push(ModelTraceEntry {
index: self.model_trace.len() + 1,
started_after_ms: model_started_after_ms,
duration_ms: model_started.elapsed().as_millis() as u64,
status: if stream_interrupted {
"interrupted".to_string()
} else if empty_completion {
"error".to_string()
} else {
model_status.to_string()
},
failure: if empty_completion {
Some(ApiFailureKind::ProtocolError)
} else {
model_failure
},
usage: model_usage,
});
if had_error {
self.checkpoint_transcript();
if let Some(delay) = pending_retry_delay.take() {
if !Self::wait_before_retry(delay, &cancel).await {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
}
continue;
}
if empty_completion {
let failure = ApiFailure::protocol_error(
"provider protocol error: response completed without assistant text or tool calls",
);
let _ = tx.send(StreamEvent::Error(failure.message.clone())).await;
self.checkpoint_transcript();
self.record_failure(&failure, transient_attempts + 1);
return Err(anyhow::Error::new(failure));
}
turn_had_meaningful_response = true;
malformed_tool_retries = 0;
retry_prompt = None;
let mut blocks = Vec::new();
if !reasoning_text.is_empty() || !reasoning_details.is_empty() {
blocks.push(ContentBlock::Reasoning {
text: (!reasoning_text.is_empty()).then_some(reasoning_text),
details: reasoning_details,
});
}
if !text_buf.is_empty() {
blocks.push(ContentBlock::Text {
text: text_buf.clone(),
});
}
for (id, name, input) in &tool_uses {
blocks.push(ContentBlock::ToolUse {
id: id.clone(),
name: name.clone(),
input: input.clone(),
});
}
if !blocks.is_empty() {
self.messages.push(Message::assistant_blocks(blocks));
}
self.checkpoint_transcript();
if stream_interrupted {
if !tool_uses.is_empty() {
let mut result_blocks = Vec::with_capacity(tool_uses.len());
for (id, name, input) in &tool_uses {
self.fire_hook(&HookTrigger::OnToolComplete).await;
let _ = tx
.send(StreamEvent::ToolResult {
is_error: true,
content: Self::INTERRUPTED_BY_USER.to_string(),
})
.await;
self.tool_trace.push(ToolTraceEntry {
id: id.clone(),
name: name.clone(),
input: input.clone(),
output: Self::INTERRUPTED_BY_USER.to_string(),
is_error: true,
read_only: self.tools.is_read_only(name),
started_after_ms: self.trace_offset_ms(),
duration_ms: 0,
});
result_blocks.push(ContentBlock::ToolResult {
tool_use_id: id.clone(),
content: Self::INTERRUPTED_BY_USER.to_string(),
is_error: Some(true),
});
}
self.messages.push(Message::tool_results(result_blocks));
self.checkpoint_transcript();
}
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
if tool_uses.is_empty() {
let _ = tx.send(StreamEvent::Done).await;
break;
}
let (result_blocks, interrupted) = self
.execute_tool_batch(&tool_uses, &tx, interactive, &cancel)
.await;
self.messages.push(Message::tool_results(result_blocks));
self.checkpoint_transcript();
if interrupted {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
let compact_notice = match self
.maybe_auto_compact_with_cancel(&cancel, Continuation::ResumeTask)
.await
{
Ok(notice) => notice,
Err(_) if cancel.is_cancelled() => {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
Err(error) => return Err(error),
};
if let Some(notice) = compact_notice {
let _ = tx.send(StreamEvent::Notice(notice)).await;
let _ = tx
.send(StreamEvent::ContextUsage(self.context_usage()))
.await;
self.checkpoint_transcript();
}
}
Ok(())
}
async fn execute_tool_batch(
&mut self,
tool_uses: &[(String, String, serde_json::Value)],
tx: &mpsc::Sender<StreamEvent>,
interactive: bool,
cancel: &tokio_util::sync::CancellationToken,
) -> (Vec<ContentBlock>, bool) {
let mut outputs: Vec<Option<TimedToolOutput>> =
(0..tool_uses.len()).map(|_| None).collect();
let mut interrupted = false;
let mut pending_permission = None;
for (idx, (_, name, input)) in tool_uses.iter().enumerate() {
if outputs[idx].is_some() {
continue;
}
if cancel.is_cancelled() {
interrupted = true;
outputs[idx] = Some(TimedToolOutput {
output: crate::tools::ToolOutput {
content: Self::INTERRUPTED_BY_USER.to_string(),
is_error: true,
},
started_after_ms: self.trace_offset_ms(),
duration_ms: 0,
});
continue;
}
if self.steering_pending() {
outputs[idx] = Some(TimedToolOutput {
output: crate::tools::ToolOutput {
content: Self::SKIPPED_FOR_STEERING.to_string(),
is_error: true,
},
started_after_ms: self.trace_offset_ms(),
duration_ms: 0,
});
continue;
}
let is_read_only = self.tools.is_read_only(name);
let perm = match pending_permission.take() {
Some(permission) => permission,
None => self.decide_permission(name, input, is_read_only).await,
};
if is_read_only && matches!(perm, PermissionResult::Allow) {
let mut end = idx + 1;
while end < tool_uses.len() && end - idx < MAX_PARALLEL_TOOLS {
let (_, next_name, next_input) = &tool_uses[end];
if !self.tools.is_read_only(next_name)
|| cancel.is_cancelled()
|| self.steering_pending()
{
break;
}
let permission = self.decide_permission(next_name, next_input, true).await;
if !matches!(permission, PermissionResult::Allow) {
pending_permission = Some(permission);
break;
}
end += 1;
}
let this: &Self = &*self;
let futures = tool_uses[idx..end].iter().enumerate().map(
|(offset, (_, name, input))| async move {
let started_after_ms = this.trace_offset_ms();
let started = Instant::now();
let output = if cancel.is_cancelled() || this.steering_pending() {
crate::tools::ToolOutput {
content: if cancel.is_cancelled() {
Self::INTERRUPTED_BY_USER
} else {
Self::SKIPPED_FOR_STEERING
}
.to_string(),
is_error: true,
}
} else {
this.execute_tool_reporting(idx + offset, name, input, tx, cancel)
.await
};
TimedToolOutput {
output,
started_after_ms,
duration_ms: started.elapsed().as_millis() as u64,
}
},
);
for (slot, output) in outputs[idx..end]
.iter_mut()
.zip(futures_util::future::join_all(futures).await)
{
*slot = Some(output);
}
continue;
}
let started_after_ms = self.trace_offset_ms();
let started = Instant::now();
let output = match perm {
PermissionResult::Allow => {
self.execute_tool_reporting(idx, name, input, tx, cancel)
.await
}
PermissionResult::Deny(reason) => crate::tools::ToolOutput {
content: format!("Permission denied: {reason}"),
is_error: true,
},
PermissionResult::Ask { message, diff } => {
if !interactive {
crate::tools::ToolOutput {
content: format!(
"Permission denied: {message} (one-shot mode has no prompt; set permission_mode in config.toml to allow)"
),
is_error: true,
}
} else {
self.ask_permission(name, input, message, diff, (idx, tx), cancel)
.await
}
}
};
outputs[idx] = Some(TimedToolOutput {
output,
started_after_ms,
duration_ms: started.elapsed().as_millis() as u64,
});
}
if cancel.is_cancelled() {
interrupted = true;
}
let mut result_blocks = Vec::with_capacity(tool_uses.len());
for (idx, (id, name, input)) in tool_uses.iter().enumerate() {
let timed = outputs[idx].take().expect("every tool got an output");
let output = timed.output;
let (content, was_truncated) = compact::truncate_tool_output(&output.content);
if was_truncated {
tracing::debug!("Truncated tool output for {}", name);
}
self.fire_hook(&HookTrigger::OnToolComplete).await;
let _ = tx
.send(StreamEvent::ToolResult {
is_error: output.is_error,
content: output.content.clone(),
})
.await;
self.tool_trace.push(ToolTraceEntry {
id: id.clone(),
name: name.clone(),
input: input.clone(),
output: content.clone(),
is_error: output.is_error,
read_only: self.tools.is_read_only(name),
started_after_ms: timed.started_after_ms,
duration_ms: timed.duration_ms,
});
result_blocks.push(ContentBlock::ToolResult {
tool_use_id: id.clone(),
content,
is_error: if output.is_error { Some(true) } else { None },
});
}
(result_blocks, interrupted)
}
async fn decide_permission(
&self,
tool_name: &str,
input: &serde_json::Value,
is_read_only: bool,
) -> PermissionResult {
let result = self.permissions.check(tool_name, input, is_read_only);
let Some(plugins) = &self.plugins else {
return result;
};
if plugins.get_by_trigger(&HookTrigger::OnPermissionCheck) == 0 {
return result;
}
let mode = serde_json::to_value(self.permissions.mode())
.ok()
.and_then(|value| value.as_str().map(str::to_string))
.unwrap_or_default();
let env: std::collections::HashMap<String, String> = [
("CLAUX_TOOL_NAME".to_string(), tool_name.to_string()),
("CLAUX_TOOL_INPUT".to_string(), input.to_string()),
("CLAUX_TOOL_READ_ONLY".to_string(), is_read_only.to_string()),
("CLAUX_PERMISSION_MODE".to_string(), mode),
(
"CLAUX_PERMISSION_DECISION".to_string(),
crate::permissions::proposed_decision(&result).to_string(),
),
]
.into_iter()
.collect();
let mut verdicts = Vec::new();
for (name, output) in plugins
.execute_decisions(&HookTrigger::OnPermissionCheck, Some(&env))
.await
{
match crate::permissions::parse_hook_verdict(&output) {
Ok(Some((decision, reason))) => verdicts.push((name, decision, reason)),
Ok(None) => {}
Err(error) => tracing::warn!("permission hook {name} ignored: {error}"),
}
}
crate::permissions::apply_hook_verdicts(tool_name, input, result, &verdicts)
}
async fn execute_tool_reporting(
&self,
index: usize,
name: &str,
input: &serde_json::Value,
tx: &mpsc::Sender<StreamEvent>,
cancel: &tokio_util::sync::CancellationToken,
) -> crate::tools::ToolOutput {
let _ = tx.send(StreamEvent::ToolRunning { index }).await;
let (progress, mut updates) = tokio::sync::watch::channel(String::new());
let execution = self.execute_tool_steerable(name, input.clone(), cancel, progress);
tokio::pin!(execution);
let mut tick = tokio::time::interval(std::time::Duration::from_millis(100));
let output = loop {
tokio::select! {
output = &mut execution => break output,
_ = tick.tick() => {
if updates.has_changed().unwrap_or(false) {
let content = updates.borrow_and_update().clone();
let _ = tx.try_send(StreamEvent::ToolOutput { index, content });
}
}
}
};
let _ = tx
.send(StreamEvent::ToolFinished {
index,
is_error: output.is_error,
content: output.content.clone(),
})
.await;
output
}
async fn ask_permission(
&mut self,
name: &str,
input: &serde_json::Value,
message: String,
diff: Option<String>,
progress: (usize, &mpsc::Sender<StreamEvent>),
cancel: &tokio_util::sync::CancellationToken,
) -> crate::tools::ToolOutput {
let (index, tx) = progress;
self.fire_hook(&HookTrigger::OnPermissionRequest).await;
let (resp_tx, resp_rx) = oneshot::channel();
let event = if let Some(d) = diff {
StreamEvent::PermissionRequestWithDiff {
tool_name: name.to_string(),
summary: message,
diff: d,
input: input.clone(),
respond: resp_tx,
}
} else {
StreamEvent::PermissionRequest {
tool_name: name.to_string(),
summary: message,
input: input.clone(),
respond: resp_tx,
}
};
let _ = tx.send(event).await;
let response = tokio::select! {
biased;
response = resp_rx => response,
_ = cancel.cancelled() => {
return crate::tools::ToolOutput {
content: "Permission request cancelled by user.".to_string(),
is_error: true,
};
}
};
match response {
Ok(PermissionResponse::Allow) => {
self.execute_tool_reporting(index, name, input, tx, cancel)
.await
}
Ok(PermissionResponse::AlwaysAllow) => {
match PermissionResponse::always_allow_for(name, input) {
PermissionResponse::AlwaysAllow => self.permissions.always_allow(name),
PermissionResponse::AlwaysAllowCommand(command) => {
self.permissions.always_allow_command(&command);
}
_ => {}
}
self.execute_tool_reporting(index, name, input, tx, cancel)
.await
}
Ok(PermissionResponse::AlwaysAllowCommand(ref cmd)) => {
self.permissions.always_allow_command(cmd);
self.execute_tool_reporting(index, name, input, tx, cancel)
.await
}
Ok(PermissionResponse::Deny) | Ok(PermissionResponse::DenyAndCancel) | Err(_) => {
crate::tools::ToolOutput {
content: "Permission denied by user.".to_string(),
is_error: true,
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::api::{MessageContent, ToolDefinition};
use crate::permissions::PermissionMode;
use crate::plugin::Plugin;
use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Instant;
struct MockProvider;
#[async_trait::async_trait]
impl Provider for MockProvider {
fn name(&self) -> &str {
"mock"
}
fn set_model(&mut self, _model: &str) {
}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let (tx, rx) = mpsc::channel(10);
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct TruncatedProvider;
struct HangingProvider;
#[async_trait::async_trait]
impl Provider for HangingProvider {
fn name(&self) -> &str {
"hanging"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let (tx, rx) = mpsc::channel(1);
let stream_cancel = cancel.child_token();
let wait_cancel = stream_cancel.clone();
tokio::spawn(async move {
wait_cancel.cancelled().await;
drop(tx);
});
Ok(ProviderStream::new(rx, stream_cancel))
}
}
#[async_trait::async_trait]
impl Provider for TruncatedProvider {
fn name(&self) -> &str {
"truncated"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let (tx, rx) = mpsc::channel(10);
let _ = tx
.send(ApiEvent::Text("partial response".to_string()))
.await;
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct EmptyCompletionProvider;
#[async_trait::async_trait]
impl Provider for EmptyCompletionProvider {
fn name(&self) -> &str {
"empty-completion"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let (tx, rx) = mpsc::channel(2);
tx.send(ApiEvent::Usage(Default::default())).await.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct ReasoningProvider;
#[async_trait::async_trait]
impl Provider for ReasoningProvider {
fn name(&self) -> &str {
"reasoning"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let (tx, rx) = mpsc::channel(4);
tx.send(ApiEvent::Reasoning {
text: Some("private thought".to_string()),
details: vec![serde_json::json!({
"type": "reasoning.text",
"text": "preserve me",
"index": 0
})],
})
.await
.unwrap();
tx.send(ApiEvent::Text("answer".to_string())).await.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct MalformedToolProvider {
calls: Arc<AtomicUsize>,
systems: Arc<Mutex<Vec<String>>>,
recover: bool,
}
struct DuplicateToolIdProvider {
calls: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl Provider for DuplicateToolIdProvider {
fn name(&self) -> &str {
"duplicate-tool-id"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
let (tx, rx) = mpsc::channel(8);
if call == 0 {
for id in ["dup", "dup", ""] {
tx.send(ApiEvent::ToolUse {
id: id.to_string(),
name: "Read".to_string(),
input: serde_json::json!({"file_path": "/dev/null"}),
})
.await
.unwrap();
}
} else {
tx.send(ApiEvent::Text("done".to_string())).await.unwrap();
}
tx.send(ApiEvent::Done).await.unwrap();
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
#[tokio::test]
async fn duplicate_and_missing_tool_use_ids_are_made_unique() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(DuplicateToolIdProvider {
calls: calls.clone(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let result = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(result, "done");
let MessageContent::Blocks(uses) = &engine.messages()[1].content else {
panic!("expected assistant tool_use blocks");
};
let use_ids: Vec<&str> = uses
.iter()
.filter_map(|block| match block {
ContentBlock::ToolUse { id, .. } => Some(id.as_str()),
_ => None,
})
.collect();
assert_eq!(use_ids, vec!["dup", "dup#2", "call_3"]);
let MessageContent::Blocks(results) = &engine.messages()[2].content else {
panic!("expected tool_result blocks");
};
let result_ids: Vec<&str> = results
.iter()
.filter_map(|block| match block {
ContentBlock::ToolResult { tool_use_id, .. } => Some(tool_use_id.as_str()),
_ => None,
})
.collect();
assert_eq!(
result_ids, use_ids,
"every renamed use must pair with its result"
);
}
struct MidStreamOutputRetryProvider {
calls: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl Provider for MidStreamOutputRetryProvider {
fn name(&self) -> &str {
"mid-stream-output-retry"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let attempt = self.calls.fetch_add(1, Ordering::SeqCst);
let (tx, rx) = mpsc::channel(4);
if attempt == 0 {
tx.send(ApiEvent::Text("rejected preamble".to_string()))
.await
.unwrap();
tx.send(ApiEvent::Error(ApiFailure::new(
ApiFailureKind::OutputLimitExceeded,
"output limit",
)))
.await
.unwrap();
} else {
tx.send(ApiEvent::Text("recovered response".to_string()))
.await
.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
}
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
#[async_trait::async_trait]
impl Provider for MalformedToolProvider {
fn name(&self) -> &str {
"malformed-tool"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
self.systems.lock().unwrap().push(system.to_string());
let attempt = self.calls.fetch_add(1, Ordering::SeqCst);
let (tx, rx) = mpsc::channel(10);
if attempt == 0 {
tx.send(ApiEvent::Text("rejected preamble".to_string()))
.await
.unwrap();
}
if self.recover && attempt > 0 {
tx.send(ApiEvent::Text("recovered response".to_string()))
.await
.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
} else {
tx.send(ApiEvent::Error(ApiFailure::malformed_tool_arguments(
"OpenAI SSE stream error: invalid arguments for tool call Read \
(call_3): EOF while parsing a value",
)))
.await
.unwrap();
}
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct ResetTrackingProvider {
resets: Arc<std::sync::atomic::AtomicUsize>,
}
struct CompactionTrackingProvider {
resets: Arc<AtomicUsize>,
complete: bool,
}
struct RecordingSummaryProvider {
resets: Arc<AtomicUsize>,
requests: Arc<Mutex<Vec<Vec<Message>>>>,
summary: String,
}
#[async_trait::async_trait]
impl Provider for RecordingSummaryProvider {
fn name(&self) -> &str {
"recording-summary"
}
fn set_model(&mut self, _model: &str) {}
fn reset_session(&mut self) {
self.resets.fetch_add(1, Ordering::SeqCst);
}
async fn stream(
&self,
messages: &[Message],
_system: &str,
tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
assert!(tools.is_empty(), "summarization must not execute tools");
self.requests.lock().unwrap().push(messages.to_vec());
let (tx, rx) = mpsc::channel(2);
tx.send(ApiEvent::Text(self.summary.clone())).await.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct WithinTurnCompactionProvider {
calls: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl Provider for WithinTurnCompactionProvider {
fn name(&self) -> &str {
"within-turn-compaction"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
let (tx, rx) = mpsc::channel(4);
match call {
0 => {
tx.send(ApiEvent::Reasoning {
text: Some("investigating the host configuration ".repeat(200)),
details: Vec::new(),
})
.await
.unwrap();
tx.send(ApiEvent::Usage(crate::api::types::Usage {
input_tokens: 110_000,
..Default::default()
}))
.await
.unwrap();
tx.send(ApiEvent::ToolUse {
id: "read-1".to_string(),
name: "Read".to_string(),
input: serde_json::json!({"file_path": "/dev/null"}),
})
.await
.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
}
1 => {
tx.send(ApiEvent::Text("current task and progress".to_string()))
.await
.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
}
_ => {
tx.send(ApiEvent::Text("finished".to_string()))
.await
.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
}
}
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
#[async_trait::async_trait]
impl Provider for CompactionTrackingProvider {
fn name(&self) -> &str {
"compaction-tracking"
}
fn set_model(&mut self, _model: &str) {}
fn reset_session(&mut self) {
self.resets.fetch_add(1, Ordering::SeqCst);
}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let (tx, rx) = mpsc::channel(2);
tx.send(ApiEvent::Text("compacted summary".to_string()))
.await
.unwrap();
if self.complete {
tx.send(ApiEvent::Done).await.unwrap();
}
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct FailingProvider {
failure: ApiFailure,
calls: Arc<AtomicUsize>,
max_tokens_seen: Arc<Mutex<Vec<u32>>>,
}
#[async_trait::async_trait]
impl Provider for FailingProvider {
fn name(&self) -> &str {
"failing"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
self.calls.fetch_add(1, Ordering::SeqCst);
self.max_tokens_seen.lock().unwrap().push(max_tokens);
let (tx, rx) = mpsc::channel(2);
tx.send(ApiEvent::Error(self.failure.clone()))
.await
.unwrap();
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
struct ToolThenFailProvider {
failure: ApiFailure,
calls: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl Provider for ToolThenFailProvider {
fn name(&self) -> &str {
"tool-then-fail"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
self.calls.fetch_add(1, Ordering::SeqCst);
let (tx, rx) = mpsc::channel(4);
tx.send(ApiEvent::ToolUse {
id: "tu_1".to_string(),
name: "Read".to_string(),
input: serde_json::json!({"file_path": "/dev/null"}),
})
.await
.unwrap();
tx.send(ApiEvent::Error(self.failure.clone()))
.await
.unwrap();
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
async fn committed_attempts_for(failure: ApiFailure) -> usize {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(ToolThenFailProvider {
failure,
calls: calls.clone(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let _ = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await;
calls.load(Ordering::SeqCst)
}
#[tokio::test]
async fn a_committed_attempt_is_not_retried_for_malformed_tool_arguments() {
let attempts =
committed_attempts_for(ApiFailure::malformed_tool_arguments("bad args")).await;
assert_eq!(attempts, 1, "committed attempt must not retry");
}
#[tokio::test]
async fn a_committed_attempt_is_not_retried_for_an_output_limit() {
let attempts =
committed_attempts_for(ApiFailure::output_limit_exceeded("output limit")).await;
assert_eq!(
attempts, 1,
"escalating max_tokens must not reissue a committed batch"
);
}
#[tokio::test]
async fn a_committed_attempt_is_not_retried_for_a_context_overflow() {
let attempts =
committed_attempts_for(ApiFailure::new(ApiFailureKind::ContextExceeded, "too long"))
.await;
assert_eq!(
attempts, 1,
"compaction recovery must not reissue a committed batch"
);
}
async fn recovery_attempts_for(failure: ApiFailure) -> (usize, Vec<u32>) {
let calls = Arc::new(AtomicUsize::new(0));
let max_tokens_seen = Arc::new(Mutex::new(Vec::new()));
let provider = Box::new(FailingProvider {
failure,
calls: calls.clone(),
max_tokens_seen: max_tokens_seen.clone(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let _ = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await;
let seen = max_tokens_seen.lock().unwrap().clone();
(calls.load(Ordering::SeqCst), seen)
}
#[tokio::test]
async fn an_output_limit_escalates_max_tokens_without_compacting() {
let (attempts, max_tokens) =
recovery_attempts_for(ApiFailure::output_limit_exceeded("output limit")).await;
assert!(attempts > 1, "the turn should retry with a larger budget");
assert!(
max_tokens.windows(2).all(|pair| pair[1] > pair[0]),
"max_tokens must escalate on each retry, got {max_tokens:?}"
);
assert_eq!(
*max_tokens.last().unwrap(),
64_000,
"escalation stops at the ceiling"
);
}
#[tokio::test]
async fn mid_stream_output_retry_discards_rejected_text() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(MidStreamOutputRetryProvider {
calls: calls.clone(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let result = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 2);
assert_eq!(result, "recovered response");
}
struct TransientProvider {
failure: ApiFailure,
failures_before_success: usize,
calls: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl Provider for TransientProvider {
fn name(&self) -> &str {
"transient"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
if call < self.failures_before_success {
if call % 2 == 0 {
return Err(anyhow::Error::new(self.failure.clone()));
}
let (tx, rx) = mpsc::channel(4);
tx.send(ApiEvent::Text("partial".to_string()))
.await
.unwrap();
tx.send(ApiEvent::Error(self.failure.clone()))
.await
.unwrap();
drop(tx);
return Ok(ProviderStream::new(rx, cancel.child_token()));
}
let (tx, rx) = mpsc::channel(4);
tx.send(ApiEvent::Text("recovered".to_string()))
.await
.unwrap();
tx.send(ApiEvent::Done).await.unwrap();
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
fn transient_engine(
failure: ApiFailure,
failures_before_success: usize,
) -> (Engine, Arc<AtomicUsize>) {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(TransientProvider {
failure,
failures_before_success,
calls: calls.clone(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
engine.set_retry_backoff_base(std::time::Duration::from_millis(5));
(engine, calls)
}
#[tokio::test]
async fn transient_failures_are_retried_then_succeed() {
let (mut engine, calls) = transient_engine(
ApiFailure::new(ApiFailureKind::RateLimited, "429")
.with_status(Some(reqwest::StatusCode::TOO_MANY_REQUESTS)),
2,
);
let result = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(
result, "recovered",
"rejected partial text must be discarded"
);
assert_eq!(calls.load(Ordering::SeqCst), 3);
let timing = engine.execution_timing();
let statuses: Vec<&str> = timing
.model_rounds
.iter()
.map(|round| round.status.as_str())
.collect();
assert_eq!(statuses, vec!["retry", "retry", "completed"]);
assert_eq!(
timing.model_rounds[0].failure,
Some(ApiFailureKind::RateLimited)
);
assert!(engine.last_failure().is_none());
}
#[tokio::test]
async fn transient_failures_stop_after_the_retry_budget() {
let (mut engine, calls) = transient_engine(
ApiFailure::new(ApiFailureKind::Unavailable, "503")
.with_status(Some(reqwest::StatusCode::SERVICE_UNAVAILABLE)),
usize::MAX,
);
let error = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap_err();
assert_eq!(
calls.load(Ordering::SeqCst),
1 + MAX_TRANSIENT_RETRIES as usize
);
assert_eq!(
error.downcast_ref::<ApiFailure>().map(|f| f.kind),
Some(ApiFailureKind::Unavailable)
);
let failure = engine.last_failure().expect("failure recorded");
assert_eq!(failure.kind, ApiFailureKind::Unavailable);
assert!(failure.retryable);
assert_eq!(failure.http_status, Some(503));
assert_eq!(failure.attempts, 1 + MAX_TRANSIENT_RETRIES);
}
#[tokio::test]
async fn a_committed_attempt_is_not_retried_for_a_transient_failure() {
let attempts =
committed_attempts_for(ApiFailure::new(ApiFailureKind::RateLimited, "429")).await;
assert_eq!(attempts, 1);
}
#[tokio::test]
async fn retry_backoff_observes_cancellation() {
let (mut engine, calls) = transient_engine(
ApiFailure::new(ApiFailureKind::RateLimited, "429")
.with_retry_after(Some(std::time::Duration::from_secs(30))),
usize::MAX,
);
let cancel = tokio_util::sync::CancellationToken::new();
let canceller = cancel.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
canceller.cancel();
});
let started = Instant::now();
let result = engine.submit("go", cancel).await;
assert!(result.is_ok(), "cancellation ends the turn cleanly");
assert!(
started.elapsed() < std::time::Duration::from_secs(5),
"the 30s Retry-After must not be awaited past cancellation"
);
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn non_retryable_failures_are_recorded_without_retry() {
let (mut engine, calls) = transient_engine(
ApiFailure::new(ApiFailureKind::Authentication, "401")
.with_status(Some(reqwest::StatusCode::UNAUTHORIZED)),
usize::MAX,
);
let _ = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await;
assert_eq!(calls.load(Ordering::SeqCst), 1);
let failure = engine.last_failure().expect("failure recorded");
assert_eq!(failure.kind, ApiFailureKind::Authentication);
assert!(!failure.retryable);
assert_eq!(failure.attempts, 1);
}
#[tokio::test]
async fn an_unclassified_failure_triggers_no_recovery() {
let (attempts, _) = recovery_attempts_for(ApiFailure::other(
"internal error (request req_413_88): invalid max_tokens parameter",
))
.await;
assert_eq!(
attempts, 1,
"an unclassified failure must not trigger compaction or escalation"
);
}
#[tokio::test]
async fn a_context_overflow_attempts_compaction() {
let (attempts, _) =
recovery_attempts_for(ApiFailure::new(ApiFailureKind::ContextExceeded, "too long"))
.await;
assert_eq!(
attempts, 2,
"a context overflow must trigger a compaction attempt"
);
}
#[test]
fn stream_errors_carry_their_classification_to_the_turn_loop() {
let error = anyhow::Error::new(ApiFailure::malformed_tool_arguments("bad args"));
assert_eq!(
Engine::failure_kind(&error),
ApiFailureKind::MalformedToolArguments
);
let untyped = anyhow::anyhow!("invalid arguments for tool call Read (call_3)");
assert_eq!(
Engine::failure_kind(&untyped),
ApiFailureKind::Other,
"prose alone must not be treated as a classification"
);
}
#[tokio::test]
async fn malformed_tool_arguments_retry_once_without_persisting_rejected_text() {
let calls = Arc::new(AtomicUsize::new(0));
let systems = Arc::new(Mutex::new(Vec::new()));
let provider = Box::new(MalformedToolProvider {
calls: calls.clone(),
systems: systems.clone(),
recover: true,
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Default);
engine.set_system_prompt("base system prompt".to_string());
let response = engine
.submit("hello", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(response, "recovered response");
assert_eq!(calls.load(Ordering::SeqCst), 2);
let systems = systems.lock().unwrap();
assert_eq!(systems[0], "base system prompt");
assert!(systems[1].starts_with("base system prompt\n\n"));
assert!(systems[1].contains("before any tools executed"));
assert!(systems[1].contains("Reissue the entire intended tool-call batch"));
let MessageContent::Blocks(blocks) = &engine.messages()[1].content else {
panic!("expected assistant blocks");
};
assert!(matches!(
blocks.as_slice(),
[ContentBlock::Text { text }] if text == "recovered response"
));
}
#[tokio::test]
async fn malformed_tool_arguments_stop_after_one_retry() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(MalformedToolProvider {
calls: calls.clone(),
systems: Arc::new(Mutex::new(Vec::new())),
recover: false,
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Default);
let error = engine
.submit("hello", tokio_util::sync::CancellationToken::new())
.await
.unwrap_err();
assert_eq!(calls.load(Ordering::SeqCst), 2);
assert!(error
.to_string()
.contains("invalid arguments for tool call"));
assert_eq!(
engine.messages().len(),
1,
"rejected assistant attempts must not enter conversation history"
);
}
#[async_trait::async_trait]
impl Provider for ResetTrackingProvider {
fn name(&self) -> &str {
"reset-tracking"
}
fn set_model(&mut self, _model: &str) {}
fn reset_session(&mut self) {
self.resets
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let (_tx, rx) = mpsc::channel(1);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
#[test]
fn set_messages_resets_session_scoped_engine_state() {
let resets = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let provider = Box::new(ResetTrackingProvider {
resets: resets.clone(),
});
let mut engine = Engine::new(
provider,
ToolRegistry::without_agent_for_tests(),
PermissionChecker::new(PermissionMode::Default),
"private-model",
);
engine
.cost
.set_pricing_override(Some(crate::cost::ModelPricing {
input: 2.0,
output: 4.0,
cache_read: 0.5,
cache_write: 1.0,
}));
engine.cost.add_usage(&crate::api::types::Usage {
input_tokens: 500,
output_tokens: 200,
cache_read_tokens: 100,
cache_creation_tokens: 50,
provider_cost_usd: None,
});
engine
.steering_queue()
.lock()
.unwrap()
.push_back("stale steering".to_string());
engine.permissions.always_allow("Write");
engine.permissions.always_allow_command("cargo test");
engine.tool_trace.push(ToolTraceEntry {
id: "old-tool".to_string(),
name: "Bash".to_string(),
input: serde_json::json!({"command": "true"}),
output: String::new(),
is_error: false,
read_only: true,
started_after_ms: 0,
duration_ms: 0,
});
engine.set_messages(vec![Message::user("loaded session")]);
assert_eq!(resets.load(std::sync::atomic::Ordering::SeqCst), 1);
assert_eq!(engine.message_count(), 1);
assert!(engine.tool_trace().is_empty());
assert!(engine.steering_queue().lock().unwrap().is_empty());
assert_eq!(engine.cost.input_tokens, 0);
assert_eq!(engine.cost.output_tokens, 0);
assert!(matches!(
engine.permissions.check(
"Write",
&serde_json::json!({"file_path": "/tmp/test"}),
false
),
PermissionResult::Ask { .. }
));
assert!(matches!(
engine
.permissions
.check("Bash", &serde_json::json!({"command": "cargo test"}), false),
PermissionResult::Ask { .. }
));
engine.cost.add_usage(&crate::api::types::Usage {
input_tokens: 1_000_000,
output_tokens: 0,
cache_read_tokens: 0,
cache_creation_tokens: 0,
provider_cost_usd: None,
});
assert_eq!(engine.cost.total_cost_usd(), 2.0);
}
fn usage(input: u32, cache_read: u32) -> crate::api::types::Usage {
crate::api::types::Usage {
input_tokens: input,
output_tokens: 0,
cache_read_tokens: cache_read,
cache_creation_tokens: 0,
provider_cost_usd: None,
}
}
#[test]
fn context_estimate_falls_back_to_message_scan_without_a_baseline() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine.messages_mut().push(Message::user("hello there"));
assert_eq!(
engine.estimated_context_tokens(),
compact::estimate_tokens(engine.messages())
);
let snapshot = engine.context_usage();
assert!(!snapshot.provider_anchored);
assert!(snapshot.short_status().starts_with("ctx ~"));
}
#[test]
fn context_estimate_anchors_to_provider_reported_usage() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine.messages_mut().push(Message::user("first"));
engine.record_request_usage(&usage(9_000, 3_000), 1);
assert_eq!(engine.estimated_context_tokens(), 12_000);
engine.messages_mut().push(Message::user("second message"));
let delta = compact::estimate_tokens(&engine.messages()[1..]);
assert!(delta > 0);
assert_eq!(engine.estimated_context_tokens(), 12_000 + delta);
}
#[test]
fn context_snapshot_reports_utilization_threshold_and_headroom() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine.context_window = 20_000;
engine.auto_compact_threshold = 0.8;
engine.messages_mut().push(Message::user("first"));
engine.record_request_usage(&usage(10_000, 0), 1);
let snapshot = engine.context_usage();
assert_eq!(snapshot.estimated_tokens, 10_000);
assert_eq!(snapshot.context_window, 20_000);
assert_eq!(snapshot.compact_threshold_tokens, 16_000);
assert!(snapshot.provider_anchored);
assert_eq!(snapshot.utilization_percent(), 50);
assert_eq!(snapshot.compact_headroom_tokens(), 6_000);
assert_eq!(snapshot.headroom_tokens(), 10_000);
assert_eq!(snapshot.short_status(), "ctx 10k/20k (50%)");
assert!(engine
.context_report()
.contains("6000 tokens until threshold"));
}
#[test]
fn compaction_clears_the_baseline_so_it_cannot_re_trigger() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine
.messages_mut()
.push(Message::user("a long conversation"));
engine.record_request_usage(&usage(150_000, 0), 1);
assert_eq!(engine.estimated_context_tokens(), 150_000);
engine.commit_compacted_messages(vec![Message::user("summary")]);
assert!(
engine.estimated_context_tokens() < 1_000,
"post-compaction estimate must not inherit the old total"
);
}
#[tokio::test]
async fn long_tool_loop_compacts_between_model_rounds() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(WithinTurnCompactionProvider {
calls: calls.clone(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let result = engine
.submit(
"repair the host",
tokio_util::sync::CancellationToken::new(),
)
.await
.unwrap();
assert_eq!(result, "finished");
assert_eq!(calls.load(Ordering::SeqCst), 3);
assert!(engine.messages().iter().any(|message| {
matches!(
&message.content,
MessageContent::Text(text) if text == "repair the host"
)
}));
}
#[test]
fn a_baseline_covering_more_messages_than_history_is_discarded() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine.messages_mut().push(Message::user("one"));
engine.messages_mut().push(Message::user("two"));
engine.record_request_usage(&usage(5_000, 0), 2);
engine.messages_mut().pop();
assert_eq!(
engine.estimated_context_tokens(),
compact::estimate_tokens(engine.messages())
);
}
#[test]
fn zero_usage_does_not_replace_a_good_baseline() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine.messages_mut().push(Message::user("first"));
engine.record_request_usage(&usage(20_000, 0), 1);
engine.record_request_usage(&usage(0, 0), 1);
assert_eq!(engine.estimated_context_tokens(), 20_000);
}
#[test]
fn resolved_model_metadata_configures_compaction_and_cost() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Default,
);
engine.set_model_metadata(crate::model::ModelMetadata {
context_window: 64_000,
pricing: Some(crate::cost::ModelPricing {
input: 2.0,
output: 4.0,
cache_read: 0.5,
cache_write: 1.0,
}),
});
engine.cost.add_usage(&crate::api::types::Usage {
input_tokens: 1_000_000,
output_tokens: 0,
cache_read_tokens: 0,
cache_creation_tokens: 0,
provider_cost_usd: None,
});
assert_eq!(engine.context_window, 64_000);
assert_eq!(engine.cost.total_cost_usd(), 2.0);
}
#[test]
fn transcript_checkpoint_preserves_running_engine_state() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("transcript.json");
engine.set_transcript_checkpoint(path.clone());
engine.start_recording();
engine.messages_mut().push(Message::user("repair it"));
engine.model_trace.push(ModelTraceEntry {
index: 1,
started_after_ms: 0,
duration_ms: 10,
failure: None,
status: "completed".to_string(),
usage: None,
});
engine.checkpoint_transcript();
let value: serde_json::Value =
serde_json::from_slice(&std::fs::read(path).unwrap()).unwrap();
assert_eq!(value["outcome"]["status"], "running");
assert_eq!(value["messages"][0]["content"], "repair it");
assert_eq!(value["timing"]["model_rounds"][0]["status"], "completed");
}
#[tokio::test]
async fn test_parallel_tool_execution() {
let provider = Box::new(MockProvider);
let tools = ToolRegistry::without_agent_for_tests();
let permissions = PermissionChecker::new(PermissionMode::Bypass);
let mut engine = Engine {
provider,
tools,
permissions,
messages: vec![],
system_prompt: String::new(),
model: "test".to_string(),
model_binding: None,
max_tokens: 1000,
context_window: 128_000,
auto_compact_threshold: 0.8,
steering: SteeringQueue::default(),
pending_images: Vec::new(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
tool_trace: Vec::new(),
model_trace: Vec::new(),
trace_started_at: Some(Instant::now()),
trace_duration_ms: None,
transcript_checkpoint: None,
cost: CostTracker::new("test"),
last_request_usage: None,
last_compaction_notice: None,
last_failure: None,
retry_backoff_base: DEFAULT_RETRY_BACKOFF_BASE,
};
let tool_uses = vec![
(
"test1".to_string(),
"Read".to_string(),
serde_json::json!({"file_path": "/dev/null"}),
),
(
"test2".to_string(),
"Glob".to_string(),
serde_json::json!({"pattern": "*.rs"}),
),
(
"test3".to_string(),
"Read".to_string(),
serde_json::json!({"file_path": "/dev/null"}),
),
];
let start = Instant::now();
let (batch_tx, mut batch_rx) = mpsc::channel(64);
let drain = tokio::spawn(async move { while batch_rx.recv().await.is_some() {} });
let (blocks, _interrupted) = engine
.execute_tool_batch(
&tool_uses,
&batch_tx,
false,
&tokio_util::sync::CancellationToken::new(),
)
.await;
drop(batch_tx);
drain.await.unwrap();
let duration = start.elapsed();
assert_eq!(blocks.len(), 3, "Should have 3 result blocks");
for (i, block) in blocks.iter().enumerate() {
if let ContentBlock::ToolResult { tool_use_id, .. } = block {
let expected_id = format!("test{}", i + 1);
assert_eq!(
tool_use_id, &expected_id,
"Results should be in original order"
);
} else {
panic!("Expected ToolResult block");
}
}
println!("Parallel execution took: {duration:?}");
}
#[tokio::test]
async fn test_mixed_readonly_and_write_tools() {
let provider = Box::new(MockProvider);
let tools = ToolRegistry::without_agent_for_tests();
let permissions = PermissionChecker::new(PermissionMode::Bypass);
let mut engine = Engine {
provider,
tools,
permissions,
messages: vec![],
system_prompt: String::new(),
model: "test".to_string(),
model_binding: None,
max_tokens: 1000,
context_window: 128_000,
auto_compact_threshold: 0.8,
steering: SteeringQueue::default(),
pending_images: Vec::new(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
tool_trace: Vec::new(),
model_trace: Vec::new(),
trace_started_at: Some(Instant::now()),
trace_duration_ms: None,
transcript_checkpoint: None,
cost: CostTracker::new("test"),
last_request_usage: None,
last_compaction_notice: None,
last_failure: None,
retry_backoff_base: DEFAULT_RETRY_BACKOFF_BASE,
};
let tool_uses = vec![
(
"test1".to_string(),
"Read".to_string(), serde_json::json!({"file_path": "/dev/null"}),
),
(
"test2".to_string(),
"Bash".to_string(), serde_json::json!({"command": "echo test"}),
),
(
"test3".to_string(),
"Glob".to_string(), serde_json::json!({"pattern": "*.rs"}),
),
];
let (batch_tx, mut batch_rx) = mpsc::channel(64);
let drain = tokio::spawn(async move { while batch_rx.recv().await.is_some() {} });
let (blocks, _interrupted) = engine
.execute_tool_batch(
&tool_uses,
&batch_tx,
false,
&tokio_util::sync::CancellationToken::new(),
)
.await;
drop(batch_tx);
drain.await.unwrap();
assert_eq!(blocks.len(), 3, "Should have 3 result blocks");
for (i, block) in blocks.iter().enumerate() {
if let ContentBlock::ToolResult { tool_use_id, .. } = block {
let expected_id = format!("test{}", i + 1);
assert_eq!(tool_use_id, &expected_id, "Results should maintain order");
}
}
}
#[tokio::test]
async fn tool_batch_reads_observe_preceding_writes() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("ordered.txt");
std::fs::write(&path, "before").unwrap();
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
let calls = vec![
(
"before".into(),
"Read".into(),
serde_json::json!({"file_path": path}),
),
(
"write".into(),
"Write".into(),
serde_json::json!({"file_path": path, "content": "after"}),
),
(
"after".into(),
"Read".into(),
serde_json::json!({"file_path": path}),
),
(
"write-again".into(),
"Write".into(),
serde_json::json!({"file_path": path, "content": "final"}),
),
(
"final".into(),
"Read".into(),
serde_json::json!({"file_path": path}),
),
];
let (tx, _rx) = mpsc::channel(64);
let (results, interrupted) = engine
.execute_tool_batch(
&calls,
&tx,
false,
&tokio_util::sync::CancellationToken::new(),
)
.await;
assert!(!interrupted);
for (index, expected) in [(0, "before"), (2, "after"), (4, "final")] {
let ContentBlock::ToolResult {
tool_use_id,
content,
is_error,
} = &results[index]
else {
panic!("expected tool result");
};
assert_eq!(tool_use_id, expected);
assert_ne!(*is_error, Some(true));
assert!(
content.contains(expected),
"{content:?} should contain {expected}"
);
}
assert_eq!(std::fs::read_to_string(path).unwrap(), "final");
}
struct ConcurrencyProbe {
active: Arc<AtomicUsize>,
peak: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl crate::tools::Tool for ConcurrencyProbe {
fn name(&self) -> &str {
"ConcurrencyProbe"
}
fn description(&self) -> &str {
"Measure concurrent calls"
}
fn input_schema(&self) -> serde_json::Value {
serde_json::json!({"type": "object"})
}
fn is_read_only(&self) -> bool {
true
}
async fn execute(
&self,
input: serde_json::Value,
_cancel: tokio_util::sync::CancellationToken,
) -> Result<crate::tools::ToolOutput> {
let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
self.peak.fetch_max(active, Ordering::SeqCst);
tokio::task::yield_now().await;
self.active.fetch_sub(1, Ordering::SeqCst);
Ok(crate::tools::ToolOutput {
content: input.to_string(),
is_error: false,
})
}
}
#[tokio::test]
async fn tool_batch_parallelism_is_bounded_and_results_stay_ordered() {
let active = Arc::new(AtomicUsize::new(0));
let peak = Arc::new(AtomicUsize::new(0));
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine.tools.add_tools(vec![Box::new(ConcurrencyProbe {
active: active.clone(),
peak: peak.clone(),
})]);
let calls: Vec<_> = (0..MAX_PARALLEL_TOOLS * 2 + 1)
.map(|i| {
(
i.to_string(),
"ConcurrencyProbe".into(),
serde_json::json!({"index": i}),
)
})
.collect();
let (tx, _rx) = mpsc::channel(64);
let (results, interrupted) = engine
.execute_tool_batch(
&calls,
&tx,
false,
&tokio_util::sync::CancellationToken::new(),
)
.await;
assert!(!interrupted);
assert_eq!(active.load(Ordering::SeqCst), 0);
assert_eq!(peak.load(Ordering::SeqCst), MAX_PARALLEL_TOOLS);
assert_eq!(results.len(), calls.len());
for (index, result) in results.iter().enumerate() {
let ContentBlock::ToolResult {
tool_use_id,
content,
is_error,
} = result
else {
panic!("expected tool result");
};
assert_eq!(tool_use_id, &index.to_string());
assert_eq!(content, &calls[index].2.to_string());
assert_ne!(*is_error, Some(true));
}
}
fn steering_engine(
first_round: Vec<(String, String, serde_json::Value)>,
push_on_first_call: Option<String>,
) -> Engine {
crate::test_support::scripted_engine(
first_round,
push_on_first_call,
PermissionMode::Bypass,
)
}
struct CountingPlugin {
trigger: HookTrigger,
count: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl Plugin for CountingPlugin {
fn name(&self) -> &str {
"counter"
}
fn trigger(&self) -> &HookTrigger {
&self.trigger
}
async fn execute(
&self,
_env_vars: Option<&HashMap<String, String>>,
) -> Result<Option<String>> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(None)
}
}
struct VerdictPlugin {
output: String,
seen_env: Arc<std::sync::Mutex<Option<HashMap<String, String>>>>,
}
#[async_trait::async_trait]
impl Plugin for VerdictPlugin {
fn name(&self) -> &str {
"verdict"
}
fn trigger(&self) -> &HookTrigger {
&HookTrigger::OnPermissionCheck
}
async fn execute(
&self,
env_vars: Option<&HashMap<String, String>>,
) -> Result<Option<String>> {
*self.seen_env.lock().unwrap() = env_vars.cloned();
Ok(Some(self.output.clone()))
}
}
async fn tool_result_after_verdict(
mode: PermissionMode,
verdict: &str,
) -> (String, HashMap<String, String>) {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(ReadThenDoneProvider {
calls: calls.clone(),
});
let mut engine = Engine::for_tests(provider, SteeringQueue::default(), mode);
let seen_env = Arc::new(std::sync::Mutex::new(None));
let mut registry = PluginRegistry::new();
registry.add(Box::new(VerdictPlugin {
output: verdict.to_string(),
seen_env: seen_env.clone(),
}));
engine.set_plugins(Arc::new(registry));
engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
let MessageContent::Blocks(blocks) = &engine.messages()[2].content else {
panic!("expected tool_result blocks");
};
let content = blocks
.iter()
.find_map(|block| match block {
ContentBlock::ToolResult { content, .. } => Some(content.clone()),
_ => None,
})
.expect("tool result");
let env = seen_env.lock().unwrap().clone().expect("hook ran");
(content, env)
}
struct ReadThenDoneProvider {
calls: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl Provider for ReadThenDoneProvider {
fn name(&self) -> &str {
"read-then-done"
}
fn set_model(&mut self, _model: &str) {}
async fn stream(
&self,
_messages: &[Message],
_system: &str,
_tools: &[ToolDefinition],
_max_tokens: u32,
cancel: tokio_util::sync::CancellationToken,
) -> Result<ProviderStream> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
let (tx, rx) = mpsc::channel(4);
if call == 0 {
tx.send(ApiEvent::ToolUse {
id: "read-1".to_string(),
name: "Read".to_string(),
input: serde_json::json!({"file_path": "/dev/null"}),
})
.await
.unwrap();
} else {
tx.send(ApiEvent::Text("done".to_string())).await.unwrap();
}
tx.send(ApiEvent::Done).await.unwrap();
drop(tx);
Ok(ProviderStream::new(rx, cancel.child_token()))
}
}
#[tokio::test]
async fn permission_hook_can_deny_a_tool_the_mode_allows() {
let (content, env) = tool_result_after_verdict(
PermissionMode::Bypass,
r#"{"decision":"deny","reason":"policy says no"}"#,
)
.await;
assert!(content.contains("Permission denied"), "{content}");
assert!(
content.contains("blocked by hook verdict: policy says no"),
"{content}"
);
assert_eq!(env["CLAUX_TOOL_NAME"], "Read");
assert_eq!(env["CLAUX_PERMISSION_DECISION"], "allow");
assert_eq!(env["CLAUX_PERMISSION_MODE"], "bypass");
assert_eq!(env["CLAUX_TOOL_READ_ONLY"], "true");
assert!(env["CLAUX_TOOL_INPUT"].contains("/dev/null"));
}
#[tokio::test]
async fn permission_hook_ask_denies_in_non_interactive_mode() {
let (content, _) =
tool_result_after_verdict(PermissionMode::Bypass, r#"{"decision":"ask"}"#).await;
assert!(content.contains("Permission denied"), "{content}");
assert!(content.contains("one-shot mode has no prompt"), "{content}");
}
#[tokio::test]
async fn permission_hook_with_no_opinion_leaves_the_decision_alone() {
let (content, _) = tool_result_after_verdict(PermissionMode::Bypass, "").await;
assert!(!content.contains("Permission denied"), "{content}");
}
#[tokio::test]
async fn one_shot_submit_fires_tool_and_turn_hooks() {
let starts = Arc::new(AtomicUsize::new(0));
let completes = Arc::new(AtomicUsize::new(0));
let turns = Arc::new(AtomicUsize::new(0));
let mut plugins = PluginRegistry::new();
for (trigger, count) in [
(HookTrigger::OnToolStart, starts.clone()),
(HookTrigger::OnToolComplete, completes.clone()),
(HookTrigger::OnTurnEnd, turns.clone()),
] {
plugins.add(Box::new(CountingPlugin { trigger, count }));
}
let mut engine = steering_engine(
vec![crate::test_support::tool_use(
"read-1",
"Read",
serde_json::json!({"file_path": "/dev/null"}),
)],
None,
);
engine.set_plugins(Arc::new(plugins));
engine
.submit("read it", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(starts.load(Ordering::SeqCst), 1);
assert_eq!(completes.load(Ordering::SeqCst), 1);
assert_eq!(turns.load(Ordering::SeqCst), 1);
}
async fn run_streaming(engine: &mut Engine, prompt: &str) {
let (tx, mut rx) = mpsc::channel(64);
let drain = tokio::spawn(async move { while rx.recv().await.is_some() {} });
engine
.submit_streaming(prompt, tx, tokio_util::sync::CancellationToken::new())
.await
.unwrap();
drain.await.unwrap();
}
#[tokio::test]
async fn test_steering_message_injected_after_tool_results() {
let mut engine = steering_engine(
vec![(
"tu_1".to_string(),
"Glob".to_string(),
serde_json::json!({"pattern": "*.does-not-exist"}),
)],
Some("also check the auth module".to_string()),
);
run_streaming(&mut engine, "do a deep review").await;
let msgs = engine.messages();
assert_eq!(msgs.len(), 4, "got: {msgs:?}");
assert_eq!(msgs[0].role, "user");
assert_eq!(msgs[1].role, "assistant");
assert_eq!(msgs[2].role, "user"); assert_eq!(msgs[3].role, "user");
match &msgs[3].content {
crate::api::MessageContent::Text(t) => {
assert_eq!(t, "also check the auth module")
}
other => panic!("expected steering text message, got {other:?}"),
}
assert!(engine.steering_queue().lock().unwrap().is_empty());
}
#[tokio::test]
async fn test_pending_steering_skips_whole_batch() {
let mut engine = steering_engine(
vec![
(
"tu_1".to_string(),
"Glob".to_string(),
serde_json::json!({"pattern": "*.a"}),
),
(
"tu_2".to_string(),
"Glob".to_string(),
serde_json::json!({"pattern": "*.b"}),
),
],
Some("wrong direction, stop".to_string()),
);
run_streaming(&mut engine, "explore").await;
let msgs = engine.messages();
let crate::api::MessageContent::Blocks(blocks) = &msgs[2].content else {
panic!("expected tool results, got {msgs:?}");
};
assert_eq!(blocks.len(), 2);
for block in blocks {
match block {
ContentBlock::ToolResult {
content, is_error, ..
} => {
assert_eq!(content, Engine::SKIPPED_FOR_STEERING);
assert_eq!(*is_error, Some(true));
}
other => panic!("expected ToolResult, got {other:?}"),
}
}
}
#[tokio::test]
async fn test_steering_cancels_running_tool() {
let mut engine = steering_engine(
vec![(
"tu_1".to_string(),
"Bash".to_string(),
serde_json::json!({"command": "sleep 5"}),
)],
None,
);
let steering = engine.steering_queue();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
steering
.lock()
.unwrap()
.push_back("no, run it in nix-shell instead".to_string());
});
let start = std::time::Instant::now();
run_streaming(&mut engine, "run the tests").await;
assert!(
start.elapsed() < std::time::Duration::from_secs(3),
"steering should cancel the running tool, not wait it out (took {:?})",
start.elapsed()
);
let last = engine.messages().last().unwrap();
match &last.content {
crate::api::MessageContent::Text(t) => {
assert_eq!(t, "no, run it in nix-shell instead")
}
other => panic!("expected steering message last, got {other:?}"),
}
}
#[tokio::test]
async fn test_cancellation_ends_turn_with_paired_results() {
let mut engine = steering_engine(
vec![(
"tu_1".to_string(),
"Bash".to_string(),
serde_json::json!({"command": "sleep 5"}),
)],
None,
);
let cancel = tokio_util::sync::CancellationToken::new();
let canceller = cancel.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
canceller.cancel();
});
let (tx, mut rx) = mpsc::channel(64);
let events = tokio::spawn(async move {
let mut interrupted = false;
while let Some(ev) = rx.recv().await {
if matches!(ev, StreamEvent::Interrupted) {
interrupted = true;
}
}
interrupted
});
let start = std::time::Instant::now();
engine.submit_streaming("run it", tx, cancel).await.unwrap();
assert!(
start.elapsed() < std::time::Duration::from_secs(3),
"cancellation should not wait out the tool (took {:?})",
start.elapsed()
);
assert!(events.await.unwrap(), "Interrupted event must be emitted");
let msgs = engine.messages();
let crate::api::MessageContent::Blocks(blocks) = &msgs.last().unwrap().content else {
panic!("expected tool results last, got {msgs:?}");
};
assert!(matches!(
&blocks[0],
ContentBlock::ToolResult {
is_error: Some(true),
..
}
));
}
#[tokio::test]
async fn test_submit_returns_text_without_notices() {
let mut engine = steering_engine(
vec![(
"tu_1".to_string(),
"Glob".to_string(),
serde_json::json!({"pattern": "*.x"}),
)],
Some("check auth too".to_string()),
);
let text = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(text, "working on it", "notices must not leak into text");
let last = engine.messages().last().unwrap();
match &last.content {
crate::api::MessageContent::Text(t) => assert_eq!(t, "check auth too"),
other => panic!("expected steering message last, got {other:?}"),
}
}
#[tokio::test]
async fn reasoning_is_hidden_from_output_and_retained_in_history() {
let mut engine = Engine::for_tests(
Box::new(ReasoningProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
let text = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(text, "answer");
let MessageContent::Blocks(blocks) = &engine.messages()[1].content else {
panic!("expected assistant blocks");
};
assert!(matches!(
&blocks[0],
ContentBlock::Reasoning { text: Some(text), details }
if text == "private thought" && details[0]["text"] == "preserve me"
));
assert!(matches!(
&blocks[1],
ContentBlock::Text { text } if text == "answer"
));
}
#[tokio::test]
async fn test_submit_rejects_stream_closed_without_done() {
let mut engine = Engine::for_tests(
Box::new(TruncatedProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
let error = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap_err();
assert!(error.to_string().contains("without completion"));
assert_eq!(
engine.messages().len(),
1,
"partial assistant content must not be committed to history"
);
assert_eq!(engine.messages()[0].role, "user");
}
#[tokio::test]
async fn test_submit_rejects_empty_completed_response() {
let mut engine = Engine::for_tests(
Box::new(EmptyCompletionProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
let error = engine
.submit("go", tokio_util::sync::CancellationToken::new())
.await
.unwrap_err();
assert_eq!(
error.to_string(),
"provider protocol error: response completed without assistant text or tool calls"
);
assert_eq!(
engine.messages().len(),
1,
"empty assistant content must not be committed to history"
);
assert_eq!(engine.execution_timing().model_rounds[0].status, "error");
}
#[tokio::test]
async fn test_compact_rejects_stream_closed_without_done() {
let mut engine = Engine::for_tests(
Box::new(TruncatedProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine
.messages_mut()
.push(Message::user("important context"));
let error = engine.compact().await.unwrap_err();
assert!(error.to_string().contains("without completion"));
assert_eq!(
engine.messages().len(),
1,
"failed compaction must preserve the original history"
);
}
#[tokio::test]
async fn summary_compaction_observes_turn_cancellation() {
let mut engine = Engine::for_tests(
Box::new(HangingProvider),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine
.messages_mut()
.push(Message::user("important context"));
let cancel = tokio_util::sync::CancellationToken::new();
let cancel_from_task = cancel.clone();
tokio::spawn(async move {
tokio::task::yield_now().await;
cancel_from_task.cancel();
});
let error = tokio::time::timeout(
std::time::Duration::from_secs(1),
engine.compact_with_cancel(&cancel, Continuation::ResumeTask),
)
.await
.expect("compaction should stop promptly")
.unwrap_err();
assert!(error.to_string().contains("cancelled"));
assert_eq!(engine.messages().len(), 1);
assert!(matches!(
&engine.messages()[0].content,
MessageContent::Text(text) if text == "important context"
));
}
#[tokio::test]
async fn old_tool_rounds_are_summarized_before_history_is_replaced() {
let resets = Arc::new(AtomicUsize::new(0));
let requests = Arc::new(Mutex::new(Vec::new()));
let provider = Box::new(RecordingSummaryProvider {
resets: resets.clone(),
requests: requests.clone(),
summary: "Objective: finish the original task. Progress: README.md was read.".into(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let old_content = "old context ".repeat(1_000);
engine.messages_mut().extend([
Message::user(&old_content),
Message::assistant_text(&old_content),
Message::assistant_blocks(vec![ContentBlock::ToolUse {
id: "call_1".to_string(),
name: "Read".to_string(),
input: serde_json::json!({"file_path": "README.md"}),
}]),
Message::tool_results(vec![ContentBlock::ToolResult {
tool_use_id: "call_1".to_string(),
content: "contents".to_string(),
is_error: None,
}]),
]);
for index in 4..13 {
engine
.messages_mut()
.push(Message::user(&format!("recent message {index}")));
}
let original = serde_json::to_value(engine.messages()).unwrap();
let result = engine.compact().await.unwrap();
assert_eq!(engine.messages().len(), 2);
let requests = requests.lock().unwrap();
assert_eq!(serde_json::to_value(&requests[0][..13]).unwrap(), original);
assert_eq!(
serde_json::to_value(&engine.messages()[0]).unwrap(),
original[0]
);
assert!(result.contains("Compacted via summary:"));
assert!(result.contains("tokens"));
assert_eq!(resets.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn summary_compaction_resets_provider_cursor_after_rewriting_history() {
let resets = Arc::new(AtomicUsize::new(0));
let provider = Box::new(CompactionTrackingProvider {
resets: resets.clone(),
complete: true,
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let large_message = "context ".repeat(10_000);
for index in 0..13 {
engine
.messages_mut()
.push(Message::user(&format!("{index}: {large_message}")));
}
engine.compact().await.unwrap();
assert_eq!(engine.messages().len(), 2);
assert!(!has_consecutive_user_messages(engine.messages()));
assert_eq!(resets.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn repeated_compaction_and_resume_retain_request_and_task_handoff() {
let requests = Arc::new(Mutex::new(Vec::new()));
let resets = Arc::new(AtomicUsize::new(0));
let summary = "Objective: fix parsing. Constraints: keep the API; use branches only. \
Decisions: preserve UTF-8. Progress: src/parser.rs updated; tests pass. \
Outstanding work: add malformed-input coverage.";
let mut engine = Engine::for_tests(
Box::new(RecordingSummaryProvider {
resets,
requests: requests.clone(),
summary: summary.into(),
}),
SteeringQueue::default(),
PermissionMode::Bypass,
);
let request = Message::user("Fix parsing; preserve the public API. Work in a worktree.");
engine.messages_mut().extend([
request.clone(),
Message::assistant_text("I will inspect the parser."),
Message::user("Correction: use branches only, no worktrees."),
Message::assistant_text(&"parser investigation and logs ".repeat(1_000)),
]);
engine.compact().await.unwrap();
assert_eq!(
serde_json::to_value(&engine.messages()[0]).unwrap(),
serde_json::to_value(&request).unwrap()
);
assert!(
matches!(&engine.messages()[1].content, MessageContent::Text(text) if text == summary)
);
let saved = serde_json::to_string(engine.messages()).unwrap();
let restored = crate::session::repair_history(serde_json::from_str(&saved).unwrap());
engine.set_messages(restored);
engine.messages_mut().extend([
Message::user("Continue with malformed-input coverage."),
Message::assistant_text(&"more investigation and logs ".repeat(1_000)),
]);
engine.compact().await.unwrap();
assert_eq!(
serde_json::to_value(&engine.messages()[0]).unwrap(),
serde_json::to_value(&request).unwrap()
);
let requests = requests.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(
matches!(&requests[0][2].content, MessageContent::Text(text) if text.contains("no worktrees"))
);
assert!(matches!(&requests[1][1].content, MessageContent::Text(text) if text == summary));
assert!(
matches!(&requests[1].last().unwrap().content, MessageContent::Text(text) if text.contains("Later user instructions supersede earlier"))
);
assert!(!has_consecutive_user_messages(engine.messages()));
}
#[tokio::test]
async fn empty_or_oversized_summary_preserves_history_and_resets_cursor() {
for summary in [" ".to_string(), "irrelevant expansion ".repeat(1_000)] {
let resets = Arc::new(AtomicUsize::new(0));
let mut engine = Engine::for_tests(
Box::new(RecordingSummaryProvider {
resets: resets.clone(),
requests: Arc::new(Mutex::new(Vec::new())),
summary,
}),
SteeringQueue::default(),
PermissionMode::Bypass,
);
engine.messages_mut().extend([
Message::user("Fix parsing; keep the public API."),
Message::assistant_text(
"Investigated src/parser.rs; next add regression coverage.",
),
]);
let original = serde_json::to_value(engine.messages()).unwrap();
let error = engine.compact().await.unwrap_err();
assert!(error.to_string().contains("history preserved"));
assert_eq!(serde_json::to_value(engine.messages()).unwrap(), original);
assert_eq!(
resets.load(Ordering::SeqCst),
1,
"a completed but rejected summary must not leave a continuation cursor"
);
}
}
#[tokio::test]
async fn turn_start_compaction_does_not_produce_consecutive_user_messages() {
let resets = Arc::new(AtomicUsize::new(0));
let provider = Box::new(CompactionTrackingProvider {
resets: resets.clone(),
complete: true,
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let large_message = "context ".repeat(10_000);
for index in 0..13 {
engine
.messages_mut()
.push(Message::user(&format!("{index}: {large_message}")));
}
engine
.submit("next prompt", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert_eq!(
resets.load(Ordering::SeqCst),
1,
"auto-compact must have run"
);
let messages = engine.messages();
assert!(matches!(
&messages[0].content,
MessageContent::Text(text) if text == &format!("0: {large_message}")
));
assert!(
!has_consecutive_user_messages(messages),
"turn-start compaction must let the incoming prompt be the user turn"
);
assert!(messages.iter().any(|message| {
matches!(&message.content, MessageContent::Text(text) if text == "next prompt")
}));
}
#[tokio::test]
async fn mid_turn_compaction_keeps_a_continuation_marker() {
let calls = Arc::new(AtomicUsize::new(0));
let provider = Box::new(WithinTurnCompactionProvider {
calls: calls.clone(),
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
engine
.submit(
"repair the host",
tokio_util::sync::CancellationToken::new(),
)
.await
.unwrap();
assert!(engine.messages().iter().any(|message| {
matches!(
&message.content,
MessageContent::Text(text)
if text == "Continue with the outstanding task described above."
)
}));
assert!(!has_consecutive_user_messages(engine.messages()));
}
fn has_consecutive_user_messages(messages: &[Message]) -> bool {
messages
.windows(2)
.any(|pair| pair[0].role == "user" && pair[1].role == "user")
}
#[tokio::test]
async fn short_messages_are_summarized_without_snipping() {
let resets = Arc::new(AtomicUsize::new(0));
let provider = Box::new(CompactionTrackingProvider {
resets: resets.clone(),
complete: true,
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
for index in 0..13 {
engine
.messages_mut()
.push(Message::user(&index.to_string()));
}
engine.compact().await.unwrap();
assert_eq!(
engine.messages().len(),
2,
"compaction retains the original request and a summary"
);
assert_eq!(resets.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn failed_summary_preserves_entire_history() {
let resets = Arc::new(AtomicUsize::new(0));
let provider = Box::new(CompactionTrackingProvider {
resets: resets.clone(),
complete: false,
});
let mut engine =
Engine::for_tests(provider, SteeringQueue::default(), PermissionMode::Bypass);
let large_message = "context ".repeat(10_000);
for index in 0..13 {
engine
.messages_mut()
.push(Message::user(&format!("{index}: {large_message}")));
}
let original = serde_json::to_value(engine.messages()).unwrap();
let error = engine.compact().await.unwrap_err();
assert!(error.to_string().contains("without completion"));
assert_eq!(
serde_json::to_value(engine.messages()).unwrap(),
original,
"failed summarization must not discard any history"
);
assert_eq!(
resets.load(Ordering::SeqCst),
0,
"failed compaction must preserve provider continuation state"
);
}
#[tokio::test]
async fn test_steering_preempts_in_non_streaming_submit() {
let mut engine = steering_engine(
vec![(
"tu_1".to_string(),
"Bash".to_string(),
serde_json::json!({"command": "sleep 5"}),
)],
None,
);
let steering = engine.steering_queue();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
steering
.lock()
.unwrap()
.push_back("stop, wrong command".to_string());
});
let start = std::time::Instant::now();
engine
.submit("run it", tokio_util::sync::CancellationToken::new())
.await
.unwrap();
assert!(
start.elapsed() < std::time::Duration::from_secs(3),
"steering should cancel the running tool via submit() too (took {:?})",
start.elapsed()
);
}
#[tokio::test]
async fn test_unknown_tool_yields_error_block_not_abort() {
let provider = Box::new(MockProvider);
let tools = ToolRegistry::without_agent_for_tests();
let permissions = PermissionChecker::new(PermissionMode::Bypass);
let mut engine = Engine {
provider,
tools,
permissions,
messages: vec![],
system_prompt: String::new(),
model: "test".to_string(),
model_binding: None,
max_tokens: 1000,
context_window: 128_000,
auto_compact_threshold: 0.8,
steering: SteeringQueue::default(),
pending_images: Vec::new(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
tool_trace: Vec::new(),
model_trace: Vec::new(),
trace_started_at: Some(Instant::now()),
trace_duration_ms: None,
transcript_checkpoint: None,
cost: CostTracker::new("test"),
last_request_usage: None,
last_compaction_notice: None,
last_failure: None,
retry_backoff_base: DEFAULT_RETRY_BACKOFF_BASE,
};
let tool_uses = vec![(
"test1".to_string(),
"TaskCreate".to_string(), serde_json::json!({"subject": "x"}),
)];
let (batch_tx, mut batch_rx) = mpsc::channel(64);
let drain = tokio::spawn(async move { while batch_rx.recv().await.is_some() {} });
let (blocks, _interrupted) = engine
.execute_tool_batch(
&tool_uses,
&batch_tx,
false,
&tokio_util::sync::CancellationToken::new(),
)
.await;
drop(batch_tx);
drain.await.unwrap();
assert_eq!(blocks.len(), 1, "every tool_use must get a tool_result");
match &blocks[0] {
ContentBlock::ToolResult {
tool_use_id,
is_error,
content,
} => {
assert_eq!(tool_use_id, "test1");
assert_eq!(*is_error, Some(true));
assert!(content.contains("Unknown tool"));
}
_ => panic!("Expected ToolResult block"),
}
assert_eq!(engine.tool_trace().len(), 1);
assert_eq!(engine.tool_trace()[0].id, "test1");
assert_eq!(engine.tool_trace()[0].name, "TaskCreate");
assert_eq!(engine.tool_trace()[0].input["subject"], "x");
assert!(engine.tool_trace()[0].is_error);
assert!(engine.tool_trace()[0].output.contains("Unknown tool"));
}
#[tokio::test]
async fn test_ask_permission_denies_in_non_streaming_mode() {
let provider = Box::new(MockProvider);
let tools = ToolRegistry::without_agent_for_tests();
let permissions = PermissionChecker::new(PermissionMode::Default);
let mut engine = Engine {
provider,
tools,
permissions,
messages: vec![],
system_prompt: String::new(),
model: "test".to_string(),
model_binding: None,
max_tokens: 1000,
context_window: 128_000,
auto_compact_threshold: 0.8,
steering: SteeringQueue::default(),
pending_images: Vec::new(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
tool_trace: Vec::new(),
model_trace: Vec::new(),
trace_started_at: Some(Instant::now()),
trace_duration_ms: None,
transcript_checkpoint: None,
cost: CostTracker::new("test"),
last_request_usage: None,
last_compaction_notice: None,
last_failure: None,
retry_backoff_base: DEFAULT_RETRY_BACKOFF_BASE,
};
let tool_uses = vec![(
"test1".to_string(),
"WebFetch".to_string(),
serde_json::json!({"url": "https://example.com/private"}),
)];
let (batch_tx, mut batch_rx) = mpsc::channel(64);
let drain = tokio::spawn(async move { while batch_rx.recv().await.is_some() {} });
let (blocks, _interrupted) = engine
.execute_tool_batch(
&tool_uses,
&batch_tx,
false,
&tokio_util::sync::CancellationToken::new(),
)
.await;
drop(batch_tx);
drain.await.unwrap();
assert_eq!(blocks.len(), 1);
match &blocks[0] {
ContentBlock::ToolResult {
is_error, content, ..
} => {
assert_eq!(
*is_error,
Some(true),
"Ask-permission tool must be denied, not executed, in non-streaming mode"
);
assert!(
content.contains("Permission denied"),
"expected a permission-denied message, got: {content}"
);
}
_ => panic!("Expected ToolResult block"),
}
}
#[tokio::test]
async fn pending_permission_request_observes_turn_cancellation() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Default,
);
let (tx, _rx) = mpsc::channel(1);
let cancel = tokio_util::sync::CancellationToken::new();
cancel.cancel();
let output = engine
.ask_permission(
"Write",
&serde_json::json!({"file_path": "/tmp/x", "content": "x"}),
"write /tmp/x".to_string(),
None,
(0, &tx),
&cancel,
)
.await;
assert!(output.is_error);
assert!(output.content.contains("cancelled"));
}
#[tokio::test]
async fn tool_completion_is_visible_before_a_later_permission_wait() {
let mut engine = Engine::for_tests(
Box::new(MockProvider),
SteeringQueue::default(),
PermissionMode::Default,
);
let dir = tempfile::tempdir().unwrap();
let tools = vec![
(
"read".to_string(),
"Glob".to_string(),
serde_json::json!({"pattern": "*.missing", "path": dir.path()}),
),
(
"bash".to_string(),
"Bash".to_string(),
serde_json::json!({"command": "echo approved"}),
),
];
let (tx, mut rx) = mpsc::channel(16);
let cancel = tokio_util::sync::CancellationToken::new();
let batch = engine.execute_tool_batch(&tools, &tx, true, &cancel);
let observer = async {
let mut running = Vec::new();
let mut finished = Vec::new();
let mut results = 0;
while let Some(event) = rx.recv().await {
match event {
StreamEvent::ToolRunning { index } => running.push(index),
StreamEvent::ToolFinished {
index,
is_error,
content,
} => {
assert!(!is_error);
if index == 1 {
assert!(content.contains("approved"));
}
finished.push(index);
}
StreamEvent::PermissionRequest { respond, .. } => {
assert_eq!(running, vec![0]);
assert_eq!(finished, vec![0]);
assert_eq!(results, 0, "ordered results wait for the batch");
respond.send(PermissionResponse::Allow).unwrap();
}
StreamEvent::ToolResult { .. } => {
results += 1;
if results == 2 {
break;
}
}
_ => {}
}
}
assert_eq!(running, vec![0, 1]);
assert_eq!(finished, vec![0, 1]);
};
let ((blocks, interrupted), ()) =
tokio::time::timeout(std::time::Duration::from_secs(5), async {
tokio::join!(batch, observer)
})
.await
.expect("execution updates must not wait for the full batch");
assert!(!interrupted);
assert_eq!(blocks.len(), 2);
}
}