From bd893f271ac2c77933607dd65f875aff369d0181 Mon Sep 17 00:00:00 2001 From: Hare Date: Thu, 6 Aug 2026 09:49:48 +0900 Subject: [PATCH] worker: treat rolled-back extraction as cancelled --- crates/worker/src/internal_worker.rs | 68 ++++++++++++++++++++++++++++ crates/worker/src/worker.rs | 37 ++++++++++++++- 2 files changed, 104 insertions(+), 1 deletion(-) diff --git a/crates/worker/src/internal_worker.rs b/crates/worker/src/internal_worker.rs index 027a2ee7..65d0dcfd 100644 --- a/crates/worker/src/internal_worker.rs +++ b/crates/worker/src/internal_worker.rs @@ -72,6 +72,16 @@ pub(crate) struct InternalWorkerError { pub(crate) async fn run_internal_worker( spec: InternalWorkerSpec, ) -> Result { + run_internal_worker_with_prepare(spec, |_| {}).await +} + +async fn run_internal_worker_with_prepare( + spec: InternalWorkerSpec, + prepare: F, +) -> Result +where + F: FnOnce(&mut Worker, EphemeralSessionStore>), +{ let InternalWorkerSpec { identity, mut manifest, @@ -163,6 +173,7 @@ pub(crate) async fn run_internal_worker( } let session_id = worker.session_id(); let segment_id = worker.segment_id(); + prepare(&mut worker); match worker.run_text(&input).await { Ok(lifecycle) => Ok(InternalWorkerResult { @@ -355,6 +366,35 @@ mod tests { } } + #[derive(Clone)] + struct CancelBeforeAiClient { + calls: Arc, + cancel_sender: Arc>>>, + } + + #[async_trait] + impl LlmClient for CancelBeforeAiClient { + fn clone_boxed(&self) -> Box { + Box::new(self.clone()) + } + + async fn stream( + &self, + _request: Request, + ) -> Result> + Send>>, ClientError> + { + self.calls.fetch_add(1, Ordering::SeqCst); + let sender = self + .cancel_sender + .lock() + .expect("cancel sender lock") + .clone() + .expect("cancel sender installed before run"); + sender.send(()).await.expect("internal engine is live"); + Ok(Box::pin(futures::stream::pending())) + } + } + fn manifest() -> WorkerManifest { WorkerManifest::from_toml( r#" @@ -414,6 +454,34 @@ permission = "write" assert_eq!(result.identity.kind, "test"); } + #[tokio::test] + async fn cancellation_before_ai_item_returns_rolled_back_lifecycle() { + let calls = Arc::new(AtomicUsize::new(0)); + let cancel_sender = Arc::new(Mutex::new(None)); + let mut internal_spec = spec(calls.clone(), &[]); + internal_spec.client = Box::new(CancelBeforeAiClient { + calls: calls.clone(), + cancel_sender: cancel_sender.clone(), + }); + let prepare_sender = cancel_sender.clone(); + + let result = match run_internal_worker_with_prepare(internal_spec, move |worker| { + *prepare_sender.lock().expect("cancel sender lock") = + Some(worker.engine_mut().cancel_sender()); + }) + .await + { + Ok(result) => result, + Err(error) => panic!( + "Worker rollback should remain a lifecycle result: {:?}", + error.source + ), + }; + + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert!(matches!(result.lifecycle, WorkerRunResult::RolledBack)); + } + #[tokio::test] async fn rejects_missing_explicit_tools_before_model_execution() { let calls = Arc::new(AtomicUsize::new(0)); diff --git a/crates/worker/src/worker.rs b/crates/worker/src/worker.rs index b3e2ccec..ee056d99 100644 --- a/crates/worker/src/worker.rs +++ b/crates/worker/src/worker.rs @@ -3475,7 +3475,22 @@ impl Worker { lifecycle = ?result.lifecycle, "internal Worker execution completed" ); - result.usage.as_ref().map(usage_audit_from_event) + let usage = result.usage.as_ref().map(usage_audit_from_event); + if let Some(error) = extract_internal_worker_lifecycle_error(&result.lifecycle) { + audit + .emit( + self.workspace_client(), + event_tx, + memory::audit::WorkerLifecycleStatus::Cancelled, + "worker_cancelled: internal Worker run rolled back before AI output", + usage, + Some(extract_audit_base), + None, + ) + .await; + return Err(error); + } + usage } Err(err) => { tracing::debug!( @@ -3629,6 +3644,13 @@ impl Worker { } } +fn extract_internal_worker_lifecycle_error(lifecycle: &WorkerRunResult) -> Option { + match lifecycle { + WorkerRunResult::RolledBack => Some(WorkerError::Engine(EngineError::Cancelled)), + WorkerRunResult::Finished | WorkerRunResult::Paused | WorkerRunResult::LimitReached => None, + } +} + fn lifecycle_status_for_worker_error(err: &WorkerError) -> memory::audit::WorkerLifecycleStatus { if matches!(err, WorkerError::Engine(EngineError::Cancelled)) { memory::audit::WorkerLifecycleStatus::Cancelled @@ -6323,6 +6345,19 @@ mod build_summary_prompt_tests { server.join().unwrap(); } + #[test] + fn rolled_back_internal_extract_is_cancelled_before_pointer_commit() { + let error = extract_internal_worker_lifecycle_error(&WorkerRunResult::RolledBack) + .expect("rolled-back extract must not enter the success path"); + + assert!(matches!(error, WorkerError::Engine(EngineError::Cancelled))); + assert!(matches!( + lifecycle_status_for_worker_error(&error), + memory::audit::WorkerLifecycleStatus::Cancelled + )); + assert!(extract_internal_worker_lifecycle_error(&WorkerRunResult::Finished).is_none()); + } + fn minimal_manifest() -> WorkerManifest { let toml_str = r#" [worker]