diff --git a/crates/protocol/src/lib.rs b/crates/protocol/src/lib.rs index ec96e226..f66be5dc 100644 --- a/crates/protocol/src/lib.rs +++ b/crates/protocol/src/lib.rs @@ -203,6 +203,53 @@ impl WorkerStateSnapshot { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerStateSnapshotApply { + Applied, + Duplicate, + Stale, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct WorkerStateSnapshotConflict { + pub execution_generation: u64, + pub revision: u64, +} + +impl std::fmt::Display for WorkerStateSnapshotConflict { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!( + formatter, + "conflicting worker state snapshots at generation {} revision {}", + self.execution_generation, self.revision + ) + } +} + +impl std::error::Error for WorkerStateSnapshotConflict {} + +pub fn apply_worker_state_snapshot( + current: &mut WorkerStateSnapshot, + incoming: &WorkerStateSnapshot, +) -> Result { + use std::cmp::Ordering; + + let ordering = (incoming.execution_generation, incoming.revision) + .cmp(&(current.execution_generation, current.revision)); + match ordering { + Ordering::Greater => { + *current = incoming.clone(); + Ok(WorkerStateSnapshotApply::Applied) + } + Ordering::Less => Ok(WorkerStateSnapshotApply::Stale), + Ordering::Equal if incoming == current => Ok(WorkerStateSnapshotApply::Duplicate), + Ordering::Equal => Err(WorkerStateSnapshotConflict { + execution_generation: incoming.execution_generation, + revision: incoming.revision, + }), + } +} + impl From for WorkerStateSnapshot { fn from(status: WorkerStatus) -> Self { let state = match status { @@ -1574,6 +1621,58 @@ pub enum Permission { mod tests { use super::*; + #[test] + fn worker_state_snapshot_apply_is_monotonic_and_detects_conflicts() { + let mut current = WorkerStateSnapshot::initial(4); + let mut newer = current.clone(); + newer.revision = 1; + newer.state = WorkerState::Busy(WorkerBusyState::Run(WorkerRunState::Running)); + + assert_eq!( + apply_worker_state_snapshot(&mut current, &newer), + Ok(WorkerStateSnapshotApply::Applied) + ); + assert_eq!( + apply_worker_state_snapshot(&mut current, &newer), + Ok(WorkerStateSnapshotApply::Duplicate) + ); + + let stale_revision = WorkerStateSnapshot::initial(4); + assert_eq!( + apply_worker_state_snapshot(&mut current, &stale_revision), + Ok(WorkerStateSnapshotApply::Stale) + ); + let stale_generation = WorkerStateSnapshot { + execution_generation: 3, + revision: u64::MAX, + ..newer.clone() + }; + assert_eq!( + apply_worker_state_snapshot(&mut current, &stale_generation), + Ok(WorkerStateSnapshotApply::Stale) + ); + + let conflicting = WorkerStateSnapshot { + state: WorkerState::Idle, + ..newer.clone() + }; + assert_eq!( + apply_worker_state_snapshot(&mut current, &conflicting), + Err(WorkerStateSnapshotConflict { + execution_generation: 4, + revision: 1, + }) + ); + assert_eq!(current, newer); + + let next_generation = WorkerStateSnapshot::initial(5); + assert_eq!( + apply_worker_state_snapshot(&mut current, &next_generation), + Ok(WorkerStateSnapshotApply::Applied) + ); + assert_eq!(current, next_generation); + } + #[test] fn method_submit_json_roundtrip_and_run_is_rejected() { let json = r#"{"method":"submit","params":{"submission_request_id":"request-1","input":[{"kind":"text","content":"Hello"}]}}"#; diff --git a/crates/tui/src/app.rs b/crates/tui/src/app.rs index a76d32b5..c543f991 100644 --- a/crates/tui/src/app.rs +++ b/crates/tui/src/app.rs @@ -1128,6 +1128,22 @@ impl App { command } + fn apply_worker_state_snapshot(&mut self, snapshot: &WorkerStateSnapshot) { + match protocol::apply_worker_state_snapshot(&mut self.worker_state, snapshot) { + Ok(protocol::WorkerStateSnapshotApply::Applied) => { + self.set_worker_status(self.worker_state.catalog_status()); + } + Ok( + protocol::WorkerStateSnapshotApply::Duplicate + | protocol::WorkerStateSnapshotApply::Stale, + ) => {} + Err(error) => self.handle_error( + ErrorCode::Internal, + format!("worker state stream rejected: {error}"), + ), + } + } + pub fn handle_worker_event(&mut self, event: Event) -> Option { if self.rewind_refresh_fence && event_is_stale_after_rewind(&event) { return None; @@ -1465,8 +1481,7 @@ impl App { self.pending_submissions = session.pending_submissions.clone(); self.restore_snapshot(&session, greeting, in_flight); self.replace_internal_worker_snapshots(internal_workers); - self.worker_state = state.clone(); - self.set_worker_status(state.catalog_status()); + self.apply_worker_state_snapshot(&state); } Event::InternalWorker { worker, @@ -1478,12 +1493,10 @@ impl App { } Event::WorkerState { snapshot } => { self.rewind_refresh_fence = false; - self.worker_state = snapshot.clone(); - self.set_worker_status(snapshot.catalog_status()); + self.apply_worker_state_snapshot(&snapshot); } Event::CommandAcknowledged { acknowledgement } => { - self.worker_state = acknowledgement.state.clone(); - self.set_worker_status(acknowledgement.state.catalog_status()); + self.apply_worker_state_snapshot(&acknowledgement.state); } // Command telemetry is an operational Web Console surface. The // TUI continues to render the final Bash ToolResult from history. @@ -3559,7 +3572,7 @@ mod completion_flow_tests { app.handle_worker_event(Event::Snapshot { greeting: test_greeting(), session: public_session(vec![session_start_value]), - state: WorkerStatus::Running.into(), + state: test_worker_state(WorkerStatus::Running), in_flight: Default::default(), internal_workers: Vec::new(), }); @@ -3570,6 +3583,59 @@ mod completion_flow_tests { assert!(matches!(app.blocks.first(), Some(Block::Greeting(_)))); } + #[test] + fn worker_state_events_and_acknowledgements_share_monotonic_application() { + let mut app = App::new("test".into()); + let running = WorkerStateSnapshot { + execution_generation: 4, + revision: 3, + state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run( + protocol::WorkerRunState::Running, + )), + last_command_id: 2, + }; + app.handle_worker_event(Event::WorkerState { + snapshot: running.clone(), + }); + app.handle_worker_event(Event::WorkerState { + snapshot: WorkerStateSnapshot { + revision: 2, + state: protocol::WorkerState::Idle, + ..running.clone() + }, + }); + assert_eq!(app.worker_state, running); + + let paused = WorkerStateSnapshot { + revision: 4, + state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run( + protocol::WorkerRunState::Paused, + )), + last_command_id: 3, + ..running.clone() + }; + app.handle_worker_event(Event::CommandAcknowledged { + acknowledgement: protocol::WorkerCommandAcknowledgement { + command_id: 3, + command: protocol::WorkerCommandKind::Pause, + disposition: protocol::WorkerCommandDisposition::Accepted, + state: paused.clone(), + }, + }); + assert_eq!(app.worker_state, paused); + + app.handle_worker_event(Event::WorkerState { + snapshot: WorkerStateSnapshot { + state: protocol::WorkerState::Idle, + ..paused.clone() + }, + }); + assert_eq!(app.worker_state, paused); + assert!(app.run_error_messages.iter().any(|message| { + message.contains("conflicting worker state snapshots at generation 4 revision 4") + })); + } + #[test] fn snapshot_replaces_live_error_with_one_durable_run_error_block() { let mut app = App::new("test".into()); @@ -3603,7 +3669,7 @@ mod completion_flow_tests { app.handle_worker_event(Event::Snapshot { greeting: test_greeting(), session: public_session(vec![serde_json::to_value(run_errored).unwrap()]), - state: WorkerStatus::Idle.into(), + state: test_worker_state(WorkerStatus::Idle), in_flight: Default::default(), internal_workers: Vec::new(), }); @@ -3667,7 +3733,7 @@ mod completion_flow_tests { pending_submissions: protocol::PendingSubmissionsSnapshot::default(), entries: Vec::new(), }, - state: WorkerStatus::Running.into(), + state: test_worker_state(WorkerStatus::Running), in_flight: InFlightSnapshot { blocks: vec![ InFlightBlock::Thinking { @@ -3994,7 +4060,7 @@ mod completion_flow_tests { pending_submissions: protocol::PendingSubmissionsSnapshot::default(), entries: Vec::new(), }, - state: WorkerStatus::Idle.into(), + state: test_worker_state(WorkerStatus::Idle), in_flight: Default::default(), internal_workers: Vec::new(), }); @@ -4046,7 +4112,7 @@ mod completion_flow_tests { pending_submissions: protocol::PendingSubmissionsSnapshot::default(), entries: Vec::new(), }, - state: WorkerStatus::Idle.into(), + state: test_worker_state(WorkerStatus::Idle), in_flight: Default::default(), internal_workers: vec![InternalWorkerSnapshot { worker: InternalWorkerRef { @@ -4194,6 +4260,13 @@ mod completion_flow_tests { .count() } + fn test_worker_state(status: WorkerStatus) -> WorkerStateSnapshot { + let mut snapshot = WorkerStateSnapshot::from(status); + snapshot.execution_generation = 1; + snapshot.revision = 1; + snapshot + } + fn test_greeting() -> protocol::Greeting { protocol::Greeting { worker_name: "test".into(), @@ -4220,7 +4293,7 @@ mod completion_flow_tests { entries: Vec::new(), }, greeting, - state: WorkerStatus::Idle.into(), + state: test_worker_state(WorkerStatus::Idle), in_flight: Default::default(), internal_workers: Vec::new(), }); @@ -4419,7 +4492,7 @@ mod completion_flow_tests { app.handle_worker_event(Event::Snapshot { greeting: test_greeting(), session: public_session(assistant_item_entries), - state: WorkerStatus::Running.into(), + state: test_worker_state(WorkerStatus::Running), in_flight: Default::default(), internal_workers: Vec::new(), }); diff --git a/crates/worker-runtime/src/worker_backend.rs b/crates/worker-runtime/src/worker_backend.rs index 6c414874..20b05aed 100644 --- a/crates/worker-runtime/src/worker_backend.rs +++ b/crates/worker-runtime/src/worker_backend.rs @@ -10,7 +10,7 @@ use std::collections::HashMap; use std::future::Future; use std::path::{Path, PathBuf}; -use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, Mutex, RwLock, mpsc}; use std::time::Duration; @@ -38,7 +38,9 @@ use crate::working_directory::{ WorkingDirectoryBinding, WorkingDirectoryDiagnostic, WorkingDirectoryMaterializer, }; use async_trait::async_trait; -use protocol::{Event, Method, Segment, WorkerCommandEnvelope, WorkerStatus}; +#[cfg(test)] +use protocol::WorkerStatus; +use protocol::{Event, Method, Segment, WorkerCommandEnvelope}; static NEXT_INTERNAL_COMMAND_ID: AtomicU64 = AtomicU64::new(1); @@ -1197,7 +1199,6 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory { struct RuntimeWorkerExecution { handle: WorkerHandle, shutdown: Arc>>, - busy: Arc, worker_state: Arc>, workspace_client: Option>, } @@ -1296,7 +1297,6 @@ where ) -> Result< ( WorkerHandle, - Arc, Arc>, Option>, ), @@ -1323,7 +1323,6 @@ where .map(|execution| { ( execution.handle.clone(), - execution.busy.clone(), execution.worker_state.clone(), execution.workspace_client.clone(), ) @@ -1434,48 +1433,31 @@ where working_directory: Option, workspace_client: Option>, ) -> WorkerExecutionSpawnResult { - let busy = Arc::new(AtomicBool::new(false)); let worker_state = Arc::new(RwLock::new(handle.shared_state.snapshot())); #[cfg(feature = "ws-server")] { let streams = subscribe_worker_protocol_session(&handle); let mut events = streams.events; let mut entry_events = streams.log_entries; - let bridge_busy = busy.clone(); let bridge_worker_state = worker_state.clone(); if let Err(message) = self.spawn_on_adapter_runtime(async move { loop { tokio::select! { event = events.recv() => { match event { - Ok(event) => { - let next_state = match &event { - Event::WorkerState { snapshot } - | Event::Snapshot { state: snapshot, .. } => { - Some(snapshot.clone()) + Ok(mut event) => { + match apply_protocol_worker_state(&bridge_worker_state, &mut event) { + Ok(true) => { + let _ = bridge_context.publish_protocol_event(event); } - Event::CommandAcknowledged { acknowledgement } => { - Some(acknowledgement.state.clone()) + Ok(false) => {} + Err(message) => { + let _ = bridge_context.publish_protocol_event(Event::Error { + code: protocol::ErrorCode::Internal, + message: format!("worker state stream rejected: {message}"), + }); + break; } - _ => None, - }; - let next_busy = next_state - .as_ref() - .map(worker_state_is_executing) - .or_else(|| matches!(event, Event::Shutdown).then_some(false)); - let _ = bridge_context.publish_protocol_event(event); - if let Some(next_state) = next_state { - if let Ok(mut current) = bridge_worker_state.write() { - if next_state.execution_generation > current.execution_generation - || (next_state.execution_generation == current.execution_generation - && next_state.revision >= current.revision) - { - *current = next_state; - } - } - } - if let Some(next_busy) = next_busy { - bridge_busy.store(next_busy, Ordering::SeqCst); } } Err(broadcast::error::RecvError::Lagged(_)) => continue, @@ -1520,7 +1502,6 @@ where RuntimeWorkerExecution { handle, shutdown, - busy, worker_state, workspace_client, }, @@ -1543,32 +1524,28 @@ impl Drop for WorkerRuntimeExecutionBackend { } } -fn worker_state_is_executing(snapshot: &protocol::WorkerStateSnapshot) -> bool { - matches!( - snapshot.state, - protocol::WorkerState::Busy(protocol::WorkerBusyState::Run( - protocol::WorkerRunState::Running - | protocol::WorkerRunState::Pausing - | protocol::WorkerRunState::Cancelling - )) | protocol::WorkerState::Busy(protocol::WorkerBusyState::Maintenance(_)) - ) -} - -fn method_starts_turn(method: &Method) -> bool { - matches!( - method, - Method::Submit { .. } - | Method::SubmitTracked { .. } - | Method::Notify { auto_run: true, .. } - | Method::NotifyTracked { auto_run: true, .. } - | Method::Resume { .. } - ) -} - -fn method_can_start_turn_from_status(method: &Method, status: WorkerStatus) -> bool { - match method { - Method::Resume { .. } => matches!(status, WorkerStatus::Idle | WorkerStatus::Paused), - _ => status == WorkerStatus::Idle, +fn apply_protocol_worker_state( + current: &Arc>, + event: &mut Event, +) -> Result { + let (incoming, replace_stale) = match event { + Event::WorkerState { snapshot } => (snapshot, false), + Event::Snapshot { state, .. } => (state, true), + Event::CommandAcknowledged { acknowledgement } => (&mut acknowledgement.state, true), + _ => return Ok(true), + }; + let mut current = current + .write() + .map_err(|_| "worker state projection lock is poisoned".to_string())?; + match protocol::apply_worker_state_snapshot(&mut current, incoming) { + Ok(protocol::WorkerStateSnapshotApply::Applied) + | Ok(protocol::WorkerStateSnapshotApply::Duplicate) => Ok(true), + Ok(protocol::WorkerStateSnapshotApply::Stale) if replace_stale => { + *incoming = current.clone(); + Ok(true) + } + Ok(protocol::WorkerStateSnapshotApply::Stale) => Ok(false), + Err(error) => Err(error.to_string()), } } @@ -1897,7 +1874,7 @@ where handle: &WorkerExecutionHandle, input: WorkerInput, ) -> WorkerExecutionResult { - let (worker, busy, worker_state, _workspace_client) = match self.get_execution(handle) { + let (worker, worker_state, _workspace_client) = match self.get_execution(handle) { Ok(execution) => execution, Err(mut result) => { result.operation = WorkerExecutionOperation::Input; @@ -1906,15 +1883,10 @@ where }; if input.kind == WorkerInputKind::Notify { - let status = worker.shared_state.catalog_status(); - let claimed_here = status == WorkerStatus::Idle - && busy - .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst) - .is_ok(); let notification_request_id = input .submission_request_id .unwrap_or_else(protocol::new_submission_request_id); - let result = self.send_method( + return self.send_method( WorkerExecutionOperation::Input, worker, Method::NotifyTracked { @@ -1926,11 +1898,6 @@ where }, }, ); - if claimed_here && result.outcome != crate::execution::WorkerExecutionOutcome::Accepted - { - busy.store(false, Ordering::SeqCst); - } - return result; } if input.kind == WorkerInputKind::Compact { @@ -1947,26 +1914,12 @@ where ); } - let is_user_submit = input.kind == WorkerInputKind::User; - let status = worker.shared_state.catalog_status(); - let claimed_here = status == WorkerStatus::Idle - && busy - .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst) - .is_ok(); - if !is_user_submit && !claimed_here { - return WorkerExecutionResult::busy( - WorkerExecutionOperation::Input, - "Worker is already running", - ); - } - let (method, submission_request_id) = match input.kind { WorkerInputKind::User => { let Some(submission_id) = input .submission_request_id .filter(|submission_id| !submission_id.trim().is_empty()) else { - busy.store(false, Ordering::SeqCst); return WorkerExecutionResult::rejected( WorkerExecutionOperation::Input, "Runtime user input is missing its internal submission id", @@ -1986,7 +1939,7 @@ where ) } WorkerInputKind::Notify => { - unreachable!("Notify input is dispatched before the turn-start busy guard") + unreachable!("Notify input is dispatched before ordinary input mapping") } WorkerInputKind::Compact => unreachable!("compact input is dispatched above"), WorkerInputKind::ListRewindTargets => (Method::ListRewindTargets, None), @@ -1999,7 +1952,7 @@ where }; let waits_for_submission_acceptance = submission_request_id.is_some(); - let result = if waits_for_submission_acceptance { + if waits_for_submission_acceptance { self.send_submit_and_wait_for_acceptance( WorkerExecutionOperation::Input, worker, @@ -2008,11 +1961,7 @@ where ) } else { self.send_method(WorkerExecutionOperation::Input, worker, method) - }; - if claimed_here && result.outcome != crate::execution::WorkerExecutionOutcome::Accepted { - busy.store(false, Ordering::SeqCst); } - result } fn upload_file( @@ -2023,7 +1972,7 @@ where content: &[u8], context: Option<&session_store::UploadedFileUploadContext>, ) -> Result { - let (worker, _, _, _) = self.get_execution(handle).map_err(|mut result| { + let (worker, _, _) = self.get_execution(handle).map_err(|mut result| { result.operation = WorkerExecutionOperation::UploadFile; result })?; @@ -2046,7 +1995,7 @@ where handle: &WorkerExecutionHandle, artifact_id: &str, ) -> WorkerExecutionResult { - let (worker, _, _, _) = match self.get_execution(handle) { + let (worker, _, _) = match self.get_execution(handle) { Ok(execution) => execution, Err(mut result) => { result.operation = WorkerExecutionOperation::DeleteUploadedFile; @@ -2067,7 +2016,7 @@ where handle: &WorkerExecutionHandle, method: Method, ) -> WorkerExecutionResult { - let (worker, busy, _worker_state, _workspace_client) = match self.get_execution(handle) { + let (worker, _worker_state, _workspace_client) = match self.get_execution(handle) { Ok(execution) => execution, Err(mut result) => { result.operation = WorkerExecutionOperation::ProtocolMethod; @@ -2075,44 +2024,7 @@ where } }; - if let Some(auto_run) = match &method { - Method::Notify { auto_run, .. } | Method::NotifyTracked { auto_run, .. } => { - Some(*auto_run) - } - _ => None, - } { - let status = worker.shared_state.catalog_status(); - let claimed_here = status == WorkerStatus::Idle - && auto_run - && busy - .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst) - .is_ok(); - let result = self.send_method(WorkerExecutionOperation::ProtocolMethod, worker, method); - if claimed_here && result.outcome != crate::execution::WorkerExecutionOutcome::Accepted - { - busy.store(false, Ordering::SeqCst); - } - return result; - } - - let starts_turn = method_starts_turn(&method); - if starts_turn - && (!method_can_start_turn_from_status(&method, worker.shared_state.catalog_status()) - || busy - .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst) - .is_err()) - { - return WorkerExecutionResult::busy( - WorkerExecutionOperation::ProtocolMethod, - "Worker is already running; runtime adapter v0 does not queue protocol methods", - ); - } - - let result = self.send_method(WorkerExecutionOperation::ProtocolMethod, worker, method); - if starts_turn && result.outcome != crate::execution::WorkerExecutionOutcome::Accepted { - busy.store(false, Ordering::SeqCst); - } - result + self.send_method(WorkerExecutionOperation::ProtocolMethod, worker, method) } fn stop_worker(&self, handle: &WorkerExecutionHandle) -> WorkerExecutionResult { @@ -2193,7 +2105,7 @@ where } fn cancel_worker(&self, handle: &WorkerExecutionHandle) -> WorkerExecutionResult { - let (worker, _busy, worker_state, _workspace_client) = match self.get_execution(handle) { + let (worker, worker_state, _workspace_client) = match self.get_execution(handle) { Ok(execution) => execution, Err(mut result) => { result.operation = WorkerExecutionOperation::Cancel; @@ -2297,6 +2209,56 @@ mod tests { WorkerCommandEnvelope::for_snapshot(state.last_command_id.saturating_add(1), &state) } + #[test] + fn protocol_bridge_applies_state_and_acknowledgement_monotonically() { + let running = protocol::WorkerStateSnapshot { + execution_generation: 4, + revision: 3, + state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run( + protocol::WorkerRunState::Running, + )), + last_command_id: 2, + }; + let current = Arc::new(RwLock::new(running.clone())); + let mut stale = Event::WorkerState { + snapshot: protocol::WorkerStateSnapshot { + revision: 2, + state: protocol::WorkerState::Idle, + ..running.clone() + }, + }; + assert!(!apply_protocol_worker_state(¤t, &mut stale).unwrap()); + assert_eq!(*current.read().unwrap(), running); + + let paused = protocol::WorkerStateSnapshot { + revision: 4, + state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run( + protocol::WorkerRunState::Paused, + )), + last_command_id: 3, + ..running.clone() + }; + let mut acknowledgement = Event::CommandAcknowledged { + acknowledgement: protocol::WorkerCommandAcknowledgement { + command_id: 3, + command: protocol::WorkerCommandKind::Pause, + disposition: protocol::WorkerCommandDisposition::Accepted, + state: paused.clone(), + }, + }; + assert!(apply_protocol_worker_state(¤t, &mut acknowledgement).unwrap()); + assert_eq!(*current.read().unwrap(), paused); + + let mut conflict = Event::WorkerState { + snapshot: protocol::WorkerStateSnapshot { + state: protocol::WorkerState::Idle, + ..paused.clone() + }, + }; + assert!(apply_protocol_worker_state(¤t, &mut conflict).is_err()); + assert_eq!(*current.read().unwrap(), paused); + } + #[test] fn workspace_prompt_projection_notification_advances_shared_cache() { let cache = WorkspacePromptProjectionCache::default(); @@ -2443,44 +2405,6 @@ mod tests { assert_eq!(after_restore_workspace_id.as_deref(), Some("workspace-a")); } - #[test] - fn compact_is_maintenance_not_a_turn_start() { - assert!(!method_starts_turn(&Method::Compact { - command: test_command(), - })); - assert!(method_starts_turn(&Method::Resume { - command: test_command(), - })); - } - - #[test] - fn resume_turn_claim_accepts_paused_and_idle_but_not_running_status() { - assert!(method_can_start_turn_from_status( - &Method::Resume { - command: test_command() - }, - WorkerStatus::Paused - )); - assert!(method_can_start_turn_from_status( - &Method::Resume { - command: test_command() - }, - WorkerStatus::Idle - )); - assert!(!method_can_start_turn_from_status( - &Method::Resume { - command: test_command() - }, - WorkerStatus::Running - )); - assert!(!method_can_start_turn_from_status( - &Method::Compact { - command: test_command() - }, - WorkerStatus::Paused - )); - } - #[derive(Clone)] enum MockResponse { Complete(Vec), @@ -2681,11 +2605,38 @@ mod tests { .collect() } + fn wait_for_adapter_command( + backend: &WorkerRuntimeExecutionBackend, + worker_ref: &WorkerRef, + expected_command_id: u64, + ) { + let deadline = std::time::Instant::now() + Duration::from_secs(5); + loop { + let observed = { + let workers = backend.workers.lock().unwrap(); + workers + .get(worker_ref) + .expect("live Worker execution") + .worker_state + .read() + .unwrap() + .last_command_id + }; + if observed >= expected_command_id { + return; + } + assert!( + std::time::Instant::now() < deadline, + "timed out waiting for adapter command {expected_command_id}; last observed={observed}", + ); + std::thread::sleep(Duration::from_millis(10)); + } + } + fn wait_for_adapter_state( backend: &WorkerRuntimeExecutionBackend, worker_ref: &WorkerRef, expected_status: WorkerStatus, - expected_busy: bool, ) { let deadline = std::time::Instant::now() + Duration::from_secs(5); loop { @@ -2693,21 +2644,16 @@ mod tests { let workers = backend.workers.lock().unwrap(); let execution = workers.get(worker_ref).expect("live Worker execution"); let projected = execution.worker_state.read().unwrap().catalog_status(); - ( - execution.handle.shared_state.catalog_status(), - projected, - execution.busy.load(Ordering::SeqCst), - ) + (execution.handle.shared_state.catalog_status(), projected) }; - if observed == (expected_status, expected_status, expected_busy) { + if observed == (expected_status, expected_status) { return; } assert!( std::time::Instant::now() < deadline, - "timed out waiting for adapter state {expected_status:?}, busy={expected_busy}; last observed controller={:?}, projected={:?}, busy={}", + "timed out waiting for adapter state {expected_status:?}; last observed controller={:?}, projected={:?}", observed.0, observed.1, - observed.2, ); std::thread::sleep(Duration::from_millis(10)); } @@ -3694,22 +3640,18 @@ mod tests { runtime .send_input(&detail.worker_ref, WorkerInput::user("pause and resume")) .expect("start initial turn"); - wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Running, true); + wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Running); - let running_resume = runtime + let running_resume = adapter_command(&backend, &detail.worker_ref); + runtime .send_protocol_method( &detail.worker_ref, Method::Resume { - command: adapter_command(&backend, &detail.worker_ref), + command: running_resume, }, ) - .expect_err("Resume while Running must be rejected"); - assert!( - running_resume - .to_string() - .contains("does not queue protocol methods"), - "unexpected Running Resume error: {running_resume}" - ); + .expect("running Resume is forwarded for controller admission"); + wait_for_adapter_command(&backend, &detail.worker_ref, running_resume.command_id); runtime .send_protocol_method( @@ -3719,7 +3661,7 @@ mod tests { }, ) .expect("pause initial turn"); - wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Paused, false); + wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Paused); runtime .send_protocol_method( @@ -3729,22 +3671,18 @@ mod tests { }, ) .expect("resume paused turn"); - wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Running, true); + wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Running); - let duplicate_resume = runtime + let duplicate_resume = adapter_command(&backend, &detail.worker_ref); + runtime .send_protocol_method( &detail.worker_ref, Method::Resume { - command: adapter_command(&backend, &detail.worker_ref), + command: duplicate_resume, }, ) - .expect_err("duplicate Resume must be rejected"); - assert!( - duplicate_resume - .to_string() - .contains("does not queue protocol methods"), - "unexpected duplicate Resume error: {duplicate_resume}" - ); + .expect("duplicate Resume is forwarded for controller admission"); + wait_for_adapter_command(&backend, &detail.worker_ref, duplicate_resume.command_id); runtime .send_protocol_method( @@ -3754,7 +3692,7 @@ mod tests { }, ) .expect("pause resumed turn"); - wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Paused, false); + wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Paused); runtime .send_protocol_method( &detail.worker_ref, @@ -3763,18 +3701,20 @@ mod tests { }, ) .expect("resume paused turn a second time"); - wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Idle, false); + wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Idle); assert_eq!(call_count.load(Ordering::SeqCst), 3); + let idle_resume = adapter_command(&backend, &detail.worker_ref); runtime .send_protocol_method( &detail.worker_ref, Method::Resume { - command: adapter_command(&backend, &detail.worker_ref), + command: idle_resume, }, ) .expect("Idle Resume preserves controller NotPaused semantics"); - wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Idle, false); + wait_for_adapter_command(&backend, &detail.worker_ref, idle_resume.command_id); + wait_for_adapter_state(&backend, &detail.worker_ref, WorkerStatus::Idle); let events = runtime .read_worker_observation_events(&detail.worker_ref, WorkerObservationCursor::zero()) .expect("read protocol events"); diff --git a/crates/worker/src/controller.rs b/crates/worker/src/controller.rs index f9cc2388..beeb4791 100644 --- a/crates/worker/src/controller.rs +++ b/crates/worker/src/controller.rs @@ -3708,7 +3708,7 @@ mod tests { WorkerCommandEnvelope { command_id: 1, expected_execution_generation: 9, - expected_worker_state_revision: 0, + expected_worker_state_revision: 1, }, &shared, ), @@ -3719,7 +3719,7 @@ mod tests { WorkerCommandEnvelope { command_id: 2, expected_execution_generation: 9, - expected_worker_state_revision: 0, + expected_worker_state_revision: 1, }, &shared, ) diff --git a/crates/worker/src/shared_state.rs b/crates/worker/src/shared_state.rs index 58c5700b..573f2bb5 100644 --- a/crates/worker/src/shared_state.rs +++ b/crates/worker/src/shared_state.rs @@ -1,6 +1,6 @@ use std::sync::{ OnceLock, RwLock, - atomic::{AtomicBool, AtomicU64, Ordering}, + atomic::{AtomicBool, Ordering}, }; use protocol::{ @@ -23,7 +23,6 @@ pub struct WorkerSharedState { pub manifest_toml: String, pub greeting: protocol::Greeting, state: RwLock, - last_command_id: AtomicU64, /// Worker-from-the-inside view of the filesystem. Set once in /// `WorkerController::start` after the local WorkdirSession provider is /// materialised, and read from the IPC server layer to answer @@ -56,7 +55,6 @@ impl WorkerSharedState { manifest_toml, greeting, state: RwLock::new(WorkerStateSnapshot::initial(execution_generation)), - last_command_id: AtomicU64::new(0), fs_view: OnceLock::new(), flow_transition_enabled: AtomicBool::new(false), } @@ -91,26 +89,27 @@ impl WorkerSharedState { snapshot.revision = snapshot.revision.saturating_add(1); snapshot.state = state; } - snapshot.last_command_id = self.last_command_id.load(Ordering::Acquire); snapshot.clone() } pub fn accept_command_id(&self, command_id: u64) -> bool { - self.last_command_id - .fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| { - (command_id > current).then_some(command_id) - }) - .is_ok() + let mut snapshot = self + .state + .write() + .expect("worker state lock poisoned; refusing command admission"); + if command_id <= snapshot.last_command_id { + return false; + } + snapshot.last_command_id = command_id; + snapshot.revision = snapshot.revision.saturating_add(1); + true } pub fn snapshot(&self) -> WorkerStateSnapshot { - let mut snapshot = self - .state + self.state .read() .expect("worker state lock poisoned; refusing an inferred fallback state") - .clone(); - snapshot.last_command_id = self.last_command_id.load(Ordering::Acquire); - snapshot + .clone() } /// Runtime catalog projection. This must not be used as live command @@ -190,6 +189,23 @@ mod tests { assert_eq!(state.catalog_status(), WorkerStatus::Paused); } + #[test] + fn accepted_command_id_advances_the_snapshot_revision_atomically() { + let state = test_state(); + assert!(state.accept_command_id(9)); + assert_eq!( + state.snapshot(), + WorkerStateSnapshot { + execution_generation: 7, + revision: 1, + last_command_id: 9, + state: WorkerState::Idle, + } + ); + assert!(!state.accept_command_id(9)); + assert_eq!(state.snapshot().revision, 1); + } + #[test] fn status_json_contains_full_snapshot_and_catalog_projection() { let state = test_state(); diff --git a/web/workspace/src/lib/workspace/console/model.test.ts b/web/workspace/src/lib/workspace/console/model.test.ts index c1f1b72a..de1e46b3 100644 --- a/web/workspace/src/lib/workspace/console/model.test.ts +++ b/web/workspace/src/lib/workspace/console/model.test.ts @@ -218,6 +218,66 @@ Deno.test("console routing projects live errors but not completion replies", () ); }); +Deno.test("Worker state events and acknowledgements apply monotonically", () => { + const projector = createConsoleProjector(); + const running: WorkerStateSnapshot = { + execution_generation: 4, + revision: 3, + last_command_id: 2, + state: { kind: "busy", state: { kind: "run", state: "running" } }, + }; + const paused: WorkerStateSnapshot = { + ...running, + revision: 4, + last_command_id: 3, + state: { kind: "busy", state: { kind: "run", state: "paused" } }, + }; + let projection = projector.append([ + { + eventId: "running", + event: { event: "worker_state", data: { snapshot: running } }, + }, + { + eventId: "stale", + event: { + event: "worker_state", + data: { snapshot: { ...running, revision: 2, state: { kind: "idle" } } }, + }, + }, + { + eventId: "pause-ack", + event: { + event: "command_acknowledged", + data: { + acknowledgement: { + command_id: 3, + command: "pause", + disposition: "accepted", + state: paused, + }, + }, + }, + }, + ]); + assertEquals(projection.workerState, paused); + assertEquals(projection.status, "paused"); + + projection = projector.append([{ + eventId: "conflict", + event: { + event: "worker_state", + data: { snapshot: { ...paused, state: { kind: "idle" } } }, + }, + }]); + assertEquals(projection.workerState, paused); + assert( + projection.lines.some((line) => + line.eventId === "conflict:worker-state-conflict" && line.error + ), + "conflicting equal-version snapshots must fail closed", + ); +}); + Deno.test("snapshot replaces a live error with one durable run_errored row", () => { const projector = createConsoleProjector(); let projection = projector.append([ diff --git a/web/workspace/src/lib/workspace/console/model.ts b/web/workspace/src/lib/workspace/console/model.ts index 857e1b30..95578cbd 100644 --- a/web/workspace/src/lib/workspace/console/model.ts +++ b/web/workspace/src/lib/workspace/console/model.ts @@ -10,6 +10,7 @@ import type { InternalWorkerRef, InternalWorkerSnapshot, Segment, + WorkerState, WorkerStateSnapshot, WorkerStatus, } from "$lib/generated/protocol"; @@ -796,6 +797,60 @@ function refreshCompactionActivity( return changed ? { ...projection, lines } : projection; } +function workerStateEqual(left: WorkerState, right: WorkerState): boolean { + if (left.kind !== right.kind) return false; + if (left.kind === "idle" || right.kind === "idle") return true; + return left.state.kind === right.state.kind && + left.state.state === right.state.state; +} + +function workerStateSnapshotEqual( + left: WorkerStateSnapshot, + right: WorkerStateSnapshot, +): boolean { + return left.execution_generation === right.execution_generation && + left.revision === right.revision && + left.last_command_id === right.last_command_id && + workerStateEqual(left.state, right.state); +} + +function applyWorkerStateSnapshot( + projection: ConsoleProjection, + incoming: WorkerStateSnapshot, + eventId: string, +): void { + const current = projection.workerState; + if (!current) { + projection.workerState = incoming; + projection.status = workerStatusFromState(incoming); + return; + } + const generationOrder = incoming.execution_generation - + current.execution_generation; + const revisionOrder = incoming.revision - current.revision; + if (generationOrder > 0 || (generationOrder === 0 && revisionOrder > 0)) { + projection.workerState = incoming; + projection.status = workerStatusFromState(incoming); + return; + } + if (generationOrder < 0 || (generationOrder === 0 && revisionOrder < 0)) { + return; + } + if (!workerStateSnapshotEqual(current, incoming)) { + projection.lines.push( + line( + `${eventId}:worker-state-conflict`, + "error", + "error ยท internal", + `worker state stream rejected: conflicting snapshots at generation ${incoming.execution_generation} revision ${incoming.revision}`, + undefined, + false, + true, + ), + ); + } +} + export function applyProtocolEvent( projection: ConsoleProjection, envelope: ConsoleEventInput, @@ -917,8 +972,6 @@ export function applyProtocolEvent( ); break; case "snapshot": { - next.workerState = event.data.state; - next.status = workerStatusFromState(event.data.state); next.cwd = event.data.greeting.cwd; const snapshot = snapshotProjectionFromSession( envelope.eventId, @@ -968,6 +1021,7 @@ export function applyProtocolEvent( }; } } + applyWorkerStateSnapshot(next, event.data.state, envelope.eventId); break; } case "internal_worker": { @@ -1016,12 +1070,14 @@ export function applyProtocolEvent( break; } case "worker_state": - next.workerState = event.data.snapshot; - next.status = workerStatusFromState(event.data.snapshot); + applyWorkerStateSnapshot(next, event.data.snapshot, envelope.eventId); break; case "command_acknowledged": - next.workerState = event.data.acknowledgement.state; - next.status = workerStatusFromState(event.data.acknowledgement.state); + applyWorkerStateSnapshot( + next, + event.data.acknowledgement.state, + envelope.eventId, + ); break; case "command": applyCommandEvent(next, envelope.eventId, event.data.event);