Implement turn cancellation
This commit is contained in:
+1
-1
@@ -6,7 +6,7 @@ This release focuses on making Cass easier to interrupt, easier to audit, and sa
|
||||
|
||||
### Agent Control
|
||||
|
||||
- [ ] **Turn cancellation.** Allow the user to stop a running turn without exiting Cass.
|
||||
- [x] **Turn cancellation.** Allow the user to stop a running turn without exiting Cass.
|
||||
- While the agent is busy, the first `Ctrl-C` should cancel the active turn.
|
||||
- A second `Ctrl-C` can retain the existing exit behavior.
|
||||
- You should also be able to press `Esc` to cancel but not exit.
|
||||
|
||||
+227
-6
@@ -1,7 +1,7 @@
|
||||
use crate::agent::{self, AgentEvent, AgentSettings};
|
||||
use crate::cli::{self, Command};
|
||||
use crate::config::{Config, ModelDefinition, ReasoningEffort};
|
||||
use crate::conversation::{self, Conversation};
|
||||
use crate::conversation::{self, Conversation, Record};
|
||||
use crate::prompt;
|
||||
use crate::ui::autofill::{AutoFillItem, AutoFillMenu};
|
||||
use crate::ui::events::poll_event;
|
||||
@@ -16,6 +16,9 @@ use std::time::{Duration, Instant};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
const TURN_CANCELLED_MESSAGE: &str = "Turn cancelled by user.";
|
||||
const TOOL_CANCELLED_MESSAGE: &str = "Tool execution cancelled by user.";
|
||||
|
||||
pub async fn run() -> Result<()> {
|
||||
let cli = cli::parse();
|
||||
if matches!(cli.command, Some(Command::Check)) {
|
||||
@@ -74,6 +77,70 @@ fn list_chats(config: &Config, cwd: &std::path::Path) -> Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn finalize_cancelled_turn(
|
||||
config: &Config,
|
||||
chat_id: &str,
|
||||
turn_start_len: Option<usize>,
|
||||
turn_message: Option<&str>,
|
||||
) -> Result<Conversation> {
|
||||
let (mut conversation, _) = Conversation::load(&config.conversations_dir(), chat_id)?;
|
||||
|
||||
if let (Some(start_len), Some(message)) = (turn_start_len, turn_message) {
|
||||
if conversation.records.len() <= start_len {
|
||||
conversation.append(Record::User {
|
||||
content: message.to_string(),
|
||||
ts: conversation::now_ts(),
|
||||
})?;
|
||||
}
|
||||
}
|
||||
|
||||
for (id, name) in pending_tool_calls(&conversation.records) {
|
||||
conversation.append(Record::Tool {
|
||||
tool_call_id: id,
|
||||
name,
|
||||
ok: false,
|
||||
content: TOOL_CANCELLED_MESSAGE.to_string(),
|
||||
ts: conversation::now_ts(),
|
||||
})?;
|
||||
}
|
||||
|
||||
if !matches!(
|
||||
conversation.records.last(),
|
||||
Some(Record::Assistant { content, tool_calls, .. })
|
||||
if content == TURN_CANCELLED_MESSAGE && tool_calls.is_empty()
|
||||
) {
|
||||
conversation.append(Record::Assistant {
|
||||
content: TURN_CANCELLED_MESSAGE.to_string(),
|
||||
reasoning: String::new(),
|
||||
reasoning_field: None,
|
||||
tool_calls: Vec::new(),
|
||||
ts: conversation::now_ts(),
|
||||
})?;
|
||||
}
|
||||
|
||||
Ok(conversation)
|
||||
}
|
||||
|
||||
fn pending_tool_calls(records: &[Record]) -> Vec<(String, String)> {
|
||||
let mut pending = Vec::new();
|
||||
for record in records {
|
||||
match record {
|
||||
Record::Assistant { tool_calls, .. } => {
|
||||
pending = tool_calls
|
||||
.iter()
|
||||
.map(|call| (call.id.clone(), call.name.clone()))
|
||||
.collect();
|
||||
}
|
||||
Record::Tool { tool_call_id, .. } => {
|
||||
pending.retain(|(id, _)| id != tool_call_id);
|
||||
}
|
||||
Record::User { .. } => pending.clear(),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
pending
|
||||
}
|
||||
|
||||
async fn run_tui(
|
||||
mut config: Config,
|
||||
cwd: PathBuf,
|
||||
@@ -101,6 +168,9 @@ async fn run_tui(
|
||||
let mut scroll: u16 = 0;
|
||||
let mut last_ctrl_c: Option<Instant> = None;
|
||||
let mut handle: Option<JoinHandle<Result<Conversation>>> = None;
|
||||
let mut cancel_requested = false;
|
||||
let mut current_turn_start_len: Option<usize> = None;
|
||||
let mut current_turn_message: Option<String> = None;
|
||||
let mut active_assistant: Option<usize> = None;
|
||||
let mut active_reasoning: Option<usize> = None;
|
||||
let mut active_tools: HashMap<String, usize> = HashMap::new();
|
||||
@@ -145,6 +215,7 @@ async fn run_tui(
|
||||
scroll: &mut scroll,
|
||||
},
|
||||
)?;
|
||||
let mut finished_status = "idle".to_string();
|
||||
match result {
|
||||
Ok(Ok(updated)) => {
|
||||
conversation = updated;
|
||||
@@ -155,6 +226,30 @@ async fn run_tui(
|
||||
title: "agent error".into(),
|
||||
content: err.to_string(),
|
||||
}),
|
||||
Err(err) if err.is_cancelled() && cancel_requested => {
|
||||
mark_active_tool_blocks_cancelled(&mut transcript, &active_tools);
|
||||
match finalize_cancelled_turn(
|
||||
&config,
|
||||
&chat_id,
|
||||
current_turn_start_len,
|
||||
current_turn_message.as_deref(),
|
||||
) {
|
||||
Ok(updated) => conversation = updated,
|
||||
Err(err) => transcript.push(TranscriptBlock {
|
||||
kind: TranscriptKind::Error,
|
||||
title: "cancel".into(),
|
||||
content: format!(
|
||||
"turn cancelled, but updating the conversation failed: {err}"
|
||||
),
|
||||
}),
|
||||
}
|
||||
transcript.push(TranscriptBlock {
|
||||
kind: TranscriptKind::Status,
|
||||
title: "cancelled".into(),
|
||||
content: TURN_CANCELLED_MESSAGE.to_string(),
|
||||
});
|
||||
finished_status = "turn cancelled".into();
|
||||
}
|
||||
Err(err) => transcript.push(TranscriptBlock {
|
||||
kind: TranscriptKind::Error,
|
||||
title: "agent task error".into(),
|
||||
@@ -164,7 +259,10 @@ async fn run_tui(
|
||||
active_assistant = None;
|
||||
active_reasoning = None;
|
||||
active_tools.clear();
|
||||
status = "idle".into();
|
||||
cancel_requested = false;
|
||||
current_turn_start_len = None;
|
||||
current_turn_message = None;
|
||||
status = finished_status;
|
||||
}
|
||||
|
||||
let autofill = build_autofill(&input, autofill_selected, &config, &cwd)?;
|
||||
@@ -217,18 +315,48 @@ async fn run_tui(
|
||||
match (key.code, key.modifiers) {
|
||||
(KeyCode::Char('c'), m) if m.contains(KeyModifiers::CONTROL) => {
|
||||
let now = Instant::now();
|
||||
input.clear();
|
||||
autofill_selected = 0;
|
||||
if last_ctrl_c
|
||||
.map(|t| now.duration_since(t) <= Duration::from_millis(1500))
|
||||
.unwrap_or(false)
|
||||
{
|
||||
if busy {
|
||||
if let Some(handle) = &handle {
|
||||
handle.abort();
|
||||
}
|
||||
if cancel_requested {
|
||||
let _ = finalize_cancelled_turn(
|
||||
&config,
|
||||
&chat_id,
|
||||
current_turn_start_len,
|
||||
current_turn_message.as_deref(),
|
||||
);
|
||||
}
|
||||
}
|
||||
terminal::leave(terminal)?;
|
||||
println!("Resume this chat with: cass --resume {}", chat_id);
|
||||
return Ok(());
|
||||
}
|
||||
last_ctrl_c = Some(now);
|
||||
status = "press Ctrl-C again within 1.5s to exit".into();
|
||||
if busy {
|
||||
if let Some(handle) = &handle {
|
||||
handle.abort();
|
||||
}
|
||||
cancel_requested = true;
|
||||
last_ctrl_c = Some(now);
|
||||
status = "turn cancellation requested; press Ctrl-C again within 1.5s to exit".into();
|
||||
} else {
|
||||
input.clear();
|
||||
autofill_selected = 0;
|
||||
last_ctrl_c = Some(now);
|
||||
status = "press Ctrl-C again within 1.5s to exit".into();
|
||||
}
|
||||
}
|
||||
(KeyCode::Esc, _) if busy => {
|
||||
if let Some(handle) = &handle {
|
||||
handle.abort();
|
||||
}
|
||||
cancel_requested = true;
|
||||
last_ctrl_c = None;
|
||||
status = "turn cancellation requested".into();
|
||||
}
|
||||
(KeyCode::BackTab, _) => {
|
||||
if busy {
|
||||
@@ -529,6 +657,10 @@ async fn run_tui(
|
||||
status = "agent is still running".into();
|
||||
} else {
|
||||
let msg = input.trim_end().to_string();
|
||||
current_turn_start_len = Some(conversation.records.len());
|
||||
current_turn_message = Some(msg.clone());
|
||||
cancel_requested = false;
|
||||
last_ctrl_c = None;
|
||||
input.clear();
|
||||
autofill_selected = 0;
|
||||
transcript.push(TranscriptBlock {
|
||||
@@ -809,6 +941,23 @@ fn active_tool_block(ctx: &mut AgentEventContext<'_>, id: &str, name: &str) -> u
|
||||
idx
|
||||
}
|
||||
|
||||
fn mark_active_tool_blocks_cancelled(
|
||||
transcript: &mut [TranscriptBlock],
|
||||
active_tools: &HashMap<String, usize>,
|
||||
) {
|
||||
for idx in active_tools.values().copied() {
|
||||
let Some(block) = transcript.get_mut(idx) else {
|
||||
continue;
|
||||
};
|
||||
block.kind = TranscriptKind::Error;
|
||||
block.title = block.title.replace('…', "cancelled");
|
||||
if !block.content.ends_with('\n') && !block.content.is_empty() {
|
||||
block.content.push('\n');
|
||||
}
|
||||
block.content.push_str(TOOL_CANCELLED_MESSAGE);
|
||||
}
|
||||
}
|
||||
|
||||
fn append_tool_output_chunk(existing: &mut String, stream: &str, chunk: &str) {
|
||||
if !existing.contains("streamed output:\n") {
|
||||
if !existing.trim().is_empty() {
|
||||
@@ -1354,6 +1503,78 @@ mod tests {
|
||||
assert_eq!(parse_local_command("/new extra"), Err("usage: /new".into()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cancelled_turn_repairs_missing_tool_results() {
|
||||
let root = tempdir().unwrap();
|
||||
let cwd = tempdir().unwrap();
|
||||
let config = Config {
|
||||
root: root.path().to_path_buf(),
|
||||
model: "test-model".into(),
|
||||
..Config::default()
|
||||
};
|
||||
let mut conversation = Conversation::create(
|
||||
&config.conversations_dir(),
|
||||
&config.model,
|
||||
cwd.path(),
|
||||
"base prompt".into(),
|
||||
)
|
||||
.unwrap();
|
||||
let chat_id = conversation.id.clone();
|
||||
let start_len = conversation.records.len();
|
||||
conversation
|
||||
.append(Record::User {
|
||||
content: "run tools".into(),
|
||||
ts: conversation::now_ts(),
|
||||
})
|
||||
.unwrap();
|
||||
conversation
|
||||
.append(Record::Assistant {
|
||||
content: String::new(),
|
||||
reasoning: String::new(),
|
||||
reasoning_field: None,
|
||||
tool_calls: vec![
|
||||
conversation::StoredToolCall {
|
||||
id: "call_done".into(),
|
||||
name: "read".into(),
|
||||
arguments: serde_json::json!({"path":"a"}),
|
||||
},
|
||||
conversation::StoredToolCall {
|
||||
id: "call_pending".into(),
|
||||
name: "shell".into(),
|
||||
arguments: serde_json::json!({"command":"sleep 60"}),
|
||||
},
|
||||
],
|
||||
ts: conversation::now_ts(),
|
||||
})
|
||||
.unwrap();
|
||||
conversation
|
||||
.append(Record::Tool {
|
||||
tool_call_id: "call_done".into(),
|
||||
name: "read".into(),
|
||||
ok: true,
|
||||
content: "ok".into(),
|
||||
ts: conversation::now_ts(),
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let updated =
|
||||
finalize_cancelled_turn(&config, &chat_id, Some(start_len), Some("run tools")).unwrap();
|
||||
|
||||
assert!(matches!(
|
||||
updated.records.get(updated.records.len() - 2),
|
||||
Some(Record::Tool { tool_call_id, name, ok, content, .. })
|
||||
if tool_call_id == "call_pending"
|
||||
&& name == "shell"
|
||||
&& !ok
|
||||
&& content == TOOL_CANCELLED_MESSAGE
|
||||
));
|
||||
assert!(matches!(
|
||||
updated.records.last(),
|
||||
Some(Record::Assistant { content, tool_calls, .. })
|
||||
if content == TURN_CANCELLED_MESSAGE && tool_calls.is_empty()
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_autofill_lists_models_from_models_json() {
|
||||
let (_root, config) = config_with_models(
|
||||
|
||||
@@ -39,6 +39,7 @@ pub async fn run(args: Value, ctx: &ToolContext) -> Result<String> {
|
||||
cmd.current_dir(&ctx.cwd);
|
||||
cmd.stdout(Stdio::piped());
|
||||
cmd.stderr(Stdio::piped());
|
||||
cmd.kill_on_drop(true);
|
||||
|
||||
let mut child = cmd.spawn().context("spawning shell command")?;
|
||||
let stdout = child.stdout.take().context("capturing command stdout")?;
|
||||
|
||||
@@ -201,6 +201,29 @@ async fn shell_streams_output_chunks() {
|
||||
assert!(streamed.contains("stderr:err"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn aborting_shell_tool_stops_command_before_completion() {
|
||||
let dir = tempdir().unwrap();
|
||||
let marker = dir.path().join("marker");
|
||||
let context = ctx(dir.path(), AccessMode::FullAccess);
|
||||
let command = "sleep 1; touch marker";
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
tools::execute(
|
||||
"shell",
|
||||
json!({"command": command, "timeout": 10}),
|
||||
&context,
|
||||
)
|
||||
.await
|
||||
});
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
handle.abort();
|
||||
assert!(handle.await.unwrap_err().is_cancelled());
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(1200)).await;
|
||||
assert!(!marker.exists());
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn full_access_blocks_writes_through_symlinked_docs_dir() {
|
||||
|
||||
Reference in New Issue
Block a user