runtime: serve selective event subscriptions

This commit is contained in:
2026-08-01 18:34:32 +09:00
parent 21cd672f64
commit ec262f0238
2 changed files with 951 additions and 14 deletions
+418 -5
View File
@@ -21,6 +21,8 @@ use crate::interaction::{WorkerInput, WorkerInteractionAck};
use crate::management::{RuntimeLimits, RuntimeSummary, WorkerDeleteResult}; use crate::management::{RuntimeLimits, RuntimeSummary, WorkerDeleteResult};
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
use crate::observation::WorkerObservationCursor; use crate::observation::WorkerObservationCursor;
#[cfg(feature = "ws-server")]
use crate::runtime::RuntimeSubscriptionRecvError;
use crate::{Runtime, RuntimeWorkspaceScope}; use crate::{Runtime, RuntimeWorkspaceScope};
use axum::body::{Body, Bytes}; use axum::body::{Body, Bytes};
use axum::extract::rejection::{JsonRejection, QueryRejection}; use axum::extract::rejection::{JsonRejection, QueryRejection};
@@ -33,10 +35,18 @@ use axum::response::{IntoResponse, Response};
use axum::routing::{get, post}; use axum::routing::{get, post};
use axum::{Json, Router}; use axum::{Json, Router};
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
use futures::StreamExt; use futures::{SinkExt, StreamExt};
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
use protocol::stream::{decode_method, encode_event}; 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}; use serde::{Deserialize, Serialize};
#[cfg(feature = "ws-server")]
use std::collections::HashMap;
use std::fmt; use std::fmt;
use std::net::SocketAddr; use std::net::SocketAddr;
#[cfg(feature = "fs-store")] #[cfg(feature = "fs-store")]
@@ -205,10 +215,12 @@ fn runtime_http_router_with_optional_auth(
.route("/v1/workers/{worker_id}/cancel", post(cancel_worker)); .route("/v1/workers/{worker_id}/cancel", post(cancel_worker));
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
let router = router.route( let router = router
"/v1/workers/{worker_id}/protocol/ws", .route("/v1/protocol/ws", get(runtime_protocol_ws))
get(worker_protocol_ws), .route(
); "/v1/workers/{worker_id}/protocol/ws",
get(worker_protocol_ws),
);
router router
.with_state(state.clone()) .with_state(state.clone())
@@ -555,6 +567,272 @@ async fn restore_worker(
Ok(Json(RuntimeHttpWorkerResponse { 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<RuntimeHttpState>,
auth: Option<Extension<RuntimeAuthContext>>,
ws: axum::extract::ws::WebSocketUpgrade,
) -> Result<Response, Response> {
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<RuntimeWorkspaceScope>,
socket: axum::extract::ws::WebSocket,
) {
let (mut socket_sender, mut socket_receiver) = socket.split();
let (outbound, mut outbound_receiver) = tokio::sync::mpsc::channel::<axum::extract::ws::Message>(
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::<SubscriptionId, tokio::task::JoinHandle<()>>::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::<SubscriptionFrame>(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<protocol::subscription::SubscriptionRequestId> {
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<axum::extract::ws::Message>,
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")] #[cfg(feature = "ws-server")]
async fn worker_protocol_ws( async fn worker_protocol_ws(
State(state): State<RuntimeHttpState>, State(state): State<RuntimeHttpState>,
@@ -1001,6 +1279,9 @@ fn required_runtime_permission(method: &Method, path: &str) -> Option<&'static s
if path.ends_with("/stop") || path.ends_with("/cancel") { if path.ends_with("/stop") || path.ends_with("/cancel") {
return Some("workers:stop"); return Some("workers:stop");
} }
if path == "/v1/protocol/ws" {
return Some("workers:list");
}
if path.ends_with("/protocol") || path.ends_with("/protocol/ws") { if path.ends_with("/protocol") || path.ends_with("/protocol/ws") {
return Some("workers:protocol"); return Some("workers:protocol");
} }
@@ -1965,6 +2246,138 @@ mod ws_tests {
serde_json::from_str(&text).unwrap() serde_json::from_str(&text).unwrap()
} }
async fn next_subscription_frame(
stream: &mut tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
) -> 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<tokio::net::TcpStream>,
>,
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] #[tokio::test]
async fn protocol_ws_connect_sends_snapshot_and_live_worker_events() { async fn protocol_ws_connect_sends_snapshot_and_live_worker_events() {
let (runtime, worker_ref, url) = spawn_runtime_server().await; let (runtime, worker_ref, url) = spawn_runtime_server().await;
+533 -9
View File
@@ -1,6 +1,6 @@
use crate::catalog::{ use crate::catalog::{
ConfigBundleRef, CreateWorkerRequest, WorkerDetail, WorkerLifecycleAck, WorkerStatus, ConfigBundleRef, CreateWorkerRequest, ProfileSelector, WorkerDetail, WorkerLifecycleAck,
WorkerSummary, WorkingDirectoryRequest, WorkerStatus, WorkerSummary, WorkingDirectoryRequest,
WorkingDirectoryStatus as CatalogWorkingDirectoryStatus, WorkspaceApiRef, WorkingDirectoryStatus as CatalogWorkingDirectoryStatus, WorkspaceApiRef,
}; };
use crate::config_bundle::{ use crate::config_bundle::{
@@ -31,13 +31,20 @@ use crate::observation::{
}; };
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
use crate::observation::{WorkerObservationCursor, WorkerObservationEvent}; use crate::observation::{WorkerObservationCursor, WorkerObservationEvent};
use protocol::subscription::{
EventSubscriptionSelector, SubscriptionEventPayload, SubscriptionSnapshot,
SubscriptionValidationError, SubscriptionWorkdirId, SubscriptionWorker, SubscriptionWorkerId,
SubscriptionWorkerState,
};
use protocol::{Event, Method}; use protocol::{Event, Method};
use std::collections::BTreeMap; use std::collections::BTreeMap;
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
use std::collections::VecDeque; 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")] #[cfg(feature = "ws-server")]
use tokio::sync::broadcast; use tokio::sync::broadcast;
use tokio::sync::mpsc;
/// Workspace-scoped Runtime authorization context supplied by a trusted backend. /// Workspace-scoped Runtime authorization context supplied by a trusted backend.
#[derive(Clone, Debug, Eq, PartialEq)] #[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<RuntimeSubscriptionUpdate>,
lagged: Arc<AtomicBool>,
runtime: Weak<Mutex<RuntimeState>>,
}
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<RuntimeSubscriptionUpdate, RuntimeSubscriptionRecvError> {
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. /// Concrete embedded Runtime domain entity.
/// ///
/// The default implementation is memory-backed and tools/provider-less by /// 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()) Ok(state.workers.values().map(WorkerRecord::summary).collect())
} }
pub fn subscribe_event_selector(
&self,
selector: EventSubscriptionSelector,
) -> Result<RuntimeEventSelectorSubscription, RuntimeError> {
self.subscribe_event_selector_for_workspace(None, selector)
}
pub fn subscribe_event_selector_scoped(
&self,
scope: &RuntimeWorkspaceScope,
selector: EventSubscriptionSelector,
) -> Result<RuntimeEventSelectorSubscription, RuntimeError> {
self.subscribe_event_selector_for_workspace(Some(scope), selector)
}
fn subscribe_event_selector_for_workspace(
&self,
scope: Option<&RuntimeWorkspaceScope>,
selector: EventSubscriptionSelector,
) -> Result<RuntimeEventSelectorSubscription, RuntimeError> {
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. /// List stopped Workers known to this Runtime.
pub fn list_stopped_workers(&self) -> Result<Vec<WorkerSummary>, RuntimeError> { pub fn list_stopped_workers(&self) -> Result<Vec<WorkerSummary>, RuntimeError> {
let state = self.lock()?; let state = self.lock()?;
@@ -846,6 +991,7 @@ impl Runtime {
let payload = input_protocol_event(&input); let payload = input_protocol_event(&input);
state.push_worker_observation_event(worker_ref.clone(), payload); state.push_worker_observation_event(worker_ref.clone(), payload);
} }
state.publish_worker_upsert(worker_ref.worker_id)?;
state.persist_runtime_snapshot()?; state.persist_runtime_snapshot()?;
state.persist_worker(&worker_ref.worker_id)?; state.persist_worker(&worker_ref.worker_id)?;
state.persist_event_by_id(event_id)?; state.persist_event_by_id(event_id)?;
@@ -979,6 +1125,7 @@ impl Runtime {
worker.working_directory = working_directory; worker.working_directory = working_directory;
worker.detail() worker.detail()
}; };
state.publish_worker_upsert(worker_ref.worker_id)?;
state.persist_runtime_snapshot()?; state.persist_runtime_snapshot()?;
state.persist_worker(&worker_ref.worker_id)?; state.persist_worker(&worker_ref.worker_id)?;
state.persist_event_by_id(detail.last_event_id)?; state.persist_event_by_id(detail.last_event_id)?;
@@ -995,6 +1142,7 @@ impl Runtime {
state.events.retain(|event| { state.events.retain(|event| {
event.id != record.last_event_id || event.worker_ref.as_ref() != Some(worker_ref) 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(()) Ok(())
} }
@@ -1005,9 +1153,9 @@ impl Runtime {
result: WorkerExecutionResult, result: WorkerExecutionResult,
) -> Result<(), RuntimeError> { ) -> Result<(), RuntimeError> {
let mut state = self.lock()?; let mut state = self.lock()?;
let worker = state.worker_mut(worker_ref)?;
if result.is_accepted() { 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(()) Ok(())
} }
@@ -1138,7 +1286,8 @@ impl Runtime {
worker_id: worker_ref.worker_id, 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); state.forget_workspace_owner_if_unused(workspace_id);
} }
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
@@ -1150,6 +1299,7 @@ impl Runtime {
RuntimeEventKind::WorkerDeleted, RuntimeEventKind::WorkerDeleted,
"worker deleted", "worker deleted",
); );
state.publish_worker_removed(worker_ref.worker_id, removed_workspace_id.as_deref())?;
state.persist_runtime_snapshot()?; state.persist_runtime_snapshot()?;
state.delete_worker_snapshot(&worker_ref.worker_id)?; state.delete_worker_snapshot(&worker_ref.worker_id)?;
state.persist_event_by_id(event_id)?; state.persist_event_by_id(event_id)?;
@@ -1304,7 +1454,10 @@ impl Runtime {
) -> Result<WorkerObservationEvent, RuntimeError> { ) -> Result<WorkerObservationEvent, RuntimeError> {
let mut state = self.lock()?; let mut state = self.lock()?;
state.ensure_worker_ref(worker_ref)?; 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); let event = state.push_worker_observation_event(worker_ref.clone(), payload);
Ok(event) Ok(event)
} }
@@ -1363,6 +1516,7 @@ impl Runtime {
worker.execution_handle = None; worker.execution_handle = None;
worker.last_event_id = event_id; worker.last_event_id = event_id;
let status = worker.status; let status = worker.status;
state.publish_worker_upsert(worker_ref.worker_id)?;
state.persist_runtime_snapshot()?; state.persist_runtime_snapshot()?;
state.persist_worker(&worker_ref.worker_id)?; state.persist_worker(&worker_ref.worker_id)?;
state.persist_event_by_id(event_id)?; state.persist_event_by_id(event_id)?;
@@ -1491,6 +1645,7 @@ impl Runtime {
worker.working_directory = working_directory; worker.working_directory = working_directory;
worker.last_event_id = event_id; worker.last_event_id = event_id;
} }
state.publish_worker_upsert(worker_ref.worker_id)?;
state.persist_runtime_snapshot()?; state.persist_runtime_snapshot()?;
state.persist_worker(&worker_ref.worker_id)?; state.persist_worker(&worker_ref.worker_id)?;
state.persist_event_by_id(event_id)?; state.persist_event_by_id(event_id)?;
@@ -1510,6 +1665,14 @@ enum RuntimePersistence {
Fs(FsRuntimeStore), Fs(FsRuntimeStore),
} }
#[derive(Debug)]
struct RuntimeEventSubscriptionSink {
selector: EventSubscriptionSelector,
workspace_id: Option<String>,
sender: mpsc::Sender<RuntimeSubscriptionUpdate>,
lagged: Arc<AtomicBool>,
}
#[derive(Debug)] #[derive(Debug)]
struct RuntimeState { struct RuntimeState {
display_name: Option<String>, display_name: Option<String>,
@@ -1528,6 +1691,10 @@ struct RuntimeState {
config_bundles: BTreeMap<String, ConfigBundle>, config_bundles: BTreeMap<String, ConfigBundle>,
events: Vec<RuntimeEvent>, events: Vec<RuntimeEvent>,
diagnostics: Vec<RuntimeDiagnostic>, diagnostics: Vec<RuntimeDiagnostic>,
subscription_revision: u64,
worker_subject_revisions: BTreeMap<WorkerId, u64>,
next_event_subscription_id: u64,
event_subscriptions: BTreeMap<u64, RuntimeEventSubscriptionSink>,
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
next_observation_sequence: u64, next_observation_sequence: u64,
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
@@ -1554,6 +1721,10 @@ impl RuntimeState {
config_bundles: BTreeMap::new(), config_bundles: BTreeMap::new(),
events: Vec::new(), events: Vec::new(),
diagnostics: 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")] #[cfg(feature = "ws-server")]
next_observation_sequence: 1, next_observation_sequence: 1,
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
@@ -1585,6 +1756,10 @@ impl RuntimeState {
config_bundles: BTreeMap::new(), config_bundles: BTreeMap::new(),
events: Vec::new(), events: Vec::new(),
diagnostics: 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")] #[cfg(feature = "ws-server")]
next_observation_sequence: 1, next_observation_sequence: 1,
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
@@ -1633,6 +1808,10 @@ impl RuntimeState {
workspace_owners: persisted.workspace_owners, workspace_owners: persisted.workspace_owners,
events: persisted.events, events: persisted.events,
diagnostics, diagnostics,
subscription_revision: 0,
worker_subject_revisions: BTreeMap::new(),
next_event_subscription_id: 1,
event_subscriptions: BTreeMap::new(),
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
next_observation_sequence: 1, next_observation_sequence: 1,
#[cfg(feature = "ws-server")] #[cfg(feature = "ws-server")]
@@ -1871,6 +2050,179 @@ impl RuntimeState {
}) })
} }
fn subscription_snapshot(
&self,
scope: Option<&RuntimeWorkspaceScope>,
selector: &EventSubscriptionSelector,
) -> Result<SubscriptionSnapshot, RuntimeError> {
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::<Result<Vec<_>, _>>()?,
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<SubscriptionWorker, RuntimeError> {
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<WorkerId> { fn primary_worker_id_for_workdir(&self, working_directory_id: &str) -> Option<WorkerId> {
self.workers.values().find_map(|worker| { self.workers.values().find_map(|worker| {
if worker if worker
@@ -1911,6 +2263,7 @@ impl RuntimeState {
let worker = self.worker_mut(worker_ref)?; let worker = self.worker_mut(worker_ref)?;
worker.execution_handle = None; worker.execution_handle = None;
worker.status = WorkerStatus::Stopped; worker.status = WorkerStatus::Stopped;
self.publish_worker_upsert(worker_ref.worker_id)?;
self.persist_runtime_snapshot()?; self.persist_runtime_snapshot()?;
Ok(()) Ok(())
} }
@@ -1989,9 +2342,9 @@ impl RuntimeState {
&mut self, &mut self,
worker_ref: &WorkerRef, worker_ref: &WorkerRef,
event: &protocol::Event, event: &protocol::Event,
) { ) -> bool {
let Some(worker) = self.workers.get_mut(&worker_ref.worker_id) else { let Some(worker) = self.workers.get_mut(&worker_ref.worker_id) else {
return; return false;
}; };
let next_status = match event { let next_status = match event {
protocol::Event::Status { protocol::Event::Status {
@@ -2018,7 +2371,11 @@ impl RuntimeState {
_ => None, _ => None,
}; };
if let Some(next_status) = next_status { if let Some(next_status) = next_status {
let changed = worker.status != next_status;
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)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -2476,6 +2847,159 @@ mod tests {
request request
} }
fn receive_subscription_update(
subscription: &mut RuntimeEventSelectorSubscription,
) -> Result<RuntimeSubscriptionUpdate, RuntimeSubscriptionRecvError> {
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] #[test]
fn scoped_worker_access_hides_other_workspace_workers() { fn scoped_worker_access_hides_other_workspace_workers() {
let runtime = runtime_with_backend(); let runtime = runtime_with_backend();