use anyhow::Result;
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use tokio::sync::{mpsc, oneshot};
#[cfg(test)]
use crate::api::ProviderStream;
use crate::api::{ApiEvent, 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>>>;
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,
plugins: Option<Arc<PluginRegistry>>,
checkpoint_enabled: bool,
pending_checkpoint: Option<PendingCheckpoint>,
last_checkpoint: Option<TurnCheckpoint>,
pub cost: CostTracker,
}
pub enum StreamEvent {
Text(String),
Retry(String),
Notice(String),
SteeringSent(String),
ToolStart {
name: String,
summary: String,
},
ToolResult {
is_error: bool,
},
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,
}
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(),
plugins: None,
checkpoint_enabled: true,
pending_checkpoint: None,
last_checkpoint: None,
cost: CostTracker::new(model),
}
}
#[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,
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
cost: CostTracker::new("test"),
}
}
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> {
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 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,
) -> 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(name, input, token).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 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
}
#[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.pending_checkpoint = None;
self.last_checkpoint = 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()
}
pub async fn maybe_auto_compact(&mut self) -> Result<bool> {
if self.auto_compact_threshold <= 0.0 {
return Ok(false);
}
let current_tokens = compact::estimate_tokens(&self.messages);
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
);
let result = self.compact().await?;
tracing::info!("Auto-compact completed: {}", result);
Ok(true)
} else {
Ok(false)
}
}
pub async fn compact(&mut self) -> Result<String> {
if self.messages.is_empty() {
return Ok("Nothing to compact.".to_string());
}
let old_count = self.messages.len();
let old_tokens = compact::estimate_tokens(&self.messages);
let mut summary_source = None;
if let Some(snipped) = compact::snip_old_messages(&self.messages, 10) {
let new_tokens = compact::estimate_tokens(&snipped);
tracing::info!(
"Snip compaction: {} msgs → {}, ~{} → ~{} tokens",
old_count,
snipped.len(),
old_tokens,
new_tokens
);
if new_tokens < old_tokens && new_tokens < self.context_window * 70 / 100 {
let new_count = snipped.len();
self.commit_compacted_messages(snipped);
return Ok(format!(
"Snipped {} old messages (~{} tokens freed)",
old_count - new_count + 1, old_tokens - new_tokens
));
}
summary_source = Some(snipped);
}
let summary_source = summary_source.unwrap_or_else(|| self.messages.clone());
self.summarize_conversation(summary_source).await
}
async fn summarize_conversation(&mut self, messages: Vec<Message>) -> Result<String> {
let summary_prompt = "Summarize the conversation so far in a concise paragraph. \
Focus on what was discussed, what decisions were made, what files were modified, \
and any outstanding tasks. Be specific about file paths and changes.";
let old_count = messages.len();
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,
tokio_util::sync::CancellationToken::new(),
)
.await?;
let mut summary = String::new();
let mut completed = false;
while let Some(event) = rx.recv().await {
match event {
ApiEvent::Text(t) => summary.push_str(&t),
ApiEvent::Usage(usage) => self.cost.add_usage(&usage),
ApiEvent::Done => {
completed = true;
break;
}
ApiEvent::Error(e) => return Err(anyhow::anyhow!("Compact error: {e}")),
_ => {}
}
}
if !completed {
anyhow::bail!("Compact error: API stream ended without completion");
}
self.commit_compacted_messages(vec![
Message::user("Here is a summary of our conversation so far:"),
Message::assistant_text(&summary),
]);
Ok(format!(
"Compacted {old_count} messages into summary.\n\n\x1b[2m{summary}\x1b[0m"
))
}
fn commit_compacted_messages(&mut self, messages: Vec<Message>) {
self.messages = messages;
self.provider.reset_session();
}
fn is_prompt_too_long(err: &str) -> bool {
let err = err.to_ascii_lowercase();
err.contains("413")
|| err.contains("prompt is too long")
|| err.contains("maximum context length")
|| err.contains("context_length_exceeded")
}
fn is_max_output_tokens(err: &str) -> bool {
let err = err.to_ascii_lowercase();
err.contains("max_output_tokens") || err.contains("max_tokens_exceeded")
}
fn is_malformed_tool_arguments(err: &str) -> bool {
err.to_ascii_lowercase()
.contains("invalid arguments for tool call")
}
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.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(user_input, tx, false, cancel).await;
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.begin_checkpoint();
let result = self.run_turn(user_input, tx, true, cancel).await;
self.finish_checkpoint();
self.fire_hook(&HookTrigger::OnTurnEnd).await;
result
}
async fn run_turn(
&mut self,
user_input: &str,
tx: mpsc::Sender<StreamEvent>,
interactive: bool,
cancel: tokio_util::sync::CancellationToken,
) -> Result<()> {
let compacted = self.maybe_auto_compact().await?;
if compacted {
let _ = tx
.send(StreamEvent::Notice(
"conversation auto-compacted to free context".to_string(),
))
.await;
}
self.messages.push(Message::user(user_input));
let mut recovery_attempts = 0;
const MAX_RECOVERY: u32 = 3;
let mut malformed_tool_retries = 0;
const MAX_MALFORMED_TOOL_RETRIES: u32 = 1;
let mut retry_prompt: Option<String> = None;
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 effective_system_prompt = retry_prompt
.as_ref()
.map(|prompt| format!("{}\n\n{prompt}", self.system_prompt));
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 mut rx = match stream_result {
Ok(rx) => rx,
Err(e) => {
if cancel.is_cancelled() {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
let err_str = e.to_string();
if Self::is_malformed_tool_arguments(&err_str)
&& 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;
}
if Self::is_max_output_tokens(&err_str) && self.max_tokens < 64_000 {
self.max_tokens = (self.max_tokens * 2).min(64_000);
continue;
}
if Self::is_prompt_too_long(&err_str) && recovery_attempts < MAX_RECOVERY {
recovery_attempts += 1;
let _ = tx
.send(StreamEvent::Notice(
"compacting conversation...".to_string(),
))
.await;
self.compact().await?;
continue;
}
let _ = tx.send(StreamEvent::Error(err_str.clone())).await;
return Err(e);
}
};
let mut text_buf = String::new();
let mut tool_uses: Vec<(String, String, serde_json::Value)> = Vec::new();
let mut had_error = false;
let mut stream_interrupted = false;
loop {
let event = tokio::select! {
event = rx.recv() => match event {
Some(event) => event,
None => {
let error = "API stream ended without completion".to_string();
let _ = tx.send(StreamEvent::Error(error.clone())).await;
return Err(anyhow::anyhow!(error));
}
},
_ = 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::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,
})
.await;
tool_uses.push((id, name, input));
}
ApiEvent::Usage(usage) => {
self.cost.add_usage(&usage);
}
ApiEvent::Done => break,
ApiEvent::Error(e) => {
if Self::is_malformed_tool_arguments(&e)
&& tool_uses.is_empty()
&& malformed_tool_retries < MAX_MALFORMED_TOOL_RETRIES
{
malformed_tool_retries += 1;
retry_prompt = Some(Self::malformed_tool_retry_prompt(&e));
let _ = tx
.send(StreamEvent::Retry(
"model returned malformed tool arguments; retrying once"
.to_string(),
))
.await;
had_error = true;
break;
}
if Self::is_max_output_tokens(&e) && self.max_tokens < 64_000 {
self.max_tokens = (self.max_tokens * 2).min(64_000);
had_error = true;
break;
}
if Self::is_prompt_too_long(&e) && recovery_attempts < MAX_RECOVERY {
recovery_attempts += 1;
let _ = tx
.send(StreamEvent::Notice(
"compacting conversation...".to_string(),
))
.await;
self.compact().await?;
had_error = true;
break;
}
let _ = tx.send(StreamEvent::Error(e.clone())).await;
return Err(anyhow::anyhow!("API error: {e}"));
}
}
}
if had_error {
continue;
}
malformed_tool_retries = 0;
retry_prompt = None;
let mut blocks = Vec::new();
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));
}
if stream_interrupted {
if !tool_uses.is_empty() {
let mut result_blocks = Vec::with_capacity(tool_uses.len());
for (id, _, _) in &tool_uses {
self.fire_hook(&HookTrigger::OnToolComplete).await;
let _ = tx.send(StreamEvent::ToolResult { is_error: true }).await;
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));
}
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));
if interrupted {
let _ = tx.send(StreamEvent::Interrupted).await;
return Ok(());
}
}
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<crate::tools::ToolOutput>> =
(0..tool_uses.len()).map(|_| None).collect();
let parallel: Vec<usize> = tool_uses
.iter()
.enumerate()
.filter(|(_, (_, name, input))| {
let ro = self.tools.is_read_only(name);
ro && matches!(
self.permissions.check(name, input, ro),
PermissionResult::Allow
)
})
.map(|(idx, _)| idx)
.collect();
let mut interrupted = false;
if !self.steering_pending() && !cancel.is_cancelled() && !parallel.is_empty() {
let this: &Self = &*self;
let futures: Vec<_> = parallel
.iter()
.map(|&idx| {
let (_, name, input) = &tool_uses[idx];
async move {
(
idx,
this.execute_tool_steerable(name, input.clone(), cancel)
.await,
)
}
})
.collect();
for (idx, output) in futures_util::future::join_all(futures).await {
outputs[idx] = Some(output);
}
}
for (idx, (_, name, input)) in tool_uses.iter().enumerate() {
if outputs[idx].is_some() {
continue;
}
if cancel.is_cancelled() {
interrupted = true;
outputs[idx] = Some(crate::tools::ToolOutput {
content: Self::INTERRUPTED_BY_USER.to_string(),
is_error: true,
});
continue;
}
if self.steering_pending() {
outputs[idx] = Some(crate::tools::ToolOutput {
content: Self::SKIPPED_FOR_STEERING.to_string(),
is_error: true,
});
continue;
}
let is_read_only = self.tools.is_read_only(name);
let perm = self.permissions.check(name, input, is_read_only);
let output = match perm {
PermissionResult::Allow => {
self.execute_tool_steerable(name, input.clone(), 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, tx, cancel)
.await
}
}
};
outputs[idx] = Some(output);
}
if cancel.is_cancelled() {
interrupted = true;
}
let mut result_blocks = Vec::with_capacity(tool_uses.len());
for (idx, (id, name, _)) in tool_uses.iter().enumerate() {
let output = outputs[idx].take().expect("every tool got an 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,
})
.await;
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 ask_permission(
&mut self,
name: &str,
input: &serde_json::Value,
message: String,
diff: Option<String>,
tx: &mpsc::Sender<StreamEvent>,
cancel: &tokio_util::sync::CancellationToken,
) -> crate::tools::ToolOutput {
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;
match resp_rx.await {
Ok(PermissionResponse::Allow) => {
self.execute_tool_steerable(name, input.clone(), 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_steerable(name, input.clone(), cancel)
.await
}
Ok(PermissionResponse::AlwaysAllowCommand(ref cmd)) => {
self.permissions.always_allow_command(cmd);
self.execute_tool_steerable(name, input.clone(), 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;
#[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 MalformedToolProvider {
calls: Arc<AtomicUsize>,
systems: Arc<Mutex<Vec<String>>>,
recover: bool,
}
#[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(
"OpenAI SSE stream error: invalid arguments for tool call Read \
(call_3): EOF while parsing a value"
.to_string(),
))
.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,
}
#[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()))
}
}
#[test]
fn output_token_limit_is_not_misclassified_as_prompt_too_long() {
let error = "max_tokens_exceeded: response reached max output tokens";
assert!(Engine::is_max_output_tokens(error));
assert!(!Engine::is_prompt_too_long(error));
}
#[test]
fn max_tokens_parameter_error_does_not_trigger_compaction() {
let error = "invalid max_tokens parameter";
assert!(!Engine::is_prompt_too_long(error));
assert!(!Engine::is_max_output_tokens(error));
}
#[test]
fn context_limit_errors_are_classified_case_insensitively() {
assert!(Engine::is_prompt_too_long(
"Maximum Context Length exceeded"
));
assert!(Engine::is_prompt_too_long("CONTEXT_LENGTH_EXCEEDED"));
}
#[test]
fn malformed_tool_argument_errors_are_classified_narrowly() {
assert!(Engine::is_malformed_tool_arguments(
"OpenAI SSE stream error: invalid arguments for tool call Read (call_3)"
));
assert!(!Engine::is_malformed_tool_arguments(
"invalid arguments for request"
));
}
#[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"));
assert_eq!(engine.messages().len(), 2);
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.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.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);
}
#[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);
}
#[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(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
cost: CostTracker::new("test"),
};
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(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
cost: CostTracker::new("test"),
};
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");
}
}
}
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)
}
}
#[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 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_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 snip_compaction_resets_provider_cursor_after_rewriting_history() {
let resets = Arc::new(AtomicUsize::new(0));
let provider = Box::new(ResetTrackingProvider {
resets: resets.clone(),
});
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}")));
}
engine.compact().await.unwrap();
assert_eq!(
engine.messages().len(),
12,
"tool-result boundary backoff reproduces the stale cursor index shape"
);
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_eq!(resets.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn snip_candidate_that_increases_tokens_is_not_committed() {
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,
"a larger snip candidate should fall back to summary compaction"
);
assert_eq!(resets.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn failed_summary_preserves_history_after_snipping_candidate() {
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 commit the snipped candidate"
);
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(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
cost: CostTracker::new("test"),
};
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"),
}
}
#[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(),
plugins: None,
checkpoint_enabled: false,
pending_checkpoint: None,
last_checkpoint: None,
cost: CostTracker::new("test"),
};
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"),
}
}
}