diff --git a/crates/worker-runtime/src/http_server.rs b/crates/worker-runtime/src/http_server.rs index f44f52a8..8812390c 100644 --- a/crates/worker-runtime/src/http_server.rs +++ b/crates/worker-runtime/src/http_server.rs @@ -21,6 +21,8 @@ use crate::interaction::{WorkerInput, WorkerInteractionAck}; use crate::management::{RuntimeLimits, RuntimeSummary, WorkerDeleteResult}; #[cfg(feature = "ws-server")] use crate::observation::WorkerObservationCursor; +#[cfg(feature = "ws-server")] +use crate::runtime::RuntimeSubscriptionRecvError; use crate::{Runtime, RuntimeWorkspaceScope}; use axum::body::{Body, Bytes}; use axum::extract::rejection::{JsonRejection, QueryRejection}; @@ -33,10 +35,18 @@ use axum::response::{IntoResponse, Response}; use axum::routing::{get, post}; use axum::{Json, Router}; #[cfg(feature = "ws-server")] -use futures::StreamExt; +use futures::{SinkExt, StreamExt}; #[cfg(feature = "ws-server")] use protocol::stream::{decode_method, encode_event}; +#[cfg(feature = "ws-server")] +use protocol::subscription::{ + SubscriptionEvent, SubscriptionFrame, SubscriptionFramePayload, SubscriptionId, + SubscriptionRejectionCode, SubscriptionRequest, SubscriptionResponse, + SubscriptionTerminationCode, +}; use serde::{Deserialize, Serialize}; +#[cfg(feature = "ws-server")] +use std::collections::HashMap; use std::fmt; use std::net::SocketAddr; #[cfg(feature = "fs-store")] @@ -205,10 +215,12 @@ fn runtime_http_router_with_optional_auth( .route("/v1/workers/{worker_id}/cancel", post(cancel_worker)); #[cfg(feature = "ws-server")] - let router = router.route( - "/v1/workers/{worker_id}/protocol/ws", - get(worker_protocol_ws), - ); + let router = router + .route("/v1/protocol/ws", get(runtime_protocol_ws)) + .route( + "/v1/workers/{worker_id}/protocol/ws", + get(worker_protocol_ws), + ); router .with_state(state.clone()) @@ -555,6 +567,272 @@ async fn restore_worker( Ok(Json(RuntimeHttpWorkerResponse { worker })) } +#[cfg(feature = "ws-server")] +const RUNTIME_PROTOCOL_OUTBOUND_CAPACITY: usize = 256; + +#[cfg(feature = "ws-server")] +async fn runtime_protocol_ws( + State(state): State, + auth: Option>, + ws: axum::extract::ws::WebSocketUpgrade, +) -> Result { + let scope = + auth_workspace_scope(&state, auth.as_ref()).map_err(|error| error.into_response())?; + Ok(ws + .on_upgrade(move |socket| runtime_protocol_ws_session(state.runtime, scope, socket)) + .into_response()) +} + +#[cfg(feature = "ws-server")] +async fn runtime_protocol_ws_session( + runtime: Runtime, + scope: Option, + socket: axum::extract::ws::WebSocket, +) { + let (mut socket_sender, mut socket_receiver) = socket.split(); + let (outbound, mut outbound_receiver) = tokio::sync::mpsc::channel::( + RUNTIME_PROTOCOL_OUTBOUND_CAPACITY, + ); + let writer = tokio::spawn(async move { + while let Some(message) = outbound_receiver.recv().await { + if socket_sender.send(message).await.is_err() { + break; + } + } + }); + + let mut next_subscription_id = 1_u64; + let mut subscriptions = HashMap::>::new(); + while let Some(message) = socket_receiver.next().await { + let Ok(message) = message else { + break; + }; + match message { + axum::extract::ws::Message::Text(text) => { + let Ok(frame) = serde_json::from_str::(text.as_str()) else { + break; + }; + let request_id = subscription_frame_request_id(&frame); + if let Err(error) = frame.validate() { + let Some(request_id) = request_id else { + break; + }; + let code = if frame.protocol_version + != protocol::subscription::SUBSCRIPTION_PROTOCOL_VERSION + { + SubscriptionRejectionCode::UnsupportedProtocolVersion + } else { + SubscriptionRejectionCode::InvalidRequest + }; + if send_runtime_subscription_frame( + &outbound, + SubscriptionFrame::new(SubscriptionFramePayload::Response( + SubscriptionResponse::SubscriptionRejected { + request_id, + subscription_id: None, + code, + message: error.to_string(), + }, + )), + ) + .await + .is_err() + { + break; + } + continue; + } + + subscriptions.retain(|_, task| !task.is_finished()); + let SubscriptionFramePayload::Request(request) = frame.payload else { + break; + }; + match request { + SubscriptionRequest::SubscribeEvents { + request_id, + selector, + } => { + let subscription = match scope.as_ref() { + Some(scope) => { + runtime.subscribe_event_selector_scoped(scope, selector.clone()) + } + None => runtime.subscribe_event_selector(selector.clone()), + }; + let mut subscription = match subscription { + Ok(subscription) => subscription, + Err(error) => { + if send_runtime_subscription_frame( + &outbound, + SubscriptionFrame::new(SubscriptionFramePayload::Response( + SubscriptionResponse::SubscriptionRejected { + request_id, + subscription_id: None, + code: runtime_subscription_rejection_code(&error), + message: error.to_string(), + }, + )), + ) + .await + .is_err() + { + break; + } + continue; + } + }; + let subscription_id = SubscriptionId::new(format!( + "runtime-subscription-{next_subscription_id}" + )) + .expect("generated Runtime subscription id is valid"); + next_subscription_id = next_subscription_id.saturating_add(1); + let response = SubscriptionFrame::new(SubscriptionFramePayload::Response( + SubscriptionResponse::Subscribed { + request_id, + subscription_id: subscription_id.clone(), + selector: selector.clone(), + snapshot_revision: subscription.snapshot_revision(), + snapshot: subscription.snapshot().clone(), + }, + )); + if send_runtime_subscription_frame(&outbound, response) + .await + .is_err() + { + break; + } + + let event_outbound = outbound.clone(); + let event_subscription_id = subscription_id.clone(); + let task = tokio::spawn(async move { + loop { + let frame = match subscription.recv().await { + Ok(update) => SubscriptionFrame::new( + SubscriptionFramePayload::Event( + SubscriptionEvent::Event { + subscription_id: event_subscription_id.clone(), + subject_revision: update.subject_revision, + payload: update.payload, + }, + ), + ), + Err(RuntimeSubscriptionRecvError::Lagged) => { + SubscriptionFrame::new(SubscriptionFramePayload::Event( + SubscriptionEvent::SubscriptionClosed { + subscription_id: event_subscription_id.clone(), + code: SubscriptionTerminationCode::Lagged, + message: "Runtime subscription lagged; resubscribe for a fresh snapshot" + .to_string(), + }, + )) + } + Err(RuntimeSubscriptionRecvError::Closed) => { + SubscriptionFrame::new(SubscriptionFramePayload::Event( + SubscriptionEvent::SubscriptionClosed { + subscription_id: event_subscription_id.clone(), + code: SubscriptionTerminationCode::ServerShutdown, + message: "Runtime subscription closed".to_string(), + }, + )) + } + }; + let terminal = matches!( + &frame.payload, + SubscriptionFramePayload::Event( + SubscriptionEvent::SubscriptionClosed { .. } + ) + ); + if send_runtime_subscription_frame(&event_outbound, frame) + .await + .is_err() + || terminal + { + break; + } + } + }); + subscriptions.insert(subscription_id, task); + } + SubscriptionRequest::UnsubscribeEvents { + request_id, + subscription_id, + } => { + if let Some(task) = subscriptions.remove(&subscription_id) { + task.abort(); + } + if send_runtime_subscription_frame( + &outbound, + SubscriptionFrame::new(SubscriptionFramePayload::Response( + SubscriptionResponse::Unsubscribed { + request_id, + subscription_id, + }, + )), + ) + .await + .is_err() + { + break; + } + } + } + } + axum::extract::ws::Message::Ping(payload) => { + if outbound + .send(axum::extract::ws::Message::Pong(payload)) + .await + .is_err() + { + break; + } + } + axum::extract::ws::Message::Pong(_) => {} + axum::extract::ws::Message::Close(_) => break, + axum::extract::ws::Message::Binary(_) => break, + } + } + + for (_, task) in subscriptions { + task.abort(); + } + drop(outbound); + let _ = writer.await; +} + +#[cfg(feature = "ws-server")] +fn subscription_frame_request_id( + frame: &SubscriptionFrame, +) -> Option { + let SubscriptionFramePayload::Request(request) = &frame.payload else { + return None; + }; + Some(match request { + SubscriptionRequest::SubscribeEvents { request_id, .. } + | SubscriptionRequest::UnsubscribeEvents { request_id, .. } => request_id.clone(), + }) +} + +#[cfg(feature = "ws-server")] +fn runtime_subscription_rejection_code(error: &RuntimeError) -> SubscriptionRejectionCode { + match error { + RuntimeError::WorkerNotFound { .. } => SubscriptionRejectionCode::ResourceNotFound, + RuntimeError::InvalidRequest(_) => SubscriptionRejectionCode::UnsupportedSelector, + _ => SubscriptionRejectionCode::Internal, + } +} + +#[cfg(feature = "ws-server")] +async fn send_runtime_subscription_frame( + outbound: &tokio::sync::mpsc::Sender, + frame: SubscriptionFrame, +) -> Result<(), ()> { + frame.validate().map_err(|_| ())?; + let text = serde_json::to_string(&frame).map_err(|_| ())?; + outbound + .send(axum::extract::ws::Message::Text(text.into())) + .await + .map_err(|_| ()) +} + #[cfg(feature = "ws-server")] async fn worker_protocol_ws( State(state): State, @@ -1001,6 +1279,9 @@ fn required_runtime_permission(method: &Method, path: &str) -> Option<&'static s if path.ends_with("/stop") || path.ends_with("/cancel") { return Some("workers:stop"); } + if path == "/v1/protocol/ws" { + return Some("workers:list"); + } if path.ends_with("/protocol") || path.ends_with("/protocol/ws") { return Some("workers:protocol"); } @@ -1965,6 +2246,138 @@ mod ws_tests { serde_json::from_str(&text).unwrap() } + async fn next_subscription_frame( + stream: &mut tokio_tungstenite::WebSocketStream< + tokio_tungstenite::MaybeTlsStream, + >, + ) -> SubscriptionFrame { + let message = stream.next().await.unwrap().unwrap(); + let Message::Text(text) = message else { + panic!("expected text frame"); + }; + serde_json::from_str(&text).unwrap() + } + + async fn send_subscription_request( + stream: &mut tokio_tungstenite::WebSocketStream< + tokio_tungstenite::MaybeTlsStream, + >, + request: SubscriptionRequest, + ) { + let frame = SubscriptionFrame::new(SubscriptionFramePayload::Request(request)); + stream + .send(Message::Text(serde_json::to_string(&frame).unwrap().into())) + .await + .unwrap(); + } + + fn runtime_protocol_url(worker_protocol_url: &str) -> String { + let base = worker_protocol_url + .split_once("/v1/workers/") + .map(|(base, _)| base) + .unwrap(); + format!("{base}/v1/protocol/ws") + } + + #[tokio::test] + async fn runtime_protocol_ws_subscribes_and_filters_worker_lifecycle() { + let (runtime, worker_ref, worker_url) = spawn_runtime_server().await; + let other = runtime + .create_worker_scoped( + &RuntimeWorkspaceScope::new("local", "local-token"), + ws_create_request(), + ) + .unwrap(); + let url = runtime_protocol_url(&worker_url); + let (mut stream, _) = connect_async(authed_ws_request(&url)).await.unwrap(); + let request_id = protocol::subscription::SubscriptionRequestId::new("subscribe-1").unwrap(); + send_subscription_request( + &mut stream, + SubscriptionRequest::SubscribeEvents { + request_id: request_id.clone(), + selector: protocol::subscription::EventSubscriptionSelector::WorkerLifecycle { + worker_ids: protocol::subscription::SubscriptionWorkerIds::new([ + protocol::subscription::SubscriptionWorkerId::new( + worker_ref.worker_id.to_string(), + ) + .unwrap(), + ]) + .unwrap(), + }, + }, + ) + .await; + let subscribed = next_subscription_frame(&mut stream).await; + let SubscriptionFramePayload::Response(SubscriptionResponse::Subscribed { + request_id: response_request_id, + subscription_id, + snapshot, + .. + }) = subscribed.payload + else { + panic!("expected subscribed response"); + }; + assert_eq!(response_request_id, request_id); + let protocol::subscription::SubscriptionSnapshot::Workers { workers } = snapshot else { + panic!("expected Worker snapshot"); + }; + assert_eq!(workers.len(), 1); + assert_eq!( + workers[0].worker_id.as_str(), + worker_ref.worker_id.to_string() + ); + + runtime + .observe_worker_event( + &other.worker_ref, + protocol::Event::Status { + status: protocol::WorkerStatus::Running, + }, + ) + .unwrap(); + runtime + .observe_worker_event( + &worker_ref, + protocol::Event::Status { + status: protocol::WorkerStatus::Running, + }, + ) + .unwrap(); + let event = next_subscription_frame(&mut stream).await; + assert!(matches!( + event.payload, + SubscriptionFramePayload::Event(SubscriptionEvent::Event { + subscription_id: delivered_subscription_id, + payload: protocol::subscription::SubscriptionEventPayload::WorkerUpserted { + worker + }, + .. + }) if delivered_subscription_id == subscription_id + && worker.worker_id.as_str() == worker_ref.worker_id.to_string() + && worker.state == protocol::subscription::SubscriptionWorkerState::Running + )); + + let unsubscribe_request_id = + protocol::subscription::SubscriptionRequestId::new("unsubscribe-1").unwrap(); + send_subscription_request( + &mut stream, + SubscriptionRequest::UnsubscribeEvents { + request_id: unsubscribe_request_id.clone(), + subscription_id: subscription_id.clone(), + }, + ) + .await; + let unsubscribed = next_subscription_frame(&mut stream).await; + assert!(matches!( + unsubscribed.payload, + SubscriptionFramePayload::Response(SubscriptionResponse::Unsubscribed { + request_id, + subscription_id: response_subscription_id, + }) if request_id == unsubscribe_request_id + && response_subscription_id == subscription_id + )); + } + #[tokio::test] async fn protocol_ws_connect_sends_snapshot_and_live_worker_events() { let (runtime, worker_ref, url) = spawn_runtime_server().await; diff --git a/crates/worker-runtime/src/runtime.rs b/crates/worker-runtime/src/runtime.rs index 5a885106..9c90fd29 100644 --- a/crates/worker-runtime/src/runtime.rs +++ b/crates/worker-runtime/src/runtime.rs @@ -1,6 +1,6 @@ use crate::catalog::{ - ConfigBundleRef, CreateWorkerRequest, WorkerDetail, WorkerLifecycleAck, WorkerStatus, - WorkerSummary, WorkingDirectoryRequest, + ConfigBundleRef, CreateWorkerRequest, ProfileSelector, WorkerDetail, WorkerLifecycleAck, + WorkerStatus, WorkerSummary, WorkingDirectoryRequest, WorkingDirectoryStatus as CatalogWorkingDirectoryStatus, WorkspaceApiRef, }; use crate::config_bundle::{ @@ -31,13 +31,20 @@ use crate::observation::{ }; #[cfg(feature = "ws-server")] use crate::observation::{WorkerObservationCursor, WorkerObservationEvent}; +use protocol::subscription::{ + EventSubscriptionSelector, SubscriptionEventPayload, SubscriptionSnapshot, + SubscriptionValidationError, SubscriptionWorkdirId, SubscriptionWorker, SubscriptionWorkerId, + SubscriptionWorkerState, +}; use protocol::{Event, Method}; use std::collections::BTreeMap; #[cfg(feature = "ws-server")] use std::collections::VecDeque; -use std::sync::{Arc, Mutex, MutexGuard}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex, MutexGuard, Weak}; #[cfg(feature = "ws-server")] use tokio::sync::broadcast; +use tokio::sync::mpsc; /// Workspace-scoped Runtime authorization context supplied by a trusted backend. #[derive(Clone, Debug, Eq, PartialEq)] @@ -55,6 +62,79 @@ impl RuntimeWorkspaceScope { } } +const RUNTIME_EVENT_SUBSCRIPTION_QUEUE_CAPACITY: usize = 256; + +#[derive(Clone, Debug)] +pub struct RuntimeSubscriptionUpdate { + pub subject_revision: u64, + pub payload: SubscriptionEventPayload, +} + +#[derive(Debug, thiserror::Error)] +pub enum RuntimeSubscriptionRecvError { + #[error("Runtime event subscription lagged and requires a fresh snapshot")] + Lagged, + #[error("Runtime event subscription closed")] + Closed, +} + +/// A gap-free snapshot/live subscription owned by one Runtime connection. +/// Dropping the subscription removes its bounded producer queue. +pub struct RuntimeEventSelectorSubscription { + subscription_id: u64, + selector: EventSubscriptionSelector, + snapshot_revision: u64, + snapshot: SubscriptionSnapshot, + receiver: mpsc::Receiver, + lagged: Arc, + runtime: Weak>, +} + +impl RuntimeEventSelectorSubscription { + pub fn selector(&self) -> &EventSubscriptionSelector { + &self.selector + } + + pub fn snapshot_revision(&self) -> u64 { + self.snapshot_revision + } + + pub fn snapshot(&self) -> &SubscriptionSnapshot { + &self.snapshot + } + + pub async fn recv( + &mut self, + ) -> Result { + if self.lagged.load(Ordering::Acquire) { + self.receiver.close(); + return Err(RuntimeSubscriptionRecvError::Lagged); + } + match self.receiver.recv().await { + Some(_) if self.lagged.load(Ordering::Acquire) => { + self.receiver.close(); + Err(RuntimeSubscriptionRecvError::Lagged) + } + Some(update) => Ok(update), + None if self.lagged.load(Ordering::Acquire) => { + Err(RuntimeSubscriptionRecvError::Lagged) + } + None => Err(RuntimeSubscriptionRecvError::Closed), + } + } +} + +impl Drop for RuntimeEventSelectorSubscription { + fn drop(&mut self) { + let Some(runtime) = self.runtime.upgrade() else { + return; + }; + if let Ok(mut state) = runtime.lock() { + state.event_subscriptions.remove(&self.subscription_id); + } + } +} + /// Concrete embedded Runtime domain entity. /// /// The default implementation is memory-backed and tools/provider-less by @@ -495,6 +575,71 @@ impl Runtime { Ok(state.workers.values().map(WorkerRecord::summary).collect()) } + pub fn subscribe_event_selector( + &self, + selector: EventSubscriptionSelector, + ) -> Result { + self.subscribe_event_selector_for_workspace(None, selector) + } + + pub fn subscribe_event_selector_scoped( + &self, + scope: &RuntimeWorkspaceScope, + selector: EventSubscriptionSelector, + ) -> Result { + self.subscribe_event_selector_for_workspace(Some(scope), selector) + } + + fn subscribe_event_selector_for_workspace( + &self, + scope: Option<&RuntimeWorkspaceScope>, + selector: EventSubscriptionSelector, + ) -> Result { + selector.validate().map_err(subscription_validation_error)?; + if matches!( + selector, + EventSubscriptionSelector::WorkerProtocol { .. } + | EventSubscriptionSelector::WorkspaceWorkers + | EventSubscriptionSelector::WorkspaceWorkdirs + ) { + return Err(RuntimeError::InvalidRequest( + "Runtime event subscriptions support only runtime_workers and worker_lifecycle selectors" + .to_string(), + )); + } + + let mut state = self.lock()?; + if let Some(scope) = scope { + state.ensure_workspace_owner_for_existing_workers(scope)?; + state.persist_runtime_snapshot()?; + } + let snapshot = state.subscription_snapshot(scope, &selector)?; + let snapshot_revision = state.subscription_revision; + let subscription_id = state.next_event_subscription_id; + state.next_event_subscription_id = state.next_event_subscription_id.saturating_add(1); + let (sender, receiver) = mpsc::channel(RUNTIME_EVENT_SUBSCRIPTION_QUEUE_CAPACITY); + let lagged = Arc::new(AtomicBool::new(false)); + state.event_subscriptions.insert( + subscription_id, + RuntimeEventSubscriptionSink { + selector: selector.clone(), + workspace_id: scope.map(|scope| scope.workspace_id.clone()), + sender, + lagged: lagged.clone(), + }, + ); + + Ok(RuntimeEventSelectorSubscription { + subscription_id, + selector, + snapshot_revision, + snapshot, + receiver, + lagged, + runtime: Arc::downgrade(&self.inner), + }) + } + /// List stopped Workers known to this Runtime. pub fn list_stopped_workers(&self) -> Result, RuntimeError> { let state = self.lock()?; @@ -846,6 +991,7 @@ impl Runtime { let payload = input_protocol_event(&input); state.push_worker_observation_event(worker_ref.clone(), payload); } + state.publish_worker_upsert(worker_ref.worker_id)?; state.persist_runtime_snapshot()?; state.persist_worker(&worker_ref.worker_id)?; state.persist_event_by_id(event_id)?; @@ -979,6 +1125,7 @@ impl Runtime { worker.working_directory = working_directory; worker.detail() }; + state.publish_worker_upsert(worker_ref.worker_id)?; state.persist_runtime_snapshot()?; state.persist_worker(&worker_ref.worker_id)?; state.persist_event_by_id(detail.last_event_id)?; @@ -995,6 +1142,7 @@ impl Runtime { state.events.retain(|event| { event.id != record.last_event_id || event.worker_ref.as_ref() != Some(worker_ref) }); + state.publish_worker_removed(worker_ref.worker_id, workspace_id.as_deref())?; } Ok(()) } @@ -1005,9 +1153,9 @@ impl Runtime { result: WorkerExecutionResult, ) -> Result<(), RuntimeError> { let mut state = self.lock()?; - let worker = state.worker_mut(worker_ref)?; if result.is_accepted() { - worker.status = worker_status_from_run_state(result.run_state); + state.worker_mut(worker_ref)?.status = worker_status_from_run_state(result.run_state); + state.publish_worker_upsert(worker_ref.worker_id)?; } Ok(()) } @@ -1138,7 +1286,8 @@ impl Runtime { worker_id: worker_ref.worker_id, } })?; - if let Some(workspace_id) = removed.workspace_id.as_deref() { + let removed_workspace_id = removed.workspace_id.clone(); + if let Some(workspace_id) = removed_workspace_id.as_deref() { state.forget_workspace_owner_if_unused(workspace_id); } #[cfg(feature = "ws-server")] @@ -1150,6 +1299,7 @@ impl Runtime { RuntimeEventKind::WorkerDeleted, "worker deleted", ); + state.publish_worker_removed(worker_ref.worker_id, removed_workspace_id.as_deref())?; state.persist_runtime_snapshot()?; state.delete_worker_snapshot(&worker_ref.worker_id)?; state.persist_event_by_id(event_id)?; @@ -1304,7 +1454,10 @@ impl Runtime { ) -> Result { let mut state = self.lock()?; state.ensure_worker_ref(worker_ref)?; - state.project_protocol_event_to_status(worker_ref, &payload); + let status_changed = state.project_protocol_event_to_status(worker_ref, &payload); + if status_changed { + state.publish_worker_upsert(worker_ref.worker_id)?; + } let event = state.push_worker_observation_event(worker_ref.clone(), payload); Ok(event) } @@ -1363,6 +1516,7 @@ impl Runtime { worker.execution_handle = None; worker.last_event_id = event_id; let status = worker.status; + state.publish_worker_upsert(worker_ref.worker_id)?; state.persist_runtime_snapshot()?; state.persist_worker(&worker_ref.worker_id)?; state.persist_event_by_id(event_id)?; @@ -1491,6 +1645,7 @@ impl Runtime { worker.working_directory = working_directory; worker.last_event_id = event_id; } + state.publish_worker_upsert(worker_ref.worker_id)?; state.persist_runtime_snapshot()?; state.persist_worker(&worker_ref.worker_id)?; state.persist_event_by_id(event_id)?; @@ -1510,6 +1665,14 @@ enum RuntimePersistence { Fs(FsRuntimeStore), } +#[derive(Debug)] +struct RuntimeEventSubscriptionSink { + selector: EventSubscriptionSelector, + workspace_id: Option, + sender: mpsc::Sender, + lagged: Arc, +} + #[derive(Debug)] struct RuntimeState { display_name: Option, @@ -1528,6 +1691,10 @@ struct RuntimeState { config_bundles: BTreeMap, events: Vec, diagnostics: Vec, + subscription_revision: u64, + worker_subject_revisions: BTreeMap, + next_event_subscription_id: u64, + event_subscriptions: BTreeMap, #[cfg(feature = "ws-server")] next_observation_sequence: u64, #[cfg(feature = "ws-server")] @@ -1554,6 +1721,10 @@ impl RuntimeState { config_bundles: BTreeMap::new(), events: Vec::new(), diagnostics: Vec::new(), + subscription_revision: 0, + worker_subject_revisions: BTreeMap::new(), + next_event_subscription_id: 1, + event_subscriptions: BTreeMap::new(), #[cfg(feature = "ws-server")] next_observation_sequence: 1, #[cfg(feature = "ws-server")] @@ -1585,6 +1756,10 @@ impl RuntimeState { config_bundles: BTreeMap::new(), events: Vec::new(), diagnostics: Vec::new(), + subscription_revision: 0, + worker_subject_revisions: BTreeMap::new(), + next_event_subscription_id: 1, + event_subscriptions: BTreeMap::new(), #[cfg(feature = "ws-server")] next_observation_sequence: 1, #[cfg(feature = "ws-server")] @@ -1633,6 +1808,10 @@ impl RuntimeState { workspace_owners: persisted.workspace_owners, events: persisted.events, diagnostics, + subscription_revision: 0, + worker_subject_revisions: BTreeMap::new(), + next_event_subscription_id: 1, + event_subscriptions: BTreeMap::new(), #[cfg(feature = "ws-server")] next_observation_sequence: 1, #[cfg(feature = "ws-server")] @@ -1871,6 +2050,179 @@ impl RuntimeState { }) } + fn subscription_snapshot( + &self, + scope: Option<&RuntimeWorkspaceScope>, + selector: &EventSubscriptionSelector, + ) -> Result { + let workers = match selector { + EventSubscriptionSelector::RuntimeWorkers => self + .workers + .values() + .filter(|worker| { + scope.is_none_or(|scope| worker.belongs_to_workspace(&scope.workspace_id)) + }) + .map(|worker| self.subscription_worker(worker)) + .collect::, _>>()?, + EventSubscriptionSelector::WorkerLifecycle { worker_ids } => { + let mut selected = Vec::with_capacity(worker_ids.as_slice().len()); + for worker_id in worker_ids.as_slice() { + let runtime_worker_id = WorkerId::parse(worker_id.as_str()).ok_or_else(|| { + RuntimeError::InvalidRequest(format!( + "worker_lifecycle selector contains invalid Runtime Worker id {worker_id}" + )) + })?; + let worker = self.workers.get(&runtime_worker_id).ok_or( + RuntimeError::WorkerNotFound { + worker_id: runtime_worker_id, + }, + )?; + if scope.is_some_and(|scope| !worker.belongs_to_workspace(&scope.workspace_id)) + { + return Err(RuntimeError::WorkerNotFound { + worker_id: runtime_worker_id, + }); + } + selected.push(self.subscription_worker(worker)?); + } + selected + } + EventSubscriptionSelector::WorkerProtocol { .. } + | EventSubscriptionSelector::WorkspaceWorkers + | EventSubscriptionSelector::WorkspaceWorkdirs => { + return Err(RuntimeError::InvalidRequest( + "selector is not produced by the Runtime lifecycle subscription".to_string(), + )); + } + }; + Ok(SubscriptionSnapshot::Workers { workers }) + } + + fn subscription_worker( + &self, + worker: &WorkerRecord, + ) -> Result { + let worker_id = SubscriptionWorkerId::new(worker.worker_id.to_string()) + .map_err(subscription_validation_error)?; + let working_directory_id = worker + .working_directory + .as_ref() + .map(|working_directory| { + SubscriptionWorkdirId::new(working_directory.summary.working_directory_id.clone()) + .map_err(subscription_validation_error) + }) + .transpose()?; + let profile = match &worker.request.profile { + ProfileSelector::Builtin(name) | ProfileSelector::Named(name) => Some(name.clone()), + }; + Ok(SubscriptionWorker { + worker_id, + subject_revision: self + .worker_subject_revisions + .get(&worker.worker_id) + .copied() + .unwrap_or(0), + state: subscription_worker_state(worker.status), + workspace_id: worker.workspace_id.clone(), + display_name: worker.request.display_name.clone(), + profile, + working_directory_id, + }) + } + + fn publish_worker_upsert(&mut self, worker_id: WorkerId) -> Result<(), RuntimeError> { + self.subscription_revision = self.subscription_revision.saturating_add(1); + let subject_revision = { + let revision = self.worker_subject_revisions.entry(worker_id).or_insert(0); + *revision = revision.saturating_add(1); + *revision + }; + let worker = self + .workers + .get(&worker_id) + .ok_or(RuntimeError::WorkerNotFound { worker_id })?; + let workspace_id = worker.workspace_id.clone(); + let mut projected = self.subscription_worker(worker)?; + projected.subject_revision = subject_revision; + let projected_worker_id = projected.worker_id.clone(); + self.deliver_worker_subscription_update( + &projected_worker_id, + workspace_id.as_deref(), + RuntimeSubscriptionUpdate { + subject_revision, + payload: SubscriptionEventPayload::WorkerUpserted { worker: projected }, + }, + ); + Ok(()) + } + + fn publish_worker_removed( + &mut self, + worker_id: WorkerId, + workspace_id: Option<&str>, + ) -> Result<(), RuntimeError> { + self.subscription_revision = self.subscription_revision.saturating_add(1); + let subject_revision = { + let revision = self.worker_subject_revisions.entry(worker_id).or_insert(0); + *revision = revision.saturating_add(1); + *revision + }; + let worker_id = SubscriptionWorkerId::new(worker_id.to_string()) + .map_err(subscription_validation_error)?; + self.deliver_worker_subscription_update( + &worker_id, + workspace_id, + RuntimeSubscriptionUpdate { + subject_revision, + payload: SubscriptionEventPayload::WorkerRemoved { + worker_id: worker_id.clone(), + }, + }, + ); + Ok(()) + } + + fn deliver_worker_subscription_update( + &mut self, + worker_id: &SubscriptionWorkerId, + workspace_id: Option<&str>, + update: RuntimeSubscriptionUpdate, + ) { + let mut closed = Vec::new(); + for (subscription_id, sink) in &self.event_subscriptions { + if sink + .workspace_id + .as_deref() + .is_some_and(|expected| workspace_id != Some(expected)) + { + continue; + } + let selected = match &sink.selector { + EventSubscriptionSelector::RuntimeWorkers => true, + EventSubscriptionSelector::WorkerLifecycle { worker_ids } => { + worker_ids.contains(worker_id) + } + EventSubscriptionSelector::WorkerProtocol { .. } + | EventSubscriptionSelector::WorkspaceWorkers + | EventSubscriptionSelector::WorkspaceWorkdirs => false, + }; + if !selected { + continue; + } + match sink.sender.try_send(update.clone()) { + Ok(()) => {} + Err(mpsc::error::TrySendError::Full(_)) => { + sink.lagged.store(true, Ordering::Release); + closed.push(*subscription_id); + } + Err(mpsc::error::TrySendError::Closed(_)) => closed.push(*subscription_id), + } + } + for subscription_id in closed { + self.event_subscriptions.remove(&subscription_id); + } + } + fn primary_worker_id_for_workdir(&self, working_directory_id: &str) -> Option { self.workers.values().find_map(|worker| { if worker @@ -1911,6 +2263,7 @@ impl RuntimeState { let worker = self.worker_mut(worker_ref)?; worker.execution_handle = None; worker.status = WorkerStatus::Stopped; + self.publish_worker_upsert(worker_ref.worker_id)?; self.persist_runtime_snapshot()?; Ok(()) } @@ -1989,9 +2342,9 @@ impl RuntimeState { &mut self, worker_ref: &WorkerRef, event: &protocol::Event, - ) { + ) -> bool { let Some(worker) = self.workers.get_mut(&worker_ref.worker_id) else { - return; + return false; }; let next_status = match event { protocol::Event::Status { @@ -2018,7 +2371,11 @@ impl RuntimeState { _ => None, }; if let Some(next_status) = next_status { + let changed = worker.status != next_status; worker.status = next_status; + changed + } else { + false } } } @@ -2213,6 +2570,20 @@ fn input_protocol_event(input: &WorkerInput) -> protocol::Event { } } +fn subscription_validation_error(error: SubscriptionValidationError) -> RuntimeError { + RuntimeError::InvalidRequest(format!("invalid event subscription: {error}")) +} + +fn subscription_worker_state(status: WorkerStatus) -> SubscriptionWorkerState { + match status { + WorkerStatus::Idle => SubscriptionWorkerState::Idle, + WorkerStatus::Running => SubscriptionWorkerState::Running, + WorkerStatus::Paused => SubscriptionWorkerState::Paused, + WorkerStatus::Stopped => SubscriptionWorkerState::Stopped, + WorkerStatus::Cancelled => SubscriptionWorkerState::Cancelled, + } +} + #[cfg(test)] mod tests { use super::*; @@ -2476,6 +2847,159 @@ mod tests { request } + fn receive_subscription_update( + subscription: &mut RuntimeEventSelectorSubscription, + ) -> Result { + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() + .block_on(subscription.recv()) + } + + #[test] + fn runtime_worker_subscription_has_gap_free_snapshot_and_live_updates() { + let runtime = runtime_with_backend(); + let mut subscription = runtime + .subscribe_event_selector(EventSubscriptionSelector::RuntimeWorkers) + .unwrap(); + assert_eq!(subscription.snapshot_revision(), 0); + let SubscriptionSnapshot::Workers { workers } = subscription.snapshot() else { + panic!("runtime_workers must return a Worker snapshot"); + }; + assert!(workers.is_empty()); + + let created = runtime.create_worker(task_request("live")).unwrap(); + let update = receive_subscription_update(&mut subscription).unwrap(); + assert_eq!(update.subject_revision, 1); + match update.payload { + SubscriptionEventPayload::WorkerUpserted { worker } => { + assert_eq!(worker.worker_id.as_str(), created.worker_id.to_string()); + assert_eq!(worker.subject_revision, 1); + assert_eq!(worker.state, SubscriptionWorkerState::Idle); + } + payload => panic!("unexpected subscription payload: {payload:?}"), + } + + runtime.stop_worker(&created.worker_ref, None).unwrap(); + let update = receive_subscription_update(&mut subscription).unwrap(); + assert_eq!(update.subject_revision, 2); + match update.payload { + SubscriptionEventPayload::WorkerUpserted { worker } => { + assert_eq!(worker.worker_id.as_str(), created.worker_id.to_string()); + assert_eq!(worker.state, SubscriptionWorkerState::Stopped); + } + payload => panic!("unexpected subscription payload: {payload:?}"), + } + } + + #[test] + fn worker_lifecycle_subscription_delivers_only_selected_workers() { + let runtime = runtime_with_backend(); + let first = runtime.create_worker(task_request("first")).unwrap(); + let second = runtime.create_worker(task_request("second")).unwrap(); + let selected_id = SubscriptionWorkerId::new(first.worker_id.to_string()).unwrap(); + let mut subscription = runtime + .subscribe_event_selector(EventSubscriptionSelector::WorkerLifecycle { + worker_ids: protocol::subscription::SubscriptionWorkerIds::new([ + selected_id.clone() + ]) + .unwrap(), + }) + .unwrap(); + let SubscriptionSnapshot::Workers { workers } = subscription.snapshot() else { + panic!("worker_lifecycle must return a Worker snapshot"); + }; + assert_eq!(workers.len(), 1); + assert_eq!(workers[0].worker_id, selected_id); + + runtime.stop_worker(&second.worker_ref, None).unwrap(); + assert!(matches!( + subscription.receiver.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + + runtime.stop_worker(&first.worker_ref, None).unwrap(); + let update = receive_subscription_update(&mut subscription).unwrap(); + assert!(matches!( + update.payload, + SubscriptionEventPayload::WorkerUpserted { worker } + if worker.worker_id == selected_id + )); + } + + #[test] + fn scoped_runtime_worker_subscription_hides_other_workspaces() { + let runtime = runtime_with_backend(); + let workspace_a = runtime + .create_worker_scoped( + &scope("workspace-a", "server-a"), + scoped_task_request("a", "workspace-a"), + ) + .unwrap(); + let workspace_b = runtime + .create_worker_scoped( + &scope("workspace-b", "server-b"), + scoped_task_request("b", "workspace-b"), + ) + .unwrap(); + let mut subscription = runtime + .subscribe_event_selector_scoped( + &scope("workspace-a", "server-a"), + EventSubscriptionSelector::RuntimeWorkers, + ) + .unwrap(); + let SubscriptionSnapshot::Workers { workers } = subscription.snapshot() else { + panic!("runtime_workers must return a Worker snapshot"); + }; + assert_eq!(workers.len(), 1); + assert_eq!( + workers[0].worker_id.as_str(), + workspace_a.worker_id.to_string() + ); + + runtime.stop_worker(&workspace_b.worker_ref, None).unwrap(); + assert!(matches!( + subscription.receiver.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + runtime.stop_worker(&workspace_a.worker_ref, None).unwrap(); + assert!(receive_subscription_update(&mut subscription).is_ok()); + } + + #[test] + fn lagged_runtime_subscription_closes_without_blocking_mutation() { + let runtime = runtime_with_backend(); + let created = runtime.create_worker(task_request("lag")).unwrap(); + let mut subscription = runtime + .subscribe_event_selector(EventSubscriptionSelector::RuntimeWorkers) + .unwrap(); + { + let mut state = runtime.lock().unwrap(); + for _ in 0..=RUNTIME_EVENT_SUBSCRIPTION_QUEUE_CAPACITY { + state.publish_worker_upsert(created.worker_id).unwrap(); + } + } + assert_eq!( + subscription.receiver.len(), + RUNTIME_EVENT_SUBSCRIPTION_QUEUE_CAPACITY + ); + assert!(matches!( + receive_subscription_update(&mut subscription), + Err(RuntimeSubscriptionRecvError::Lagged) + )); + } + + #[test] + fn dropping_runtime_subscription_releases_producer_state() { + let runtime = runtime_with_backend(); + let subscription = runtime + .subscribe_event_selector(EventSubscriptionSelector::RuntimeWorkers) + .unwrap(); + assert_eq!(runtime.lock().unwrap().event_subscriptions.len(), 1); + drop(subscription); + assert!(runtime.lock().unwrap().event_subscriptions.is_empty()); + } + #[test] fn scoped_worker_access_hides_other_workspace_workers() { let runtime = runtime_with_backend();