diff --git a/Cargo.lock b/Cargo.lock index 28870bb0..4b2f4b62 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -637,6 +637,7 @@ checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" name = "client" version = "0.1.0" dependencies = [ + "async-trait", "chrono", "futures", "protocol", @@ -4622,6 +4623,7 @@ version = "0.1.0" dependencies = [ "agen", "async-trait", + "client", "fs4", "futures", "manifest", diff --git a/crates/client/Cargo.toml b/crates/client/Cargo.toml index a17997e7..e3475d33 100644 --- a/crates/client/Cargo.toml +++ b/crates/client/Cargo.toml @@ -5,6 +5,7 @@ edition.workspace = true license.workspace = true [dependencies] +async-trait.workspace = true chrono = { version = "0.4", default-features = false, features = ["clock"] } protocol = { workspace = true } ticket = { workspace = true } diff --git a/crates/client/src/backend_runtime.rs b/crates/client/src/backend_runtime.rs index 8cb67625..78f61833 100644 --- a/crates/client/src/backend_runtime.rs +++ b/crates/client/src/backend_runtime.rs @@ -1,13 +1,7 @@ -use crate::{BackendApiClient, BackendApiClientError}; -use futures::{SinkExt, StreamExt}; -use protocol::stream::{decode_event, encode_method}; -use protocol::{ErrorCode, Event, Method}; +use crate::transport::websocket::{Socket as WebSocket, SocketError as WebSocketError}; +use crate::{BackendApiClient, BackendApiClientError, Client}; use reqwest::Method as HttpMethod; -use std::collections::VecDeque; use std::fmt; -use tokio::sync::mpsc; -use tokio_tungstenite::connect_async; -use tokio_tungstenite::tungstenite::Message as TungsteniteMessage; use tokio_tungstenite::tungstenite::client::IntoClientRequest; use tokio_tungstenite::tungstenite::http::HeaderValue; use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; @@ -106,20 +100,12 @@ impl BackendRuntimeListTarget { } } -#[derive(Debug)] -pub struct BackendRuntimeClient { - target: BackendRuntimeTarget, - command_tx: mpsc::UnboundedSender, - events: mpsc::UnboundedReceiver, - diagnostics: VecDeque, - _protocol_task: tokio::task::JoinHandle<()>, -} - #[derive(Debug)] pub enum BackendRuntimeClientError { InvalidTarget(String), Api(BackendApiClientError), Http(reqwest::Error), + Protocol(String), } impl fmt::Display for BackendRuntimeClientError { @@ -128,6 +114,7 @@ impl fmt::Display for BackendRuntimeClientError { Self::InvalidTarget(message) => f.write_str(message), Self::Api(error) => write!(f, "{error}"), Self::Http(error) => write!(f, "{error}"), + Self::Protocol(message) => f.write_str(message), } } } @@ -282,151 +269,22 @@ pub async fn restore_backend_worker( Ok(response.json::().await?) } -impl BackendRuntimeClient { - pub async fn connect(target: BackendRuntimeTarget) -> Result { - validate_target(&target)?; - let api = BackendApiClient::from_stored_token(&target.base_url)?; - let (event_tx, rx) = mpsc::unbounded_channel(); - let (command_tx, command_rx) = mpsc::unbounded_channel(); - - let protocol_target = target.clone(); - let protocol_event_tx = event_tx.clone(); - let protocol_task = tokio::spawn(async move { - run_worker_protocol_transport(protocol_target, api, command_rx, protocol_event_tx) - .await; - }); - - Ok(Self { - target, - command_tx, - events: rx, - diagnostics: VecDeque::new(), - _protocol_task: protocol_task, - }) - } - - pub fn try_next_event(&mut self) -> Option { - if let Some(event) = self.diagnostics.pop_front() { - return Some(event); - } - self.events.try_recv().ok() - } - - pub async fn next_event(&mut self) -> Option { - if let Some(event) = self.diagnostics.pop_front() { - return Some(event); - } - self.events.recv().await - } - - pub async fn send(&mut self, method: &Method) -> Result<(), BackendRuntimeClientError> { - self.command_tx.send(method.clone()).map_err(|_| { - BackendRuntimeClientError::InvalidTarget(format!( - "Backend protocol command stream is closed for {}", - self.target.display_label() - )) - })?; - Ok(()) - } -} - -impl Drop for BackendRuntimeClient { - fn drop(&mut self) { - self._protocol_task.abort(); - } -} - -async fn run_worker_protocol_transport( +pub async fn connect_backend_runtime( target: BackendRuntimeTarget, - api: BackendApiClient, - mut commands: mpsc::UnboundedReceiver, - tx: mpsc::UnboundedSender, -) { - let request = match protocol_ws_request(&target, &api) { - Ok(request) => request, - Err(error) => { - let _ = tx.send(diagnostic_event(format!( - "Backend protocol request could not be constructed for {}: {error}", - target.display_label() - ))); - return; - } - }; - match connect_async(request).await { - Ok((ws, _)) => { - let (mut sink, mut stream) = ws.split(); - loop { - tokio::select! { - maybe_method = commands.recv() => { - let Some(method) = maybe_method else { - break; - }; - match encode_method(&method) { - Ok(text) => { - if let Err(error) = sink.send(TungsteniteMessage::Text(text.into())).await { - let _ = tx.send(diagnostic_event(format!( - "Backend protocol command send failed for {}: {error}", - target.display_label() - ))); - break; - } - } - Err(error) => { - let _ = tx.send(diagnostic_event(format!( - "Backend protocol command could not serialize method for {}: {error}", - target.display_label() - ))); - } - } - } - frame = stream.next() => { - match frame { - Some(Ok(TungsteniteMessage::Text(text))) => { - match decode_event(&text) { - Ok(event) => { - let _ = tx.send(event); - } - Err(error) => { - let _ = tx.send(diagnostic_event(format!( - "Backend protocol response was not valid Event JSON for {}: {error}", - target.display_label() - ))); - } - } - } - Some(Ok(TungsteniteMessage::Close(_))) | None => { - let _ = tx.send(diagnostic_event(format!( - "Backend protocol command stream closed for {}", - target.display_label() - ))); - break; - } - Some(Ok(TungsteniteMessage::Ping(_))) - | Some(Ok(TungsteniteMessage::Pong(_))) - | Some(Ok(TungsteniteMessage::Binary(_))) - | Some(Ok(TungsteniteMessage::Frame(_))) => {} - Some(Err(error)) => { - let _ = tx.send(diagnostic_event(format!( - "Backend protocol WebSocket error for {}: {error}", - target.display_label() - ))); - break; - } - } - } - } - } - } - Err(error) => { - let message = protocol_connect_error_message(&target, &api, &error); - let _ = tx.send(diagnostic_event(message)); - while commands.recv().await.is_some() { - let _ = tx.send(diagnostic_event(format!( - "Backend protocol command was not sent because command stream is unavailable for {}", - target.display_label() - ))); - } - } +) -> Result, BackendRuntimeClientError> { + validate_target(&target)?; + let api = BackendApiClient::from_stored_token(&target.base_url)?; + let request = protocol_ws_request(&target, &api).map_err(|error| { + BackendRuntimeClientError::Protocol(format!( + "Backend protocol request could not be constructed for {}: {error}", + target.display_label() + )) + })?; + match WebSocket::connect(request).await { + Ok(socket) => Ok(Client::new(socket)), + Err(WebSocketError::WebSocket(error)) => Err(BackendRuntimeClientError::Protocol( + protocol_connect_error_message(&target, &api, &error), + )), } } @@ -453,13 +311,6 @@ fn protocol_connect_error_message( ) } -fn diagnostic_event(message: impl Into) -> Event { - Event::Error { - code: ErrorCode::Internal, - message: message.into(), - } -} - fn validate_target(target: &BackendRuntimeTarget) -> Result<(), BackendRuntimeClientError> { if target.base_url.trim().is_empty() { return Err(BackendRuntimeClientError::InvalidTarget( diff --git a/crates/client/src/client.rs b/crates/client/src/client.rs new file mode 100644 index 00000000..517d6a7c --- /dev/null +++ b/crates/client/src/client.rs @@ -0,0 +1,137 @@ +use std::error::Error; +use std::fmt; + +use protocol::stream::{decode_event, encode_method}; +use protocol::{Event, Method}; + +use crate::transport::Socket; + +/// Typed Worker protocol client over an injected message transport. +pub struct Client { + socket: T, +} + +#[derive(Debug)] +pub enum ClientError { + Transport(E), + Protocol(serde_json::Error), +} + +impl Client { + pub fn new(socket: T) -> Self { + Self { socket } + } + + pub fn into_inner(self) -> T { + self.socket + } +} + +impl Client { + pub async fn send(&mut self, method: &Method) -> Result<(), ClientError> { + let message = encode_method(method).map_err(ClientError::Protocol)?; + self.socket + .send(message) + .await + .map_err(ClientError::Transport) + } + + pub async fn next_event(&mut self) -> Result, ClientError> { + self.socket + .next() + .await + .map_err(ClientError::Transport)? + .map(|message| decode_event(&message).map_err(ClientError::Protocol)) + .transpose() + } + + pub fn try_next_event(&mut self) -> Result, ClientError> { + self.socket + .try_next() + .map_err(ClientError::Transport)? + .map(|message| decode_event(&message).map_err(ClientError::Protocol)) + .transpose() + } +} + +impl fmt::Display for ClientError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Transport(error) => write!(formatter, "Worker transport error: {error}"), + Self::Protocol(error) => write!(formatter, "Worker protocol error: {error}"), + } + } +} + +impl Error for ClientError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + Self::Transport(error) => Some(error), + Self::Protocol(error) => Some(error), + } + } +} + +#[cfg(test)] +mod tests { + use std::collections::VecDeque; + use std::convert::Infallible; + + use async_trait::async_trait; + use protocol::stream::{decode_method, encode_event}; + use protocol::{Event, Method, WorkerStatus}; + + use super::Client; + use crate::transport::Socket; + + #[derive(Default)] + struct TestSocket { + sent: Vec, + incoming: VecDeque, + } + + #[async_trait] + impl Socket for TestSocket { + type Error = Infallible; + + async fn send(&mut self, message: String) -> Result<(), Self::Error> { + self.sent.push(message); + Ok(()) + } + + async fn next(&mut self) -> Result, Self::Error> { + Ok(self.incoming.pop_front()) + } + + fn try_next(&mut self) -> Result, Self::Error> { + Ok(self.incoming.pop_front()) + } + } + + #[tokio::test] + async fn encodes_methods_and_decodes_events_above_transport() { + let mut socket = TestSocket::default(); + socket.incoming.push_back( + encode_event(&Event::Status { + status: WorkerStatus::Idle, + }) + .expect("encode event"), + ); + let mut client = Client::new(socket); + + client + .send(&Method::run_text("hello")) + .await + .expect("send method"); + assert!(matches!( + decode_method(&client.socket.sent[0]), + Ok(Method::Run { .. }) + )); + assert!(matches!( + client.next_event().await, + Ok(Some(Event::Status { + status: WorkerStatus::Idle + })) + )); + } +} diff --git a/crates/client/src/lib.rs b/crates/client/src/lib.rs index b373747f..6aae3f6a 100644 --- a/crates/client/src/lib.rs +++ b/crates/client/src/lib.rs @@ -7,8 +7,9 @@ pub mod backend_api; mod backend_auth; pub mod backend_runtime; pub mod backend_workspace; +mod client; pub mod target; -mod worker_client; +pub mod transport; mod workspace_product; pub use backend_api::{ @@ -20,23 +21,23 @@ pub use backend_auth::{ poll_device_login, start_device_login, wait_for_device_login, }; pub use backend_runtime::{ - BackendDiagnostic, BackendDiagnosticSeverity, BackendRuntimeClient, BackendRuntimeClientError, + BackendDiagnostic, BackendDiagnosticSeverity, BackendRuntimeClientError, BackendRuntimeListResponse, BackendRuntimeListTarget, BackendRuntimeSummary, BackendRuntimeTarget, BackendWorkerCapabilitySummary, BackendWorkerImplementationSummary, BackendWorkerRestoreResponse, BackendWorkerRestoreResult, BackendWorkerSummary, - BackendWorkerWorkspaceSummary, BackendWorkingDirectorySummary, list_backend_stopped_workers, - list_backend_workers, restore_backend_worker, + BackendWorkerWorkspaceSummary, BackendWorkingDirectorySummary, connect_backend_runtime, + list_backend_stopped_workers, list_backend_workers, restore_backend_worker, }; pub use backend_workspace::{ BackendWorkspace, BackendWorkspaceCatalogTarget, BackendWorkspaceClientError, CreateBackendWorkspaceRepository, CreateBackendWorkspaceRequest, CreateBackendWorkspaceResponse, create_backend_workspace, list_backend_workspaces, }; +pub use client::{Client, ClientError}; pub use target::{ BackendTarget, Dashboard, ResolvedTarget, StandaloneSessionListIntent, StandaloneSessionResumeIntent, StandaloneTarget, Target, TargetError, TargetKind, WorkerConnection, WorkerConnectionSelector, WorkerList, WorkerListRequest, WorkerSpawn, }; -pub use worker_client::WorkerClient; pub use workspace_api::{ObjectiveDetail, ObjectiveSummary}; pub use workspace_product::BackendWorkspaceProductClient; diff --git a/crates/client/src/transport/in_process.rs b/crates/client/src/transport/in_process.rs new file mode 100644 index 00000000..b1db8e6b --- /dev/null +++ b/crates/client/src/transport/in_process.rs @@ -0,0 +1,115 @@ +use async_trait::async_trait; +use thiserror::Error; +use tokio::sync::mpsc; + +use super::Socket as SocketContract; + +const CHANNEL_CAPACITY: usize = 256; + +pub struct Socket { + outgoing: mpsc::Sender, + incoming: mpsc::Receiver, +} + +/// Host-side endpoint paired with an in-process client transport. +pub struct Peer { + incoming: mpsc::Receiver, + outgoing: mpsc::Sender, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)] +pub enum SocketError { + #[error("in-process Worker protocol transport closed")] + Closed, +} + +impl Socket { + pub fn pair() -> (Self, Peer) { + let (client_tx, peer_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (peer_tx, client_rx) = mpsc::channel(CHANNEL_CAPACITY); + ( + Self { + outgoing: client_tx, + incoming: client_rx, + }, + Peer { + incoming: peer_rx, + outgoing: peer_tx, + }, + ) + } +} + +#[async_trait] +impl SocketContract for Socket { + type Error = SocketError; + + async fn send(&mut self, message: String) -> Result<(), Self::Error> { + self.outgoing + .send(message) + .await + .map_err(|_| SocketError::Closed) + } + + async fn next(&mut self) -> Result, Self::Error> { + Ok(self.incoming.recv().await) + } + + fn try_next(&mut self) -> Result, Self::Error> { + match self.incoming.try_recv() { + Ok(message) => Ok(Some(message)), + Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) => { + Ok(None) + } + } + } +} + +impl Peer { + pub async fn next(&mut self) -> Option { + self.incoming.recv().await + } + + pub async fn send(&self, message: String) -> Result<(), String> { + self.outgoing.send(message).await.map_err(|error| error.0) + } +} + +#[cfg(test)] +mod tests { + use protocol::stream::{decode_method, encode_event}; + use protocol::{Event, Method, WorkerStatus}; + + use super::Socket; + use crate::Client; + + #[tokio::test] + async fn pair_carries_typed_protocol_through_generic_client() { + let (socket, mut peer) = Socket::pair(); + let mut client = Client::new(socket); + + client + .send(&Method::run_text("hello")) + .await + .expect("send method"); + assert!(matches!( + peer.next().await.as_deref().map(decode_method), + Some(Ok(Method::Run { .. })) + )); + + peer.send( + encode_event(&Event::Status { + status: WorkerStatus::Idle, + }) + .expect("encode event"), + ) + .await + .expect("send event"); + assert!(matches!( + client.next_event().await, + Ok(Some(Event::Status { + status: WorkerStatus::Idle + })) + )); + } +} diff --git a/crates/client/src/transport/mod.rs b/crates/client/src/transport/mod.rs new file mode 100644 index 00000000..12d5e6d2 --- /dev/null +++ b/crates/client/src/transport/mod.rs @@ -0,0 +1,22 @@ +use std::error::Error; + +use async_trait::async_trait; + +pub mod in_process; +pub mod unix_socket; +pub mod websocket; + +/// Message-oriented transport for one Worker protocol connection. +/// +/// Implementations own physical framing. `client::Client` owns the typed +/// Method/Event protocol encoding layered on top of these UTF-8 messages. +#[async_trait] +pub trait Socket { + type Error: Error + Send + Sync + 'static; + + async fn send(&mut self, message: String) -> Result<(), Self::Error>; + + async fn next(&mut self) -> Result, Self::Error>; + + fn try_next(&mut self) -> Result, Self::Error>; +} diff --git a/crates/client/src/transport/unix_socket.rs b/crates/client/src/transport/unix_socket.rs new file mode 100644 index 00000000..0262bff2 --- /dev/null +++ b/crates/client/src/transport/unix_socket.rs @@ -0,0 +1,172 @@ +use std::io; +use std::path::Path; + +use async_trait::async_trait; +use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; +use tokio::net::UnixStream; +use tokio::sync::mpsc; +use tokio::task::JoinHandle; + +use super::Socket as SocketContract; + +pub struct Socket { + writer: tokio::io::WriteHalf, + messages: mpsc::Receiver>, + reader_task: JoinHandle<()>, +} + +impl Socket { + pub async fn connect(path: &Path) -> io::Result { + let stream = UnixStream::connect(path).await?; + let (reader, writer) = tokio::io::split(stream); + let (message_tx, messages) = mpsc::channel(256); + let reader_task = tokio::spawn(async move { + let mut lines = BufReader::new(reader).lines(); + loop { + match lines.next_line().await { + Ok(Some(message)) if message.trim().is_empty() => {} + Ok(Some(message)) => { + if message_tx.send(Ok(message)).await.is_err() { + return; + } + } + Ok(None) => return, + Err(error) => { + let _ = message_tx.send(Err(error)).await; + return; + } + } + } + }); + Ok(Self { + writer, + messages, + reader_task, + }) + } +} + +#[async_trait] +impl SocketContract for Socket { + type Error = io::Error; + + async fn send(&mut self, message: String) -> Result<(), Self::Error> { + self.writer.write_all(message.as_bytes()).await?; + self.writer.write_all(b"\n").await?; + self.writer.flush().await + } + + async fn next(&mut self) -> Result, Self::Error> { + match self.messages.recv().await { + Some(message) => message.map(Some), + None => Ok(None), + } + } + + fn try_next(&mut self) -> Result, Self::Error> { + match self.messages.try_recv() { + Ok(message) => message.map(Some), + Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) => { + Ok(None) + } + } + } +} + +impl Drop for Socket { + fn drop(&mut self) { + self.reader_task.abort(); + } +} + +#[cfg(test)] +mod tests { + use std::io::ErrorKind; + use std::time::Duration; + + use protocol::stream::{decode_method, encode_event}; + use protocol::{Event, Method, WorkerStatus}; + use tempfile::tempdir; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::UnixListener; + + use super::*; + use crate::Client; + + async fn assert_peer_closed(stream: &mut UnixStream, reason: &str) { + let mut buf = [0_u8; 1]; + match tokio::time::timeout(Duration::from_secs(1), stream.read(&mut buf)) + .await + .expect(reason) + { + Ok(0) => {} + Err(error) if error.kind() == ErrorKind::ConnectionReset => {} + Ok(n) => panic!("server should observe peer close, read {n} byte(s)"), + Err(error) => panic!("server read failed unexpectedly: {error}"), + } + } + + #[tokio::test] + async fn client_receives_events_over_unix_socket() { + let socket_dir = tempdir().unwrap(); + let socket_path = socket_dir.path().join("events.sock"); + let listener = UnixListener::bind(&socket_path).unwrap(); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let event = encode_event(&Event::Status { + status: WorkerStatus::Idle, + }) + .unwrap(); + stream.write_all(event.as_bytes()).await.unwrap(); + stream.write_all(b"\n").await.unwrap(); + }); + + let mut client = Client::new(Socket::connect(&socket_path).await.unwrap()); + let event = tokio::time::timeout(Duration::from_secs(1), client.next_event()) + .await + .expect("client should receive event while alive") + .expect("transport should succeed"); + assert!(matches!( + event, + Some(Event::Status { + status: WorkerStatus::Idle + }) + )); + server.await.unwrap(); + } + + #[tokio::test] + async fn client_sends_methods_over_unix_socket() { + let socket_dir = tempdir().unwrap(); + let socket_path = socket_dir.path().join("send.sock"); + let listener = UnixListener::bind(&socket_path).unwrap(); + let server = tokio::spawn(async move { + let (reader, _) = listener.accept().await.unwrap(); + BufReader::new(reader).lines().next_line().await.unwrap() + }); + + let mut client = Client::new(Socket::connect(&socket_path).await.unwrap()); + client + .send(&Method::run_text("hello")) + .await + .expect("send method"); + + let received = server.await.unwrap().expect("method message"); + assert!(matches!(decode_method(&received), Ok(Method::Run { .. }))); + } + + #[tokio::test] + async fn dropping_socket_closes_server_connection() { + let socket_dir = tempdir().unwrap(); + let socket_path = socket_dir.path().join("drop.sock"); + let listener = UnixListener::bind(&socket_path).unwrap(); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + assert_peer_closed(&mut stream, "dropped socket should close promptly").await; + }); + + let socket = Socket::connect(&socket_path).await.unwrap(); + drop(socket); + server.await.unwrap(); + } +} diff --git a/crates/client/src/transport/websocket.rs b/crates/client/src/transport/websocket.rs new file mode 100644 index 00000000..e8640573 --- /dev/null +++ b/crates/client/src/transport/websocket.rs @@ -0,0 +1,140 @@ +use async_trait::async_trait; +use futures::{SinkExt, StreamExt}; +use thiserror::Error; +use tokio::net::TcpStream; +use tokio::sync::mpsc; +use tokio::task::JoinHandle; +use tokio_tungstenite::tungstenite::http::Request; +use tokio_tungstenite::tungstenite::{self, Message}; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async}; + +use super::Socket as SocketContract; + +type Writer = futures::stream::SplitSink>, Message>; + +pub struct Socket { + writer: Writer, + messages: mpsc::Receiver>, + reader_task: JoinHandle<()>, +} + +#[derive(Debug, Error)] +pub enum SocketError { + #[error("WebSocket transport failed: {0}")] + WebSocket(#[from] tungstenite::Error), +} + +impl Socket { + pub async fn connect(request: Request<()>) -> Result { + let (stream, _) = connect_async(request).await?; + let (writer, mut reader) = stream.split(); + let (message_tx, messages) = mpsc::channel(256); + let reader_task = tokio::spawn(async move { + loop { + match reader.next().await { + Some(Ok(Message::Text(message))) => { + if message_tx.send(Ok(message.to_string())).await.is_err() { + return; + } + } + Some(Ok(Message::Close(_))) | None => return, + Some(Ok( + Message::Binary(_) + | Message::Ping(_) + | Message::Pong(_) + | Message::Frame(_), + )) => {} + Some(Err(error)) => { + let _ = message_tx.send(Err(SocketError::WebSocket(error))).await; + return; + } + } + } + }); + Ok(Self { + writer, + messages, + reader_task, + }) + } +} + +#[async_trait] +impl SocketContract for Socket { + type Error = SocketError; + + async fn send(&mut self, message: String) -> Result<(), Self::Error> { + self.writer.send(Message::Text(message.into())).await?; + Ok(()) + } + + async fn next(&mut self) -> Result, Self::Error> { + match self.messages.recv().await { + Some(message) => message.map(Some), + None => Ok(None), + } + } + + fn try_next(&mut self) -> Result, Self::Error> { + match self.messages.try_recv() { + Ok(message) => message.map(Some), + Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) => { + Ok(None) + } + } + } +} + +impl Drop for Socket { + fn drop(&mut self) { + self.reader_task.abort(); + } +} + +#[cfg(test)] +mod tests { + use futures::{SinkExt, StreamExt}; + use protocol::stream::{decode_method, encode_event}; + use protocol::{Event, Method, WorkerStatus}; + use tokio::net::TcpListener; + use tokio_tungstenite::accept_async; + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + + use super::*; + use crate::Client; + + #[tokio::test] + async fn carries_typed_protocol_through_generic_client() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let mut socket = accept_async(stream).await.unwrap(); + let message = socket.next().await.unwrap().unwrap(); + assert!(matches!( + message, + Message::Text(ref text) + if matches!(decode_method(text), Ok(Method::Run { .. })) + )); + let event = encode_event(&Event::Status { + status: WorkerStatus::Idle, + }) + .unwrap(); + socket.send(Message::Text(event.into())).await.unwrap(); + }); + + let request = format!("ws://{address}").into_client_request().unwrap(); + let mut client = Client::new(Socket::connect(request).await.unwrap()); + client + .send(&Method::run_text("hello")) + .await + .expect("send method"); + assert!(matches!( + client.next_event().await, + Ok(Some(Event::Status { + status: WorkerStatus::Idle + })) + )); + server.await.unwrap(); + } +} diff --git a/crates/client/src/worker_client.rs b/crates/client/src/worker_client.rs deleted file mode 100644 index a0661f41..00000000 --- a/crates/client/src/worker_client.rs +++ /dev/null @@ -1,186 +0,0 @@ -use std::io; -use std::path::Path; - -use protocol::stream::{JsonLineReader, JsonLineWriter}; -use protocol::{Event, Method}; -use tokio::net::UnixStream; -use tokio::sync::mpsc; -use tokio::task::JoinHandle; - -pub struct WorkerClient { - writer: JsonLineWriter>, - event_rx: mpsc::Receiver, - reader_task: JoinHandle<()>, -} - -impl WorkerClient { - pub async fn connect(path: &Path) -> Result { - let stream = UnixStream::connect(path).await?; - let (reader, writer) = tokio::io::split(stream); - let writer = JsonLineWriter::new(writer); - - let (event_tx, event_rx) = mpsc::channel::(256); - - let reader_task = tokio::spawn(async move { - let mut reader = JsonLineReader::new(reader); - while let Ok(Some(event)) = reader.next::().await { - if event_tx.send(event).await.is_err() { - break; - } - } - }); - - Ok(Self { - writer, - event_rx, - reader_task, - }) - } - - pub async fn send(&mut self, method: &Method) -> Result<(), io::Error> { - self.writer.write(method).await - } - - pub fn try_next_event(&mut self) -> Option { - self.event_rx.try_recv().ok() - } - - pub async fn next_event(&mut self) -> Option { - self.event_rx.recv().await - } -} - -impl Drop for WorkerClient { - fn drop(&mut self) { - self.reader_task.abort(); - } -} - -#[cfg(test)] -mod tests { - use std::io::ErrorKind; - use std::time::Duration; - - use protocol::{Segment, WorkerStatus}; - use tempfile::tempdir; - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - use tokio::net::UnixListener; - - use super::*; - - async fn assert_peer_closed(stream: &mut UnixStream, reason: &str) { - let mut buf = [0_u8; 1]; - match tokio::time::timeout(Duration::from_secs(1), stream.read(&mut buf)) - .await - .expect(reason) - { - Ok(0) => {} - Err(error) if error.kind() == ErrorKind::ConnectionReset => {} - Ok(n) => panic!("server should observe peer close, read {n} byte(s)"), - Err(error) => panic!("server read failed unexpectedly: {error}"), - } - } - - #[tokio::test] - async fn receives_events_while_client_is_alive() { - let socket_dir = tempdir().unwrap(); - let socket_path = socket_dir.path().join("events.sock"); - let listener = UnixListener::bind(&socket_path).unwrap(); - let server = tokio::spawn(async move { - let (stream, _) = listener.accept().await.unwrap(); - let mut writer = JsonLineWriter::new(stream); - writer - .write(&Event::Status { - status: WorkerStatus::Idle, - }) - .await - .unwrap(); - }); - - let mut client = WorkerClient::connect(&socket_path).await.unwrap(); - - let event = tokio::time::timeout(Duration::from_secs(1), client.next_event()) - .await - .expect("client should receive event while alive"); - assert!(matches!( - event, - Some(Event::Status { - status: WorkerStatus::Idle - }) - )); - server.await.unwrap(); - } - - #[tokio::test] - async fn send_writes_methods_while_client_is_alive() { - let socket_dir = tempdir().unwrap(); - let socket_path = socket_dir.path().join("send.sock"); - let listener = UnixListener::bind(&socket_path).unwrap(); - let server = tokio::spawn(async move { - let (stream, _) = listener.accept().await.unwrap(); - let mut reader = JsonLineReader::new(stream); - reader.next::().await.unwrap() - }); - - let mut client = WorkerClient::connect(&socket_path).await.unwrap(); - let method = Method::Run { - input: vec![Segment::text("hello")], - }; - client.send(&method).await.unwrap(); - - let received = tokio::time::timeout(Duration::from_secs(1), server) - .await - .expect("server should receive method while client is alive") - .unwrap(); - match received { - Some(Method::Run { input }) => assert_eq!(input, vec![Segment::text("hello")]), - other => panic!("expected Run method, got {other:?}"), - } - } - - #[tokio::test] - async fn dropping_repeated_clients_closes_server_connections() { - let socket_dir = tempdir().unwrap(); - let socket_path = socket_dir.path().join("drop.sock"); - let listener = UnixListener::bind(&socket_path).unwrap(); - let server = tokio::spawn(async move { - for _ in 0..16 { - let (mut stream, _) = listener.accept().await.unwrap(); - assert_peer_closed( - &mut stream, - "dropped client should close its socket promptly", - ) - .await; - } - }); - - for _ in 0..16 { - let client = WorkerClient::connect(&socket_path).await.unwrap(); - drop(client); - } - - server.await.unwrap(); - } - - #[tokio::test] - async fn dropping_client_aborts_blocked_reader_task() { - let socket_dir = tempdir().unwrap(); - let socket_path = socket_dir.path().join("blocked-reader.sock"); - let listener = UnixListener::bind(&socket_path).unwrap(); - let server = tokio::spawn(async move { - let (mut stream, _) = listener.accept().await.unwrap(); - stream.write_all(b"{\"event\"").await.unwrap(); - assert_peer_closed( - &mut stream, - "aborting the blocked client reader should close the socket", - ) - .await; - }); - - let client = WorkerClient::connect(&socket_path).await.unwrap(); - tokio::task::yield_now().await; - drop(client); - - server.await.unwrap(); - } -} diff --git a/crates/standalone/Cargo.toml b/crates/standalone/Cargo.toml index a3bb1813..8c5adc9d 100644 --- a/crates/standalone/Cargo.toml +++ b/crates/standalone/Cargo.toml @@ -7,6 +7,7 @@ license.workspace = true [dependencies] agen.workspace = true +client.workspace = true fs4.workspace = true manifest.workspace = true protocol.workspace = true diff --git a/crates/standalone/src/host.rs b/crates/standalone/src/host.rs index 031a554f..b6e2f7ba 100644 --- a/crates/standalone/src/host.rs +++ b/crates/standalone/src/host.rs @@ -2,14 +2,20 @@ use std::path::PathBuf; use std::time::Duration; use agen::llm_client::client::LlmClient; +use client::Client; +use client::transport::in_process::{Peer as InProcessPeer, Socket as InProcessSocket}; +use protocol::stream::{decode_method, encode_event}; use protocol::{Event, Method}; use session_store::{ CombinedStore, FsStore, FsWorkerStore, WorkerActiveSegmentRef, WorkerMetadataStore, }; use thiserror::Error; -use tokio::sync::broadcast; use worker::bootstrap::{WorkerBootstrap, WorkerBootstrapError, WorkerBootstrapLayout}; use worker::controller::WorkerControllerTransport; +use worker::ipc::protocol_session::{ + WorkerProtocolSessionStreams, dispatch_worker_protocol_method, live_log_entry_event, + subscribe_worker_protocol_session, +}; use worker::{BootstrappedWorker, WorkerError, WorkerFilesystemAuthority, WorkerWorkspaceContext}; use crate::launch::ResolvedStandaloneLaunch; @@ -55,12 +61,6 @@ pub enum StandaloneStartupError { Controller, } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)] -pub enum StandaloneRequestError { - #[error("the standalone Worker is no longer accepting requests")] - WorkerUnavailable, -} - #[derive(Debug, Clone, Copy, PartialEq, Eq, Error)] pub enum StandaloneShutdownError { #[error("the standalone Worker did not stop before the shutdown deadline")] @@ -277,19 +277,15 @@ impl StandaloneHost { &self.record } - pub async fn send(&self, method: Method) -> Result<(), StandaloneRequestError> { - self.handle - .send(method) - .await - .map_err(|_| StandaloneRequestError::WorkerUnavailable) - } - - pub fn subscribe(&self) -> broadcast::Receiver { - self.handle.subscribe() - } - - pub fn snapshot(&self) -> Event { - self.handle.snapshot_event() + /// Open one complete client-side Worker protocol session. + /// + /// Working events, committed session entries, alert snapshots, and the + /// initial history snapshot are merged behind the client boundary. + pub fn connect(&self) -> Client { + let streams = subscribe_worker_protocol_session(&self.handle); + let (socket, peer) = InProcessSocket::pair(); + tokio::spawn(run_protocol_session(self.handle.clone(), streams, peer)); + Client::new(socket) } pub fn with_shutdown_timeout(mut self, shutdown_timeout: Duration) -> Self { @@ -349,6 +345,111 @@ impl StandaloneHost { } } +async fn run_protocol_session( + handle: worker::WorkerHandle, + streams: WorkerProtocolSessionStreams, + mut peer: InProcessPeer, +) { + let WorkerProtocolSessionStreams { + snapshot_event, + mut log_entries, + alert_snapshot, + mut events, + } = streams; + + if !send_protocol_snapshot(&peer, alert_snapshot, snapshot_event).await { + return; + } + + loop { + tokio::select! { + message = peer.next() => { + let Some(message) = message else { + return; + }; + let Ok(method) = decode_method(&message) else { + return; + }; + if let Some(event) = dispatch_worker_protocol_method(&handle, method).await + && !send_protocol_event(&peer, event).await + { + return; + } + } + event = events.recv() => { + match event { + Ok(event) => { + if !send_protocol_event(&peer, event).await { + return; + } + } + Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { + let replacement = subscribe_worker_protocol_session(&handle); + let WorkerProtocolSessionStreams { + snapshot_event, + log_entries: replacement_log_entries, + alert_snapshot, + events: replacement_events, + } = replacement; + log_entries = replacement_log_entries; + events = replacement_events; + if !send_protocol_snapshot(&peer, alert_snapshot, snapshot_event).await { + return; + } + } + Err(tokio::sync::broadcast::error::RecvError::Closed) => return, + } + } + entry = log_entries.recv() => { + match entry { + Ok(entry) => { + if let Some(event) = live_log_entry_event(entry) + && !send_protocol_event(&peer, event).await + { + return; + } + } + Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { + let replacement = subscribe_worker_protocol_session(&handle); + let WorkerProtocolSessionStreams { + snapshot_event, + log_entries: replacement_log_entries, + alert_snapshot, + events: replacement_events, + } = replacement; + log_entries = replacement_log_entries; + events = replacement_events; + if !send_protocol_snapshot(&peer, alert_snapshot, snapshot_event).await { + return; + } + } + Err(tokio::sync::broadcast::error::RecvError::Closed) => return, + } + } + } + } +} + +async fn send_protocol_snapshot( + peer: &InProcessPeer, + alert_snapshot: Vec, + snapshot_event: Event, +) -> bool { + for alert in alert_snapshot { + if !send_protocol_event(peer, Event::Alert(alert)).await { + return false; + } + } + send_protocol_event(peer, snapshot_event).await +} + +async fn send_protocol_event(peer: &InProcessPeer, event: Event) -> bool { + let Ok(message) = encode_event(&event) else { + return false; + }; + peer.send(message).await.is_ok() +} + fn backing_store( store: &StandaloneSessionStore, id: StandaloneSessionId, diff --git a/crates/standalone/src/lib.rs b/crates/standalone/src/lib.rs index 19123cb4..8c04f863 100644 --- a/crates/standalone/src/lib.rs +++ b/crates/standalone/src/lib.rs @@ -8,9 +8,7 @@ pub mod host; pub mod launch; pub mod store; -pub use host::{ - StandaloneHost, StandaloneRequestError, StandaloneShutdownError, StandaloneStartupError, -}; +pub use host::{StandaloneHost, StandaloneShutdownError, StandaloneStartupError}; pub use launch::{ResolvedStandaloneLaunch, StandaloneLaunchConfig, StandaloneLaunchError}; pub use store::{ StaleLeasePolicy, StandaloneCwdIdentity, StandaloneListScope, StandaloneSessionId, diff --git a/crates/standalone/tests/host.rs b/crates/standalone/tests/host.rs index b01bcbfb..8198f3ac 100644 --- a/crates/standalone/tests/host.rs +++ b/crates/standalone/tests/host.rs @@ -8,6 +8,8 @@ use agen::llm_client::error::ClientError; use agen::llm_client::event::{Event as LlmEvent, StopReason}; use agen::llm_client::types::Request; use async_trait::async_trait; +use client::Client; +use client::transport::in_process::Socket as InProcessSocket; use futures::{Stream, stream}; use protocol::{Event, Method}; use standalone::{ @@ -88,17 +90,29 @@ async fn in_process_host_runs_text_and_read_tool_then_shuts_down() { let host = StandaloneHost::start_with_model_client(launch, client) .await .expect("start in-process host"); - let mut events = host.subscribe(); + let mut protocol_client = host.connect(); - host.send(Method::run_text("read the probe")) + protocol_client + .send(&Method::run_text("read the probe")) .await .expect("submit input"); tokio::time::timeout(Duration::from_secs(30), async { + let mut saw_user_message = false; let mut saw_text = false; let mut saw_tool_result = false; loop { - match events.recv().await.expect("worker event") { + match protocol_client + .next_event() + .await + .expect("protocol event") + .expect("worker event") + { + Event::UserMessage { segments } + if format!("{segments:?}").contains("read the probe") => + { + saw_user_message = true; + } Event::TextDelta { text } if text.contains("standalone response") => { saw_text = true; } @@ -106,6 +120,10 @@ async fn in_process_host_runs_text_and_read_tool_then_shuts_down() { saw_tool_result = true; } Event::RunEnd { .. } => { + assert!( + saw_user_message, + "stream must expose the committed user message" + ); assert!(saw_text, "stream must expose the model text delta"); assert!(saw_tool_result, "stream must expose the tool result"); break; @@ -251,15 +269,18 @@ async fn standalone_restore_preserves_history_tasks_notifications_and_cwd_scope( ]); let host = StandaloneHost::start_with_model_client(launch, first_client).await?; let session_id = host.session_id(); - let mut events = host.subscribe(); - host.send(Method::run_text("first request")).await?; - wait_for_run_end(&mut events).await?; - host.send(Method::Notify { - message: "persisted notification".to_string(), - auto_run: true, - }) - .await?; - wait_for_run_end(&mut events).await?; + let mut protocol_client = host.connect(); + protocol_client + .send(&Method::run_text("first request")) + .await?; + wait_for_run_end(&mut protocol_client).await?; + protocol_client + .send(&Method::Notify { + message: "persisted notification".to_string(), + auto_run: true, + }) + .await?; + wait_for_run_end(&mut protocol_client).await?; host.shutdown().await?; let store = StandaloneSessionStore::open(&state_dir)?; @@ -288,16 +309,24 @@ async fn standalone_restore_preserves_history_tasks_notifications_and_cwd_scope( let host = StandaloneHost::restore_with_model_client(state_dir.clone(), session_id, second_client) .await?; - let snapshot = format!("{:?}", host.snapshot()); + let mut protocol_client = host.connect(); + let snapshot = format!( + "{:?}", + protocol_client + .next_event() + .await + .expect("restored protocol stream") + .expect("restored snapshot") + ); assert!(snapshot.contains("first request"), "{snapshot}"); assert!(snapshot.contains("first answer"), "{snapshot}"); assert!(snapshot.contains("persisted task"), "{snapshot}"); assert!(snapshot.contains("persisted notification"), "{snapshot}"); - let mut events = host.subscribe(); - host.send(Method::run_text("continue after restore")) + protocol_client + .send(&Method::run_text("continue after restore")) .await?; - wait_for_run_end(&mut events).await?; + wait_for_run_end(&mut protocol_client).await?; let request = second_inspection .requests() .into_iter() @@ -503,10 +532,10 @@ async fn standalone_metadata_fails_closed_on_incomplete_or_newer_records() -> Te Ok(()) } -async fn wait_for_run_end(events: &mut tokio::sync::broadcast::Receiver) -> TestResult { +async fn wait_for_run_end(client: &mut Client) -> TestResult { tokio::time::timeout(Duration::from_secs(10), async { loop { - if matches!(events.recv().await, Ok(Event::RunEnd { .. })) { + if matches!(client.next_event().await, Ok(Some(Event::RunEnd { .. }))) { break; } } diff --git a/crates/tui/src/console/mod.rs b/crates/tui/src/console/mod.rs index 0d8fe7d2..541ff446 100644 --- a/crates/tui/src/console/mod.rs +++ b/crates/tui/src/console/mod.rs @@ -21,10 +21,13 @@ use protocol::{Greeting, RewindSummary, RewindTarget, RewindTargetId, Segment}; use ratatui::Terminal; use ratatui::backend::CrosstermBackend; use standalone::{StandaloneHost, StandaloneLaunchConfig}; -use tokio::sync::{broadcast, mpsc}; +use tokio::sync::mpsc; use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; -use client::{BackendRuntimeClient, BackendRuntimeTarget, StandaloneSessionResumeIntent}; +use client::transport::Socket; +use client::{ + BackendRuntimeTarget, Client, StandaloneSessionResumeIntent, connect_backend_runtime, +}; use crate::app::{ActionbarNoticeLevel, ActionbarNoticeSource, App}; use crate::composer_keys::{ComposerEditAction, composer_edit_action}; @@ -119,74 +122,40 @@ fn copy_selection_to_terminal(app: &mut App) -> bool { copy_selection_to_writer(app, &mut stdout) } -enum ConsoleConnection { - BackendRuntime(BackendRuntimeClient), - Standalone { - host: Option, - events: broadcast::Receiver, - initial_snapshot: Option, - }, +struct ConsoleConnection { + client: Client, + standalone_host: Option, } -impl ConsoleConnection { - fn standalone(host: StandaloneHost) -> Self { - let events = host.subscribe(); - let initial_snapshot = Some(host.snapshot()); - Self::Standalone { - host: Some(host), - events, - initial_snapshot, +impl ConsoleConnection { + fn new(client: Client) -> Self { + Self { + client, + standalone_host: None, } } - fn try_next_event(&mut self) -> Option { - match self { - Self::BackendRuntime(client) => client.try_next_event(), - Self::Standalone { - events, - initial_snapshot, - .. - } => initial_snapshot.take().or_else(|| events.try_recv().ok()), + fn with_standalone_host(client: Client, host: StandaloneHost) -> Self { + Self { + client, + standalone_host: Some(host), } } - async fn next_event(&mut self) -> Option { - match self { - Self::BackendRuntime(client) => client.next_event().await, - Self::Standalone { host, events, .. } => loop { - match events.recv().await { - Ok(event) => break Some(event), - Err(broadcast::error::RecvError::Lagged(_)) => { - let Some(host) = host.as_ref() else { - break None; - }; - break Some(host.snapshot()); - } - Err(broadcast::error::RecvError::Closed) => break None, - } - }, - } + fn try_next_event(&mut self) -> Result, Box> { + Ok(self.client.try_next_event()?) + } + + async fn next_event(&mut self) -> Result, Box> { + Ok(self.client.next_event().await?) } async fn send(&mut self, method: &Method) -> Result<(), Box> { - match self { - Self::BackendRuntime(client) => Ok(client.send(method).await?), - Self::Standalone { host, .. } => { - let host = host.as_ref().ok_or_else(|| { - io::Error::new( - io::ErrorKind::BrokenPipe, - "Standalone Worker has already shut down", - ) - })?; - Ok(host.send(method.clone()).await?) - } - } + Ok(self.client.send(method).await?) } async fn shutdown(&mut self) -> Result<(), Box> { - if let Self::Standalone { host, .. } = self - && let Some(host) = host.take() - { + if let Some(host) = self.standalone_host.take() { host.shutdown().await?; } Ok(()) @@ -251,7 +220,8 @@ async fn run_standalone_host( worker_label: String, history_root: PathBuf, ) -> Result<(), Box> { - let mut connection = ConsoleConnection::standalone(host); + let client = host.connect(); + let mut connection = ConsoleConnection::with_standalone_host(client, host); let mut terminal = match enter_fullscreen() { Ok(terminal) => terminal, @@ -280,12 +250,12 @@ pub(crate) async fn run_backend_runtime( target: BackendRuntimeTarget, ) -> Result<(), Box> { let worker_label = target.display_label(); - let client = BackendRuntimeClient::connect(target).await?; + let client = connect_backend_runtime(target).await?; let mut terminal = enter_fullscreen()?; let workspace_root = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")); let mut app = App::new_with_persistent_input_history(worker_label, &workspace_root); app.connected = true; - let mut connection = ConsoleConnection::BackendRuntime(client); + let mut connection = ConsoleConnection::new(client); let result = run_loop(&mut terminal, &mut app, &mut connection).await; let _ = leave_fullscreen(&mut terminal); result @@ -560,7 +530,7 @@ enum E2eRewindInput { enum LoopInput

