11use std:: collections:: BTreeMap ;
22use std:: fmt:: { Display , Formatter } ;
3+ use std:: sync:: {
4+ atomic:: { AtomicBool , Ordering } ,
5+ Arc ,
6+ } ;
37
48use serde_json:: { Map , Value } ;
59use 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`].
5796pub 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}
0 commit comments