server: stream workspace worker subscriptions

This commit is contained in:
2026-08-01 20:12:27 +09:00
parent 5d4deb258a
commit b977c4cbad
7 changed files with 613 additions and 16 deletions
+1
View File
@@ -23,6 +23,7 @@ pub mod runtime_subscription;
pub mod server;
pub mod skills;
pub mod store;
mod workspace_subscription;
pub use authority::{
MemoryAuthority, MemoryDocument, MemoryStagingEntry, MemoryStagingResolution,
@@ -242,6 +242,7 @@ impl RuntimeSubscriptionBroker {
})?;
let downstream_id = self.next_downstream.fetch_add(1, Ordering::Relaxed);
let (events, receiver) = mpsc::channel(DOWNSTREAM_QUEUE_CAPACITY);
let initial_events = events.clone();
registration
.commands
.send(Command::Subscribe {
@@ -250,6 +251,17 @@ impl RuntimeSubscriptionBroker {
events,
})
.map_err(|_| RuntimeSubscriptionBrokerError::Closed)?;
let initial_status = registration
.status
.read()
.expect("broker status poisoned")
.clone();
if !initial_status.connected {
let _ = initial_events.try_send(BrokerSubscriptionEvent::Disconnected {
connection_generation: initial_status.connection_generation,
message: "Runtime subscription connection is not currently available".to_string(),
});
}
Ok(BrokerSubscription {
downstream_id,
runtime_id: runtime_id.to_string(),
@@ -291,6 +303,7 @@ impl SelectorState {
}
struct State {
runtime_id: String,
generation: u64,
next_request: u64,
selectors: HashMap<EventSubscriptionSelector, SelectorState>,
@@ -299,8 +312,9 @@ struct State {
upstream_index: HashMap<SubscriptionId, EventSubscriptionSelector>,
}
impl State {
fn new(generation: u64) -> Self {
fn new(runtime_id: String, generation: u64) -> Self {
Self {
runtime_id,
generation,
next_request: 1,
selectors: HashMap::new(),
@@ -433,8 +447,17 @@ fn project_payload_runtime(
mut payload: SubscriptionEventPayload,
runtime_id: &str,
) -> SubscriptionEventPayload {
if let SubscriptionEventPayload::WorkerUpserted { worker } = &mut payload {
worker.runtime_id = Some(runtime_id.to_string());
match &mut payload {
SubscriptionEventPayload::WorkerUpserted { worker } => {
worker.runtime_id = Some(runtime_id.to_string());
}
SubscriptionEventPayload::WorkerRemoved {
runtime_id: projected_runtime_id,
..
} => {
*projected_runtime_id = Some(runtime_id.to_string());
}
_ => {}
}
payload
}
@@ -446,7 +469,7 @@ async fn run_connection(
mut commands: mpsc::UnboundedReceiver<Command>,
status: Arc<RwLock<RuntimeSubscriptionBrokerStatus>>,
) {
let mut state = State::new(generation);
let mut state = State::new(config.runtime_id.clone(), generation);
let mut disconnect_notified = false;
loop {
update_status(&status, &state, false);
@@ -653,6 +676,7 @@ async fn handle_frame(
if state.pending.remove(&request_id) != Some(selector.clone()) {
return Err(());
}
let snapshot = project_snapshot_runtime(snapshot, &state.runtime_id);
let entry = state.selectors.get_mut(&selector).ok_or(())?;
entry.pending = false;
entry.upstream_id = Some(subscription_id.clone());
@@ -704,6 +728,7 @@ async fn handle_frame(
.cloned()
.ok_or(())?;
payload.validate_for_selector(&selector).map_err(|_| ())?;
let payload = project_payload_runtime(payload, &state.runtime_id);
let entry = state.selectors.get_mut(&selector).ok_or(())?;
if let Some(subject) = event_subject(&payload) {
let revision = entry.revisions.entry(subject).or_insert(0);
@@ -793,7 +818,16 @@ fn snapshot_revisions(snapshot: &SubscriptionSnapshot) -> HashMap<String, u64> {
match snapshot {
SubscriptionSnapshot::Workers { workers } => workers
.iter()
.map(|worker| (worker.worker_id.to_string(), worker.subject_revision))
.map(|worker| {
(
format!(
"{}:{}",
worker.runtime_id.as_deref().unwrap_or_default(),
worker.worker_id
),
worker.subject_revision,
)
})
.collect(),
SubscriptionSnapshot::WorkerProtocol { worker_id, .. } => {
HashMap::from([(worker_id.to_string(), 0)])
@@ -803,9 +837,20 @@ fn snapshot_revisions(snapshot: &SubscriptionSnapshot) -> HashMap<String, u64> {
}
fn event_subject(payload: &SubscriptionEventPayload) -> Option<String> {
Some(match payload {
SubscriptionEventPayload::WorkerUpserted { worker } => worker.worker_id.to_string(),
SubscriptionEventPayload::WorkerRemoved { worker_id }
| SubscriptionEventPayload::WorkerProtocol { worker_id, .. } => worker_id.to_string(),
SubscriptionEventPayload::WorkerUpserted { worker } => format!(
"{}:{}",
worker.runtime_id.as_deref().unwrap_or_default(),
worker.worker_id
),
SubscriptionEventPayload::WorkerRemoved {
worker_id,
runtime_id,
} => format!(
"{}:{}",
runtime_id.as_deref().unwrap_or_default(),
worker_id
),
SubscriptionEventPayload::WorkerProtocol { worker_id, .. } => worker_id.to_string(),
SubscriptionEventPayload::WorkdirUpserted { workdir } => {
workdir.working_directory_id.to_string()
}
+139 -2
View File
@@ -867,6 +867,10 @@ pub fn build_router(api: WorkspaceApi) -> Router {
"/api/w/{workspace_id}/workers",
get(scoped_list_workers).post(scoped_create_workspace_worker),
)
.route(
"/api/w/{workspace_id}/protocol/ws",
get(scoped_workspace_protocol_ws),
)
.route(
"/api/workers/launch-options",
get(get_worker_launch_options),
@@ -3779,6 +3783,38 @@ async fn scoped_list_runtimes(
list_runtimes(State(api)).await
}
async fn scoped_workspace_protocol_ws(
State(api): State<WorkspaceApi>,
AxumPath(workspace_id): AxumPath<String>,
headers: HeaderMap,
ws: axum::extract::ws::WebSocketUpgrade,
) -> std::result::Result<Response, Response> {
validate_workspace_scope(&api, &workspace_id).map_err(|error| error.into_response())?;
let actor = resolve_actor(&api, &headers)
.await
.map_err(|error| error.into_response())?
.ok_or_else(|| StatusCode::UNAUTHORIZED.into_response())?;
let workspace = api
.store
.get_workspace(&workspace_id)
.await
.map_err(|error| ApiError::from(error).into_response())?
.ok_or_else(|| StatusCode::NOT_FOUND.into_response())?;
if workspace
.owner_account_id
.as_deref()
.is_some_and(|owner| owner != actor.account_id.as_str())
{
return Err(StatusCode::FORBIDDEN.into_response());
}
let broker = api.runtime_subscription_broker().clone();
Ok(ws
.on_upgrade(move |socket| {
crate::workspace_subscription::serve_workspace_subscription(broker, socket)
})
.into_response())
}
async fn scoped_list_workers(
State(api): State<WorkspaceApi>,
AxumPath(path): AxumPath<ScopedWorkspacePath>,
@@ -8716,6 +8752,7 @@ mod tests {
use std::{fs, sync::Arc};
use tokio_tungstenite::connect_async;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tower::ServiceExt;
use worker_runtime::resource::BackendResourceClient;
use worker_runtime::working_directory::WorkingDirectoryMaterializer;
@@ -8725,8 +8762,9 @@ mod tests {
WorkerSpawnIntent,
};
use crate::store::{
MemoryDocumentRecord, MemoryStagingRecord, ObjectiveRecord, ObjectiveResourceRecord,
ObjectiveTicketLinkRecord, SqliteWorkspaceStore, WorkspaceRecord,
AccountRecord, MemoryDocumentRecord, MemoryStagingRecord, ObjectiveRecord,
ObjectiveResourceRecord, ObjectiveTicketLinkRecord, SqliteWorkspaceStore, UserRecord,
WorkspaceRecord,
};
const TEST_WORKSPACE_ID: &str = "0192f0e8-4d84-7d6e-a000-000000000001";
@@ -12623,6 +12661,105 @@ mod tests {
)
}
#[tokio::test]
async fn workspace_subscription_requires_browser_session() {
let dir = tempfile::tempdir().unwrap();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let app = build_router(test_api(dir.path()).await);
let server = tokio::spawn(async move {
let _ = axum::serve(listener, app).await;
});
let error = tokio_tungstenite::connect_async(format!(
"ws://{address}/api/w/{TEST_WORKSPACE_ID}/protocol/ws"
))
.await
.unwrap_err();
let tokio_tungstenite::tungstenite::Error::Http(response) = error else {
panic!("expected HTTP authentication rejection");
};
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
server.abort();
}
#[tokio::test]
async fn workspace_subscription_returns_authenticated_workspace_snapshot() {
let dir = tempfile::tempdir().unwrap();
let api = test_api(dir.path()).await;
let account = AccountRecord {
account_id: "account-test".to_string(),
kind: "user".to_string(),
handle: "tester".to_string(),
display_name: "Tester".to_string(),
created_at: TEST_CREATED_AT.to_string(),
updated_at: TEST_CREATED_AT.to_string(),
};
let user = UserRecord {
user_id: "user-test".to_string(),
account_id: account.account_id.clone(),
handle: account.handle.clone(),
display_name: account.display_name.clone(),
created_at: TEST_CREATED_AT.to_string(),
updated_at: TEST_CREATED_AT.to_string(),
};
api.store.upsert_account(&account).unwrap();
api.store.upsert_user(&user).unwrap();
let session = issue_browser_session_response(&api, user).unwrap();
let cookie = session
.headers()
.get(SET_COOKIE)
.unwrap()
.to_str()
.unwrap()
.split(';')
.next()
.unwrap()
.to_string();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let app = build_router(api);
let server = tokio::spawn(async move {
let _ = axum::serve(listener, app).await;
});
let mut request = format!("ws://{address}/api/w/{TEST_WORKSPACE_ID}/protocol/ws")
.into_client_request()
.unwrap();
request
.headers_mut()
.insert(axum::http::header::COOKIE, cookie.parse().unwrap());
let (mut socket, _) = connect_async(request).await.unwrap();
let frame = protocol::subscription::SubscriptionFrame::new(
protocol::subscription::SubscriptionFramePayload::Request(
protocol::subscription::SubscriptionRequest::SubscribeEvents {
request_id: protocol::subscription::SubscriptionRequestId::new("request-1")
.unwrap(),
selector: protocol::subscription::EventSubscriptionSelector::WorkspaceWorkers,
},
),
);
socket
.send(Message::Text(serde_json::to_string(&frame).unwrap().into()))
.await
.unwrap();
let Message::Text(text) = socket.next().await.unwrap().unwrap() else {
panic!("expected subscription response");
};
let response: protocol::subscription::SubscriptionFrame =
serde_json::from_str(text.as_str()).unwrap();
assert!(matches!(
response.payload,
protocol::subscription::SubscriptionFramePayload::Response(
protocol::subscription::SubscriptionResponse::Subscribed {
selector: protocol::subscription::EventSubscriptionSelector::WorkspaceWorkers,
snapshot: protocol::subscription::SubscriptionSnapshot::Workers { .. },
..
}
)
));
server.abort();
}
#[tokio::test]
async fn passkey_registration_rejects_unverified_credential_response() {
let dir = tempfile::tempdir().unwrap();
@@ -0,0 +1,389 @@
use std::collections::{BTreeMap, HashMap, HashSet};
use axum::extract::ws::{Message as WsMessage, WebSocket};
use futures::{SinkExt, StreamExt};
use protocol::subscription::{
EventSubscriptionSelector, SubscriptionEvent, SubscriptionEventPayload, SubscriptionFrame,
SubscriptionFramePayload, SubscriptionId, SubscriptionRejectionCode, SubscriptionRequest,
SubscriptionResponse, SubscriptionSnapshot, SubscriptionTerminationCode, SubscriptionWorker,
};
use tokio::sync::mpsc;
use crate::runtime_subscription::{BrokerSubscriptionEvent, RuntimeSubscriptionBroker};
const OUTBOUND_CAPACITY: usize = 256;
pub(crate) async fn serve_workspace_subscription(
broker: RuntimeSubscriptionBroker,
socket: WebSocket,
) {
let (mut socket_sender, mut socket_receiver) = socket.split();
let (outbound, mut outbound_receiver) = mpsc::channel::<WsMessage>(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 {
WsMessage::Text(text) => {
let Ok(frame) = serde_json::from_str::<SubscriptionFrame>(text.as_str()) else {
break;
};
if frame.validate().is_err() {
break;
}
let SubscriptionFramePayload::Request(request) = frame.payload else {
break;
};
subscriptions.retain(|_, task| !task.is_finished());
match request {
SubscriptionRequest::SubscribeEvents {
request_id,
selector,
} => {
if selector != EventSubscriptionSelector::WorkspaceWorkers {
let _ = send_frame(&outbound, SubscriptionFrame::new(
SubscriptionFramePayload::Response(
SubscriptionResponse::SubscriptionRejected {
request_id,
subscription_id: None,
code: SubscriptionRejectionCode::UnsupportedSelector,
message: "Workspace clients may subscribe only to workspace_workers on this endpoint".to_string(),
},
),
)).await;
continue;
}
let subscription_id = SubscriptionId::new(format!(
"workspace-subscription-{next_subscription_id}"
))
.expect("generated Workspace subscription id is valid");
next_subscription_id = next_subscription_id.saturating_add(1);
let task = tokio::spawn(run_workspace_workers(
broker.clone(),
request_id,
subscription_id.clone(),
outbound.clone(),
));
subscriptions.insert(subscription_id, task);
}
SubscriptionRequest::UnsubscribeEvents {
request_id,
subscription_id,
} => {
if let Some(task) = subscriptions.remove(&subscription_id) {
task.abort();
}
if send_frame(
&outbound,
SubscriptionFrame::new(SubscriptionFramePayload::Response(
SubscriptionResponse::Unsubscribed {
request_id,
subscription_id,
},
)),
)
.await
.is_err()
{
break;
}
}
}
}
WsMessage::Ping(value) => {
if outbound.send(WsMessage::Pong(value)).await.is_err() {
break;
}
}
WsMessage::Pong(_) => {}
WsMessage::Close(_) | WsMessage::Binary(_) => break,
}
}
for (_, task) in subscriptions {
task.abort();
}
drop(outbound);
let _ = writer.await;
}
async fn run_workspace_workers(
broker: RuntimeSubscriptionBroker,
request_id: protocol::subscription::SubscriptionRequestId,
subscription_id: SubscriptionId,
outbound: mpsc::Sender<WsMessage>,
) {
let runtime_ids = broker.runtime_ids();
let mut pending = runtime_ids.iter().cloned().collect::<HashSet<_>>();
let (events, mut event_receiver) = mpsc::channel(OUTBOUND_CAPACITY);
let mut upstreams = tokio::task::JoinSet::new();
for runtime_id in runtime_ids {
let Ok(mut subscription) =
broker.subscribe(&runtime_id, EventSubscriptionSelector::RuntimeWorkers)
else {
pending.remove(&runtime_id);
continue;
};
let sender = events.clone();
upstreams.spawn(async move {
while let Some(event) = subscription.recv().await {
if sender.send((runtime_id.clone(), event)).await.is_err() {
break;
}
}
});
}
drop(events);
let mut workers = HashMap::<String, BTreeMap<String, SubscriptionWorker>>::new();
while !pending.is_empty() {
let Some((runtime_id, event)) = event_receiver.recv().await else {
return;
};
match event {
BrokerSubscriptionEvent::Snapshot { snapshot, .. } => {
install_snapshot(&mut workers, &runtime_id, snapshot);
pending.remove(&runtime_id);
}
BrokerSubscriptionEvent::Disconnected { .. }
| BrokerSubscriptionEvent::Rejected { .. }
| BrokerSubscriptionEvent::Closed { .. } => {
pending.remove(&runtime_id);
}
BrokerSubscriptionEvent::Event { .. } => {}
}
}
let mut revisions = HashMap::<String, u64>::new();
let mut initial_workers = workers
.values_mut()
.flat_map(|runtime| runtime.values_mut())
.map(|worker| {
let key = worker_key(worker.runtime_id.as_deref(), worker.worker_id.as_str());
worker.subject_revision = next_revision(&mut revisions, &key);
worker.clone()
})
.collect::<Vec<_>>();
sort_workers(&mut initial_workers);
if send_frame(
&outbound,
SubscriptionFrame::new(SubscriptionFramePayload::Response(
SubscriptionResponse::Subscribed {
request_id,
subscription_id: subscription_id.clone(),
selector: EventSubscriptionSelector::WorkspaceWorkers,
snapshot_revision: 1,
snapshot: SubscriptionSnapshot::Workers {
workers: initial_workers,
},
},
)),
)
.await
.is_err()
{
return;
}
while let Some((runtime_id, event)) = event_receiver.recv().await {
match event {
BrokerSubscriptionEvent::Snapshot { snapshot, .. } => {
let removed = workers.remove(&runtime_id).unwrap_or_default();
for worker in removed.values() {
let key = worker_key(Some(&runtime_id), worker.worker_id.as_str());
let revision = next_revision(&mut revisions, &key);
if send_event(
&outbound,
&subscription_id,
revision,
SubscriptionEventPayload::WorkerRemoved {
worker_id: worker.worker_id.clone(),
runtime_id: Some(runtime_id.clone()),
},
)
.await
.is_err()
{
return;
}
}
install_snapshot(&mut workers, &runtime_id, snapshot);
if let Some(current) = workers.get_mut(&runtime_id) {
for worker in current.values_mut() {
let key = worker_key(Some(&runtime_id), worker.worker_id.as_str());
let revision = next_revision(&mut revisions, &key);
worker.subject_revision = revision;
if send_event(
&outbound,
&subscription_id,
revision,
SubscriptionEventPayload::WorkerUpserted {
worker: worker.clone(),
},
)
.await
.is_err()
{
return;
}
}
}
}
BrokerSubscriptionEvent::Event { payload, .. } => match payload {
SubscriptionEventPayload::WorkerUpserted { mut worker } => {
worker.runtime_id = Some(runtime_id.clone());
let key = worker_key(Some(&runtime_id), worker.worker_id.as_str());
let revision = next_revision(&mut revisions, &key);
worker.subject_revision = revision;
workers
.entry(runtime_id)
.or_default()
.insert(worker.worker_id.to_string(), worker.clone());
if send_event(
&outbound,
&subscription_id,
revision,
SubscriptionEventPayload::WorkerUpserted { worker },
)
.await
.is_err()
{
return;
}
}
SubscriptionEventPayload::WorkerRemoved { worker_id, .. } => {
workers
.entry(runtime_id.clone())
.or_default()
.remove(worker_id.as_str());
let key = worker_key(Some(&runtime_id), worker_id.as_str());
let revision = next_revision(&mut revisions, &key);
if send_event(
&outbound,
&subscription_id,
revision,
SubscriptionEventPayload::WorkerRemoved {
worker_id,
runtime_id: Some(runtime_id),
},
)
.await
.is_err()
{
return;
}
}
_ => {}
},
BrokerSubscriptionEvent::Disconnected { .. } => {}
BrokerSubscriptionEvent::Rejected { code, message, .. } => {
let _ = send_frame(
&outbound,
SubscriptionFrame::new(SubscriptionFramePayload::Event(
SubscriptionEvent::SubscriptionClosed {
subscription_id: subscription_id.clone(),
code: rejection_termination(code),
message,
},
)),
)
.await;
return;
}
BrokerSubscriptionEvent::Closed { code, message, .. } => {
let _ = send_frame(
&outbound,
SubscriptionFrame::new(SubscriptionFramePayload::Event(
SubscriptionEvent::SubscriptionClosed {
subscription_id: subscription_id.clone(),
code,
message,
},
)),
)
.await;
return;
}
}
}
}
fn install_snapshot(
workers: &mut HashMap<String, BTreeMap<String, SubscriptionWorker>>,
runtime_id: &str,
snapshot: SubscriptionSnapshot,
) {
let SubscriptionSnapshot::Workers {
workers: snapshot_workers,
} = snapshot
else {
return;
};
let mut projected = BTreeMap::new();
for mut worker in snapshot_workers {
worker.runtime_id = Some(runtime_id.to_string());
projected.insert(worker.worker_id.to_string(), worker);
}
workers.insert(runtime_id.to_string(), projected);
}
async fn send_event(
outbound: &mpsc::Sender<WsMessage>,
subscription_id: &SubscriptionId,
subject_revision: u64,
payload: SubscriptionEventPayload,
) -> Result<(), ()> {
send_frame(
outbound,
SubscriptionFrame::new(SubscriptionFramePayload::Event(SubscriptionEvent::Event {
subscription_id: subscription_id.clone(),
subject_revision,
payload,
})),
)
.await
}
async fn send_frame(
outbound: &mpsc::Sender<WsMessage>,
frame: SubscriptionFrame,
) -> Result<(), ()> {
frame.validate().map_err(|_| ())?;
outbound
.send(WsMessage::Text(
serde_json::to_string(&frame).map_err(|_| ())?.into(),
))
.await
.map_err(|_| ())
}
fn next_revision(revisions: &mut HashMap<String, u64>, key: &str) -> u64 {
let revision = revisions.entry(key.to_string()).or_insert(0);
*revision = revision.saturating_add(1);
*revision
}
fn worker_key(runtime_id: Option<&str>, worker_id: &str) -> String {
format!("{}:{worker_id}", runtime_id.unwrap_or_default())
}
fn sort_workers(workers: &mut [SubscriptionWorker]) {
workers.sort_by(|left, right| {
left.runtime_id
.cmp(&right.runtime_id)
.then_with(|| left.worker_id.cmp(&right.worker_id))
});
}
fn rejection_termination(code: SubscriptionRejectionCode) -> SubscriptionTerminationCode {
match code {
SubscriptionRejectionCode::Unauthorized => SubscriptionTerminationCode::Unauthorized,
SubscriptionRejectionCode::ResourceNotFound => SubscriptionTerminationCode::ResourceGone,
_ => SubscriptionTerminationCode::ServerShutdown,
}
}