{ Terminal(TerminalEventResult), - Worker(Option

), + Worker(P), } async fn next_loop_input( @@ -569,7 +539,7 @@ async fn next_loop_input( pod_next: F, ) -> LoopInput

where - F: Future>, + F: Future, { tokio::select! { biased; @@ -586,9 +556,9 @@ where } } -async fn drain_terminal_events( +async fn drain_terminal_events( app: &mut App, - client: &mut ConsoleConnection, + client: &mut ConsoleConnection, term_rx: &mut mpsc::UnboundedReceiver, ) -> Result> { let mut handled = false; @@ -613,13 +583,13 @@ async fn drain_terminal_events( Ok(handled) } -async fn drain_worker_events( +async fn drain_worker_events( app: &mut App, - client: &mut ConsoleConnection, + client: &mut ConsoleConnection, ) -> Result> { let mut handled = false; for _ in 0..POD_EVENT_DRAIN_LIMIT { - match client.try_next_event() { + match client.try_next_event()? { Some(ev) => { handled = true; if let Some(method) = app.handle_worker_event(ev) { @@ -632,10 +602,10 @@ async fn drain_worker_events( Ok(handled) } -async fn run_loop( +async fn run_loop( terminal: &mut Terminal>, app: &mut App, - client: &mut ConsoleConnection, + client: &mut ConsoleConnection, ) -> Result<(), Box> { let (_terminal_reader, mut term_rx) = TerminalEventReader::spawn()?; @@ -660,7 +630,7 @@ async fn run_loop( LoopInput::Terminal(term_event) => { handle_terminal_event(app, client, term_event?).await?; } - LoopInput::Worker(event) => match event { + LoopInput::Worker(event) => match event? { Some(ev) => { if let Some(method) = app.handle_worker_event(ev) { client.send(&method).await?; @@ -680,9 +650,9 @@ async fn run_loop( Ok(()) } -async fn handle_terminal_event( +async fn handle_terminal_event( app: &mut App, - client: &mut ConsoleConnection, + client: &mut ConsoleConnection, event: TermEvent, ) -> Result<(), Box> { match event {