Skip to content

Commit 82a4ef6

Browse files
committed
Merge PR ultraworkers#3255: feat/esc-turn-interruption
2 parents b40ad9a + b0fad35 commit 82a4ef6

7 files changed

Lines changed: 698 additions & 39 deletions

File tree

‎rust/Cargo.lock‎

Lines changed: 1 addition & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎rust/crates/runtime/src/conversation.rs‎

Lines changed: 266 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,9 @@
11
use std::collections::BTreeMap;
22
use std::fmt::{Display, Formatter};
3+
use std::sync::{
4+
atomic::{AtomicBool, Ordering},
5+
Arc,
6+
};
37

48
use serde_json::{Map, Value};
59
use telemetry::SessionTracer;
@@ -53,6 +57,41 @@ pub struct PromptCacheEvent {
5357
pub token_drop: u32,
5458
}
5559

60+
/// Shared flag used to request graceful interruption of a running turn.
61+
///
62+
/// Cloning shares the underlying flag, mirroring
63+
/// [`HookAbortSignal`](crate::hooks::HookAbortSignal). An input listener
64+
/// (e.g. Esc or Ctrl+C handling in the CLI) sets the flag while the
65+
/// conversation loop and the streaming API client poll it at safe points.
66+
/// When the flag is observed, the turn winds down without treating the
67+
/// stop as a failure: pending tool calls receive synthesized error
68+
/// results so the session stays consistent, and [`TurnSummary`] reports
69+
/// `interrupted: true`.
70+
#[derive(Debug, Clone, Default)]
71+
pub struct TurnInterruptSignal {
72+
interrupted: Arc<AtomicBool>,
73+
}
74+
75+
impl TurnInterruptSignal {
76+
#[must_use]
77+
pub fn new() -> Self {
78+
Self::default()
79+
}
80+
81+
pub fn interrupt(&self) {
82+
self.interrupted.store(true, Ordering::SeqCst);
83+
}
84+
85+
#[must_use]
86+
pub fn is_interrupted(&self) -> bool {
87+
self.interrupted.load(Ordering::SeqCst)
88+
}
89+
90+
pub fn reset(&self) {
91+
self.interrupted.store(false, Ordering::SeqCst);
92+
}
93+
}
94+
5695
/// Minimal streaming API contract required by [`ConversationRuntime`].
5796
pub trait ApiClient {
5897
fn stream(&mut self, request: ApiRequest) -> Result<Vec<AssistantEvent>, RuntimeError>;
@@ -118,6 +157,7 @@ pub struct TurnSummary {
118157
pub iterations: usize,
119158
pub usage: TokenUsage,
120159
pub auto_compaction: Option<AutoCompactionEvent>,
160+
pub interrupted: bool,
121161
}
122162

123163
/// Details about automatic session compaction applied during a turn.
@@ -138,6 +178,7 @@ pub struct ConversationRuntime<C, T> {
138178
hook_runner: HookRunner,
139179
auto_compaction_input_tokens_threshold: u32,
140180
hook_abort_signal: HookAbortSignal,
181+
turn_interrupt_signal: TurnInterruptSignal,
141182
hook_progress_reporter: Option<Box<dyn HookProgressReporter>>,
142183
session_tracer: Option<SessionTracer>,
143184
}
@@ -187,6 +228,7 @@ where
187228
hook_runner: HookRunner::from_feature_config(feature_config),
188229
auto_compaction_input_tokens_threshold: auto_compaction_threshold_from_env(),
189230
hook_abort_signal: HookAbortSignal::default(),
231+
turn_interrupt_signal: TurnInterruptSignal::default(),
190232
hook_progress_reporter: None,
191233
session_tracer: None,
192234
}
@@ -217,6 +259,15 @@ where
217259
self
218260
}
219261

