server: stream workspace worker subscriptions
This commit is contained in:
@@ -643,6 +643,8 @@ pub enum SubscriptionEventPayload {
|
|||||||
},
|
},
|
||||||
WorkerRemoved {
|
WorkerRemoved {
|
||||||
worker_id: SubscriptionWorkerId,
|
worker_id: SubscriptionWorkerId,
|
||||||
|
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||||
|
runtime_id: Option<String>,
|
||||||
},
|
},
|
||||||
WorkerProtocol {
|
WorkerProtocol {
|
||||||
worker_id: SubscriptionWorkerId,
|
worker_id: SubscriptionWorkerId,
|
||||||
@@ -660,9 +662,17 @@ impl SubscriptionEventPayload {
|
|||||||
pub fn validate(&self) -> Result<(), SubscriptionValidationError> {
|
pub fn validate(&self) -> Result<(), SubscriptionValidationError> {
|
||||||
match self {
|
match self {
|
||||||
Self::WorkerUpserted { worker } => worker.validate(),
|
Self::WorkerUpserted { worker } => worker.validate(),
|
||||||
Self::WorkerRemoved { worker_id } | Self::WorkerProtocol { worker_id, .. } => {
|
Self::WorkerRemoved {
|
||||||
worker_id.validate()
|
worker_id,
|
||||||
|
runtime_id,
|
||||||
|
} => {
|
||||||
|
worker_id.validate()?;
|
||||||
|
if let Some(runtime_id) = runtime_id {
|
||||||
|
validate_identifier("runtime_id", runtime_id, MAX_RESOURCE_ID_BYTES)?;
|
||||||
}
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
Self::WorkerProtocol { worker_id, .. } => worker_id.validate(),
|
||||||
Self::WorkdirUpserted { workdir } => workdir.validate(),
|
Self::WorkdirUpserted { workdir } => workdir.validate(),
|
||||||
Self::WorkdirRemoved {
|
Self::WorkdirRemoved {
|
||||||
working_directory_id,
|
working_directory_id,
|
||||||
@@ -688,7 +698,7 @@ impl SubscriptionEventPayload {
|
|||||||
) if worker_ids.contains(&worker.worker_id) => Ok(()),
|
) if worker_ids.contains(&worker.worker_id) => Ok(()),
|
||||||
(
|
(
|
||||||
EventSubscriptionSelector::WorkerLifecycle { worker_ids },
|
EventSubscriptionSelector::WorkerLifecycle { worker_ids },
|
||||||
Self::WorkerRemoved { worker_id },
|
Self::WorkerRemoved { worker_id, .. },
|
||||||
) if worker_ids.contains(worker_id) => Ok(()),
|
) if worker_ids.contains(worker_id) => Ok(()),
|
||||||
(
|
(
|
||||||
EventSubscriptionSelector::WorkerProtocol {
|
EventSubscriptionSelector::WorkerProtocol {
|
||||||
@@ -708,7 +718,7 @@ impl SubscriptionEventPayload {
|
|||||||
}),
|
}),
|
||||||
(
|
(
|
||||||
EventSubscriptionSelector::WorkerLifecycle { .. },
|
EventSubscriptionSelector::WorkerLifecycle { .. },
|
||||||
Self::WorkerRemoved { worker_id },
|
Self::WorkerRemoved { worker_id, .. },
|
||||||
) => Err(SubscriptionValidationError::UnselectedWorker {
|
) => Err(SubscriptionValidationError::UnselectedWorker {
|
||||||
worker_id: worker_id.to_string(),
|
worker_id: worker_id.to_string(),
|
||||||
}),
|
}),
|
||||||
@@ -721,7 +731,7 @@ fn validate_workers(workers: &[SubscriptionWorker]) -> Result<(), SubscriptionVa
|
|||||||
let mut seen = HashSet::with_capacity(workers.len());
|
let mut seen = HashSet::with_capacity(workers.len());
|
||||||
for worker in workers {
|
for worker in workers {
|
||||||
worker.validate()?;
|
worker.validate()?;
|
||||||
if !seen.insert(&worker.worker_id) {
|
if !seen.insert((worker.runtime_id.as_deref(), &worker.worker_id)) {
|
||||||
return Err(SubscriptionValidationError::DuplicateWorkerId {
|
return Err(SubscriptionValidationError::DuplicateWorkerId {
|
||||||
worker_id: worker.worker_id.to_string(),
|
worker_id: worker.worker_id.to_string(),
|
||||||
});
|
});
|
||||||
@@ -868,6 +878,7 @@ mod tests {
|
|||||||
subject_revision: 8,
|
subject_revision: 8,
|
||||||
payload: SubscriptionEventPayload::WorkerRemoved {
|
payload: SubscriptionEventPayload::WorkerRemoved {
|
||||||
worker_id: worker_id("worker-2"),
|
worker_id: worker_id("worker-2"),
|
||||||
|
runtime_id: None,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
assert!(matches!(
|
assert!(matches!(
|
||||||
@@ -880,6 +891,19 @@ mod tests {
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn workspace_snapshot_allows_equal_local_worker_ids_from_distinct_runtimes() {
|
||||||
|
let mut first = worker("1");
|
||||||
|
first.runtime_id = Some("runtime-a".to_string());
|
||||||
|
let mut second = worker("1");
|
||||||
|
second.runtime_id = Some("runtime-b".to_string());
|
||||||
|
SubscriptionSnapshot::Workers {
|
||||||
|
workers: vec![first, second],
|
||||||
|
}
|
||||||
|
.validate_for_selector(&EventSubscriptionSelector::WorkspaceWorkers)
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn subscription_closed_is_a_typed_server_event() {
|
fn subscription_closed_is_a_typed_server_event() {
|
||||||
let frame = SubscriptionFrame::new(SubscriptionFramePayload::Event(
|
let frame = SubscriptionFrame::new(SubscriptionFramePayload::Event(
|
||||||
|
|||||||
@@ -2177,6 +2177,7 @@ impl RuntimeState {
|
|||||||
subject_revision,
|
subject_revision,
|
||||||
payload: SubscriptionEventPayload::WorkerRemoved {
|
payload: SubscriptionEventPayload::WorkerRemoved {
|
||||||
worker_id: worker_id.clone(),
|
worker_id: worker_id.clone(),
|
||||||
|
runtime_id: None,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ pub mod runtime_subscription;
|
|||||||
pub mod server;
|
pub mod server;
|
||||||
pub mod skills;
|
pub mod skills;
|
||||||
pub mod store;
|
pub mod store;
|
||||||
|
mod workspace_subscription;
|
||||||
|
|
||||||
pub use authority::{
|
pub use authority::{
|
||||||
MemoryAuthority, MemoryDocument, MemoryStagingEntry, MemoryStagingResolution,
|
MemoryAuthority, MemoryDocument, MemoryStagingEntry, MemoryStagingResolution,
|
||||||
|
|||||||
@@ -242,6 +242,7 @@ impl RuntimeSubscriptionBroker {
|
|||||||
})?;
|
})?;
|
||||||
let downstream_id = self.next_downstream.fetch_add(1, Ordering::Relaxed);
|
let downstream_id = self.next_downstream.fetch_add(1, Ordering::Relaxed);
|
||||||
let (events, receiver) = mpsc::channel(DOWNSTREAM_QUEUE_CAPACITY);
|
let (events, receiver) = mpsc::channel(DOWNSTREAM_QUEUE_CAPACITY);
|
||||||
|
let initial_events = events.clone();
|
||||||
registration
|
registration
|
||||||
.commands
|
.commands
|
||||||
.send(Command::Subscribe {
|
.send(Command::Subscribe {
|
||||||
@@ -250,6 +251,17 @@ impl RuntimeSubscriptionBroker {
|
|||||||
events,
|
events,
|
||||||
})
|
})
|
||||||
.map_err(|_| RuntimeSubscriptionBrokerError::Closed)?;
|
.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 {
|
Ok(BrokerSubscription {
|
||||||
downstream_id,
|
downstream_id,
|
||||||
runtime_id: runtime_id.to_string(),
|
runtime_id: runtime_id.to_string(),
|
||||||
@@ -291,6 +303,7 @@ impl SelectorState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
struct State {
|
struct State {
|
||||||
|
runtime_id: String,
|
||||||
generation: u64,
|
generation: u64,
|
||||||
next_request: u64,
|
next_request: u64,
|
||||||
selectors: HashMap<EventSubscriptionSelector, SelectorState>,
|
selectors: HashMap<EventSubscriptionSelector, SelectorState>,
|
||||||
@@ -299,8 +312,9 @@ struct State {
|
|||||||
upstream_index: HashMap<SubscriptionId, EventSubscriptionSelector>,
|
upstream_index: HashMap<SubscriptionId, EventSubscriptionSelector>,
|
||||||
}
|
}
|
||||||
impl State {
|
impl State {
|
||||||
fn new(generation: u64) -> Self {
|
fn new(runtime_id: String, generation: u64) -> Self {
|
||||||
Self {
|
Self {
|
||||||
|
runtime_id,
|
||||||
generation,
|
generation,
|
||||||
next_request: 1,
|
next_request: 1,
|
||||||
selectors: HashMap::new(),
|
selectors: HashMap::new(),
|
||||||
@@ -433,9 +447,18 @@ fn project_payload_runtime(
|
|||||||
mut payload: SubscriptionEventPayload,
|
mut payload: SubscriptionEventPayload,
|
||||||
runtime_id: &str,
|
runtime_id: &str,
|
||||||
) -> SubscriptionEventPayload {
|
) -> SubscriptionEventPayload {
|
||||||
if let SubscriptionEventPayload::WorkerUpserted { worker } = &mut payload {
|
match &mut payload {
|
||||||
|
SubscriptionEventPayload::WorkerUpserted { worker } => {
|
||||||
worker.runtime_id = Some(runtime_id.to_string());
|
worker.runtime_id = Some(runtime_id.to_string());
|
||||||
}
|
}
|
||||||
|
SubscriptionEventPayload::WorkerRemoved {
|
||||||
|
runtime_id: projected_runtime_id,
|
||||||
|
..
|
||||||
|
} => {
|
||||||
|
*projected_runtime_id = Some(runtime_id.to_string());
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
payload
|
payload
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -446,7 +469,7 @@ async fn run_connection(
|
|||||||
mut commands: mpsc::UnboundedReceiver<Command>,
|
mut commands: mpsc::UnboundedReceiver<Command>,
|
||||||
status: Arc<RwLock<RuntimeSubscriptionBrokerStatus>>,
|
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;
|
let mut disconnect_notified = false;
|
||||||
loop {
|
loop {
|
||||||
update_status(&status, &state, false);
|
update_status(&status, &state, false);
|
||||||
@@ -653,6 +676,7 @@ async fn handle_frame(
|
|||||||
if state.pending.remove(&request_id) != Some(selector.clone()) {
|
if state.pending.remove(&request_id) != Some(selector.clone()) {
|
||||||
return Err(());
|
return Err(());
|
||||||
}
|
}
|
||||||
|
let snapshot = project_snapshot_runtime(snapshot, &state.runtime_id);
|
||||||
let entry = state.selectors.get_mut(&selector).ok_or(())?;
|
let entry = state.selectors.get_mut(&selector).ok_or(())?;
|
||||||
entry.pending = false;
|
entry.pending = false;
|
||||||
entry.upstream_id = Some(subscription_id.clone());
|
entry.upstream_id = Some(subscription_id.clone());
|
||||||
@@ -704,6 +728,7 @@ async fn handle_frame(
|
|||||||
.cloned()
|
.cloned()
|
||||||
.ok_or(())?;
|
.ok_or(())?;
|
||||||
payload.validate_for_selector(&selector).map_err(|_| ())?;
|
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(())?;
|
let entry = state.selectors.get_mut(&selector).ok_or(())?;
|
||||||
if let Some(subject) = event_subject(&payload) {
|
if let Some(subject) = event_subject(&payload) {
|
||||||
let revision = entry.revisions.entry(subject).or_insert(0);
|
let revision = entry.revisions.entry(subject).or_insert(0);
|
||||||
@@ -793,7 +818,16 @@ fn snapshot_revisions(snapshot: &SubscriptionSnapshot) -> HashMap<String, u64> {
|
|||||||
match snapshot {
|
match snapshot {
|
||||||
SubscriptionSnapshot::Workers { workers } => workers
|
SubscriptionSnapshot::Workers { workers } => workers
|
||||||
.iter()
|
.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(),
|
.collect(),
|
||||||
SubscriptionSnapshot::WorkerProtocol { worker_id, .. } => {
|
SubscriptionSnapshot::WorkerProtocol { worker_id, .. } => {
|
||||||
HashMap::from([(worker_id.to_string(), 0)])
|
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> {
|
fn event_subject(payload: &SubscriptionEventPayload) -> Option<String> {
|
||||||
Some(match payload {
|
Some(match payload {
|
||||||
SubscriptionEventPayload::WorkerUpserted { worker } => worker.worker_id.to_string(),
|
SubscriptionEventPayload::WorkerUpserted { worker } => format!(
|
||||||
SubscriptionEventPayload::WorkerRemoved { worker_id }
|
"{}:{}",
|
||||||
| SubscriptionEventPayload::WorkerProtocol { worker_id, .. } => worker_id.to_string(),
|
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 } => {
|
SubscriptionEventPayload::WorkdirUpserted { workdir } => {
|
||||||
workdir.working_directory_id.to_string()
|
workdir.working_directory_id.to_string()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -867,6 +867,10 @@ pub fn build_router(api: WorkspaceApi) -> Router {
|
|||||||
"/api/w/{workspace_id}/workers",
|
"/api/w/{workspace_id}/workers",
|
||||||
get(scoped_list_workers).post(scoped_create_workspace_worker),
|
get(scoped_list_workers).post(scoped_create_workspace_worker),
|
||||||
)
|
)
|
||||||
|
.route(
|
||||||
|
"/api/w/{workspace_id}/protocol/ws",
|
||||||
|
get(scoped_workspace_protocol_ws),
|
||||||
|
)
|
||||||
.route(
|
.route(
|
||||||
"/api/workers/launch-options",
|
"/api/workers/launch-options",
|
||||||
get(get_worker_launch_options),
|
get(get_worker_launch_options),
|
||||||
@@ -3779,6 +3783,38 @@ async fn scoped_list_runtimes(
|
|||||||
list_runtimes(State(api)).await
|
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(
|
async fn scoped_list_workers(
|
||||||
State(api): State<WorkspaceApi>,
|
State(api): State<WorkspaceApi>,
|
||||||
AxumPath(path): AxumPath<ScopedWorkspacePath>,
|
AxumPath(path): AxumPath<ScopedWorkspacePath>,
|
||||||
@@ -8716,6 +8752,7 @@ mod tests {
|
|||||||
use std::{fs, sync::Arc};
|
use std::{fs, sync::Arc};
|
||||||
use tokio_tungstenite::connect_async;
|
use tokio_tungstenite::connect_async;
|
||||||
use tokio_tungstenite::tungstenite::Message;
|
use tokio_tungstenite::tungstenite::Message;
|
||||||
|
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||||
use tower::ServiceExt;
|
use tower::ServiceExt;
|
||||||
use worker_runtime::resource::BackendResourceClient;
|
use worker_runtime::resource::BackendResourceClient;
|
||||||
use worker_runtime::working_directory::WorkingDirectoryMaterializer;
|
use worker_runtime::working_directory::WorkingDirectoryMaterializer;
|
||||||
@@ -8725,8 +8762,9 @@ mod tests {
|
|||||||
WorkerSpawnIntent,
|
WorkerSpawnIntent,
|
||||||
};
|
};
|
||||||
use crate::store::{
|
use crate::store::{
|
||||||
MemoryDocumentRecord, MemoryStagingRecord, ObjectiveRecord, ObjectiveResourceRecord,
|
AccountRecord, MemoryDocumentRecord, MemoryStagingRecord, ObjectiveRecord,
|
||||||
ObjectiveTicketLinkRecord, SqliteWorkspaceStore, WorkspaceRecord,
|
ObjectiveResourceRecord, ObjectiveTicketLinkRecord, SqliteWorkspaceStore, UserRecord,
|
||||||
|
WorkspaceRecord,
|
||||||
};
|
};
|
||||||
|
|
||||||
const TEST_WORKSPACE_ID: &str = "0192f0e8-4d84-7d6e-a000-000000000001";
|
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]
|
#[tokio::test]
|
||||||
async fn passkey_registration_rejects_unverified_credential_response() {
|
async fn passkey_registration_rejects_unverified_credential_response() {
|
||||||
let dir = tempfile::tempdir().unwrap();
|
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,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -128,7 +128,7 @@ export type SubscriptionWorkdir = { working_directory_id: SubscriptionWorkdirId,
|
|||||||
|
|
||||||
export type SubscriptionSnapshot = { "topic": "workers", "data": { workers: Array<SubscriptionWorker>, } } | { "topic": "worker_protocol", "data": { worker_id: SubscriptionWorkerId, events: Array<Event>, } } | { "topic": "workspace_workdirs", "data": { workdirs: Array<SubscriptionWorkdir>, } };
|
export type SubscriptionSnapshot = { "topic": "workers", "data": { workers: Array<SubscriptionWorker>, } } | { "topic": "worker_protocol", "data": { worker_id: SubscriptionWorkerId, events: Array<Event>, } } | { "topic": "workspace_workdirs", "data": { workdirs: Array<SubscriptionWorkdir>, } };
|
||||||
|
|
||||||
export type SubscriptionEventPayload = { "event": "worker_upserted", "data": { worker: SubscriptionWorker, } } | { "event": "worker_removed", "data": { worker_id: SubscriptionWorkerId, } } | { "event": "worker_protocol", "data": { worker_id: SubscriptionWorkerId, event: Event, } } | { "event": "workdir_upserted", "data": { workdir: SubscriptionWorkdir, } } | { "event": "workdir_removed", "data": { working_directory_id: SubscriptionWorkdirId, } };
|
export type SubscriptionEventPayload = { "event": "worker_upserted", "data": { worker: SubscriptionWorker, } } | { "event": "worker_removed", "data": { worker_id: SubscriptionWorkerId, runtime_id?: string | null, } } | { "event": "worker_protocol", "data": { worker_id: SubscriptionWorkerId, event: Event, } } | { "event": "workdir_upserted", "data": { workdir: SubscriptionWorkdir, } } | { "event": "workdir_removed", "data": { working_directory_id: SubscriptionWorkdirId, } };
|
||||||
|
|
||||||
export type SubscriptionRejectionCode = "invalid_request" | "unsupported_protocol_version" | "unsupported_selector" | "unauthorized" | "resource_not_found" | "capacity_exceeded" | "internal";
|
export type SubscriptionRejectionCode = "invalid_request" | "unsupported_protocol_version" | "unsupported_selector" | "unauthorized" | "resource_not_found" | "capacity_exceeded" | "internal";
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user