feat: add worker observation websocket proxy
This commit is contained in:
@@ -14,11 +14,18 @@ required-features = ["http-server"]
|
||||
default = []
|
||||
fs-store = ["dep:serde_json"]
|
||||
http-server = ["dep:axum", "dep:serde_json", "dep:tokio", "dep:tower"]
|
||||
ws-server = ["http-server", "axum/ws", "dep:futures", "dep:protocol", "tokio/sync"]
|
||||
|
||||
[dependencies]
|
||||
axum = { workspace = true, optional = true }
|
||||
futures = { workspace = true, optional = true }
|
||||
protocol = { workspace = true, optional = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
serde_json = { workspace = true, optional = true }
|
||||
thiserror = { workspace = true }
|
||||
tokio = { workspace = true, features = ["net", "rt"], optional = true }
|
||||
tower = { workspace = true, features = ["util"], optional = true }
|
||||
|
||||
[dev-dependencies]
|
||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
|
||||
tokio-tungstenite.workspace = true
|
||||
|
||||
@@ -14,15 +14,21 @@ use crate::fs_store::FsRuntimeStoreOptions;
|
||||
use crate::identity::{RuntimeId, WorkerId, WorkerRef};
|
||||
use crate::interaction::{WorkerInput, WorkerInteractionAck};
|
||||
use crate::management::{RuntimeLimits, RuntimeOptions, RuntimeSummary};
|
||||
#[cfg(feature = "ws-server")]
|
||||
use crate::observation::WorkerObservationCursor;
|
||||
use crate::observation::{TranscriptProjection, TranscriptQuery};
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::extract::rejection::{JsonRejection, QueryRejection};
|
||||
#[cfg(feature = "ws-server")]
|
||||
use axum::extract::ws::{Message as WsMessage, WebSocket, WebSocketUpgrade};
|
||||
use axum::extract::{Path, Query, State};
|
||||
use axum::http::{Request, StatusCode, header};
|
||||
use axum::middleware::{self, Next};
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use axum::routing::{get, post};
|
||||
use axum::{Json, Router};
|
||||
#[cfg(feature = "ws-server")]
|
||||
use futures::StreamExt;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
use std::net::SocketAddr;
|
||||
@@ -157,7 +163,7 @@ pub fn runtime_http_router(runtime: Runtime, local_token: Option<String>) -> Rou
|
||||
local_token: local_token.map(Arc::<str>::from),
|
||||
};
|
||||
|
||||
Router::new()
|
||||
let router = Router::new()
|
||||
.route("/v1/runtime", get(get_runtime))
|
||||
.route("/v1/workers", get(list_workers).post(create_worker))
|
||||
.route("/v1/workers/{worker_id}", get(get_worker))
|
||||
@@ -167,7 +173,12 @@ pub fn runtime_http_router(runtime: Runtime, local_token: Option<String>) -> Rou
|
||||
.route(
|
||||
"/v1/workers/{worker_id}/transcript",
|
||||
get(get_worker_transcript),
|
||||
)
|
||||
);
|
||||
|
||||
#[cfg(feature = "ws-server")]
|
||||
let router = router.route("/v1/workers/{worker_id}/events/ws", get(worker_events_ws));
|
||||
|
||||
router
|
||||
.with_state(state.clone())
|
||||
.layer(middleware::from_fn_with_state(state, require_local_token))
|
||||
}
|
||||
@@ -255,6 +266,43 @@ 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>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
struct RuntimeHttpTranscriptQuery {
|
||||
#[serde(default)]
|
||||
@@ -267,6 +315,51 @@ fn default_transcript_limit() -> usize {
|
||||
256
|
||||
}
|
||||
|
||||
#[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(
|
||||
@@ -313,6 +406,182 @@ async fn create_worker(
|
||||
Ok(Json(RuntimeHttpWorkerResponse { worker }))
|
||||
}
|
||||
|
||||
#[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>,
|
||||
@@ -688,3 +957,159 @@ mod tests {
|
||||
assert!(error.error.message.contains("worker-missing"));
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, feature = "ws-server"))]
|
||||
mod ws_tests {
|
||||
use super::*;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use tokio_tungstenite::connect_async;
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
|
||||
async fn spawn_runtime_server() -> (Runtime, WorkerRef, String) {
|
||||
let runtime = Runtime::new_memory();
|
||||
let worker = runtime
|
||||
.create_worker(CreateWorkerRequest::default())
|
||||
.unwrap();
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
tokio::spawn({
|
||||
let runtime = runtime.clone();
|
||||
async move { serve_runtime_http(runtime, listener, None).await.unwrap() }
|
||||
});
|
||||
(
|
||||
runtime,
|
||||
worker.worker_ref.clone(),
|
||||
format!(
|
||||
"ws://{addr}/v1/workers/{}/events/ws",
|
||||
worker.worker_ref.worker_id
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
async fn next_frame(
|
||||
stream: &mut tokio_tungstenite::WebSocketStream<
|
||||
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
|
||||
>,
|
||||
) -> RuntimeWorkerEventWsFrame {
|
||||
let message = stream.next().await.unwrap().unwrap();
|
||||
let Message::Text(text) = message else {
|
||||
panic!("expected text frame");
|
||||
};
|
||||
serde_json::from_str(&text).unwrap()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_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:?}");
|
||||
}
|
||||
}
|
||||
|
||||
let stored = runtime
|
||||
.observe_worker_event(
|
||||
&worker_ref,
|
||||
protocol::Event::TextDelta {
|
||||
text: "started".into(),
|
||||
},
|
||||
)
|
||||
.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:?}");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_ws_cursor_resume_is_duplicate_safe_and_filters_workers() {
|
||||
let (runtime, worker_ref, url) = spawn_runtime_server().await;
|
||||
let other = runtime
|
||||
.create_worker(CreateWorkerRequest::default())
|
||||
.unwrap();
|
||||
let first = runtime
|
||||
.observe_worker_event(
|
||||
&worker_ref,
|
||||
protocol::Event::TextDelta {
|
||||
text: "started".into(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
runtime
|
||||
.observe_worker_event(
|
||||
&other.worker_ref,
|
||||
protocol::Event::TextDelta {
|
||||
text: "started".into(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let (mut stream, _) = connect_async(format!("{url}?cursor={}", first.cursor))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
next_frame(&mut stream).await,
|
||||
RuntimeWorkerEventWsFrame::Event { envelope } if matches!(envelope.payload, protocol::Event::Snapshot { .. })
|
||||
));
|
||||
|
||||
let second = runtime
|
||||
.observe_worker_event(
|
||||
&worker_ref,
|
||||
protocol::Event::TextDone {
|
||||
text: "done".into(),
|
||||
},
|
||||
)
|
||||
.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:?}");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_ws_reports_malformed_cursor_and_observation_only_input() {
|
||||
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:?}");
|
||||
}
|
||||
}
|
||||
|
||||
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:?}");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -93,3 +93,62 @@ pub struct RuntimeEventBatch {
|
||||
pub events: Vec<RuntimeEvent>,
|
||||
pub has_more: bool,
|
||||
}
|
||||
|
||||
/// Runtime-local cursor for worker-scoped WebSocket observation.
|
||||
#[cfg(feature = "ws-server")]
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
|
||||
pub struct WorkerObservationCursor {
|
||||
pub sequence: u64,
|
||||
}
|
||||
|
||||
#[cfg(feature = "ws-server")]
|
||||
impl WorkerObservationCursor {
|
||||
pub const PREFIX: &'static str = "wo";
|
||||
|
||||
pub fn new(sequence: u64) -> Self {
|
||||
Self { sequence }
|
||||
}
|
||||
|
||||
pub fn zero() -> Self {
|
||||
Self { sequence: 0 }
|
||||
}
|
||||
|
||||
pub fn encode(self) -> String {
|
||||
format!("{}_{:016x}", Self::PREFIX, self.sequence)
|
||||
}
|
||||
|
||||
pub fn decode(value: &str) -> Option<Self> {
|
||||
let encoded = value.strip_prefix("wo_")?;
|
||||
if encoded.len() != 16 {
|
||||
return None;
|
||||
}
|
||||
u64::from_str_radix(encoded, 16)
|
||||
.ok()
|
||||
.map(|sequence| Self { sequence })
|
||||
}
|
||||
}
|
||||
|
||||
/// One protocol event observed from a runtime Worker.
|
||||
#[cfg(feature = "ws-server")]
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub struct WorkerObservationEvent {
|
||||
pub cursor: String,
|
||||
pub event_id: String,
|
||||
pub sequence: u64,
|
||||
pub worker_ref: WorkerRef,
|
||||
pub payload: protocol::Event,
|
||||
}
|
||||
|
||||
#[cfg(feature = "ws-server")]
|
||||
impl WorkerObservationEvent {
|
||||
pub fn new(sequence: u64, worker_ref: WorkerRef, payload: protocol::Event) -> Self {
|
||||
let cursor = WorkerObservationCursor::new(sequence).encode();
|
||||
Self {
|
||||
event_id: cursor.clone(),
|
||||
cursor,
|
||||
sequence,
|
||||
worker_ref,
|
||||
payload,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,9 +16,15 @@ use crate::observation::{
|
||||
EventCursor, EventSubscription, EventSubscriptionMode, RuntimeEvent, RuntimeEventBatch,
|
||||
RuntimeEventKind, TranscriptEntry, TranscriptProjection, TranscriptQuery, TranscriptRole,
|
||||
};
|
||||
#[cfg(feature = "ws-server")]
|
||||
use crate::observation::{WorkerObservationCursor, WorkerObservationEvent};
|
||||
use std::collections::BTreeMap;
|
||||
#[cfg(feature = "ws-server")]
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex, MutexGuard};
|
||||
#[cfg(feature = "ws-server")]
|
||||
use tokio::sync::broadcast;
|
||||
|
||||
static NEXT_RUNTIME_SEQUENCE: AtomicU64 = AtomicU64::new(1);
|
||||
|
||||
@@ -395,6 +401,88 @@ impl Runtime {
|
||||
})
|
||||
}
|
||||
|
||||
/// Cursor pointing after the current worker-scoped protocol observation event.
|
||||
#[cfg(feature = "ws-server")]
|
||||
pub fn worker_observation_cursor_now(
|
||||
&self,
|
||||
worker_ref: &WorkerRef,
|
||||
) -> Result<WorkerObservationCursor, RuntimeError> {
|
||||
let state = self.lock()?;
|
||||
state.ensure_worker_ref(worker_ref)?;
|
||||
let sequence = state
|
||||
.observation_events
|
||||
.iter()
|
||||
.rev()
|
||||
.find(|event| &event.worker_ref == worker_ref)
|
||||
.map(|event| event.sequence)
|
||||
.unwrap_or(0);
|
||||
Ok(WorkerObservationCursor::new(sequence))
|
||||
}
|
||||
|
||||
/// Build the current Worker Snapshot event used as the first observation frame.
|
||||
#[cfg(feature = "ws-server")]
|
||||
pub fn worker_observation_snapshot(
|
||||
&self,
|
||||
worker_ref: &WorkerRef,
|
||||
) -> Result<protocol::Event, RuntimeError> {
|
||||
let state = self.lock()?;
|
||||
let _worker = state.worker(worker_ref)?;
|
||||
Ok(protocol::Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
greeting: protocol::Greeting {
|
||||
worker_name: worker_ref.worker_id.to_string(),
|
||||
cwd: String::new(),
|
||||
provider: "worker-runtime".to_string(),
|
||||
model: "worker-runtime".to_string(),
|
||||
scope_summary: "runtime worker observation".to_string(),
|
||||
tools: Vec::new(),
|
||||
context_window: 0,
|
||||
context_tokens: 0,
|
||||
},
|
||||
status: protocol::WorkerStatus::Idle,
|
||||
in_flight: protocol::InFlightSnapshot { blocks: Vec::new() },
|
||||
})
|
||||
}
|
||||
|
||||
/// Replay retained worker-scoped protocol observation events after a cursor.
|
||||
#[cfg(feature = "ws-server")]
|
||||
pub fn read_worker_observation_events(
|
||||
&self,
|
||||
worker_ref: &WorkerRef,
|
||||
cursor: WorkerObservationCursor,
|
||||
) -> Result<Vec<WorkerObservationEvent>, RuntimeError> {
|
||||
let state = self.lock()?;
|
||||
state.ensure_worker_ref(worker_ref)?;
|
||||
state.validate_worker_observation_cursor(worker_ref, cursor)?;
|
||||
Ok(state
|
||||
.observation_events
|
||||
.iter()
|
||||
.filter(|event| &event.worker_ref == worker_ref && event.sequence > cursor.sequence)
|
||||
.cloned()
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// Subscribe to live protocol observation events.
|
||||
#[cfg(feature = "ws-server")]
|
||||
pub fn subscribe_worker_observation(
|
||||
&self,
|
||||
) -> Result<broadcast::Receiver<WorkerObservationEvent>, RuntimeError> {
|
||||
Ok(self.lock()?.observation_tx.subscribe())
|
||||
}
|
||||
|
||||
/// Append a Worker protocol event to the observation bus.
|
||||
#[cfg(feature = "ws-server")]
|
||||
pub fn observe_worker_event(
|
||||
&self,
|
||||
worker_ref: &WorkerRef,
|
||||
payload: protocol::Event,
|
||||
) -> Result<WorkerObservationEvent, RuntimeError> {
|
||||
let mut state = self.lock()?;
|
||||
state.ensure_worker_ref(worker_ref)?;
|
||||
let event = state.push_worker_observation_event(worker_ref.clone(), payload);
|
||||
Ok(event)
|
||||
}
|
||||
|
||||
/// Snapshot current diagnostics.
|
||||
pub fn diagnostics(&self) -> Result<Vec<RuntimeDiagnostic>, RuntimeError> {
|
||||
Ok(self.lock()?.diagnostics.clone())
|
||||
@@ -465,6 +553,12 @@ struct RuntimeState {
|
||||
workers: BTreeMap<WorkerId, WorkerRecord>,
|
||||
events: Vec<RuntimeEvent>,
|
||||
diagnostics: Vec<RuntimeDiagnostic>,
|
||||
#[cfg(feature = "ws-server")]
|
||||
next_observation_sequence: u64,
|
||||
#[cfg(feature = "ws-server")]
|
||||
observation_events: VecDeque<WorkerObservationEvent>,
|
||||
#[cfg(feature = "ws-server")]
|
||||
observation_tx: broadcast::Sender<WorkerObservationEvent>,
|
||||
}
|
||||
|
||||
impl RuntimeState {
|
||||
@@ -482,6 +576,12 @@ impl RuntimeState {
|
||||
workers: BTreeMap::new(),
|
||||
events: Vec::new(),
|
||||
diagnostics: Vec::new(),
|
||||
#[cfg(feature = "ws-server")]
|
||||
next_observation_sequence: 1,
|
||||
#[cfg(feature = "ws-server")]
|
||||
observation_events: VecDeque::new(),
|
||||
#[cfg(feature = "ws-server")]
|
||||
observation_tx: broadcast::channel(256).0,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -505,6 +605,12 @@ impl RuntimeState {
|
||||
workers: BTreeMap::new(),
|
||||
events: Vec::new(),
|
||||
diagnostics: Vec::new(),
|
||||
#[cfg(feature = "ws-server")]
|
||||
next_observation_sequence: 1,
|
||||
#[cfg(feature = "ws-server")]
|
||||
observation_events: VecDeque::new(),
|
||||
#[cfg(feature = "ws-server")]
|
||||
observation_tx: broadcast::channel(256).0,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -762,6 +868,54 @@ impl RuntimeState {
|
||||
self.next_event_id.saturating_sub(1)
|
||||
}
|
||||
|
||||
#[cfg(feature = "ws-server")]
|
||||
fn validate_worker_observation_cursor(
|
||||
&self,
|
||||
worker_ref: &WorkerRef,
|
||||
cursor: WorkerObservationCursor,
|
||||
) -> Result<(), RuntimeError> {
|
||||
if let Some(first) = self
|
||||
.observation_events
|
||||
.iter()
|
||||
.find(|event| &event.worker_ref == worker_ref)
|
||||
{
|
||||
if cursor.sequence != 0 && cursor.sequence < first.sequence {
|
||||
return Err(RuntimeError::InvalidRequest(format!(
|
||||
"worker observation cursor {} is expired for worker {}",
|
||||
cursor.encode(),
|
||||
worker_ref.worker_id
|
||||
)));
|
||||
}
|
||||
}
|
||||
if cursor.sequence >= self.next_observation_sequence {
|
||||
return Err(RuntimeError::InvalidRequest(format!(
|
||||
"worker observation cursor {} is unknown for worker {}",
|
||||
cursor.encode(),
|
||||
worker_ref.worker_id
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(feature = "ws-server")]
|
||||
fn push_worker_observation_event(
|
||||
&mut self,
|
||||
worker_ref: WorkerRef,
|
||||
payload: protocol::Event,
|
||||
) -> WorkerObservationEvent {
|
||||
const MAX_OBSERVATION_BACKLOG: usize = 1024;
|
||||
|
||||
let sequence = self.next_observation_sequence;
|
||||
self.next_observation_sequence += 1;
|
||||
let event = WorkerObservationEvent::new(sequence, worker_ref, payload);
|
||||
self.observation_events.push_back(event.clone());
|
||||
while self.observation_events.len() > MAX_OBSERVATION_BACKLOG {
|
||||
self.observation_events.pop_front();
|
||||
}
|
||||
let _ = self.observation_tx.send(event.clone());
|
||||
event
|
||||
}
|
||||
|
||||
fn push_diagnostic(
|
||||
&mut self,
|
||||
severity: DiagnosticSeverity,
|
||||
|
||||
Reference in New Issue
Block a user