262+
#[must_use]
263+
pub fn with_turn_interrupt_signal(
264+
mut self,
265+
turn_interrupt_signal: TurnInterruptSignal,
266+
) -> Self {
267+
self.turn_interrupt_signal = turn_interrupt_signal;
268+
self
269+
}
270+
220271
#[must_use]
221272
pub fn with_hook_progress_reporter(
222273
mut self,
@@ -350,8 +401,14 @@ where
350401
let mut prompt_cache_events = Vec::new();
351402
let mut iterations = 0;
352403
let mut auto_compaction = None;
404+
let mut interrupted = false;
353405

354406
loop {
407+
if self.turn_interrupt_signal.is_interrupted() {
408+
self.record_turn_interrupted(iterations, "before_request");
409+
interrupted = true;
410+
break;
411+
}
355412
iterations += 1;
356413
if iterations > self.max_iterations {
357414
let error = RuntimeError::new(
@@ -368,6 +425,14 @@ where
368425
let events = match self.api_client.stream(request) {
369426
Ok(events) => events,
370427
Err(error) => {
428+
if self.turn_interrupt_signal.is_interrupted() {
429+
// The client aborted because the user interrupted the
430+
// turn; any partial response is discarded and the stop
431+
// is reported as an interruption rather than a failure.
432+
self.record_turn_interrupted(iterations, "during_request");
433+
interrupted = true;
434+
break;
435+
}
371436
self.record_turn_failed(iterations, &error);
372437
return Err(error);
373438
}
@@ -416,6 +481,25 @@ where
416481
}
417482

418483
for (tool_use_id, tool_name, input) in pending_tool_uses {
484+
if interrupted || self.turn_interrupt_signal.is_interrupted() {
485+
// Every pending tool_use must still receive a tool_result
486+
// so the session stays valid for the next request.
487+
if !interrupted {
488+
self.record_turn_interrupted(iterations, "before_tool");
489+
interrupted = true;
490+
}
491+
let result_message = ConversationMessage::tool_result(
492+
tool_use_id,
493+
tool_name,
494+
"Interrupted by user before this tool could run.",
495+
true,
496+
);
497+
self.session
498+
.push_message(result_message.clone())
499+
.map_err(|error| RuntimeError::new(error.to_string()))?;
500+
tool_results.push(result_message);
501+
continue;
502+
}
419503
let pre_hook_result = self.run_pre_tool_use_hook(&tool_name, &input);
420504
let effective_input = pre_hook_result
421505
.updated_input()
@@ -515,6 +599,10 @@ where
515599
self.record_tool_finished(iterations, &result_message);
516600
tool_results.push(result_message);
517601
}
602+
603+
if interrupted {
604+
break;
605+
}
518606
}
519607

520608
let summary = TurnSummary {
@@ -524,8 +612,11 @@ where
524612
iterations,
525613
usage: self.usage_tracker.cumulative_usage(),
526614
auto_compaction,
615+
interrupted,
527616
};
528-
self.record_turn_completed(&summary);
617+
if !interrupted {
618+
self.record_turn_completed(&summary);
619+
}
529620

530621
Ok(summary)
531622
}
@@ -689,6 +780,17 @@ where
689780
session_tracer.record("turn_completed", attributes);
690781
}
691782

783+
fn record_turn_interrupted(&self, iteration: usize, phase: &str) {
784+
let Some(session_tracer) = &self.session_tracer else {
785+
return;
786+
};
787+
788+
let mut attributes = Map::new();
789+
attributes.insert("iteration".to_string(), Value::from(iteration as u64));
790+
attributes.insert("phase".to_string(), Value::String(phase.to_string()));
791+
session_tracer.record("turn_interrupted", attributes);
792+
}
793+
692794
fn record_turn_failed(&self, iteration: usize, error: &RuntimeError) {
693795
let Some(session_tracer) = &self.session_tracer else {
694796
return;
@@ -850,7 +952,8 @@ mod tests {
850952
use super::{
851953
build_assistant_message, parse_auto_compaction_threshold, ApiClient, ApiRequest,
852954
AssistantEvent, AutoCompactionEvent, ConversationRuntime, PromptCacheEvent, RuntimeError,
853-
StaticToolExecutor, ToolExecutor, DEFAULT_AUTO_COMPACTION_INPUT_TOKENS_THRESHOLD,
955+
StaticToolExecutor, ToolExecutor, TurnInterruptSignal,
956+
DEFAULT_AUTO_COMPACTION_INPUT_TOKENS_THRESHOLD,
854957
};
855958
use crate::compact::CompactionConfig;
856959
use crate::config::{RuntimeFeatureConfig, RuntimeHookConfig};
@@ -1875,4 +1978,165 @@ mod tests {
18751978
// then
18761979
assert_eq!(error.to_string(), "upstream failed");
18771980
}
1981+
1982+
#[test]
1983+
fn interrupt_before_first_request_skips_the_api_call() {
1984+
struct UnreachableApi;
1985+
1986+
impl ApiClient for UnreachableApi {
1987+
fn stream(
1988+
&mut self,
1989+
_request: ApiRequest,
1990+
) -> Result<Vec<AssistantEvent>, RuntimeError> {
1991+
unreachable!("interrupted turn must not reach the API")
1992+
}
1993+
}
1994+
1995+
// given
1996+
let interrupt = TurnInterruptSignal::new();
1997+
interrupt.interrupt();
1998+
let mut runtime = ConversationRuntime::new(
1999+
Session::new(),
2000+
UnreachableApi,
2001+
StaticToolExecutor::new(),
2002+
PermissionPolicy::new(PermissionMode::DangerFullAccess),
2003+
vec!["system".to_string()],
2004+
)
2005+
.with_turn_interrupt_signal(interrupt);
2006+
2007+
// when
2008+
let summary = runtime
2009+
.run_turn("hello", None)
2010+
.expect("interruption should not be reported as a failure");
2011+
2012+
// then
2013+
assert!(summary.interrupted);
2014+
assert_eq!(summary.iterations, 0);
2015+
assert!(summary.assistant_messages.is_empty());
2016+
assert!(summary.tool_results.is_empty());
2017+
assert_eq!(summary.auto_compaction, None);
2018+
assert_eq!(runtime.session().messages.len(), 1);
2019+
assert_eq!(runtime.session().messages[0].role, MessageRole::User);
2020+
}
2021+
2022+
#[test]
2023+
fn interrupt_after_stream_synthesizes_results_for_pending_tools() {
2024+
struct ToolUseApi {
2025+
interrupt: TurnInterruptSignal,
2026+
}
2027+
2028+
impl ApiClient for ToolUseApi {
2029+
fn stream(
2030+
&mut self,
2031+
_request: ApiRequest,
2032+
) -> Result<Vec<AssistantEvent>, RuntimeError> {
2033+
// Simulate the user pressing Esc while the response streams in.
2034+
self.interrupt.interrupt();
2035+
Ok(vec![
2036+
AssistantEvent::TextDelta("Running the tool.".to_string()),
2037+
AssistantEvent::ToolUse {
2038+
id: "tool-1".to_string(),
2039+
name: "add".to_string(),
2040+
input: "2,2".to_string(),
2041+
},
2042+
AssistantEvent::MessageStop,
2043+
])
2044+
}
2045+
}
2046+
2047+
// given
2048+
let interrupt = TurnInterruptSignal::new();
2049+
let mut runtime = ConversationRuntime::new(
2050+
Session::new(),
2051+
ToolUseApi {
2052+
interrupt: interrupt.clone(),
2053+
},
2054+
StaticToolExecutor::new()
2055+
.register("add", |_input| panic!("interrupted tool must not run")),
2056+
PermissionPolicy::new(PermissionMode::DangerFullAccess),
2057+
vec!["system".to_string()],
2058+
)
2059+
.with_turn_interrupt_signal(interrupt);
2060+
2061+
// when
2062+
let summary = runtime
2063+
.run_turn("what is 2 + 2?", None)
2064+
.expect("interruption should not be reported as a failure");
2065+
2066+
// then
2067+
assert!(summary.interrupted);
2068+
assert_eq!(summary.iterations, 1);
2069+
assert_eq!(summary.assistant_messages.len(), 1);
2070+
assert_eq!(summary.tool_results.len(), 1);
2071+
assert!(matches!(
2072+
&summary.tool_results[0].blocks[0],
2073+
ContentBlock::ToolResult {
2074+
tool_use_id,
2075+
is_error: true,
2076+
output,
2077+
..
2078+
} if tool_use_id == "tool-1" && output.contains("Interrupted by user")
2079+
));
2080+
// user text, assistant tool_use, synthesized tool_result
2081+
assert_eq!(runtime.session().messages.len(), 3);
2082+
assert!(matches!(
2083+
runtime.session().messages[2].blocks[0],
2084+
ContentBlock::ToolResult { is_error: true, .. }
2085+
));
2086+
}
2087+
2088+
#[test]
2089+
fn stream_error_during_interrupt_is_reported_as_interruption() {
2090+
struct AbortedApi {
2091+
interrupt: TurnInterruptSignal,
2092+
}
2093+
2094+
impl ApiClient for AbortedApi {
2095+
fn stream(
2096+
&mut self,
2097+
_request: ApiRequest,
2098+
) -> Result<Vec<AssistantEvent>, RuntimeError> {
2099+
// Simulate the streaming client aborting the connection after
2100+
// observing the interrupt flag mid-stream.
2101+
self.interrupt.interrupt();
2102+
Err(RuntimeError::new("request aborted"))
2103+
}
2104+
}
2105+
2106+
// given
2107+
let sink = Arc::new(MemoryTelemetrySink::default());
2108+
let tracer = SessionTracer::new("session-interrupt", sink.clone());
2109+
let interrupt = TurnInterruptSignal::new();
2110+
let mut runtime = ConversationRuntime::new(
2111+
Session::new(),
2112+
AbortedApi {
2113+
interrupt: interrupt.clone(),
2114+
},
2115+
StaticToolExecutor::new(),
2116+
PermissionPolicy::new(PermissionMode::DangerFullAccess),
2117+
vec!["system".to_string()],
2118+
)
2119+
.with_turn_interrupt_signal(interrupt)
2120+
.with_session_tracer(tracer);
2121+
2122+
// when
2123+
let summary = runtime
2124+
.run_turn("hello", None)
2125+
.expect("interrupt-driven aborts should not surface as errors");
2126+
2127+
// then
2128+
assert!(summary.interrupted);
2129+
assert!(summary.assistant_messages.is_empty());
2130+
let trace_names = sink
2131+
.events()
2132+
.iter()
2133+
.filter_map(|event| match event {
2134+
TelemetryEvent::SessionTrace(trace) => Some(trace.name.clone()),
2135+
_ => None,
2136+
})
2137+
.collect::<Vec<_>>();
2138+
assert!(trace_names.iter().any(|name| name == "turn_interrupted"));
2139+
assert!(!trace_names.iter().any(|name| name == "turn_failed"));
2140+
assert!(!trace_names.iter().any(|name| name == "turn_completed"));
2141+
}
18782142
}

‎rust/crates/runtime/src/lib.rs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -83,7 +83,7 @@ pub use config_validate::{
8383
pub use conversation::{
8484
auto_compaction_threshold_from_env, ApiClient, ApiRequest, AssistantEvent, AutoCompactionEvent,
8585
ConversationRuntime, PromptCacheEvent, RuntimeError, StaticToolExecutor, ToolError,
86-
ToolExecutor, TurnSummary,
86+
ToolExecutor, TurnInterruptSignal, TurnSummary,
8787
};
8888
pub use file_ops::{
8989
edit_file, edit_file_in_workspace, glob_search, glob_search_in_workspace, grep_search,

0 commit comments

Comments
 (0)