fix: unify worker protocol websocket

This commit is contained in:
2026-07-21 20:29:30 +09:00
parent 70a26a3042
commit 8841d063be
10 changed files with 377 additions and 1046 deletions
+148 -342
View File
@@ -154,12 +154,10 @@ pub fn runtime_http_router(runtime: Runtime, local_token: Option<String>) -> Rou
.route("/v1/workers/{worker_id}/cancel", post(cancel_worker));
#[cfg(feature = "ws-server")]
let router = router
.route("/v1/workers/{worker_id}/events/ws", get(worker_events_ws))
.route(
"/v1/workers/{worker_id}/protocol/ws",
get(worker_protocol_ws),
);
let router = router.route(
"/v1/workers/{worker_id}/protocol/ws",
get(worker_protocol_ws),
);
router
.with_state(state.clone())
@@ -277,88 +275,12 @@ pub struct RuntimeHttpErrorDetail {
pub message: String,
}
/// Runtime-owned WebSocket frame for worker-scoped observation.
#[cfg(feature = "ws-server")]
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum RuntimeWorkerEventWsFrame {
Event {
envelope: RuntimeWorkerEventWsEnvelope,
},
Diagnostic {
diagnostic: RuntimeWorkerEventWsDiagnostic,
},
}
/// Runtime-local protocol event envelope.
#[cfg(feature = "ws-server")]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct RuntimeWorkerEventWsEnvelope {
pub cursor: String,
pub event_id: String,
pub worker_id: WorkerId,
pub payload: protocol::Event,
}
/// Runtime-local observation diagnostic.
#[cfg(feature = "ws-server")]
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct RuntimeWorkerEventWsDiagnostic {
pub code: String,
pub message: String,
}
#[cfg(feature = "ws-server")]
#[derive(Clone, Debug, Default, Deserialize)]
struct RuntimeWorkerEventsWsQuery {
cursor: Option<String>,
}
#[cfg(feature = "ws-server")]
impl RuntimeWorkerEventWsFrame {
fn event(
cursor: String,
event_id: String,
worker_id: WorkerId,
payload: protocol::Event,
) -> Self {
Self::Event {
envelope: RuntimeWorkerEventWsEnvelope {
cursor,
event_id,
worker_id,
payload,
},
}
}
fn diagnostic(code: impl Into<String>, message: impl Into<String>) -> Self {
Self::Diagnostic {
diagnostic: RuntimeWorkerEventWsDiagnostic {
code: code.into(),
message: message.into(),
},
}
}
}
#[cfg(feature = "ws-server")]
async fn send_ws_frame(socket: &mut WebSocket, frame: &RuntimeWorkerEventWsFrame) -> bool {
match serde_json::to_string(frame) {
Ok(text) => socket.send(WsMessage::Text(text.into())).await.is_ok(),
Err(error) => {
let fallback = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.serialize_failed",
format!("failed to serialize observation frame: {error}"),
);
let Ok(text) = serde_json::to_string(&fallback) else {
return false;
};
socket.send(WsMessage::Text(text.into())).await.is_ok()
}
}
}
type RestResult<T> = Result<Json<T>, RuntimeHttpRestError>;
async fn get_runtime(
@@ -515,6 +437,7 @@ async fn create_worker(
async fn worker_protocol_ws(
State(state): State<RuntimeHttpState>,
Path(worker_id): Path<String>,
Query(query): Query<RuntimeWorkerEventsWsQuery>,
ws: WebSocketUpgrade,
) -> Result<Response, RuntimeHttpRestError> {
let worker_ref = worker_ref_for(&state.runtime, worker_id)?;
@@ -523,7 +446,9 @@ async fn worker_protocol_ws(
.worker_detail(&worker_ref)
.map_err(RuntimeHttpRestError::runtime)?;
Ok(ws
.on_upgrade(move |socket| worker_protocol_ws_session(state.runtime, worker_ref, socket))
.on_upgrade(move |socket| {
worker_protocol_ws_session(state.runtime, worker_ref, query, socket)
})
.into_response())
}
@@ -531,45 +456,130 @@ async fn worker_protocol_ws(
async fn worker_protocol_ws_session(
runtime: Runtime,
worker_ref: WorkerRef,
query: RuntimeWorkerEventsWsQuery,
mut socket: WebSocket,
) {
while let Some(frame) = socket.next().await {
match frame {
Ok(WsMessage::Text(text)) => match serde_json::from_str::<protocol::Method>(&text) {
Ok(method) => match runtime.send_protocol_method(&worker_ref, method) {
Ok(events) => {
for event in events {
let mut cursor = match query.cursor.as_deref() {
Some(raw) => match WorkerObservationCursor::decode(raw) {
Some(cursor) => cursor,
None => {
let event =
protocol_error_event(format!("malformed worker observation cursor: {raw}"));
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
},
None => match runtime.worker_observation_cursor_now(&worker_ref) {
Ok(cursor) => cursor,
Err(error) => {
let event = protocol_error_event(error.to_string());
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
},
};
let mut receiver = match runtime.subscribe_worker_observation() {
Ok(receiver) => receiver,
Err(error) => {
let event =
protocol_error_event(format!("runtime observation bus unavailable: {error}"));
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
};
let snapshot = match runtime.worker_observation_snapshot(&worker_ref) {
Ok(snapshot) => snapshot,
Err(error) => {
let event = protocol_error_event(error.to_string());
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
};
if !send_protocol_event(&mut socket, &snapshot).await {
return;
}
match runtime.read_worker_observation_events(&worker_ref, cursor) {
Ok(backlog) => {
for event in backlog {
cursor = WorkerObservationCursor::new(event.sequence);
if !send_protocol_event(&mut socket, &event.payload).await {
return;
}
}
}
Err(error) => {
let event = protocol_error_event(error.to_string());
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
}
loop {
tokio::select! {
inbound = socket.next() => {
match inbound {
Some(Ok(WsMessage::Text(text))) => match serde_json::from_str::<protocol::Method>(&text) {
Ok(method) => match runtime.send_protocol_method(&worker_ref, method) {
Ok(events) => {
for event in events {
if !send_protocol_event(&mut socket, &event).await {
return;
}
}
}
Err(error) => {
let event = protocol_error_event(error.to_string());
if !send_protocol_event(&mut socket, &event).await {
return;
}
}
},
Err(error) => {
let event = protocol_error_event(format!(
"malformed protocol method frame: {error}"
));
if !send_protocol_event(&mut socket, &event).await {
return;
}
}
}
Err(error) => {
let event = protocol_error_event(error.to_string());
if !send_protocol_event(&mut socket, &event).await {
},
Some(Ok(WsMessage::Close(_))) | None => return,
Some(Ok(WsMessage::Ping(payload))) => {
if socket.send(WsMessage::Pong(payload)).await.is_err() {
return;
}
}
},
Err(error) => {
let event =
protocol_error_event(format!("malformed protocol method frame: {error}"));
if !send_protocol_event(&mut socket, &event).await {
Some(Ok(WsMessage::Pong(_))) | Some(Ok(WsMessage::Binary(_))) => {}
Some(Err(error)) => {
let event = protocol_error_event(format!("protocol WebSocket error: {error}"));
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
}
},
Ok(WsMessage::Close(_)) => return,
Ok(WsMessage::Ping(payload)) => {
if socket.send(WsMessage::Pong(payload)).await.is_err() {
return;
}
}
Ok(WsMessage::Pong(_)) | Ok(WsMessage::Binary(_)) => {}
Err(error) => {
let event = protocol_error_event(format!("protocol WebSocket error: {error}"));
let _ = send_protocol_event(&mut socket, &event).await;
return;
event = receiver.recv() => {
match event {
Ok(event) if event.worker_ref == worker_ref && event.sequence > cursor.sequence => {
cursor = WorkerObservationCursor::new(event.sequence);
if !send_protocol_event(&mut socket, &event.payload).await {
return;
}
}
Ok(_) => {}
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
let event = protocol_error_event("runtime observation backlog was overrun");
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => {
let event = protocol_error_event("runtime observation bus closed");
let _ = send_protocol_event(&mut socket, &event).await;
return;
}
}
}
}
}
@@ -599,182 +609,6 @@ fn protocol_error_event(message: impl Into<String>) -> protocol::Event {
}
}
#[cfg(feature = "ws-server")]
async fn worker_events_ws(
State(state): State<RuntimeHttpState>,
Path(worker_id): Path<String>,
Query(query): Query<RuntimeWorkerEventsWsQuery>,
ws: WebSocketUpgrade,
) -> Result<Response, RuntimeHttpRestError> {
let worker_ref = worker_ref_for(&state.runtime, worker_id)?;
state
.runtime
.worker_detail(&worker_ref)
.map_err(RuntimeHttpRestError::runtime)?;
Ok(ws
.on_upgrade(move |socket| {
worker_events_ws_session(state.runtime, worker_ref, query, socket)
})
.into_response())
}
#[cfg(feature = "ws-server")]
async fn worker_events_ws_session(
runtime: Runtime,
worker_ref: WorkerRef,
query: RuntimeWorkerEventsWsQuery,
mut socket: WebSocket,
) {
let mut cursor = match query.cursor.as_deref() {
Some(raw) => match WorkerObservationCursor::decode(raw) {
Some(cursor) => cursor,
None => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.cursor_malformed",
format!("malformed worker observation cursor: {raw}"),
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
},
None => match runtime.worker_observation_cursor_now(&worker_ref) {
Ok(cursor) => cursor,
Err(error) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.worker_not_found",
error.to_string(),
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
},
};
let mut receiver = match runtime.subscribe_worker_observation() {
Ok(receiver) => receiver,
Err(error) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.unavailable",
format!("runtime observation bus unavailable: {error}"),
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
};
let snapshot = match runtime.worker_observation_snapshot(&worker_ref) {
Ok(snapshot) => snapshot,
Err(error) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.worker_not_found",
error.to_string(),
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
};
let snapshot_cursor = cursor.encode();
let snapshot_frame = RuntimeWorkerEventWsFrame::event(
snapshot_cursor.clone(),
format!("snapshot:{snapshot_cursor}"),
worker_ref.worker_id.clone(),
snapshot,
);
if !send_ws_frame(&mut socket, &snapshot_frame).await {
return;
}
match runtime.read_worker_observation_events(&worker_ref, cursor) {
Ok(backlog) => {
for event in backlog {
cursor = WorkerObservationCursor::new(event.sequence);
let frame = RuntimeWorkerEventWsFrame::event(
event.cursor,
event.event_id,
event.worker_ref.worker_id,
event.payload,
);
if !send_ws_frame(&mut socket, &frame).await {
return;
}
}
}
Err(error) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.cursor_unknown_or_expired",
error.to_string(),
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
}
loop {
tokio::select! {
inbound = socket.next() => {
match inbound {
Some(Ok(WsMessage::Close(_))) | None => return,
Some(Ok(WsMessage::Ping(payload))) => {
if socket.send(WsMessage::Pong(payload)).await.is_err() {
return;
}
}
Some(Ok(WsMessage::Pong(_))) => {}
Some(Ok(_)) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.observation_only",
"runtime worker event WebSocket is observation-only",
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
Some(Err(error)) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.websocket_error",
format!("runtime WebSocket receive error: {error}"),
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
}
}
event = receiver.recv() => {
match event {
Ok(event) if event.worker_ref == worker_ref && event.sequence > cursor.sequence => {
cursor = WorkerObservationCursor::new(event.sequence);
let frame = RuntimeWorkerEventWsFrame::event(
event.cursor,
event.event_id,
event.worker_ref.worker_id,
event.payload,
);
if !send_ws_frame(&mut socket, &frame).await {
return;
}
}
Ok(_) => {}
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.cursor_expired",
"runtime observation backlog was overrun",
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => {
let frame = RuntimeWorkerEventWsFrame::diagnostic(
"runtime.upstream_closed",
"runtime observation bus closed",
);
let _ = send_ws_frame(&mut socket, &frame).await;
return;
}
}
}
}
}
}
async fn send_worker_input(
State(state): State<RuntimeHttpState>,
Path(worker_id): Path<String>,
@@ -1409,7 +1243,7 @@ mod ws_tests {
runtime,
worker.worker_ref.clone(),
format!(
"ws://{addr}/v1/workers/{}/events/ws",
"ws://{addr}/v1/workers/{}/protocol/ws",
worker.worker_ref.worker_id
),
)
@@ -1419,7 +1253,7 @@ mod ws_tests {
stream: &mut tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
) -> RuntimeWorkerEventWsFrame {
) -> protocol::Event {
let message = stream.next().await.unwrap().unwrap();
let Message::Text(text) = message else {
panic!("expected text frame");
@@ -1428,21 +1262,16 @@ mod ws_tests {
}
#[tokio::test]
async fn runtime_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 (mut stream, _) = connect_async(&url).await.unwrap();
match next_frame(&mut stream).await {
RuntimeWorkerEventWsFrame::Event { envelope } => {
assert_eq!(envelope.worker_id, worker_ref.worker_id);
assert!(matches!(envelope.payload, protocol::Event::Snapshot { .. }));
}
RuntimeWorkerEventWsFrame::Diagnostic { diagnostic } => {
panic!("unexpected diagnostic: {diagnostic:?}");
}
}
assert!(matches!(
next_frame(&mut stream).await,
protocol::Event::Snapshot { .. }
));
let stored = runtime
runtime
.observe_worker_event(
&worker_ref,
protocol::Event::TextDelta {
@@ -1450,23 +1279,14 @@ mod ws_tests {
},
)
.unwrap();
match next_frame(&mut stream).await {
RuntimeWorkerEventWsFrame::Event { envelope } => {
assert_eq!(envelope.worker_id, worker_ref.worker_id);
assert_eq!(envelope.cursor, stored.cursor);
assert!(matches!(
envelope.payload,
protocol::Event::TextDelta { .. }
));
}
RuntimeWorkerEventWsFrame::Diagnostic { diagnostic } => {
panic!("unexpected diagnostic: {diagnostic:?}");
}
}
assert!(matches!(
next_frame(&mut stream).await,
protocol::Event::TextDelta { .. }
));
}
#[tokio::test]
async fn runtime_ws_cursor_resume_is_duplicate_safe_and_filters_workers() {
async fn protocol_ws_cursor_resume_is_duplicate_safe_and_filters_workers() {
let (runtime, worker_ref, url) = spawn_runtime_server().await;
let other = runtime.create_worker(ws_create_request()).unwrap();
let first = runtime
@@ -1481,7 +1301,7 @@ mod ws_tests {
.observe_worker_event(
&other.worker_ref,
protocol::Event::TextDelta {
text: "started".into(),
text: "other".into(),
},
)
.unwrap();
@@ -1491,10 +1311,10 @@ mod ws_tests {
.unwrap();
assert!(matches!(
next_frame(&mut stream).await,
RuntimeWorkerEventWsFrame::Event { envelope } if matches!(envelope.payload, protocol::Event::Snapshot { .. })
protocol::Event::Snapshot { .. }
));
let second = runtime
runtime
.observe_worker_event(
&worker_ref,
protocol::Event::TextDone {
@@ -1502,41 +1322,27 @@ mod ws_tests {
},
)
.unwrap();
match next_frame(&mut stream).await {
RuntimeWorkerEventWsFrame::Event { envelope } => {
assert_eq!(envelope.cursor, second.cursor);
assert_ne!(envelope.cursor, first.cursor);
assert!(matches!(envelope.payload, protocol::Event::TextDone { .. }));
}
RuntimeWorkerEventWsFrame::Diagnostic { diagnostic } => {
panic!("unexpected diagnostic: {diagnostic:?}");
}
}
assert!(matches!(
next_frame(&mut stream).await,
protocol::Event::TextDone { .. }
));
}
#[tokio::test]
async fn runtime_ws_reports_malformed_cursor_and_observation_only_input() {
async fn protocol_ws_reports_malformed_cursor_and_method_frame() {
let (_runtime, _worker_ref, url) = spawn_runtime_server().await;
let (mut malformed, _) = connect_async(format!("{url}?cursor=bad")).await.unwrap();
match next_frame(&mut malformed).await {
RuntimeWorkerEventWsFrame::Diagnostic { diagnostic } => {
assert_eq!(diagnostic.code, "runtime.cursor_malformed");
}
RuntimeWorkerEventWsFrame::Event { envelope } => {
panic!("unexpected event: {envelope:?}");
}
}
assert!(matches!(
next_frame(&mut malformed).await,
protocol::Event::Error { .. }
));
let (mut stream, _) = connect_async(&url).await.unwrap();
let _ = next_frame(&mut stream).await;
stream.send(Message::Text("{}".into())).await.unwrap();
match next_frame(&mut stream).await {
RuntimeWorkerEventWsFrame::Diagnostic { diagnostic } => {
assert_eq!(diagnostic.code, "runtime.observation_only");
}
RuntimeWorkerEventWsFrame::Event { envelope } => {
panic!("unexpected event: {envelope:?}");
}
}
assert!(matches!(
next_frame(&mut stream).await,
protocol::Event::Error { .. }
));
}
}