diff --git a/.github/workflows/rust.yml b/.github/workflows/rust.yml index bab99db..cfd9813 100644 --- a/.github/workflows/rust.yml +++ b/.github/workflows/rust.yml @@ -1,10 +1,10 @@ name: Rust +# Manual only. Nothing here runs on a push or a pull request: the gate is the +# one an author runs locally before asking for review, and a green tick that +# nobody asked for teaches people to stop reading it. on: - push: - branches: [ "master" ] - pull_request: - branches: [ "master" ] + workflow_dispatch: env: CARGO_TERM_COLOR: always diff --git a/AGENTS.md b/AGENTS.md index 4f9da03..8ba7b61 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -9,6 +9,13 @@ natively, and Claude Code loads it through the `@AGENTS.md` import in ## Invariants (don't break these) +- **No Python.** Not a script, not `python3 -c`, not a heredoc. Reaching for it is the + tell that a step is being solved by parsing when the tool that owns the answer could + just be asked. Do not swap it for another parser either, and do not assume `jq` is + present: it does not ship with macOS. A fixed-shape field is one `sed -nE` line; + anything needing real parsing belongs in this repo's own language, where it can be + tested. If a task seems to need Python, the approach is wrong. + - **Verify the whole chain, not just this repo.** endpoint-gen, honey_id-types, endpoint-validator and the six backends all break silently when this crate moves. `./scripts/check-chain.sh` verifies all of it; run it before calling a change done. diff --git a/Cargo.toml b/Cargo.toml index 14fe1cc..d8b54cb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "endpoint-libs" -version = "3.0.3" +version = "3.1.0" edition = "2024" authors = ["Veon "] description = "Launch MCP services fast: describe endpoints once in RON, and endpoint-gen generates the Rust models, docs and MCP tool schemas that this crate serves over WebSocket RPC, with roles and typed errors built in." @@ -67,6 +67,13 @@ framed-transport = [ "wire-core", "dep:tokio-util", ] +nagoya-transport = [ + # `framed_json_neutral` over a Nagoya socket. Off by default and additive: a + # consumer that does not run Nagoya never compiles it, and the tokio path is + # untouched either way. + "framed-transport", + "dep:nagoya", +] agent-control = ["framed-transport"] ws-client = [ # WS client (WsClient, WsClientBuilder) - standalone @@ -155,6 +162,11 @@ tokio-rustls = { version = "0.26", optional = true, default-features = false, fe rustls = { version = "0.23", optional = true, default-features = false, features = ["ring", "logging", "std"] } tokio = { version = "1.39", features = ["full"] } tokio-util = { version = "0.7", features = ["codec"], optional = true } +# The I/O driver is behind Nagoya's own non-default `reactor` feature; without it +# `nagoya::reactor` does not exist and the error reads like a missing module. +# 0.1.9 is the published release that already speaks AF_UNIX and carries `TaskSet`, +# so this resolves from the registry with no path dependency and no sibling checkout. +nagoya = { version = "0.1.9", default-features = false, features = ["reactor"], optional = true } tokio-cron-scheduler = { version = "0.11", optional = true } parking_lot = { version = "0.12", optional = true } dashmap = { version = "6.0", optional = true } @@ -191,6 +203,10 @@ tonic = { version = "0.14", optional = true } [dev-dependencies] tempfile = "3.19" tokio = { version = "1.39", features = ["full", "test-util"] } +# `compat` only, and only for tests: it adapts a tokio duplex into the +# `futures-io` pair the neutral transport takes, which is what lets one test +# drive both framing paths and compare the bytes they produce. +tokio-util = { version = "0.7", features = ["codec", "compat"] } tracing-throttle = { version = "0.4", features = ["async", "test-helpers"] } uuid = { version = "1", features = ["v4", "serde"] } rcgen = "0.14" diff --git a/src/libs/ws/server.rs b/src/libs/ws/server.rs index ce1f2fd..2ae0baf 100644 --- a/src/libs/ws/server.rs +++ b/src/libs/ws/server.rs @@ -171,31 +171,58 @@ impl WebsocketServer { .upgrade_stream(stream, addr, &self.config, &cached_date) .await?; - // Loop: spawn session task for each upgrade event - while let Ok(event) = rx.recv().await { + use futures::StreamExt; + use futures::future::FutureExt; + use futures::stream::FuturesUnordered; + + // Poll sessions in place, the same one-thread model `serve_with` uses. + // The hyper upgrader still `spawn_local`s (TokioExecutor); that is why + // `run_shard` keeps a LocalSet. nago-wss is the replacement, not this loop. + let mut sessions = FuturesUnordered::new(); + loop { + let event = if sessions.is_empty() { + match rx.recv().await { + Ok(event) => event, + Err(_) => break, + } + } else { + let recv = rx.recv(); + futures::pin_mut!(recv); + match futures::future::select(recv, sessions.next()).await { + futures::future::Either::Left((Ok(event), _)) => event, + futures::future::Either::Left((Err(_), _)) => { + while sessions.next().await.is_some() {} + break; + } + futures::future::Either::Right(_) => continue, + } + }; + let this = Arc::clone(&self); let states = Arc::clone(&states); let addr_clone = addr; + sessions.push( + async move { + let ws_stream = match create_ws_stream(event.on_upgrade).await { + Ok(s) => s, + Err(e) => { + error!(ws_server = true, ?addr_clone, "on_upgrade failed: {e}"); + return; + } + }; - tokio::task::spawn_local(async move { - let ws_stream = match create_ws_stream(event.on_upgrade).await { - Ok(s) => s, - Err(e) => { - error!(ws_server = true, ?addr_clone, "on_upgrade failed: {e}"); - return; - } - }; - - debug!( - ws_server = true, - ?addr_clone, - protocol = %event.protocol, - "WsServer: upgrade succeeded, protocol received" - ); + debug!( + ws_server = true, + ?addr_clone, + protocol = %event.protocol, + "WsServer: upgrade succeeded, protocol received" + ); - this.post_upgrade_connection(addr_clone, states, ws_stream, event.protocol) - .await; - }); + this.post_upgrade_connection(addr_clone, states, ws_stream, event.protocol) + .await; + } + .boxed_local(), + ); } debug!( @@ -230,8 +257,9 @@ impl WebsocketServer { /// time — the WebSocket subprotocol string today, a handed-over token for local /// transports. It is passed to [`AuthController::auth`] unchanged. /// - /// Must be called inside a `tokio::task::LocalSet`: [`MessageStream`]'s futures - /// are not `Send`. + /// [`MessageStream`]'s futures are not `Send`, so this must be polled on + /// the thread that owns the stream. [`Self::serve_with`] does that by + /// driving connections with `FuturesUnordered` rather than `spawn_local`. pub async fn serve_connection( self: Arc, peer: PeerIdentity, @@ -351,12 +379,19 @@ impl WebsocketServer { /// runs on a single runtime — the shard-per-core model is a property of the TCP /// path and buys nothing for a 1:1 sidecar channel. /// - /// Must be called inside a `tokio::task::LocalSet` (see - /// [`Self::serve_connection`]). + /// Connections are polled in place with `FuturesUnordered`. That is the + /// same one-thread model `spawn_local` had, without tying the method to a + /// tokio `LocalSet`, so a nagoya reactor can drive it. TCP `listen` now + /// polls its connections the same way; it still needs a `LocalSet` for the + /// hyper upgrader. pub async fn serve_with(self, listener: L) -> Result<()> where L: SessionListener + 'static, { + use futures::StreamExt; + use futures::future::FutureExt; + use futures::stream::FuturesUnordered; + self.validate_protocol_mode()?; let this = Arc::new(self); let states = Arc::new(WebsocketStates::new()); @@ -366,8 +401,19 @@ impl WebsocketServer { this.config.drop_conn_on_buffer_full, ); + let mut connections = FuturesUnordered::new(); loop { - let (stream, peer) = match listener.accept().await { + let accepted = if connections.is_empty() { + listener.accept().await + } else { + let accept = listener.accept(); + futures::pin_mut!(accept); + match futures::future::select(accept, connections.next()).await { + futures::future::Either::Left((accepted, _)) => accepted, + futures::future::Either::Right(_) => continue, + } + }; + let (stream, peer) = match accepted { Ok(accepted) => accepted, Err(err) => { error!(ws_server = true, error = %err, "listener accept failed; stopping"); @@ -378,11 +424,14 @@ impl WebsocketServer { let this = Arc::clone(&this); let states = Arc::clone(&states); - tokio::task::spawn_local(async move { - // Local transports carry credentials out of band (an inherited fd is - // already a capability), so there is no subprotocol string to pass. - this.serve_connection(peer, states, stream, None).await; - }); + connections.push( + async move { + // Local transports carry credentials out of band (an inherited fd is + // already a capability), so there is no subprotocol string to pass. + this.serve_connection(peer, states, stream, None).await; + } + .boxed_local(), + ); } } @@ -495,43 +544,68 @@ impl WebsocketServer { .enable_all() .build() .expect("Failed to build shard runtime"); + // LocalSet remains only because the hyper upgrader `spawn_local`s onto + // TokioExecutor. Connection tasks themselves are polled in place. let local_set = LocalSet::new(); rt.block_on(local_set.run_until(async move { + use futures::StreamExt; + use futures::future::FutureExt; + use futures::stream::FuturesUnordered; + + let mut connections = FuturesUnordered::new(); loop { - let Some((stream, addr)) = rx.recv().await else { + let received = if connections.is_empty() { + rx.recv().await + } else { + let recv = rx.recv(); + futures::pin_mut!(recv); + match futures::future::select(recv, connections.next()).await { + futures::future::Either::Left((received, _)) => received, + futures::future::Either::Right(_) => continue, + } + }; + let Some((stream, addr)) = received else { + while connections.next().await.is_some() {} break; }; let this = Arc::clone(&this); let states = Arc::clone(&states); let listener = Arc::clone(&listener); - tokio::task::spawn_local(async move { - let stream = match listener.handshake(stream).await { - Ok(channel) => { - debug!(ws_server = true, "Accepted stream from {}", addr); - channel - } - Err(err) => { + connections.push( + async move { + let stream = match listener.handshake(stream).await { + Ok(channel) => { + debug!(ws_server = true, "Accepted stream from {}", addr); + channel + } + Err(err) => { + error!( + ws_server = true, + "Error while handshaking stream: {:?}", err + ); + return; + } + }; + if let Err(err) = TOOLBOX + .scope( + this.toolbox.clone(), + this.handle_ws_handshake_and_connection( + addr, + states, + Box::new(stream), + ), + ) + .await + { error!( ws_server = true, - "Error while handshaking stream: {:?}", err + ?addr, + "Failed to handle WS connection: {err}" ); - return; } - }; - if let Err(err) = TOOLBOX - .scope( - this.toolbox.clone(), - this.handle_ws_handshake_and_connection(addr, states, Box::new(stream)), - ) - .await - { - error!( - ws_server = true, - ?addr, - "Failed to handle WS connection: {err}" - ); } - }); + .boxed_local(), + ); } })); } diff --git a/src/libs/ws/session.rs b/src/libs/ws/session.rs index 24ea64e..a408328 100644 --- a/src/libs/ws/session.rs +++ b/src/libs/ws/session.rs @@ -1,9 +1,23 @@ use eyre::Result; +use futures::StreamExt; +use futures::future::{FutureExt, LocalBoxFuture}; +use futures::stream::FuturesUnordered; use std::collections::HashSet; use std::sync::Arc; use tokio::sync::mpsc; use tracing::*; +/// What the session loop does after one inbound frame. +/// +/// Handler bodies are polled in place on this task (`Task`) so a slow hook +/// cannot stall the read loop, without `spawn_local` and the LocalSet it +/// needs. `serve_with` already drives connections the same way. +enum Dispatch { + Close, + Keep, + Task(LocalBoxFuture<'static, ()>), +} + use crate::libs::ws::WsMessage as Message; use crate::libs::error_code::ErrorCode; @@ -58,7 +72,7 @@ impl WsClientSession { } } - fn handle_message(&mut self, msg: Message) -> Result { + fn handle_message(&mut self, msg: Message) -> Result { let addr = self.conn_info.peer.display(); let mut context = RequestContext::from_conn(&self.conn_info); @@ -73,8 +87,10 @@ impl WsClientSession { }; if let Some(frame) = payload.and_then(mcp::try_parse_jsonrpc) { let mcp = Arc::clone(mcp); - self.handle_mcp_frame(mcp, frame, context); - return Ok(true); + if let Some(task) = self.handle_mcp_frame(mcp, frame, context) { + return Ok(Dispatch::Task(task)); + } + return Ok(Dispatch::Keep); } } @@ -90,7 +106,7 @@ impl WsClientSession { ) .to_string(), ); - return Ok(true); + return Ok(Dispatch::Keep); } #[allow(unreachable_patterns)] @@ -114,14 +130,14 @@ impl WsClientSession { serde_json::from_slice(&b) } Message::Ping(_) => { - return Ok(true); + return Ok(Dispatch::Keep); } Message::Pong(_) => { - return Ok(true); + return Ok(Dispatch::Keep); } Message::Close(_) => { debug!(ws_server = true, ?addr, "Receive side terminated"); - return Ok(false); + return Ok(Dispatch::Close); } _ => { warn!( @@ -129,7 +145,7 @@ impl WsClientSession { ?addr, "Ignoring unsupported WebSocket frame" ); - return Ok(true); + return Ok(Dispatch::Keep); } }; let req = match obj { @@ -148,7 +164,7 @@ impl WsClientSession { }), }), ); - return Ok(true); + return Ok(Dispatch::Keep); } }; debug!( @@ -177,7 +193,7 @@ impl WsClientSession { }), }), ); - return Ok(true); + return Ok(Dispatch::Keep); }; if !check_roles(&context.roles, &endpoint.allowed_roles) { @@ -194,51 +210,52 @@ impl WsClientSession { }), }), ); - return Ok(true); + return Ok(Dispatch::Keep); } let handler = endpoint.handler.clone(); let toolbox = self.server.toolbox.clone(); let hooks = self.server.hooks.clone(); let schema = endpoint.schema.clone(); - tokio::task::spawn_local(async move { - let mut context = context; - // Hooks run inside the spawned task so a slow hook cannot stall the - // session loop, and after check_roles so they only see calls that were - // already allowed to reach this endpoint. - if let Err(custom) = hooks.run_before(&mut context, &schema, &req.params).await { - let code = custom.code.to_u32(); - toolbox.send( - context.connection_id, - WsResponseValue::Error(WsResponseError { - method: context.method, - code, - seq: context.seq, - log_id: context.log_id.to_string(), - params: custom.params.clone(), - }), - ); + Ok(Dispatch::Task( + async move { + let mut context = context; + // Polled alongside the session read loop so a slow hook cannot + // stall it, and after check_roles so they only see calls that + // were already allowed to reach this endpoint. + if let Err(custom) = hooks.run_before(&mut context, &schema, &req.params).await { + let code = custom.code.to_u32(); + toolbox.send( + context.connection_id, + WsResponseValue::Error(WsResponseError { + method: context.method, + code, + seq: context.seq, + log_id: context.log_id.to_string(), + params: custom.params.clone(), + }), + ); + hooks + .run_after(&context, &schema, &RequestOutcome::PublicErr { code }) + .await; + return; + } + + TOOLBOX + .scope( + toolbox.clone(), + handler.handle(&toolbox, context.clone(), req.params), + ) + .await; + + // The erased handler reports its own outcome through the toolbox, so + // AfterRequest observes completion rather than the specific result here. hooks - .run_after(&context, &schema, &RequestOutcome::PublicErr { code }) + .run_after(&context, &schema, &RequestOutcome::Ok) .await; - return; } - - TOOLBOX - .scope( - toolbox.clone(), - handler.handle(&toolbox, context.clone(), req.params), - ) - .await; - - // The erased handler reports its own outcome through the toolbox, so - // AfterRequest observes completion rather than the specific result here. - hooks - .run_after(&context, &schema, &RequestOutcome::Ok) - .await; - }); - - Ok(true) + .boxed_local(), + )) } /// Handles one parsed JSON-RPC frame: lifecycle methods are answered @@ -251,7 +268,7 @@ impl WsClientSession { mcp: Arc, frame: Result, mut context: RequestContext, - ) { + ) -> Option> { let conn_id = context.connection_id; let req = match frame { Ok(req) => req, @@ -259,7 +276,7 @@ impl WsClientSession { self.server .toolbox .send_raw(conn_id, error_frame.to_string()); - return; + return None; } }; @@ -269,8 +286,9 @@ impl WsClientSession { match mcp.route(req, &context.roles) { McpAction::Respond(frame) => { self.server.toolbox.send_raw(conn_id, frame.to_string()); + None } - McpAction::Ignore => {} + McpAction::Ignore => None, McpAction::ToolCall { id, method_code, @@ -294,52 +312,58 @@ impl WsClientSession { ) .to_string(), ); - return; + return None; }; let handler = endpoint.handler.clone(); let toolbox = self.server.toolbox.clone(); let hooks = self.server.hooks.clone(); let schema = endpoint.schema.clone(); - tokio::task::spawn_local(async move { - let mut context = context; - // Same placement as the legacy path, but the rejection has to go - // back in the MCP envelope — a tool error, not a WsResponseError. - if let Err(custom) = hooks.run_before(&mut context, &schema, &arguments).await { - let code = custom.code.to_u32(); - toolbox.send_raw( - conn_id, - jsonrpc_result(&id, encode_tool_error(custom.code, &custom.params)) - .to_string(), - ); + Some( + async move { + let mut context = context; + // Same placement as the legacy path, but the rejection has to go + // back in the MCP envelope — a tool error, not a WsResponseError. + if let Err(custom) = + hooks.run_before(&mut context, &schema, &arguments).await + { + let code = custom.code.to_u32(); + toolbox.send_raw( + conn_id, + jsonrpc_result(&id, encode_tool_error(custom.code, &custom.params)) + .to_string(), + ); + hooks + .run_after(&context, &schema, &RequestOutcome::PublicErr { code }) + .await; + return; + } + + TOOLBOX + .scope( + toolbox.clone(), + handler.handle_mcp( + &toolbox, + context.clone(), + McpCallCtx { id }, + arguments, + ), + ) + .await; + hooks - .run_after(&context, &schema, &RequestOutcome::PublicErr { code }) + .run_after(&context, &schema, &RequestOutcome::Ok) .await; - return; } - - TOOLBOX - .scope( - toolbox.clone(), - handler.handle_mcp( - &toolbox, - context.clone(), - McpCallCtx { id }, - arguments, - ), - ) - .await; - - hooks - .run_after(&context, &schema, &RequestOutcome::Ok) - .await; - }); + .boxed_local(), + ) } } } async fn run_loop(&mut self) -> Result<()> { let conn_id = self.conn_info.connection_id; + let mut handlers = FuturesUnordered::new(); loop { while let Ok(msg) = self.rx.try_recv() { if !self.send_message(msg).await { @@ -385,14 +409,17 @@ impl WsClientSession { break; } }; - if !self.handle_message(msg)? { - break; + match self.handle_message(msg)? { + Dispatch::Close => break, + Dispatch::Keep => {} + Dispatch::Task(task) => handlers.push(task), } } else { debug!(ws_server = true, ?conn_id, "Inbound stream ended"); break; } } + _ = handlers.next(), if !handlers.is_empty() => {} } } diff --git a/src/libs/ws/traits.rs b/src/libs/ws/traits.rs index a4c9746..a2bb7c5 100644 --- a/src/libs/ws/traits.rs +++ b/src/libs/ws/traits.rs @@ -54,8 +54,10 @@ impl std::error::Error for StreamError { /// implements this just as well, either directly or via /// [`TransportStream`](super::TransportStream). /// -/// Note `?Send`: implementations' futures need not be `Send`, which matches the -/// `spawn_local` dispatch model. Drivers must run inside a `LocalSet`. +/// Note `?Send`: implementations' futures need not be `Send`. `serve_with` +/// and the session loop poll them in place on one thread. TCP `listen` polls +/// connections the same way; a `LocalSet` remains only for the hyper +/// upgrader's `spawn_local`. #[async_trait(?Send)] pub trait MessageStream: Unpin + Send { async fn send(&mut self, msg: Message) -> Result<(), StreamError>; diff --git a/src/libs/ws/transport.rs b/src/libs/ws/transport.rs index 7ff59e8..41faa15 100644 --- a/src/libs/ws/transport.rs +++ b/src/libs/ws/transport.rs @@ -18,11 +18,12 @@ //! //! # Threading //! -//! [`MessageStream`] is `#[async_trait(?Send)]` — its futures are **not** `Send`, -//! matching the existing `spawn_local` dispatch model. Anything driving a session -//! (`serve_connection`, `serve_with`) must therefore run inside a -//! `tokio::task::LocalSet`. This is not an oversight; it is what lets handlers hold -//! non-`Send` state across await points. +//! [`MessageStream`] is `#[async_trait(?Send)]` — its futures are **not** `Send`. +//! `serve_connection` / `serve_with` poll the session and its request handlers +//! in place on one thread, so they do not need a `LocalSet`. TCP `listen` now +//! polls connections the same way. A `LocalSet` remains only because the hyper +//! upgrader `spawn_local`s onto `TokioExecutor`; nago-wss is the tokio-free +//! replacement for that backend, and is not wired here yet. use eyre::eyre; use futures::{Sink, SinkExt, Stream, StreamExt}; @@ -89,5 +90,11 @@ where #[cfg(feature = "framed-transport")] pub mod framed; +#[cfg(feature = "nagoya-transport")] +pub mod nagoya; + #[cfg(feature = "framed-transport")] pub use framed::{FramedError, framed_json}; + +#[cfg(feature = "nagoya-transport")] +pub use nagoya::NagoyaStream; diff --git a/src/libs/ws/transport/framed.rs b/src/libs/ws/transport/framed.rs index a3423dd..c9535f1 100644 --- a/src/libs/ws/transport/framed.rs +++ b/src/libs/ws/transport/framed.rs @@ -30,8 +30,11 @@ //! value, so the serde codec layer would buy nothing. use std::io; +use std::pin::Pin; +use std::task::{Context, Poll}; use bytes::{Buf, BufMut, Bytes, BytesMut}; +use futures::io::{AsyncRead as FuturesRead, AsyncWrite as FuturesWrite}; use futures::{Sink, Stream}; use tokio::io::{AsyncRead, AsyncWrite}; use tokio_util::codec::{Framed, LengthDelimitedCodec}; @@ -257,6 +260,228 @@ where } } +/// Length-delimited framing over a `futures-io` byte stream. +/// +/// The same wire format as [`framed_json`], carried over +/// [`futures::io::AsyncRead`]/[`AsyncWrite`] rather than tokio's. That trait pair is +/// the neutral one: tokio adapts to it through `tokio-util`'s `Compat`, and Nagoya +/// through [`NagoyaStream`](super::nagoya::NagoyaStream) behind the `nagoya-transport` +/// feature, so a runtime-agnostic caller has a path that does not name a runtime. +/// +/// `encode` and `decode` are shared with the tokio path, so the bytes on the wire +/// are identical by construction rather than by agreement. +pub fn framed_json_neutral( + io: S, +) -> impl Transport + Unpin + Send +where + S: FuturesRead + FuturesWrite + Unpin + Send + 'static, +{ + framed_json_neutral_with_max_frame(io, DEFAULT_MAX_FRAME_BYTES) +} + +/// [`framed_json_neutral`] with an explicit maximum frame length. +pub fn framed_json_neutral_with_max_frame( + io: S, + max_frame_bytes: usize, +) -> impl Transport + Unpin + Send +where + S: FuturesRead + FuturesWrite + Unpin + Send + 'static, +{ + NeutralFramed { + io, + max_frame_bytes, + read_buf: BytesMut::new(), + write_buf: BytesMut::new(), + read_eof: false, + } +} + +/// The `futures-io` counterpart of [`WireFramed`]. +/// +/// `tokio_util`'s `Framed` is what the tokio path gets for free; there is no +/// equivalent in `futures-util`, so the buffering is here. It is the same two +/// buffers `Framed` keeps: one accumulating what has been read but not yet framed, +/// one holding what has been encoded but not yet written. +struct NeutralFramed { + io: S, + max_frame_bytes: usize, + read_buf: BytesMut, + write_buf: BytesMut, + read_eof: bool, +} + +/// How much to ask the reader for at once when the buffer needs filling. +const READ_CHUNK_BYTES: usize = 8 * 1024; + +/// The length prefix itself: `u32` big-endian. +const LENGTH_PREFIX_BYTES: usize = 4; + +impl NeutralFramed { + /// Take one whole frame out of `read_buf`, if one is there. + /// + /// `Ok(None)` means "not yet", which is the ordinary case and not an error. + /// An oversized length is refused **before** the body is buffered, which is the + /// property that keeps a hostile peer from naming 4 GiB and being believed. + fn take_frame(&mut self) -> Result, FramedError> { + if self.read_buf.len() < LENGTH_PREFIX_BYTES { + return Ok(None); + } + + let mut prefix = [0u8; LENGTH_PREFIX_BYTES]; + prefix.copy_from_slice(&self.read_buf[..LENGTH_PREFIX_BYTES]); + let frame_len = u32::from_be_bytes(prefix) as usize; + + if frame_len > self.max_frame_bytes { + return Err(FramedError::Io(io::Error::new( + io::ErrorKind::InvalidData, + format!( + "frame of {frame_len} bytes exceeds the {} byte maximum", + self.max_frame_bytes + ), + ))); + } + + if self.read_buf.len() < LENGTH_PREFIX_BYTES + frame_len { + return Ok(None); + } + + self.read_buf.advance(LENGTH_PREFIX_BYTES); + Ok(Some(self.read_buf.split_to(frame_len))) + } +} + +impl Stream for NeutralFramed +where + S: FuturesRead + FuturesWrite + Unpin, +{ + type Item = Result; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + + loop { + match this.take_frame() { + Err(err) => return Poll::Ready(Some(Err(err))), + Ok(Some(frame)) => return Poll::Ready(Some(decode(frame))), + Ok(None) => {} + } + + // A frame was still arriving when the stream ended. That is a truncated + // frame and an error, not a clean close: a clean close lands on a frame + // boundary with nothing buffered. + if this.read_eof { + return if this.read_buf.is_empty() { + Poll::Ready(None) + } else { + Poll::Ready(Some(Err(FramedError::Io(io::Error::new( + io::ErrorKind::UnexpectedEof, + "the stream ended in the middle of a frame", + ))))) + }; + } + + let before = this.read_buf.len(); + this.read_buf.resize(before + READ_CHUNK_BYTES, 0); + let result = Pin::new(&mut this.io).poll_read(cx, &mut this.read_buf[before..]); + + match result { + Poll::Ready(Ok(0)) => { + this.read_buf.truncate(before); + this.read_eof = true; + } + Poll::Ready(Ok(read)) => this.read_buf.truncate(before + read), + Poll::Ready(Err(err)) => { + this.read_buf.truncate(before); + return Poll::Ready(Some(Err(FramedError::Io(err)))); + } + Poll::Pending => { + this.read_buf.truncate(before); + return Poll::Pending; + } + } + } + } +} + +impl Sink for NeutralFramed +where + S: FuturesRead + FuturesWrite + Unpin, +{ + type Error = FramedError; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + // Bound the outbound buffer the way `Framed` does: once a backlog has built + // up, make the caller wait for it to drain rather than accepting without + // limit. Below the threshold this is free. + let this = self.get_mut(); + if this.write_buf.len() >= this.max_frame_bytes { + Pin::new(&mut *this).poll_flush(cx) + } else { + Poll::Ready(Ok(())) + } + } + + fn start_send(self: Pin<&mut Self>, item: WireMessage) -> Result<(), Self::Error> { + let this = self.get_mut(); + let payload = encode(item); + + if payload.len() > this.max_frame_bytes { + return Err(FramedError::Io(io::Error::new( + io::ErrorKind::InvalidData, + format!( + "frame of {} bytes exceeds the {} byte maximum", + payload.len(), + this.max_frame_bytes + ), + ))); + } + + this.write_buf.reserve(LENGTH_PREFIX_BYTES + payload.len()); + this.write_buf + .put_u32(u32::try_from(payload.len()).map_err(|_| { + FramedError::Io(io::Error::new( + io::ErrorKind::InvalidData, + "frame length does not fit in u32", + )) + })?); + this.write_buf.put_slice(&payload); + Ok(()) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + + while !this.write_buf.is_empty() { + match Pin::new(&mut this.io).poll_write(cx, &this.write_buf) { + Poll::Ready(Ok(0)) => { + return Poll::Ready(Err(FramedError::Io(io::Error::new( + io::ErrorKind::WriteZero, + "the stream accepted no bytes", + )))); + } + Poll::Ready(Ok(written)) => this.write_buf.advance(written), + Poll::Ready(Err(err)) => return Poll::Ready(Err(FramedError::Io(err))), + Poll::Pending => return Poll::Pending, + } + } + + Pin::new(&mut this.io) + .poll_flush(cx) + .map_err(FramedError::Io) + } + + fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match self.as_mut().poll_flush(cx) { + Poll::Ready(Ok(())) => {} + other => return other, + } + let this = self.get_mut(); + Pin::new(&mut this.io) + .poll_close(cx) + .map_err(FramedError::Io) + } +} + #[cfg(test)] mod tests { use super::*; @@ -388,4 +613,139 @@ mod tests { "expected an io error for an oversized declared length, got {got:?}" ); } + /// The neutral path must put the same bytes on the wire as the tokio path. + /// Not "equivalent": identical, because the format is normative for non-Rust + /// peers and there are now two implementations that could drift. + #[tokio::test] + async fn the_neutral_path_writes_the_same_bytes_as_the_tokio_path() { + use tokio_util::compat::TokioAsyncReadCompatExt; + + let cases = vec![ + WireMessage::Text("hello".into()), + WireMessage::Binary(vec![0, 1, 2, 255].into()), + WireMessage::Ping(vec![9].into()), + WireMessage::Close(Some(CloseFrame { + code: 1000, + reason: "bye".into(), + })), + ]; + + for case in cases { + let (mut tokio_sink, tokio_reader) = tokio::io::duplex(4096); + let mut tokio_framed = framed_json(tokio_reader); + tokio_framed.send(case.clone()).await.expect("tokio send"); + let mut tokio_bytes = Vec::new(); + tokio::io::AsyncReadExt::read_buf(&mut tokio_sink, &mut tokio_bytes) + .await + .expect("tokio read"); + + let (mut neutral_sink, neutral_reader) = tokio::io::duplex(4096); + let mut neutral_framed = framed_json_neutral(neutral_reader.compat()); + neutral_framed + .send(case.clone()) + .await + .expect("neutral send"); + let mut neutral_bytes = Vec::new(); + tokio::io::AsyncReadExt::read_buf(&mut neutral_sink, &mut neutral_bytes) + .await + .expect("neutral read"); + + assert_eq!( + tokio_bytes, neutral_bytes, + "the two framing paths disagree on the wire format for {case:?}" + ); + } + } + + #[tokio::test] + async fn the_neutral_path_carries_messages_both_ways() { + use tokio_util::compat::TokioAsyncReadCompatExt; + + let (client, server) = tokio::io::duplex(4096); + let mut client = framed_json_neutral(client.compat()); + let mut server = framed_json_neutral(server.compat()); + + let sent = WireMessage::Text("ping".into()); + client.send(sent.clone()).await.expect("send"); + let got = server.next().await.expect("a frame").expect("decode"); + assert_eq!(sent, got); + + let back = WireMessage::Binary(vec![7, 7, 7].into()); + server.send(back.clone()).await.expect("send back"); + let got = client.next().await.expect("a frame").expect("decode"); + assert_eq!(back, got); + } + + /// A length prefix naming more than the maximum is refused before the body is + /// buffered. Believing it is how a peer asks for an allocation it never sends. + #[tokio::test] + async fn the_neutral_path_rejects_an_oversized_length_prefix() { + use tokio_util::compat::TokioAsyncReadCompatExt; + + let (mut writer, reader) = tokio::io::duplex(4096); + let mut framed = framed_json_neutral_with_max_frame(reader.compat(), 64); + + tokio::io::AsyncWriteExt::write_all(&mut writer, &u32::MAX.to_be_bytes()) + .await + .expect("write prefix"); + + let err = framed + .next() + .await + .expect("a result") + .expect_err("must refuse"); + assert!( + matches!(err, FramedError::Io(_)), + "expected an io error, got {err:?}" + ); + } + + /// A stream that ends mid-frame is truncation, not a clean close. A clean close + /// lands on a frame boundary with nothing buffered. + #[tokio::test] + async fn the_neutral_path_reports_a_truncated_frame() { + use tokio_util::compat::TokioAsyncReadCompatExt; + + let (mut writer, reader) = tokio::io::duplex(4096); + let mut framed = framed_json_neutral(reader.compat()); + + // Promise ten bytes, send three, then hang up. + tokio::io::AsyncWriteExt::write_all(&mut writer, &10u32.to_be_bytes()) + .await + .expect("write prefix"); + tokio::io::AsyncWriteExt::write_all(&mut writer, &[KIND_TEXT, b'h', b'i']) + .await + .expect("write partial body"); + drop(writer); + + let err = framed + .next() + .await + .expect("a result") + .expect_err("must refuse"); + assert!( + matches!(err, FramedError::Io(ref e) if e.kind() == io::ErrorKind::UnexpectedEof), + "expected UnexpectedEof, got {err:?}" + ); + } + + /// A clean close after a whole frame ends the stream rather than erroring. + #[tokio::test] + async fn the_neutral_path_ends_cleanly_on_a_frame_boundary() { + use tokio_util::compat::TokioAsyncReadCompatExt; + + let (writer, reader) = tokio::io::duplex(4096); + let mut sender = framed_json_neutral(writer.compat()); + let mut framed = framed_json_neutral(reader.compat()); + + sender + .send(WireMessage::Text("only".into())) + .await + .expect("send"); + sender.close().await.expect("close"); + + let got = framed.next().await.expect("a frame").expect("decode"); + assert_eq!(WireMessage::Text("only".into()), got); + assert!(framed.next().await.is_none(), "expected a clean end"); + } } diff --git a/src/libs/ws/transport/nagoya.rs b/src/libs/ws/transport/nagoya.rs new file mode 100644 index 0000000..60cd1eb --- /dev/null +++ b/src/libs/ws/transport/nagoya.rs @@ -0,0 +1,103 @@ +//! A Nagoya socket as a `futures-io` byte stream. +//! +//! [`framed_json_neutral`](super::framed::framed_json_neutral) asks for +//! [`futures::io::AsyncRead`]/[`AsyncWrite`], which is the trait pair that names no +//! runtime. Nagoya's socket does not implement it: it implements Nagoya's own +//! [`Stream`](nagoya::io::Stream), whose `read` and `write_all` are `async fn` rather +//! than `poll_` methods. Nagoya's `io::Compat` does not close that gap either, since it +//! runs the other way, presenting a `futures-io` stream to a Nagoya consumer. +//! +//! The adapter is thin because `TcpStream` already exposes the poll-shaped methods +//! underneath its futures: `poll_read`, `poll_write`, and `poll_flush` are public and have +//! exactly the `futures-io` signature. So this is delegation, not a bridge — no boxed +//! future, no buffer, and no second copy of the readiness logic. Had those methods been +//! private, this file would have had to poll an `async fn` borrowing `&mut self` from +//! inside `poll_read`, which is not something a wrapper can do soundly. +//! +//! `TcpStream` is the type for a Unix-domain connection too, so there is one adapter +//! rather than one per address family. +//! +//! ```no_run +//! # use endpoint_libs::libs::ws::transport::{framed::framed_json_neutral, nagoya::NagoyaStream}; +//! # fn example(socket: nagoya::reactor::TcpStream) { +//! let transport = framed_json_neutral(NagoyaStream::new(socket)); +//! # let _ = transport; +//! # } +//! ``` + +use std::io; +use std::pin::Pin; +use std::task::{Context, Poll}; + +use futures::io::{AsyncRead, AsyncWrite}; +use nagoya::io::StreamError; +use nagoya::reactor::TcpStream; + +/// A Nagoya [`TcpStream`] presented as a `futures-io` byte stream. +#[derive(Debug)] +pub struct NagoyaStream(TcpStream); + +impl NagoyaStream { + pub const fn new(stream: TcpStream) -> Self { + Self(stream) + } + + /// The socket back, for a caller that wants the Nagoya API again. + pub fn into_inner(self) -> TcpStream { + self.0 + } + + pub const fn get_ref(&self) -> &TcpStream { + &self.0 + } + + pub const fn get_mut(&mut self) -> &mut TcpStream { + &mut self.0 + } +} + +impl From for NagoyaStream { + fn from(stream: TcpStream) -> Self { + Self::new(stream) + } +} + +/// The platform's own number, which is what Nagoya carries and what `std` wants. +fn io_error(error: StreamError) -> io::Error { + io::Error::from_raw_os_error(error.0) +} + +impl AsyncRead for NagoyaStream { + fn poll_read( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buffer: &mut [u8], + ) -> Poll> { + self.get_mut().0.poll_read(cx, buffer).map_err(io_error) + } +} + +impl AsyncWrite for NagoyaStream { + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buffer: &[u8], + ) -> Poll> { + self.get_mut().0.poll_write(cx, buffer).map_err(io_error) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.get_mut().0.poll_flush(cx).map_err(io_error) + } + + /// Flush, and leave the close to the drop. + /// + /// Nagoya exposes no half-close, so this cannot shut the write side down while + /// leaving the read side open. Reporting success after a flush is the honest answer + /// available: everything written has been handed to the kernel, and the descriptor is + /// closed when the stream is dropped. A peer that needs to see end-of-stream before + /// the drop needs a half-close on the Nagoya side first. + fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.poll_flush(cx) + } +} diff --git a/tests/nagoya_transport.rs b/tests/nagoya_transport.rs new file mode 100644 index 0000000..92a6452 --- /dev/null +++ b/tests/nagoya_transport.rs @@ -0,0 +1,84 @@ +//! The framed transport carried over a Nagoya Unix socket, with no tokio anywhere. +//! +//! The claim under test is that `framed_json_neutral` is genuinely runtime-neutral: the +//! same wire format reaches a peer over a socket driven by Nagoya's reactor, on Nagoya's +//! executor, in a build that does not start a tokio runtime. A compile alone would not +//! show that, because the adapter's whole job is readiness, so the message has to travel. +//! +//! If this file fails to compile, the neutral seam has regressed to naming a runtime. + +#![cfg(feature = "nagoya-transport")] + +use std::os::unix::ffi::OsStrExt; + +use endpoint_libs::libs::ws::WireMessage; +use endpoint_libs::libs::ws::transport::framed::framed_json_neutral; +use endpoint_libs::libs::ws::transport::nagoya::NagoyaStream; +use futures::{SinkExt, StreamExt}; +use nagoya::reactor::{Addr, Reactor, TaskSet, TcpListener, TcpStream, block_on_with}; + +/// A socket path that removes itself, so a failed run does not poison the next one. +struct SocketPath(std::path::PathBuf); + +impl SocketPath { + fn new(name: &str) -> Self { + let path = std::env::temp_dir().join(format!("elibs-nagoya-{name}-{}", std::process::id())); + let _ = std::fs::remove_file(&path); + Self(path) + } +} + +impl Drop for SocketPath { + fn drop(&mut self) { + let _ = std::fs::remove_file(&self.0); + } +} + +#[test] +fn a_wire_message_crosses_a_nagoya_unix_socket_in_both_directions() { + let path = SocketPath::new("roundtrip"); + let addr = Addr::path(path.0.as_os_str().as_bytes()).expect("the path fits a sun_path"); + + let reactor = Reactor::local().expect("a reactor"); + let handle = reactor.handle(); + let listener = TcpListener::bind(addr, &handle).expect("bind"); + + // Connected before the executor starts, so neither task has to poll for the other's + // existence: a Unix-domain connect finishes inside the call rather than by becoming + // writable later. + let client_socket = + nagoya::reactor::socket::TcpSocket::connect(addr).expect("connect to the listener"); + let client = TcpStream::from_socket(client_socket, &handle).expect("register the client"); + + let request = WireMessage::Text("what did you index".into()); + let reply = WireMessage::Binary(vec![0, 1, 2, 255].into()); + + let expected_request = request.clone(); + let expected_reply = reply.clone(); + + let mut tasks = TaskSet::new(); + tasks.push(async move { + let (server, _) = listener.accept().await.expect("accept"); + let mut framed = framed_json_neutral(NagoyaStream::new(server)); + let received = framed + .next() + .await + .expect("a frame arrives") + .expect("it decodes"); + assert_eq!(received, expected_request); + framed.send(expected_reply).await.expect("send the reply"); + framed.close().await.expect("close cleanly"); + }); + tasks.push(async move { + let mut framed = framed_json_neutral(NagoyaStream::new(client)); + framed.send(request).await.expect("send the request"); + let received = framed + .next() + .await + .expect("a reply arrives") + .expect("it decodes"); + assert_eq!(received, reply); + }); + + block_on_with(&reactor, tasks); +}