From 63b5dc893634f038d4afd309d162a35328a09c56 Mon Sep 17 00:00:00 2001 From: wsp Date: Tue, 8 Sep 2026 19:43:20 +0800 Subject: [PATCH] fix(agentic): make context compaction respond to cancellation - Cancel automatic and manual compaction preparation, hooks, and summary requests using the dialog turn cancellation token - Propagate cancellation without fallback compression or further rounds - Discard cancelled preparation results while preserving context commits - Cover early and in-flight cancellation with regression tests - Document focused compaction verification commands --- .../src/agentic/coordination/coordinator.rs | 186 +++++-- .../core/src/agentic/execution/AGENTS.md | 7 + .../src/agentic/execution/execution_engine.rs | 519 ++++++++++++------ .../src/agentic/execution/round_executor.rs | 19 + 4 files changed, 500 insertions(+), 231 deletions(-) diff --git a/src/crates/assembly/core/src/agentic/coordination/coordinator.rs b/src/crates/assembly/core/src/agentic/coordination/coordinator.rs index 999b70b08d..960b41820d 100644 --- a/src/crates/assembly/core/src/agentic/coordination/coordinator.rs +++ b/src/crates/assembly/core/src/agentic/coordination/coordinator.rs @@ -25,8 +25,8 @@ use crate::agentic::events::{ AgenticEvent, DeepReviewQueueState, EventPriority, EventQueue, EventRouter, EventSubscriber, }; use crate::agentic::execution::{ - ContextCompactionOutcome, ExecutionContext, ExecutionEngine, ExecutionResult, - ManualCompactionCommitGate, + prepare_compression_cancellable, ContextCompactionOutcome, ExecutionContext, ExecutionEngine, + ExecutionResult, ManualCompactionCommitGate, }; use crate::agentic::fork_agent::ForkAgentContextSnapshot; use crate::agentic::goal_mode::{ @@ -5693,61 +5693,73 @@ Update the persona files and delete BOOTSTRAP.md as soon as bootstrap is complet cancellation_token: CancellationToken, commit_gate: Arc, ) -> OpenBitFunResult<()> { - let manual_workspace_services = Self::build_workspace_services(&manual_workspace).await?; - let manual_execution_context = ExecutionContext { - session_id: session_id.clone(), - dialog_turn_id: turn_id.clone(), - turn_index, - agent_type: runtime_agent_type, - workspace: manual_workspace, - context: HashMap::from([( - "cancel_lifecycle_owner".to_string(), - "coordinator".to_string(), - )]), - subagent_parent_info: None, - permission_delegation: None, - permission_runtime_ceiling: None, - delegation_policy: DelegationPolicy::top_level(), - runtime_tool_restrictions: ToolRuntimeRestrictions::default(), - workspace_services: manual_workspace_services, - terminal_port, - remote_exec_port, - round_injection: None, - emit_lifecycle_events: false, - recover_partial_on_cancel: false, - }; - let session_max_tokens = session.config.max_context_tokens; - - // Unify context_window: min(model capability, session config) - let model_context_window = - match crate::infrastructure::ai::get_global_ai_client_factory().await { - Ok(factory) => { - let model_id = session.config.model_id.as_deref().unwrap_or("default"); - match factory.get_client_resolved(model_id).await { - Ok(client) => Some(client.config.context_window as usize), - Err(_) => None, - } - } - Err(_) => None, - }; - let context_window = match model_context_window { - Some(mcw) => mcw.min(session_max_tokens), - None => session_max_tokens, - }; let compression_id = format!("compression_{}", uuid::Uuid::new_v4()); - match execution_engine - .compact_session_context( - session_id.clone(), - turn_id.clone(), - compression_id.clone(), - manual_execution_context, - context_messages, - "manual", - cancellation_token, - commit_gate, - ) - .await - { + let mut context_window = session.config.max_context_tokens; + let result = async { + let (manual_execution_context, resolved_context_window) = + prepare_compression_cancellable(&cancellation_token, async { + let manual_workspace_services = + Self::build_workspace_services(&manual_workspace).await?; + let manual_execution_context = ExecutionContext { + session_id: session_id.clone(), + dialog_turn_id: turn_id.clone(), + turn_index, + agent_type: runtime_agent_type, + workspace: manual_workspace, + context: HashMap::from([( + "cancel_lifecycle_owner".to_string(), + "coordinator".to_string(), + )]), + subagent_parent_info: None, + permission_delegation: None, + permission_runtime_ceiling: None, + delegation_policy: DelegationPolicy::top_level(), + runtime_tool_restrictions: ToolRuntimeRestrictions::default(), + workspace_services: manual_workspace_services, + terminal_port, + remote_exec_port, + round_injection: None, + emit_lifecycle_events: false, + recover_partial_on_cancel: false, + }; + let session_max_tokens = session.config.max_context_tokens; + + // Unify context_window: min(model capability, session config) + let model_context_window = + match crate::infrastructure::ai::get_global_ai_client_factory().await { + Ok(factory) => { + let model_id = + session.config.model_id.as_deref().unwrap_or("default"); + match factory.get_client_resolved(model_id).await { + Ok(client) => Some(client.config.context_window as usize), + Err(_) => None, + } + } + Err(_) => None, + }; + let context_window = match model_context_window { + Some(mcw) => mcw.min(session_max_tokens), + None => session_max_tokens, + }; + Ok((manual_execution_context, context_window)) + }) + .await?; + context_window = resolved_context_window; + execution_engine + .compact_session_context( + session_id.clone(), + turn_id.clone(), + compression_id.clone(), + manual_execution_context, + context_messages, + "manual", + cancellation_token, + commit_gate, + ) + .await + } + .await; + match result { Ok(outcome) => { Self::finalize_manual_compaction_success( session_manager.as_ref(), @@ -15753,6 +15765,68 @@ mod tests { assert_eq!(result.error.as_deref(), Some("summary request failed")); } + #[tokio::test] + async fn manual_compaction_cancelled_before_setup_preserves_context_and_settles_turn() { + let (coordinator, session_manager) = test_coordinator(); + let workspace = tempfile::tempdir().unwrap(); + let session = session_manager + .create_session( + "Cancelled compaction".to_string(), + "agentic".to_string(), + SessionConfig { + workspace_path: Some(workspace.path().to_string_lossy().into_owned()), + ..SessionConfig::default() + }, + ) + .await + .unwrap(); + let session_id = session.session_id.clone(); + let original = vec![Message::user("Keep this context".to_string())]; + session_manager + .replace_context_messages(&session_id, original.clone()) + .await; + let turn_id = session_manager + .start_maintenance_turn( + &session_id, + "/compact".to_string(), + None, + Some(ConversationCoordinator::manual_compaction_metadata()), + ) + .await + .unwrap(); + let token = CancellationToken::new(); + token.cancel(); + let result = ConversationCoordinator::execute_manual_compaction_task( + session_manager.clone(), + coordinator.execution_engine.clone(), + coordinator.event_queue.clone(), + session, + original, + session_id.clone(), + turn_id.clone(), + 0, + "agentic".to_string(), + None, + None, + None, + token, + Arc::new(ManualCompactionCommitGate::planning()), + ) + .await; + assert!(matches!(result, Err(OpenBitFunError::Cancelled(_)))); + let session = session_manager.get_session(&session_id).unwrap(); + assert!(matches!(session.state, SessionState::Idle)); + assert_eq!(session.compression_state.compression_count, 0); + let messages = session_manager + .get_context_messages(&session_id) + .await + .unwrap(); + assert_eq!(messages.len(), 1); + assert!( + matches!(&messages[0].content, MessageContent::Text(text) if text == "Keep this context") + ); + } + #[tokio::test] async fn applied_manual_compaction_emits_failed_terminal_when_turn_persistence_fails() { let root = tempfile::tempdir().expect("test root"); diff --git a/src/crates/assembly/core/src/agentic/execution/AGENTS.md b/src/crates/assembly/core/src/agentic/execution/AGENTS.md index c632a73721..f92420edfe 100644 --- a/src/crates/assembly/core/src/agentic/execution/AGENTS.md +++ b/src/crates/assembly/core/src/agentic/execution/AGENTS.md @@ -21,3 +21,10 @@ tests take absolute JSONL input/report paths in `OPENBITFUN_SHELL_REPLAY_INPUT` `OPENBITFUN_SHELL_NORMAL_OUTPUT`. They only analyze strings; never execute archived commands. The ignored Bash append integration test requires `OPENBITFUN_SHELL_TEST_BASH` to name a trusted Bash 4+ executable and uses only isolated synthetic commands. + +For automatic/manual context compaction cancellation, preparation, and commit races, use: + +```bash +cargo test --locked -p openbitfun-core --no-default-features --features agent-runtime,git --lib compression +cargo test --locked -p openbitfun-core --no-default-features --features agent-runtime,git --lib compaction +``` diff --git a/src/crates/assembly/core/src/agentic/execution/execution_engine.rs b/src/crates/assembly/core/src/agentic/execution/execution_engine.rs index 8b85197453..38b63afadc 100644 --- a/src/crates/assembly/core/src/agentic/execution/execution_engine.rs +++ b/src/crates/assembly/core/src/agentic/execution/execution_engine.rs @@ -247,6 +247,28 @@ impl ManualCompactionCommitGate { } } +/// Cancel preparation as a unit, including provider retries and pre-compact hooks. +/// Recheck after completion because synchronous planning may observe cancellation +/// in the same poll that produces a result. Callers must commit outside this future. +pub(crate) async fn prepare_compression_cancellable( + cancellation_token: &CancellationToken, + preparation: impl std::future::Future>, +) -> OpenBitFunResult { + let result = tokio::select! { + biased; + _ = cancellation_token.cancelled() => { + return Err(OpenBitFunError::Cancelled("Context compaction cancelled".to_string())); + } + result = preparation => result, + }; + if cancellation_token.is_cancelled() { + return Err(OpenBitFunError::Cancelled( + "Context compaction cancelled".to_string(), + )); + } + result +} + fn manual_compaction_terminal_error(error: OpenBitFunError) -> OpenBitFunError { match error { error @ OpenBitFunError::Cancelled(_) => error, @@ -2495,6 +2517,7 @@ impl ExecutionEngine { } break; } + Err(err @ OpenBitFunError::Cancelled(_)) => return Err(err), Err(err) => { warn!( "Model-based compression failed, falling back to structured local compression: {}", @@ -2800,92 +2823,103 @@ impl ExecutionEngine { // Captured before `ai_client` is consumed by summary generation. let ai_client_model = ai_client.config.model.clone(); - native_hooks::dispatch_pre_compact( - Self::native_hook_facts(session_id, dialog_turn_id, workspace, &ai_client_model), - trigger, - ) - .await; - - // Emit compression started event - self.emit_event( - AgenticEvent::ContextCompressionStarted { - session_id: session_id.to_string(), - turn_id: dialog_turn_id.to_string(), - compression_id: compression_id.clone(), - trigger: trigger.to_string(), - tokens_before: before_pressure.total_tokens, - context_window, - }, - EventPriority::Normal, - ) - .await; + let cancellation_token = self.round_executor.ensure_cancel_token(dialog_turn_id); + let planned_result = prepare_compression_cancellable(&cancellation_token, async { + native_hooks::dispatch_pre_compact( + Self::native_hook_facts(session_id, dialog_turn_id, workspace, &ai_client_model), + trigger, + ) + .await; - // Execute compression - let compression_contract = self - .session_manager - .compression_contract_for_session(session_id, compression_contract_limit); - let model_exchange_trace_dir = self - .session_manager - .persistent_model_exchange_trace_dir(session_id) + // Emit compression started event + self.emit_event( + AgenticEvent::ContextCompressionStarted { + session_id: session_id.to_string(), + turn_id: dialog_turn_id.to_string(), + compression_id: compression_id.clone(), + trigger: trigger.to_string(), + tokens_before: before_pressure.total_tokens, + context_window, + }, + EventPriority::Normal, + ) .await; - let trace_config = prepare_model_exchange_trace_for_workspace( - session_id, - dialog_turn_id, - workspace, - model_exchange_trace_dir.as_deref(), - ModelExchangeTraceOperation { - kind: "context_compression", - id: &compression_id, - trigger: Some(trigger), - }, - ai_client.as_ref(), - ) - .await; - let planned_result = self - .build_planned_compression_result( + + // Execute compression + let compression_contract = self + .session_manager + .compression_contract_for_session(session_id, compression_contract_limit); + let model_exchange_trace_dir = self + .session_manager + .persistent_model_exchange_trace_dir(session_id) + .await; + let trace_config = prepare_model_exchange_trace_for_workspace( session_id, dialog_turn_id, - &runtime_messages, - context_window, - compression_contract, - ai_client, - model_request_context, - tool_definitions, - prepended_prompt_reminders, - primary_supports_image_understanding, workspace, - trace_config, + model_exchange_trace_dir.as_deref(), + ModelExchangeTraceOperation { + kind: "context_compression", + id: &compression_id, + trigger: Some(trigger), + }, + ai_client.as_ref(), ) .await; - match planned_result { - Ok(Some(mut compression_result)) => { - let boundary_turn_index = self - .session_manager - .get_turn_count(session_id) - .saturating_sub(1); - match self - .session_manager - .create_compression_transcript_reference( - session_id, - boundary_turn_index, - &compression_id, - trigger, - ) - .await - { - Ok(Some(reference)) => { - self.context_compressor.append_transcript_reference( - &mut compression_result, - &reference.uri, - &reference.index_range, - ); - } - Ok(None) => {} - Err(error) => warn!( - "Failed to create automatic compression transcript; continuing without reference: session_id={}, turn_id={}, error={}", - session_id, dialog_turn_id, error - ), + let planned_result = self + .build_planned_compression_result( + session_id, + dialog_turn_id, + &runtime_messages, + context_window, + compression_contract, + ai_client, + model_request_context, + tool_definitions, + prepended_prompt_reminders, + primary_supports_image_understanding, + workspace, + trace_config, + ) + .await; + let mut compression_result = match planned_result? { + Some(result) => result, + None => return Ok(None), + }; + let boundary_turn_index = self + .session_manager + .get_turn_count(session_id) + .saturating_sub(1); + match self + .session_manager + .create_compression_transcript_reference( + session_id, + boundary_turn_index, + &compression_id, + trigger, + ) + .await + { + Ok(Some(reference)) => { + self.context_compressor.append_transcript_reference( + &mut compression_result, + &reference.uri, + &reference.index_range, + ); } + Ok(None) => {} + Err(error) => warn!( + "Failed to create automatic compression transcript; continuing without reference: session_id={}, turn_id={}, error={}", + session_id, dialog_turn_id, error + ), + } + Ok(Some(compression_result)) + }) + .await; + // Preparation has no context writes. Once admitted here, finish the + // context commit without dropping it halfway through persistence. + match planned_result { + Ok(Some(compression_result)) => { self.session_manager .replace_context_messages(session_id, compression_result.messages.clone()) .await; @@ -2992,15 +3026,19 @@ impl ExecutionEngine { ) .await; - native_hooks::dispatch_post_compact( - Self::native_hook_facts( - session_id, - dialog_turn_id, - workspace, - &ai_client_model, - ), - trigger, - ) + let _ = prepare_compression_cancellable(&cancellation_token, async { + native_hooks::dispatch_post_compact( + Self::native_hook_facts( + session_id, + dialog_turn_id, + workspace, + &ai_client_model, + ), + trigger, + ) + .await; + Ok(()) + }) .await; Ok(Some((compressed_tokens, new_messages))) @@ -3019,7 +3057,7 @@ impl ExecutionEngine { ) .await; - Err(OpenBitFunError::Session(e.to_string())) + Err(manual_compaction_terminal_error(e)) } } } @@ -3045,98 +3083,89 @@ impl ExecutionEngine { OpenBitFunError::NotFound(format!("Session not found: {}", session_id)) })?; let start_time = std::time::Instant::now(); - let scaffold = self - .resolve_compression_runtime_scaffold(&session, &context) - .await?; - native_hooks::dispatch_pre_compact( - Self::native_hook_facts( - &session_id, - &dialog_turn_id, - context.workspace.as_ref(), - &scaffold.ai_client.config.model, - ), - trigger, - ) - .await; - let context_window = (scaffold.ai_client.config.context_window as usize) - .min(session.config.max_context_tokens); - let prepended_reminders = scaffold.prepended_prompt_reminders.ordered_reminders(); - let prepended_reminder_tokens = - Self::prepended_reminder_tokens_for_pressure(&prepended_reminders); - let compression_trigger_budget = - Self::compression_trigger_budget(context_window, scaffold.ai_client.config.max_tokens); - let mut runtime_messages = vec![scaffold.system_prompt_message.clone()]; - runtime_messages.extend(messages.clone()); - let before_pressure = Self::estimate_auto_compression_pressure( - &runtime_messages, - scaffold.tool_definitions.as_deref(), - context_window, - compression_trigger_budget, - prepended_reminder_tokens, - ); - - self.emit_event( - AgenticEvent::ContextCompressionStarted { - session_id: session_id.to_string(), - turn_id: dialog_turn_id.to_string(), - compression_id: compression_id.clone(), - trigger: trigger.to_string(), - tokens_before: before_pressure.total_tokens, + let preparation = prepare_compression_cancellable(&cancellation_token, async { + let scaffold = self + .resolve_compression_runtime_scaffold(&session, &context) + .await?; + native_hooks::dispatch_pre_compact( + Self::native_hook_facts( + &session_id, + &dialog_turn_id, + context.workspace.as_ref(), + &scaffold.ai_client.config.model, + ), + trigger, + ) + .await; + let context_window = (scaffold.ai_client.config.context_window as usize) + .min(session.config.max_context_tokens); + let prepended_reminders = scaffold.prepended_prompt_reminders.ordered_reminders(); + let prepended_reminder_tokens = + Self::prepended_reminder_tokens_for_pressure(&prepended_reminders); + let compression_trigger_budget = Self::compression_trigger_budget( context_window, - }, - EventPriority::Normal, - ) - .await; + scaffold.ai_client.config.max_tokens, + ); + let mut runtime_messages = vec![scaffold.system_prompt_message.clone()]; + runtime_messages.extend(messages.clone()); + let before_pressure = Self::estimate_auto_compression_pressure( + &runtime_messages, + scaffold.tool_definitions.as_deref(), + context_window, + compression_trigger_budget, + prepended_reminder_tokens, + ); - let compression_contract = self - .session_manager - .compression_contract_for_session(&session_id, scaffold.compression_contract_limit); - let model_exchange_trace_dir = self - .session_manager - .persistent_model_exchange_trace_dir(&session_id) + self.emit_event( + AgenticEvent::ContextCompressionStarted { + session_id: session_id.to_string(), + turn_id: dialog_turn_id.to_string(), + compression_id: compression_id.clone(), + trigger: trigger.to_string(), + tokens_before: before_pressure.total_tokens, + context_window, + }, + EventPriority::Normal, + ) .await; - let trace_config = prepare_model_exchange_trace_for_workspace( - &session_id, - &dialog_turn_id, - context.workspace.as_ref(), - model_exchange_trace_dir.as_deref(), - ModelExchangeTraceOperation { - kind: "context_compression", - id: &compression_id, - trigger: Some(trigger), - }, - scaffold.ai_client.as_ref(), - ) - .await; - let planned_result = tokio::select! { - biased; - _ = cancellation_token.cancelled() => { - Err(OpenBitFunError::Cancelled("Manual context compaction cancelled".to_string())) - } - result = self.build_planned_compression_result( + + let compression_contract = self + .session_manager + .compression_contract_for_session(&session_id, scaffold.compression_contract_limit); + let model_exchange_trace_dir = self + .session_manager + .persistent_model_exchange_trace_dir(&session_id) + .await; + let trace_config = prepare_model_exchange_trace_for_workspace( &session_id, &dialog_turn_id, - &runtime_messages, - context_window, - compression_contract, - scaffold.ai_client.clone(), - &scaffold.model_request_context, - &scaffold.tool_definitions, - &scaffold.prepended_prompt_reminders, - scaffold.primary_supports_image_understanding, context.workspace.as_ref(), - trace_config, - ) => result, - }; - let planned_result = match planned_result { - Ok(result) if commit_gate.try_begin_commit() => Ok(result), - Ok(_) => Err(OpenBitFunError::Cancelled( - "Manual context compaction cancelled".to_string(), - )), - Err(error) => Err(error), - }; - match planned_result { - Ok(Some(mut compression_result)) => { + model_exchange_trace_dir.as_deref(), + ModelExchangeTraceOperation { + kind: "context_compression", + id: &compression_id, + trigger: Some(trigger), + }, + scaffold.ai_client.as_ref(), + ) + .await; + let mut planned_result = self + .build_planned_compression_result( + &session_id, + &dialog_turn_id, + &runtime_messages, + context_window, + compression_contract, + scaffold.ai_client.clone(), + &scaffold.model_request_context, + &scaffold.tool_definitions, + &scaffold.prepended_prompt_reminders, + scaffold.primary_supports_image_understanding, + context.workspace.as_ref(), + trace_config, + ) + .await?; + if let Some(compression_result) = planned_result.as_mut() { let boundary_turn_index = self .session_manager .get_turn_count(&session_id) @@ -3153,7 +3182,7 @@ impl ExecutionEngine { { Ok(Some(reference)) => { self.context_compressor.append_transcript_reference( - &mut compression_result, + compression_result, &reference.uri, &reference.index_range, ); @@ -3164,6 +3193,41 @@ impl ExecutionEngine { session_id, dialog_turn_id, error ), } + } + Ok((scaffold, before_pressure, planned_result)) + }) + .await; + let (scaffold, before_pressure, planned_result) = match preparation { + Ok(result) => result, + Err(err) => { + self.emit_event( + AgenticEvent::ContextCompressionFailed { + session_id: session_id.clone(), + turn_id: dialog_turn_id.clone(), + compression_id: compression_id.clone(), + error: err.to_string(), + }, + EventPriority::High, + ) + .await; + return Err(manual_compaction_terminal_error(err)); + } + }; + let context_window = before_pressure.context_window; + let compression_trigger_budget = + Self::compression_trigger_budget(context_window, scaffold.ai_client.config.max_tokens); + let prepended_reminders = scaffold.prepended_prompt_reminders.ordered_reminders(); + let prepended_reminder_tokens = + Self::prepended_reminder_tokens_for_pressure(&prepended_reminders); + let planned_result = if commit_gate.try_begin_commit() { + Ok(planned_result) + } else { + Err(OpenBitFunError::Cancelled( + "Manual context compaction cancelled".to_string(), + )) + }; + match planned_result { + Ok(Some(compression_result)) => { let compressed_messages = compression_result.messages; self.session_manager .replace_context_messages(&session_id, compressed_messages.clone()) @@ -3802,6 +3866,12 @@ impl ExecutionEngine { // Loop to execute model rounds loop { + if self + .round_executor + .is_dialog_turn_cancelled(&dialog_turn_id) + { + return Err(OpenBitFunError::Cancelled("Dialog cancelled".to_string())); + } if reached_fixed_model_round_limit(self.config.max_rounds, completed_rounds) { warn!( "Reached max rounds limit: {}, stopping execution", @@ -3982,6 +4052,12 @@ impl ExecutionEngine { compressed_tokens, ); + if self + .round_executor + .is_dialog_turn_cancelled(&dialog_turn_id) + { + return Err(OpenBitFunError::Cancelled("Dialog cancelled".to_string())); + } messages = compressed_messages; turn_prompt_scaffold = self .resolve_turn_prompt_scaffold(TurnPromptScaffoldInput { @@ -4006,6 +4082,7 @@ impl ExecutionEngine { debug!("No eligible multi-turn context available for compression"); consecutive_compression_failures = 0; } + Err(err @ OpenBitFunError::Cancelled(_)) => return Err(err), Err(e) => { consecutive_compression_failures += 1; compression_failure_count += 1; @@ -4216,6 +4293,14 @@ impl ExecutionEngine { send_pressure.total_tokens, compressed_tokens ); + if self + .round_executor + .is_dialog_turn_cancelled(&dialog_turn_id) + { + return Err(OpenBitFunError::Cancelled( + "Dialog cancelled".to_string(), + )); + } messages = compressed_messages; turn_prompt_scaffold = self .resolve_turn_prompt_scaffold(TurnPromptScaffoldInput { @@ -4252,6 +4337,7 @@ impl ExecutionEngine { ); return Err(err); } + Err(err @ OpenBitFunError::Cancelled(_)) => return Err(err), Err(compression_error) => { error!( "Context-overflow recovery compression failed: session_id={}, turn_id={}, round_index={}, error={}", @@ -5205,6 +5291,89 @@ mod tests { use std::sync::Arc; use std::time::Duration; + #[tokio::test] + async fn compression_cancellation_before_preparation_does_not_start_work() { + let token = tokio_util::sync::CancellationToken::new(); + token.cancel(); + let result = super::prepare_compression_cancellable(&token, async { + panic!("Cancelled preparation must not be polled"); + #[allow(unreachable_code)] + Ok(()) + }) + .await; + assert!(matches!(result, Err(crate::OpenBitFunError::Cancelled(_)))); + } + + #[tokio::test] + async fn compression_cancellation_drops_pending_work_without_commit() { + struct DropProbe(Arc); + impl Drop for DropProbe { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + let token = tokio_util::sync::CancellationToken::new(); + let dropped = Arc::new(AtomicBool::new(false)); + let committed = Arc::new(AtomicBool::new(false)); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let work_token = token.clone(); + let work_dropped = dropped.clone(); + let work_committed = committed.clone(); + let task = tokio::spawn(async move { + let result = super::prepare_compression_cancellable(&work_token, async { + let _probe = DropProbe(work_dropped); + started_tx.send(()).unwrap(); + std::future::pending::<()>().await; + Ok(()) + }) + .await; + if result.is_ok() { + work_committed.store(true, Ordering::SeqCst); + } + result + }); + started_rx.await.unwrap(); + token.cancel(); + let result = tokio::time::timeout(Duration::from_secs(1), task) + .await + .expect("Cancellation must not wait for pending work") + .unwrap(); + assert!(matches!(result, Err(crate::OpenBitFunError::Cancelled(_)))); + assert!(dropped.load(Ordering::SeqCst)); + assert!(!committed.load(Ordering::SeqCst)); + } + + #[tokio::test] + async fn compression_cancellation_in_completion_poll_rejects_result() { + let token = tokio_util::sync::CancellationToken::new(); + let result = super::prepare_compression_cancellable(&token, async { + token.cancel(); + Ok("summary that must not be committed") + }) + .await; + assert!(matches!(result, Err(crate::OpenBitFunError::Cancelled(_)))); + } + + #[tokio::test] + async fn compression_preparation_preserves_success_and_failure() { + let token = tokio_util::sync::CancellationToken::new(); + assert_eq!( + super::prepare_compression_cancellable(&token, async { Ok(42) }) + .await + .unwrap(), + 42 + ); + let result = super::prepare_compression_cancellable::<()>(&token, async { + Err(crate::OpenBitFunError::Session( + "preparation failed".to_string(), + )) + }) + .await; + assert!( + matches!(result, Err(crate::OpenBitFunError::Session(message)) if message == "preparation failed") + ); + } + fn resolved_tool_manifest( allowed_tool_names: &[&str], tool_definition_names: &[&str], diff --git a/src/crates/assembly/core/src/agentic/execution/round_executor.rs b/src/crates/assembly/core/src/agentic/execution/round_executor.rs index 2a55a2542d..11c0ace1a7 100644 --- a/src/crates/assembly/core/src/agentic/execution/round_executor.rs +++ b/src/crates/assembly/core/src/agentic/execution/round_executor.rs @@ -1326,6 +1326,11 @@ impl RoundExecutor { self.cancellation_tokens.insert(dialog_turn_id, token); } + /// Reuse an early registered token, including its already-cancelled state. + pub(crate) fn ensure_cancel_token(&self, dialog_turn_id: &str) -> CancellationToken { + self.cancellation_tokens.get_or_insert_new(dialog_turn_id) + } + /// Return a clone of the cancellation token registered for a dialog turn. pub fn cancel_token_for_dialog_turn(&self, dialog_turn_id: &str) -> Option { self.cancellation_tokens.token(dialog_turn_id) @@ -2128,6 +2133,20 @@ mod tests { assert!(executor.cancel_token_for_dialog_turn("missing").is_none()); } + #[tokio::test] + async fn compression_token_is_cancellable_before_first_round_and_keeps_early_cancel() { + let executor = test_round_executor(); + let token = executor.ensure_cancel_token("turn-1"); + executor.cancel_dialog_turn("turn-1").await.unwrap(); + assert!(token.is_cancelled()); + assert!(executor.ensure_cancel_token("turn-1").is_cancelled()); + + let early = CancellationToken::new(); + early.cancel(); + executor.register_cancel_token("turn-2", early); + assert!(executor.ensure_cancel_token("turn-2").is_cancelled()); + } + #[tokio::test] async fn cancel_keeps_token_registered_until_cleanup() { let executor = test_round_executor();