From 97516c6416c109875bf7c7e95a791f25f5636422 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 15:04:57 +0900 Subject: [PATCH 01/17] fix: enforce websocket resource limits across transports --- websock-proto/src/lib.rs | 2 +- websock-proto/src/options.rs | 71 ++++++++++++++++++++- websock-tungstenite/src/builder.rs | 14 +++- websock-tungstenite/src/connection.rs | 15 ++++- websock-tungstenite/src/server.rs | 28 +++++--- websock-tungstenite/tests/connection.rs | 85 +++++++++++++++++++++++++ websock-wasm/src/builder.rs | 8 ++- websock-wasm/src/connection.rs | 83 ++++++++++++++++++++---- websock-wasm/src/stream.rs | 27 ++++++-- 9 files changed, 295 insertions(+), 38 deletions(-) create mode 100644 websock-tungstenite/tests/connection.rs diff --git a/websock-proto/src/lib.rs b/websock-proto/src/lib.rs index fd9770e..b99a29b 100644 --- a/websock-proto/src/lib.rs +++ b/websock-proto/src/lib.rs @@ -12,5 +12,5 @@ mod transport; pub use bytes::Bytes; pub use error::{Error, Result}; pub use message::{CloseFrame, Message}; -pub use options::{ConnectOptions, ServerOptions, default_ws_alpn}; +pub use options::{ConnectOptions, ServerOptions, WebSocketLimits, default_ws_alpn}; pub use transport::{LocalBoxFuture, WebSocketConnection}; diff --git a/websock-proto/src/options.rs b/websock-proto/src/options.rs index 187e372..0699912 100644 --- a/websock-proto/src/options.rs +++ b/websock-proto/src/options.rs @@ -1,5 +1,56 @@ pub const ALPN_HTTP_1_1: &[u8] = b"http/1.1"; -pub const ALPN_H2: &[u8] = b"h2"; + +/// Resource limits applied by WebSocket transports. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct WebSocketLimits { + /// Maximum accepted WebSocket message size in bytes. + pub max_message_size: usize, + /// Maximum accepted WebSocket frame payload size in bytes. + /// + /// Browsers do not expose individual frames, so this limit is native-only. + pub max_frame_size: usize, + /// Maximum native transport write-buffer size in bytes. + /// + /// Browsers manage their own write buffers and ignore this value. + pub max_write_buffer_size: usize, +} + +impl Default for WebSocketLimits { + fn default() -> Self { + Self { + max_message_size: 16 * 1024 * 1024, + max_frame_size: 4 * 1024 * 1024, + max_write_buffer_size: 16 * 1024 * 1024, + } + } +} + +impl WebSocketLimits { + /// Validate that all limits are non-zero and internally consistent. + pub fn validate(&self) -> crate::Result<()> { + if self.max_message_size == 0 { + return Err(crate::Error::Protocol( + "max_message_size must be greater than zero".into(), + )); + } + if self.max_frame_size == 0 { + return Err(crate::Error::Protocol( + "max_frame_size must be greater than zero".into(), + )); + } + if self.max_frame_size > self.max_message_size { + return Err(crate::Error::Protocol( + "max_frame_size must not exceed max_message_size".into(), + )); + } + if self.max_write_buffer_size == 0 { + return Err(crate::Error::Protocol( + "max_write_buffer_size must be greater than zero".into(), + )); + } + Ok(()) + } +} /// Connection configuration shared by native and WebAssembly transports. #[derive(Debug, Clone, Default)] @@ -9,6 +60,9 @@ pub struct ConnectOptions { /// Additional headers (native only). Browser WebSockets typically ignore these by design. pub headers: Vec<(String, String)>, + + /// Resource limits for the WebSocket transport. + pub limits: WebSocketLimits, } /// Server configuration shared by native transports. @@ -19,8 +73,21 @@ pub struct ServerOptions { /// Additional response headers (native only). pub headers: Vec<(String, String)>, + + /// Resource limits for accepted WebSocket connections. + pub limits: WebSocketLimits, } pub fn default_ws_alpn() -> Vec> { - vec![ALPN_HTTP_1_1.to_vec(), ALPN_H2.to_vec()] + vec![ALPN_HTTP_1_1.to_vec()] +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn default_alpn_only_advertises_supported_http_version() { + assert_eq!(default_ws_alpn(), vec![ALPN_HTTP_1_1.to_vec()]); + } } diff --git a/websock-tungstenite/src/builder.rs b/websock-tungstenite/src/builder.rs index ec71c4c..27e3fdf 100644 --- a/websock-tungstenite/src/builder.rs +++ b/websock-tungstenite/src/builder.rs @@ -6,7 +6,7 @@ use rustls::{ClientConfig, RootCertStore, ServerConfig}; use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; use std::sync::Arc; use websock_proto::default_ws_alpn; -use websock_proto::{ConnectOptions, Error, Result, ServerOptions}; +use websock_proto::{ConnectOptions, Error, Result, ServerOptions, WebSocketLimits}; /// Builder for creating a WebSocket client. /// @@ -64,6 +64,12 @@ impl ClientBuilder { self } + /// Configure WebSocket message, frame, and write-buffer limits. + pub fn with_limits(mut self, limits: WebSocketLimits) -> Self { + self.opts.limits = limits; + self + } + /// Add a single subprotocol. pub fn with_protocol(mut self, protocol: impl Into) -> Self { self.opts.protocols.push(protocol.into()); @@ -238,6 +244,12 @@ impl ServerBuilder { self } + /// Configure accepted WebSocket message, frame, and write-buffer limits. + pub fn with_limits(mut self, limits: WebSocketLimits) -> Self { + self.opts.limits = limits; + self + } + /// Add a single subprotocol to the allowed set. pub fn with_protocol(mut self, protocol: impl Into) -> Self { self.opts.protocols.push(protocol.into()); diff --git a/websock-tungstenite/src/connection.rs b/websock-tungstenite/src/connection.rs index 757fce8..867a7f6 100644 --- a/websock-tungstenite/src/connection.rs +++ b/websock-tungstenite/src/connection.rs @@ -29,6 +29,14 @@ pub async fn connect_with_tls( opts: ConnectOptions, tls: Option>, ) -> Result { + opts.limits.validate()?; + let write_buffer_size = (128 * 1024).min(opts.limits.max_write_buffer_size.saturating_sub(1)); + let config = tungstenite::protocol::WebSocketConfig::default() + .read_buffer_size((128 * 1024).min(opts.limits.max_frame_size)) + .write_buffer_size(write_buffer_size) + .max_write_buffer_size(opts.limits.max_write_buffer_size) + .max_message_size(Some(opts.limits.max_message_size)) + .max_frame_size(Some(opts.limits.max_frame_size)); let mut req = url .into_client_request() .map_err(|e| Error::InvalidUrl(e.to_string()))?; @@ -54,9 +62,10 @@ pub async fn connect_with_tls( } let connector = tls.map(Connector::Rustls); - let (ws, _resp) = tokio_tungstenite::connect_async_tls_with_config(req, None, false, connector) - .await - .map_err(map_tungstenite_err)?; + let (ws, _resp) = + tokio_tungstenite::connect_async_tls_with_config(req, Some(config), false, connector) + .await + .map_err(map_tungstenite_err)?; let info = ConnectionInfo { peer: ws diff --git a/websock-tungstenite/src/server.rs b/websock-tungstenite/src/server.rs index 0d04dbe..0a379d2 100644 --- a/websock-tungstenite/src/server.rs +++ b/websock-tungstenite/src/server.rs @@ -5,7 +5,7 @@ use std::sync::Arc; use tokio::io::{AsyncRead, AsyncWrite}; use tokio::net::{TcpListener, ToSocketAddrs}; use tokio_rustls::TlsAcceptor; -use tokio_tungstenite::tungstenite; +use tokio_tungstenite::{accept_hdr_async_with_config, tungstenite}; use tungstenite::handshake::server::{Request, Response}; use tungstenite::http::header::{HeaderName, HeaderValue, SEC_WEBSOCKET_PROTOCOL}; use websock_proto::{Error, Result, ServerOptions}; @@ -22,6 +22,7 @@ pub async fn bind( where A: ToSocketAddrs, { + opts.limits.validate()?; let listener = TcpListener::bind(addr) .await .map_err(|e| Error::Io(e.to_string()))?; @@ -29,12 +30,20 @@ where validate_protocols(&opts)?; let acceptor = tls.map(|cfg| TlsAcceptor::from(Arc::new(cfg))); + let write_buffer_size = (128 * 1024).min(opts.limits.max_write_buffer_size.saturating_sub(1)); + let config = tungstenite::protocol::WebSocketConfig::default() + .read_buffer_size((128 * 1024).min(opts.limits.max_frame_size)) + .write_buffer_size(write_buffer_size) + .max_write_buffer_size(opts.limits.max_write_buffer_size) + .max_message_size(Some(opts.limits.max_message_size)) + .max_frame_size(Some(opts.limits.max_frame_size)); Ok(Server { listener, protocols: Arc::new(opts.protocols.into_iter().collect()), headers: Arc::new(headers), acceptor, + config, }) } @@ -52,6 +61,7 @@ pub struct Server { protocols: Arc>, headers: Arc>, acceptor: Option, + config: tungstenite::protocol::WebSocketConfig, } impl Server { @@ -85,7 +95,7 @@ impl Server { let headers = Arc::clone(&self.headers); let protocols = Arc::clone(&self.protocols); - let ws = tokio_tungstenite::accept_hdr_async( + let ws = accept_hdr_async_with_config( stream, move |req: &Request, mut resp: Response| { for (name, value) in headers.iter() { @@ -100,6 +110,7 @@ impl Server { Ok(resp) }, + Some(self.config), ) .await .map_err(map_tungstenite_err)?; @@ -142,7 +153,7 @@ impl Server { is_tls: true, }; - let ws = tokio_tungstenite::accept_hdr_async( + let ws = accept_hdr_async_with_config( tls_stream, move |req: &Request, mut resp: Response| { for (name, value) in headers.iter() { @@ -157,6 +168,7 @@ impl Server { Ok(resp) }, + Some(self.config), ) .await .map_err(map_tungstenite_err)?; @@ -201,10 +213,8 @@ fn select_protocol<'a>(req: &'a Request, allowed: &HashSet) -> Option<&' } let header = req.headers().get(SEC_WEBSOCKET_PROTOCOL)?; let header = header.to_str().ok()?; - for candidate in header.split(',').map(|s| s.trim()) { - if allowed.contains(candidate) { - return Some(candidate); - } - } - None + header + .split(',') + .map(str::trim) + .find(|candidate| allowed.contains(*candidate)) } diff --git a/websock-tungstenite/tests/connection.rs b/websock-tungstenite/tests/connection.rs new file mode 100644 index 0000000..a06822e --- /dev/null +++ b/websock-tungstenite/tests/connection.rs @@ -0,0 +1,85 @@ +use websock_proto::{Error, Message, WebSocketLimits}; +use websock_tungstenite::{ClientBuilder, ServerBuilder}; + +#[tokio::test] +async fn client_server_round_trip_text_and_binary() { + let server = ServerBuilder::new() + .with_protocol("test.v1") + .build() + .await + .expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let mut connection = server.accept().await.expect("accept connection"); + assert_eq!( + connection.recv().await.expect("receive text"), + Message::Text("hello".into()) + ); + connection + .send(Message::Binary(b"world".as_slice().into())) + .await + .expect("send binary"); + connection.close().await.expect("close connection"); + }); + + let client = ClientBuilder::new().with_protocol("test.v1").build(); + let mut connection = client + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + connection + .send(Message::Text("hello".into())) + .await + .expect("send text"); + assert_eq!( + connection.recv().await.expect("receive binary"), + Message::Binary(b"world".as_slice().into()) + ); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn server_rejects_message_above_configured_limit() { + let limits = WebSocketLimits { + max_message_size: 64, + max_frame_size: 64, + max_write_buffer_size: 1024, + }; + let server = ServerBuilder::new() + .with_limits(limits) + .build() + .await + .expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let mut connection = server.accept().await.expect("accept connection"); + let err = connection + .recv() + .await + .expect_err("oversized message must fail"); + assert!(matches!(err, Error::Protocol(_))); + }); + + let client = ClientBuilder::new().build(); + let mut connection = client + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + let _ = connection.send(Message::Binary(vec![0; 65].into())).await; + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn invalid_server_limits_return_an_error() { + let limits = WebSocketLimits { + max_message_size: 0, + ..WebSocketLimits::default() + }; + let err = match ServerBuilder::new().with_limits(limits).build().await { + Ok(_) => panic!("invalid limits must fail"), + Err(err) => err, + }; + assert!(matches!(err, Error::Protocol(_))); +} diff --git a/websock-wasm/src/builder.rs b/websock-wasm/src/builder.rs index 56bf44d..f619193 100644 --- a/websock-wasm/src/builder.rs +++ b/websock-wasm/src/builder.rs @@ -1,6 +1,6 @@ //! Builders for browser WebSocket clients. -use websock_proto::{ConnectOptions, Error, Result}; +use websock_proto::{ConnectOptions, Error, Result, WebSocketLimits}; use crate::Connection; use crate::connection::connect; @@ -57,6 +57,12 @@ impl ClientBuilder { self } + /// Configure WebSocket message resource limits. + pub fn with_limits(mut self, limits: WebSocketLimits) -> Self { + self.opts.limits = limits; + self + } + /// Add a single subprotocol. pub fn with_protocol(mut self, protocol: impl Into) -> Self { self.opts.protocols.push(protocol.into()); diff --git a/websock-wasm/src/connection.rs b/websock-wasm/src/connection.rs index e503090..d31d564 100644 --- a/websock-wasm/src/connection.rs +++ b/websock-wasm/src/connection.rs @@ -12,6 +12,8 @@ use wasm_bindgen::prelude::*; /// Establish a browser WebSocket connection. pub async fn connect(url: &str, opts: ConnectOptions) -> Result { + opts.limits.validate()?; + let max_message_size = opts.limits.max_message_size; let ws = if opts.protocols.is_empty() { web_sys::WebSocket::new(url).map_err(js_err)? } else { @@ -25,7 +27,7 @@ pub async fn connect(url: &str, opts: ConnectOptions) -> Result { ws.set_binary_type(web_sys::BinaryType::Arraybuffer); // Channel used to deliver messages to the consumer. - let (tx, rx) = mpsc::unbounded::>(); + let (tx, rx) = mpsc::channel::>(64); // Handle the connection process. let (open_tx, open_rx) = oneshot::channel::>(); @@ -77,44 +79,79 @@ pub async fn connect(url: &str, opts: ConnectOptions) -> Result { } // Set up message/error/close handlers. - let tx_msg = tx.clone(); + let mut tx_msg = tx.clone(); + let ws_onmessage = ws.clone(); let onmessage = Closure::::new(move |e: web_sys::MessageEvent| { let data = e.data(); if let Some(s) = data.as_string() { - let _ = tx_msg.unbounded_send(Ok(Message::Text(s))); + if s.len() > max_message_size { + let _ = tx_msg.try_send(Err(Error::Protocol( + "websocket message exceeds max_message_size".into(), + ))); + tx_msg.close_channel(); + let _ = ws_onmessage.close(); + return; + } + if tx_msg.try_send(Ok(Message::Text(s))).is_err() { + tx_msg.close_channel(); + let _ = ws_onmessage.close(); + } return; } if data.is_instance_of::() { let ab: js_sys::ArrayBuffer = data.unchecked_into(); let u8arr = js_sys::Uint8Array::new(&ab); + if u8arr.length() as usize > max_message_size { + let _ = tx_msg.try_send(Err(Error::Protocol( + "websocket message exceeds max_message_size".into(), + ))); + tx_msg.close_channel(); + let _ = ws_onmessage.close(); + return; + } let mut buf = vec![0u8; u8arr.length() as usize]; u8arr.copy_to(&mut buf); - let _ = tx_msg.unbounded_send(Ok(Message::Binary(Bytes::from(buf)))); + if tx_msg + .try_send(Ok(Message::Binary(Bytes::from(buf)))) + .is_err() + { + tx_msg.close_channel(); + let _ = ws_onmessage.close(); + } return; } - let _ = tx_msg.unbounded_send(Err(Error::Protocol("unsupported message type".into()))); + let _ = tx_msg.try_send(Err(Error::Protocol("unsupported message type".into()))); + tx_msg.close_channel(); + let _ = ws_onmessage.close(); }); ws.set_onmessage(Some(onmessage.as_ref().unchecked_ref())); - let tx_err = tx.clone(); + let mut tx_err = tx.clone(); let onerror = Closure::::new(move |_e: web_sys::Event| { - let _ = tx_err.unbounded_send(Err(Error::Other("websocket error".into()))); + if tx_err + .try_send(Err(Error::Other("websocket error".into()))) + .is_err() + { + tx_err.close_channel(); + } }); ws.set_onerror(Some(onerror.as_ref().unchecked_ref())); - let tx_close = tx.clone(); + let mut tx_close = tx; let onclose = Closure::::new(move |_e: web_sys::CloseEvent| { - let _ = tx_close.unbounded_send(Err(Error::Closed)); + let _ = tx_close.try_send(Err(Error::Closed)); + tx_close.close_channel(); }); ws.set_onclose(Some(onclose.as_ref().unchecked_ref())); Ok(Connection { ws: Rc::new(ws), rx: Some(rx), + max_write_buffer_size: opts.limits.max_write_buffer_size, _onmessage: Some(onmessage), _onerror: Some(onerror), _onclose: Some(onclose), @@ -124,7 +161,8 @@ pub async fn connect(url: &str, opts: ConnectOptions) -> Result { /// WebSocket connection wrapper for browser WebSockets. pub struct Connection { pub(crate) ws: Rc, - pub(crate) rx: Option>>, + pub(crate) rx: Option>>, + pub(crate) max_write_buffer_size: usize, pub(crate) _onmessage: Option>, pub(crate) _onerror: Option>, @@ -135,8 +173,14 @@ impl Connection { /// Send a text or binary message. pub async fn send(&mut self, msg: Message) -> Result<()> { match msg { - Message::Text(s) => self.ws.send_with_str(&s).map_err(js_err)?, - Message::Binary(b) => self.ws.send_with_u8_array(b.as_ref()).map_err(js_err)?, + Message::Text(s) => { + check_send_capacity(&self.ws, s.len(), self.max_write_buffer_size)?; + self.ws.send_with_str(&s).map_err(js_err)?; + } + Message::Binary(b) => { + check_send_capacity(&self.ws, b.len(), self.max_write_buffer_size)?; + self.ws.send_with_u8_array(b.as_ref()).map_err(js_err)?; + } } Ok(()) } @@ -144,8 +188,7 @@ impl Connection { /// Receive the next text or binary message. pub async fn recv(&mut self) -> Result { let rx = self.rx.as_mut().ok_or(Error::Closed)?; - let item = rx.next().await.ok_or(Error::Closed)?; - item + rx.next().await.ok_or(Error::Closed)? } /// Close the WebSocket connection. @@ -195,3 +238,15 @@ impl Drop for Connection { pub(crate) fn js_err(e: JsValue) -> Error { Error::Other(format!("{e:?}")) } + +pub(crate) fn check_send_capacity( + ws: &web_sys::WebSocket, + message_len: usize, + max_write_buffer_size: usize, +) -> Result<()> { + let buffered = usize::try_from(ws.buffered_amount()).unwrap_or(usize::MAX); + if buffered.saturating_add(message_len) > max_write_buffer_size { + return Err(Error::Other("websocket write buffer limit exceeded".into())); + } + Ok(()) +} diff --git a/websock-wasm/src/stream.rs b/websock-wasm/src/stream.rs index 090f184..b7b968a 100644 --- a/websock-wasm/src/stream.rs +++ b/websock-wasm/src/stream.rs @@ -1,7 +1,7 @@ //! Sink/Stream split helpers for browser WebSocket connections. use crate::Connection; -use crate::connection::js_err; +use crate::connection::{check_send_capacity, js_err}; use futures_channel::mpsc; use futures_core::Stream; use futures_sink::Sink; @@ -17,15 +17,21 @@ pub struct ConnectionSink { ws: Rc, closed: bool, ref_count: Rc>, + max_write_buffer_size: usize, } impl ConnectionSink { /// Create a sink backed by the provided WebSocket instance. - fn new(ws: Rc, ref_count: Rc>) -> Self { + fn new( + ws: Rc, + ref_count: Rc>, + max_write_buffer_size: usize, + ) -> Self { Self { ws, closed: false, ref_count, + max_write_buffer_size, } } } @@ -50,8 +56,14 @@ impl Sink for ConnectionSink { return Err(Error::Closed); } match item { - Message::Text(s) => this.ws.send_with_str(&s).map_err(js_err)?, - Message::Binary(b) => this.ws.send_with_u8_array(b.as_ref()).map_err(js_err)?, + Message::Text(s) => { + check_send_capacity(&this.ws, s.len(), this.max_write_buffer_size)?; + this.ws.send_with_str(&s).map_err(js_err)?; + } + Message::Binary(b) => { + check_send_capacity(&this.ws, b.len(), this.max_write_buffer_size)?; + this.ws.send_with_u8_array(b.as_ref()).map_err(js_err)?; + } } Ok(()) } @@ -90,7 +102,7 @@ impl Drop for ConnectionSink { /// Stream wrapper for receiving messages from a browser WebSocket. pub struct ConnectionStream { ws: Rc, - rx: mpsc::UnboundedReceiver>, + rx: mpsc::Receiver>, terminated: bool, ref_count: Rc>, _onmessage: Closure, @@ -102,7 +114,7 @@ impl ConnectionStream { /// Create a stream backed by the provided WebSocket and receiver. fn new( ws: Rc, - rx: mpsc::UnboundedReceiver>, + rx: mpsc::Receiver>, ref_count: Rc>, onmessage: Closure, onerror: Closure, @@ -164,10 +176,11 @@ pub fn split(mut conn: Connection) -> (ConnectionSink, ConnectionStream) { let ws_for_sink = Rc::clone(&conn.ws); let ws_for_stream = Rc::clone(&conn.ws); + let max_write_buffer_size = conn.max_write_buffer_size; let ref_count = Rc::new(Cell::new(2)); ( - ConnectionSink::new(ws_for_sink, Rc::clone(&ref_count)), + ConnectionSink::new(ws_for_sink, Rc::clone(&ref_count), max_write_buffer_size), ConnectionStream::new(ws_for_stream, rx, ref_count, onmessage, onerror, onclose), ) } From fd0728c179f411d87d6de684cc6db75347653998 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 15:05:54 +0900 Subject: [PATCH 02/17] fix: harden mux flow control and stream lifecycle --- websock-mux-proto/src/stream.rs | 7 +- websock-mux-proto/tests/frame.rs | 6 + websock-tungstenite-mux/src/client.rs | 14 +- websock-tungstenite-mux/src/server.rs | 6 +- websock-tungstenite-mux/src/session.rs | 559 +++++++++++++++++++---- websock-tungstenite-mux/tests/session.rs | 64 +++ websock-wasm-demo/src/echo_mux.rs | 6 +- websock-wasm-mux/Cargo.toml | 4 +- websock-wasm-mux/src/client.rs | 3 +- websock-wasm-mux/src/session.rs | 458 ++++++++++++++----- 10 files changed, 928 insertions(+), 199 deletions(-) create mode 100644 websock-tungstenite-mux/tests/session.rs diff --git a/websock-mux-proto/src/stream.rs b/websock-mux-proto/src/stream.rs index 33ebb20..152be0a 100644 --- a/websock-mux-proto/src/stream.rs +++ b/websock-mux-proto/src/stream.rs @@ -22,7 +22,7 @@ impl StreamId { StreamDir::Bi => 0, StreamDir::Uni => 1, }; - let value = counter.checked_shl(2).ok_or(VarIntBoundsExceeded)? + let value = counter.checked_mul(4).ok_or(VarIntBoundsExceeded)? | ((dir_bit as u64) << 1) | (initiator as u64); VarInt::from_u64(value)?; @@ -40,6 +40,11 @@ impl StreamId { pub fn initiator_is_server(self) -> bool { self.0 & 1 == 1 } + + /// Return the stream sequence number encoded in this ID. + pub fn counter(self) -> u64 { + self.0 >> 2 + } } #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/websock-mux-proto/tests/frame.rs b/websock-mux-proto/tests/frame.rs index d03e9a2..5597b46 100644 --- a/websock-mux-proto/tests/frame.rs +++ b/websock-mux-proto/tests/frame.rs @@ -73,6 +73,12 @@ fn frame_encoded_len_matches_actual_size() { assert_eq!(frame.encoded_len(), encoded.len()); } +#[test] +fn stream_id_rejects_counter_overflow() { + assert!(StreamId::new(u64::MAX, false, StreamDir::Bi).is_err()); + assert!(StreamId::new(1 << 62, true, StreamDir::Uni).is_err()); +} + #[test] fn frame_decode_unknown_tag() { let mut buf = BytesMut::new(); diff --git a/websock-tungstenite-mux/src/client.rs b/websock-tungstenite-mux/src/client.rs index 7688cd9..bc395e7 100644 --- a/websock-tungstenite-mux/src/client.rs +++ b/websock-tungstenite-mux/src/client.rs @@ -64,6 +64,7 @@ impl Client { url: &str, tls: Option>, ) -> Result { + self.limits.validate()?; validate_client_protocols(&self.opts)?; let mut request = url @@ -88,10 +89,15 @@ impl Client { } let connector = tls.map(Connector::Rustls); - let (stream, response) = - tokio_tungstenite::connect_async_tls_with_config(request, None, false, connector) - .await - .map_err(map_tungstenite_err)?; + let config = self.limits.websocket_config(); + let (stream, response) = tokio_tungstenite::connect_async_tls_with_config( + request, + Some(config), + false, + connector, + ) + .await + .map_err(map_tungstenite_err)?; let proto = negotiated_protocol(&response) .ok_or_else(|| Error::Protocol("missing SEC_WEBSOCKET_PROTOCOL in response".into()))?; diff --git a/websock-tungstenite-mux/src/server.rs b/websock-tungstenite-mux/src/server.rs index 09cc048..6bcfee9 100644 --- a/websock-tungstenite-mux/src/server.rs +++ b/websock-tungstenite-mux/src/server.rs @@ -3,7 +3,7 @@ use std::sync::Arc; use tokio::io::{AsyncRead, AsyncWrite}; use tokio::net::{TcpListener, ToSocketAddrs}; use tokio_rustls::TlsAcceptor; -use tokio_tungstenite::accept_hdr_async; +use tokio_tungstenite::accept_hdr_async_with_config; use tokio_tungstenite::tungstenite; use tokio_tungstenite::tungstenite::handshake::server; use tungstenite::http; @@ -69,6 +69,7 @@ pub async fn bind( where A: ToSocketAddrs, { + limits.validate()?; let listener = TcpListener::bind(addr) .await .map_err(|e| Error::Io(e.to_string()))?; @@ -122,7 +123,7 @@ impl Server { let headers = Arc::clone(&self.headers); let allowed = Arc::clone(&self.allowed); - let ws = accept_hdr_async( + let ws = accept_hdr_async_with_config( stream, move |req: &server::Request, mut resp: server::Response| { // Additional headers from configuration @@ -153,6 +154,7 @@ impl Server { Ok(resp) }, + Some(self.limits.websocket_config()), ) .await .map_err(map_tungstenite_err)?; diff --git a/websock-tungstenite-mux/src/session.rs b/websock-tungstenite-mux/src/session.rs index 5be3723..cba71c8 100644 --- a/websock-tungstenite-mux/src/session.rs +++ b/websock-tungstenite-mux/src/session.rs @@ -2,7 +2,7 @@ use std::collections::HashMap; use std::io; use std::pin::Pin; use std::sync::{ - Arc, + Arc, Mutex as StdMutex, MutexGuard as StdMutexGuard, atomic::{AtomicBool, AtomicU64, Ordering}, }; use std::task::{Context, Poll}; @@ -17,6 +17,7 @@ use tokio_tungstenite::tungstenite; use tokio_util::sync::PollSender; use websock_proto::{Error, Result}; +use websock_mux_proto::VarInt; use websock_mux_proto::stream::{Frame, StreamDir, StreamId}; const MAX_WRITE_CHUNK: usize = 16 * 1024; @@ -58,7 +59,7 @@ impl Default for Limits { recv_event_queue_len: 128, outbound_queue_len: 256, max_batch_frames: 64, - max_batch_bytes: 256 * 1024, + max_batch_bytes: 512 * 1024, initial_stream_window: 512 * 1024, stream_window_update_threshold: 256 * 1024, accept_uni_queue_len: 128, @@ -67,6 +68,69 @@ impl Default for Limits { } } +impl Limits { + /// Validate that the limits are non-zero and internally consistent. + pub fn validate(&self) -> Result<()> { + let non_zero = [ + ("max_ws_message_size", self.max_ws_message_size), + ("max_stream_data_per_frame", self.max_stream_data_per_frame), + ("max_open_streams", self.max_open_streams), + ("recv_event_queue_len", self.recv_event_queue_len), + ("outbound_queue_len", self.outbound_queue_len), + ("max_batch_frames", self.max_batch_frames), + ("max_batch_bytes", self.max_batch_bytes), + ("initial_stream_window", self.initial_stream_window), + ( + "stream_window_update_threshold", + self.stream_window_update_threshold, + ), + ("accept_uni_queue_len", self.accept_uni_queue_len), + ("accept_bi_queue_len", self.accept_bi_queue_len), + ]; + if let Some((name, _)) = non_zero.into_iter().find(|(_, value)| *value == 0) { + return Err(Error::Protocol(format!("{name} must be greater than zero"))); + } + if self.max_stream_data_per_frame > self.max_ws_message_size { + return Err(Error::Protocol( + "max_stream_data_per_frame must not exceed max_ws_message_size".into(), + )); + } + if self.max_batch_bytes > self.max_ws_message_size { + return Err(Error::Protocol( + "max_batch_bytes must not exceed max_ws_message_size".into(), + )); + } + if self.max_batch_bytes < self.max_stream_data_per_frame.saturating_add(33) { + return Err(Error::Protocol( + "max_batch_bytes must accommodate a maximum-size stream frame".into(), + )); + } + if self.stream_window_update_threshold > self.initial_stream_window { + return Err(Error::Protocol( + "stream_window_update_threshold must not exceed initial_stream_window".into(), + )); + } + if self.max_ws_message_size as u64 > VarInt::MAX.into_inner() + || self.initial_stream_window as u64 > VarInt::MAX.into_inner() + { + return Err(Error::Protocol( + "byte and flow-control limits must fit in a mux varint".into(), + )); + } + Ok(()) + } + + pub(crate) fn websocket_config(&self) -> tungstenite::protocol::WebSocketConfig { + let buffer_size = self.max_ws_message_size.min(128 * 1024); + tungstenite::protocol::WebSocketConfig::default() + .read_buffer_size(buffer_size) + .write_buffer_size(buffer_size) + .max_write_buffer_size(buffer_size.saturating_add(self.max_ws_message_size)) + .max_message_size(Some(self.max_ws_message_size)) + .max_frame_size(Some(self.max_ws_message_size)) + } +} + #[derive(Clone)] pub struct Session { inner: Arc, @@ -83,6 +147,7 @@ impl Session { where S: AsyncRead + AsyncWrite + Unpin + Send + 'static, { + limits.validate()?; let (outbound_tx, outbound_rx) = mpsc::channel(limits.outbound_queue_len); let (accept_uni_tx, accept_uni_rx) = mpsc::channel(limits.accept_uni_queue_len); let (accept_bi_tx, accept_bi_rx) = mpsc::channel(limits.accept_bi_queue_len); @@ -107,23 +172,28 @@ impl Session { pub async fn open_uni(&self) -> Result { let id = self.inner.next_stream_id(StreamDir::Uni)?; - let flow = self - .inner - .register_send_flow(id, self.inner.limits.initial_stream_window as u64) - .await; - self.inner.send_frame(Frame::OpenUni { id }).await?; + let flow = self.inner.register_send_flow(id, 0).await?; + if let Err(err) = self.inner.send_frame(Frame::OpenUni { id }).await { + self.inner.remove_send_flow(id); + return Err(err); + } Ok(SendStream::new(id, self.inner.clone(), flow)) } pub async fn open_bi(&self) -> Result<(SendStream, RecvStream)> { let id = self.inner.next_stream_id(StreamDir::Bi)?; - let flow = self - .inner - .register_send_flow(id, self.inner.limits.initial_stream_window as u64) - .await; + let flow = self.inner.register_send_flow(id, 0).await?; let recv = self.inner.register_recv_stream(id).await; - self.inner.send_frame(Frame::OpenBi { id }).await?; - self.inner.send_initial_credit(id).await?; + if let Err(err) = self.inner.send_frame(Frame::OpenBi { id }).await { + self.inner.remove_stream(id).await; + self.inner.remove_send_flow(id); + return Err(err); + } + if let Err(err) = self.inner.send_initial_credit(id).await { + self.inner.remove_stream(id).await; + self.inner.remove_send_flow(id); + return Err(err); + } Ok((SendStream::new(id, self.inner.clone(), flow), recv)) } @@ -141,6 +211,7 @@ impl Session { struct SendFlowState { max_data: AtomicU64, sent_data: AtomicU64, + closed: AtomicBool, waker: AtomicWaker, } @@ -149,6 +220,7 @@ impl SendFlowState { Self { max_data: AtomicU64::new(initial_max), sent_data: AtomicU64::new(0), + closed: AtomicBool::new(false), waker: AtomicWaker::new(), } } @@ -185,6 +257,9 @@ impl SendFlowState { } fn update_max(&self, max: u64) { + if self.closed.load(Ordering::Acquire) { + return; + } let mut current = self.max_data.load(Ordering::Acquire); while max > current { match self @@ -199,6 +274,15 @@ impl SendFlowState { } } } + + fn close(&self) { + self.closed.store(true, Ordering::Release); + self.waker.wake(); + } + + fn is_closed(&self) -> bool { + self.closed.load(Ordering::Acquire) + } } pub struct SendStream { @@ -248,7 +332,10 @@ impl SendStream { self.flow.waker.register(cx.waker()); let n = self.flow.try_reserve(wanted); if n == 0 { - if self.finished.load(Ordering::SeqCst) || self.session.is_closed() { + if self.finished.load(Ordering::SeqCst) + || self.flow.is_closed() + || self.session.is_closed() + { Poll::Ready(Err(Error::Closed)) } else { Poll::Pending @@ -282,6 +369,12 @@ impl SendStream { } pub async fn finish(&self) -> Result<()> { + if self.finished.load(Ordering::SeqCst) { + return Ok(()); + } + if self.flow.is_closed() || self.session.is_closed() { + return Err(Error::Closed); + } if self .finished .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst) @@ -300,6 +393,8 @@ impl SendStream { } pub async fn reset(&self, code: u64) -> Result<()> { + VarInt::from_u64(code) + .map_err(|_| Error::Protocol("reset code exceeds mux varint range".into()))?; self.finished.store(true, Ordering::SeqCst); self.session.remove_send_flow(self.id); self.session @@ -308,7 +403,7 @@ impl SendStream { } pub fn closed(&self) -> bool { - self.finished.load(Ordering::SeqCst) + self.finished.load(Ordering::SeqCst) || self.flow.is_closed() || self.session.is_closed() } } @@ -337,7 +432,8 @@ impl AsyncWrite for SendStream { if buf.is_empty() { return Poll::Ready(Ok(0)); } - if this.finished.load(Ordering::SeqCst) || this.session.is_closed() { + if this.finished.load(Ordering::SeqCst) || this.flow.is_closed() || this.session.is_closed() + { return Poll::Ready(Err(io_closed())); } @@ -351,7 +447,11 @@ impl AsyncWrite for SendStream { } let chunk_len = this.flow.try_reserve(wanted_len); if chunk_len == 0 { - return Poll::Pending; + return if this.flow.is_closed() || this.session.is_closed() { + Poll::Ready(Err(io_closed())) + } else { + Poll::Pending + }; } match this.outbound.poll_reserve(cx) { @@ -422,6 +522,11 @@ impl AsyncWrite for SendStream { fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { let this = self.get_mut(); + if !this.finished.load(Ordering::SeqCst) + && (this.flow.is_closed() || this.session.is_closed()) + { + return Poll::Ready(Err(io_closed())); + } if !this.finished.load(Ordering::SeqCst) && !this.fin_queued { match this.outbound.poll_reserve(cx) { @@ -449,6 +554,9 @@ impl AsyncWrite for SendStream { impl Drop for SendStream { fn drop(&mut self) { + if Arc::strong_count(&self.finished) != 1 { + return; + } self.session.remove_send_flow(self.id); if !self.finished.load(Ordering::SeqCst) { let _ = self.session.try_send_frame(Frame::ResetStream { @@ -465,6 +573,12 @@ struct RecvEvent { fin: bool, } +struct RecvState { + sender: mpsc::Sender, + received: u64, + max_data: Arc, +} + pub struct RecvStream { id: StreamId, session: Arc, @@ -475,6 +589,8 @@ pub struct RecvStream { granted: u64, initial_window: u64, update_threshold: u64, + max_data: Arc, + stop_sent: AtomicBool, } impl RecvStream { @@ -484,6 +600,7 @@ impl RecvStream { receiver: mpsc::Receiver, initial_window: u64, update_threshold: u64, + max_data: Arc, ) -> Self { Self { id, @@ -495,6 +612,8 @@ impl RecvStream { granted: initial_window, initial_window, update_threshold, + max_data, + stop_sent: AtomicBool::new(false), } } @@ -503,7 +622,10 @@ impl RecvStream { return; } self.consumed = self.consumed.saturating_add(n as u64); - let target = self.consumed.saturating_add(self.initial_window); + let target = self + .consumed + .saturating_add(self.initial_window) + .min(VarInt::MAX.into_inner()); if target <= self.granted { return; } @@ -519,14 +641,15 @@ impl RecvStream { .is_ok() { self.granted = target; + self.max_data.store(target, Ordering::Release); } } pub async fn read(&mut self, buf: &mut [u8]) -> Result> { - if self.finished { - return Ok(None); - } if self.pending.is_empty() { + if self.finished { + return Ok(None); + } if let Some(chunk) = self.read_chunk_internal().await? { self.pending = chunk; } else { @@ -541,13 +664,13 @@ impl RecvStream { } pub async fn read_buf(&mut self, buf: &mut B) -> Result> { - if self.finished { - return Ok(None); - } if buf.remaining_mut() == 0 { return Ok(Some(0)); } if self.pending.is_empty() { + if self.finished { + return Ok(None); + } if let Some(chunk) = self.read_chunk_internal().await? { self.pending = chunk; } else { @@ -582,6 +705,20 @@ impl RecvStream { } pub async fn read_chunk(&mut self, max: usize) -> Result> { + if max == 0 { + return Err(Error::Protocol( + "read_chunk max must be greater than zero".into(), + )); + } + if !self.pending.is_empty() { + let amount = self.pending.len().min(max); + let chunk = self.pending.split_to(amount); + self.on_bytes_consumed(chunk.len()); + return Ok(Some(chunk)); + } + if self.finished { + return Ok(None); + } match self.receiver.recv().await { Some(mut event) => { if event.fin { @@ -607,6 +744,12 @@ impl RecvStream { } pub async fn stop(&self, code: u64) -> Result<()> { + VarInt::from_u64(code) + .map_err(|_| Error::Protocol("stop code exceeds mux varint range".into()))?; + if self.stop_sent.swap(true, Ordering::SeqCst) { + return Ok(()); + } + self.session.remove_recv_stream(self.id); self.session .send_frame(Frame::StopSending { id: self.id, code }) .await @@ -619,7 +762,8 @@ impl RecvStream { impl Drop for RecvStream { fn drop(&mut self) { - if !self.finished { + self.session.remove_recv_stream(self.id); + if !self.finished && !self.stop_sent.swap(true, Ordering::SeqCst) { let _ = self.session.try_send_frame(Frame::StopSending { id: self.id, code: 0, @@ -694,10 +838,12 @@ pub(crate) struct SessionInner { outbound_tx: mpsc::Sender, accept_uni_tx: Mutex>>, accept_bi_tx: Mutex>>, - streams: Mutex>>, - send_flows: Mutex>>, + streams: StdMutex>, + send_flows: StdMutex>>, next_uni: AtomicU64, next_bi: AtomicU64, + next_peer_uni: AtomicU64, + next_peer_bi: AtomicU64, closed: AtomicBool, } @@ -715,10 +861,12 @@ impl SessionInner { outbound_tx, accept_uni_tx: Mutex::new(Some(accept_uni_tx)), accept_bi_tx: Mutex::new(Some(accept_bi_tx)), - streams: Mutex::new(HashMap::new()), - send_flows: Mutex::new(HashMap::new()), + streams: StdMutex::new(HashMap::new()), + send_flows: StdMutex::new(HashMap::new()), next_uni: AtomicU64::new(0), next_bi: AtomicU64::new(0), + next_peer_uni: AtomicU64::new(0), + next_peer_bi: AtomicU64::new(0), closed: AtomicBool::new(false), } } @@ -865,6 +1013,11 @@ impl SessionInner { if id.initiator_is_server() != self.peer_is_server() { return self.protocol_error(1, "OpenUni with wrong initiator").await; } + if !self.validate_peer_stream_id(id) { + return self + .protocol_error(1, "OpenUni with non-monotonic StreamId") + .await; + } let recv = match self.try_register_inbound_recv_stream(id).await { Ok(recv) => recv, @@ -878,7 +1031,7 @@ impl SessionInner { Err(mpsc::error::TrySendError::Full(_)) => { // Application is not accepting inbound streams fast enough. // Reset this stream and keep the connection alive. - let mut streams = self.streams.lock().await; + let mut streams = self.lock_streams(); streams.remove(&id); let _ = self.try_send_frame(Frame::ResetStream { id, code: 3 }); return Ok(()); @@ -896,14 +1049,20 @@ impl SessionInner { if id.initiator_is_server() != self.peer_is_server() { return self.protocol_error(1, "OpenBi with wrong initiator").await; } + if !self.validate_peer_stream_id(id) { + return self + .protocol_error(1, "OpenBi with non-monotonic StreamId") + .await; + } let recv = match self.try_register_inbound_recv_stream(id).await { Ok(recv) => recv, Err((code, reason)) => return self.protocol_error(code, reason).await, }; - let flow = self - .register_send_flow(id, self.limits.initial_stream_window as u64) - .await; + let flow = match self.register_send_flow(id, 0).await { + Ok(flow) => flow, + Err(err) => return self.protocol_error(3, err.to_string()).await, + }; let send = SendStream::new(id, self.clone(), flow); let tx = { self.accept_bi_tx.lock().await.clone() }; @@ -911,7 +1070,7 @@ impl SessionInner { match tx.try_send((send, recv)) { Ok(()) => {} Err(mpsc::error::TrySendError::Full(_)) => { - let mut streams = self.streams.lock().await; + let mut streams = self.lock_streams(); streams.remove(&id); let _ = self.try_send_frame(Frame::ResetStream { id, code: 3 }); return Ok(()); @@ -927,47 +1086,63 @@ impl SessionInner { return self.protocol_error(2, "stream data too large").await; } let tx = { - let streams = self.streams.lock().await; - streams.get(&id).cloned() + let mut streams = self.lock_streams(); + match streams.get_mut(&id) { + None => Err("Stream data on unknown stream"), + Some(state) => match state.received.checked_add(data.len() as u64) { + None => Err("stream data overflow"), + Some(received) if received > state.max_data.load(Ordering::Acquire) => { + Err("stream flow-control limit exceeded") + } + Some(received) => { + state.received = received; + Ok(state.sender.clone()) + } + }, + } }; - let Some(tx) = tx else { - return self - .protocol_error(1, "Stream data on unknown stream") - .await; + let tx = match tx { + Ok(tx) => tx, + Err(reason) => return self.protocol_error(2, reason).await, }; - if tx.send(RecvEvent { data, fin }).await.is_err() { - let mut streams = self.streams.lock().await; - streams.remove(&id); - return Ok(()); + match tx.try_send(RecvEvent { data, fin }) { + Ok(()) => {} + Err(mpsc::error::TrySendError::Full(_)) => { + self.remove_recv_stream(id); + let _ = self.try_send_frame(Frame::ResetStream { id, code: 3 }); + return Ok(()); + } + Err(mpsc::error::TrySendError::Closed(_)) => { + self.remove_recv_stream(id); + return Ok(()); + } } if fin { - self.remove_send_flow(id); - let mut streams = self.streams.lock().await; + let mut streams = self.lock_streams(); streams.remove(&id); } } - Frame::ResetStream { id, .. } | Frame::StopSending { id, .. } => { - let mut streams = self.streams.lock().await; - let removed_recv = streams.remove(&id).is_some(); - drop(streams); - let mut send_flows = self.send_flows.lock().await; - let removed_send = send_flows.remove(&id).is_some(); - if !removed_recv && !removed_send { - return self.protocol_error(1, "reset/stop on unknown stream").await; + Frame::ResetStream { id, .. } => { + let removed = { self.lock_streams().remove(&id).is_some() }; + if !removed { + return self.protocol_error(1, "reset on unknown stream").await; + } + } + Frame::StopSending { id, .. } => { + if !self.remove_send_flow(id) { + return self.protocol_error(1, "stop on unknown stream").await; } } Frame::MaxStreamData { id, max } => { - let flows = self.send_flows.lock().await; + let flows = self.lock_send_flows(); if let Some(flow) = flows.get(&id) { flow.update_max(max); } } Frame::ConnectionClose { .. } => { - let mut streams = self.streams.lock().await; - streams.clear(); - let mut send_flows = self.send_flows.lock().await; - send_flows.clear(); + self.close_all().await; + return Err(Error::Closed); } } Ok(()) @@ -978,15 +1153,24 @@ impl SessionInner { id: StreamId, ) -> std::result::Result { let (tx, rx) = mpsc::channel(self.limits.recv_event_queue_len); - let mut streams = self.streams.lock().await; - if streams.len() >= self.limits.max_open_streams { - return Err((3, "too many open streams")); - } - if streams.contains_key(&id) { - return Err((1, "duplicate stream open")); + let max_data = Arc::new(AtomicU64::new(self.limits.initial_stream_window as u64)); + { + let mut streams = self.lock_streams(); + if streams.len() >= self.limits.max_open_streams { + return Err((3, "too many open streams")); + } + if streams.contains_key(&id) { + return Err((1, "duplicate stream open")); + } + streams.insert( + id, + RecvState { + sender: tx, + received: 0, + max_data: max_data.clone(), + }, + ); } - streams.insert(id, tx); - drop(streams); let initial_window = self.limits.initial_stream_window as u64; let recv = RecvStream::new( @@ -995,6 +1179,7 @@ impl SessionInner { rx, initial_window, self.limits.stream_window_update_threshold as u64, + max_data, ); self.send_frame(Frame::MaxStreamData { id, @@ -1007,8 +1192,16 @@ impl SessionInner { pub(crate) async fn register_recv_stream(self: &Arc, id: StreamId) -> RecvStream { let (tx, rx) = mpsc::channel(self.limits.recv_event_queue_len); - let mut streams = self.streams.lock().await; - streams.insert(id, tx); + let max_data = Arc::new(AtomicU64::new(self.limits.initial_stream_window as u64)); + let mut streams = self.lock_streams(); + streams.insert( + id, + RecvState { + sender: tx, + received: 0, + max_data: max_data.clone(), + }, + ); drop(streams); let initial_window = self.limits.initial_stream_window as u64; RecvStream::new( @@ -1017,6 +1210,7 @@ impl SessionInner { rx, initial_window, self.limits.stream_window_update_threshold as u64, + max_data, ) } @@ -1036,19 +1230,40 @@ impl SessionInner { StreamId::new(counter, self.is_server, dir).map_err(|e| Error::StreamId(e.to_string())) } - async fn register_send_flow(&self, id: StreamId, initial_max: u64) -> Arc { + async fn register_send_flow( + &self, + id: StreamId, + initial_max: u64, + ) -> Result> { let flow = Arc::new(SendFlowState::new(initial_max)); - let mut send_flows = self.send_flows.lock().await; + let mut send_flows = self.lock_send_flows(); + if send_flows.len() >= self.limits.max_open_streams { + return Err(Error::Protocol("too many open send streams".into())); + } + if send_flows.contains_key(&id) { + return Err(Error::Protocol("duplicate send stream".into())); + } send_flows.insert(id, flow.clone()); - flow + Ok(flow) } - fn remove_send_flow(&self, id: StreamId) { - if let Ok(mut send_flows) = self.send_flows.try_lock() { - send_flows.remove(&id); + fn remove_send_flow(&self, id: StreamId) -> bool { + if let Some(flow) = self.lock_send_flows().remove(&id) { + flow.close(); + true + } else { + false } } + async fn remove_stream(&self, id: StreamId) { + self.remove_recv_stream(id); + } + + fn remove_recv_stream(&self, id: StreamId) -> bool { + self.lock_streams().remove(&id).is_some() + } + pub(crate) async fn send_frame(&self, frame: Frame) -> Result<()> { if self.is_closed() { return Err(Error::Closed); @@ -1083,11 +1298,14 @@ impl SessionInner { } // Close existing streams { - let mut streams = self.streams.lock().await; + let mut streams = self.lock_streams(); streams.clear(); } { - let mut send_flows = self.send_flows.lock().await; + let mut send_flows = self.lock_send_flows(); + for flow in send_flows.values() { + flow.close(); + } send_flows.clear(); } } @@ -1112,6 +1330,40 @@ impl SessionInner { fn peer_is_server(&self) -> bool { !self.is_server } + + fn lock_send_flows(&self) -> StdMutexGuard<'_, HashMap>> { + self.send_flows + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + + fn lock_streams(&self) -> StdMutexGuard<'_, HashMap> { + self.streams + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + + fn validate_peer_stream_id(&self, id: StreamId) -> bool { + let next = match id.dir() { + StreamDir::Uni => &self.next_peer_uni, + StreamDir::Bi => &self.next_peer_bi, + }; + let mut current = next.load(Ordering::SeqCst); + loop { + if id.counter() < current { + return false; + } + match next.compare_exchange( + current, + id.counter().saturating_add(1), + Ordering::SeqCst, + Ordering::SeqCst, + ) { + Ok(_) => return true, + Err(actual) => current = actual, + } + } + } } impl websock_mux_proto::MuxSendStream for SendStream { @@ -1258,7 +1510,10 @@ mod tests { accept_bi_tx, )); let id = StreamId::new(0, false, StreamDir::Bi).expect("stream id"); - let flow = inner.register_send_flow(id, u64::MAX).await; + let flow = inner + .register_send_flow(id, u64::MAX) + .await + .expect("register flow"); let mut send = SendStream::new(id, inner.clone(), flow); let mut recv = inner.register_recv_stream(id).await; @@ -1355,7 +1610,10 @@ mod tests { accept_bi_tx, )); let id = StreamId::new(0, false, StreamDir::Bi).expect("stream id"); - let flow = inner.register_send_flow(id, u64::MAX).await; + let flow = inner + .register_send_flow(id, u64::MAX) + .await + .expect("register flow"); let mut send = SendStream::new(id, inner, flow); let first = AsyncWriteExt::write(&mut send, b"x") @@ -1377,4 +1635,147 @@ mod tests { .expect("second write result"); assert_eq!(second_ok, 1); } + + #[test] + fn limits_reject_zero_and_inconsistent_values() { + let limits = Limits { + outbound_queue_len: 0, + ..Limits::default() + }; + assert!(matches!(limits.validate(), Err(Error::Protocol(_)))); + + let mut limits = Limits::default(); + limits.stream_window_update_threshold = limits.initial_stream_window + 1; + assert!(matches!(limits.validate(), Err(Error::Protocol(_)))); + } + + #[tokio::test] + async fn receive_flow_control_violation_closes_session() { + let limits = Limits::default(); + let (outbound_tx, _outbound_rx) = mpsc::channel(8); + let (accept_uni_tx, _accept_uni_rx) = mpsc::channel(8); + let (accept_bi_tx, _accept_bi_rx) = mpsc::channel(8); + let inner = Arc::new(SessionInner::new( + false, + limits.clone(), + outbound_tx, + accept_uni_tx, + accept_bi_tx, + )); + let id = StreamId::new(0, false, StreamDir::Bi).expect("stream id"); + let _recv = inner.register_recv_stream(id).await; + + for _ in 0..2 { + inner + .handle_frame(Frame::Stream { + id, + data: Bytes::from(vec![0; limits.max_stream_data_per_frame]), + fin: false, + }) + .await + .expect("data within advertised credit"); + } + + let err = inner + .handle_frame(Frame::Stream { + id, + data: Bytes::from_static(b"x"), + fin: false, + }) + .await + .expect_err("credit violation must fail"); + assert!(matches!(err, Error::Protocol(_))); + assert!(inner.is_closed()); + } + + #[tokio::test] + async fn peer_stream_ids_must_be_monotonic() { + let limits = Limits::default(); + let (outbound_tx, _outbound_rx) = mpsc::channel(8); + let (accept_uni_tx, _accept_uni_rx) = mpsc::channel(8); + let (accept_bi_tx, _accept_bi_rx) = mpsc::channel(8); + let inner = Arc::new(SessionInner::new( + false, + limits, + outbound_tx, + accept_uni_tx, + accept_bi_tx, + )); + let first = StreamId::new(1, true, StreamDir::Uni).expect("stream id"); + inner + .handle_frame(Frame::OpenUni { id: first }) + .await + .expect("forward stream id is valid"); + let stale = StreamId::new(0, true, StreamDir::Uni).expect("stream id"); + let err = inner + .handle_frame(Frame::OpenUni { id: stale }) + .await + .expect_err("stale stream id must fail"); + assert!(matches!(err, Error::Protocol(_))); + } + + #[tokio::test] + async fn final_frame_data_survives_partial_chunk_reads() { + let limits = Limits::default(); + let (outbound_tx, _outbound_rx) = mpsc::channel(8); + let (accept_uni_tx, _accept_uni_rx) = mpsc::channel(8); + let (accept_bi_tx, _accept_bi_rx) = mpsc::channel(8); + let inner = Arc::new(SessionInner::new( + false, + limits, + outbound_tx, + accept_uni_tx, + accept_bi_tx, + )); + let id = StreamId::new(0, false, StreamDir::Bi).expect("stream id"); + let mut recv = inner.register_recv_stream(id).await; + inner + .handle_frame(Frame::Stream { + id, + data: Bytes::from_static(b"hello"), + fin: true, + }) + .await + .expect("final frame"); + + assert_eq!( + recv.read_chunk(2).await.expect("first chunk").as_deref(), + Some(b"he".as_slice()) + ); + assert_eq!( + recv.read_chunk(2).await.expect("second chunk").as_deref(), + Some(b"ll".as_slice()) + ); + assert_eq!( + recv.read_chunk(2).await.expect("third chunk").as_deref(), + Some(b"o".as_slice()) + ); + assert!(recv.read_chunk(2).await.expect("end of stream").is_none()); + } + + #[tokio::test] + async fn dropping_a_send_stream_clone_does_not_reset_the_stream() { + let limits = Limits::default(); + let (outbound_tx, mut outbound_rx) = mpsc::channel(8); + let (accept_uni_tx, _accept_uni_rx) = mpsc::channel(8); + let (accept_bi_tx, _accept_bi_rx) = mpsc::channel(8); + let inner = Arc::new(SessionInner::new( + false, + limits, + outbound_tx, + accept_uni_tx, + accept_bi_tx, + )); + let id = StreamId::new(0, false, StreamDir::Uni).expect("stream id"); + let flow = inner + .register_send_flow(id, 64) + .await + .expect("register flow"); + let send = SendStream::new(id, inner.clone(), flow); + drop(send.clone()); + + assert!(inner.lock_send_flows().contains_key(&id)); + assert!(outbound_rx.try_recv().is_err()); + send.write(b"x").await.expect("remaining clone is usable"); + } } diff --git a/websock-tungstenite-mux/tests/session.rs b/websock-tungstenite-mux/tests/session.rs new file mode 100644 index 0000000..9f22582 --- /dev/null +++ b/websock-tungstenite-mux/tests/session.rs @@ -0,0 +1,64 @@ +use websock_tungstenite_mux::{ClientBuilder, Limits, ServerBuilder}; + +fn limits(window: usize) -> Limits { + Limits { + max_ws_message_size: 1024, + max_stream_data_per_frame: 64, + max_open_streams: 8, + recv_event_queue_len: 4, + outbound_queue_len: 8, + max_batch_frames: 4, + max_batch_bytes: 128, + initial_stream_window: window, + stream_window_update_threshold: window / 2, + accept_uni_queue_len: 4, + accept_bi_queue_len: 4, + } +} + +#[tokio::test] +async fn different_peer_windows_interoperate_with_backpressure() { + let server = ServerBuilder::new() + .with_limits(limits(64)) + .build() + .await + .expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let session = server.accept().await.expect("accept session"); + let mut stream = session.accept_uni().await.expect("accept uni stream"); + let mut received = Vec::new(); + let mut buffer = [0_u8; 31]; + while let Some(read) = stream.read(&mut buffer).await.expect("read stream") { + received.extend_from_slice(&buffer[..read]); + } + received + }); + + let client = ClientBuilder::new().with_limits(limits(128)).build(); + let session = client + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + let stream = session.open_uni().await.expect("open uni stream"); + let payload = vec![42_u8; 256]; + stream.write_all(&payload).await.expect("write payload"); + stream.finish().await.expect("finish stream"); + + assert_eq!(server_task.await.expect("server task"), payload); +} + +#[tokio::test] +async fn invalid_limits_fail_before_connecting() { + let invalid = Limits { + outbound_queue_len: 0, + ..Limits::default() + }; + let client = ClientBuilder::new().with_limits(invalid).build(); + let err = match client.connect("ws://127.0.0.1:1").await { + Ok(_) => panic!("invalid limits must fail"), + Err(err) => err, + }; + assert!(matches!(err, websock_proto::Error::Protocol(_))); +} diff --git a/websock-wasm-demo/src/echo_mux.rs b/websock-wasm-demo/src/echo_mux.rs index 6fd5dc6..f6f5daa 100644 --- a/websock-wasm-demo/src/echo_mux.rs +++ b/websock-wasm-demo/src/echo_mux.rs @@ -26,8 +26,10 @@ pub async fn run_mux_bi_demo(url: &str, log: impl Fn(&str)) { let (send, mut recv) = session.open_bi().expect("open_bi failed"); - send.write(b"hello mux from wasm").expect("send failed"); - send.finish().expect("finish failed"); + send.write(b"hello mux from wasm") + .await + .expect("send failed"); + send.finish().await.expect("finish failed"); let mut buf = vec![0u8; 1024]; while let Some(n) = recv.read(&mut buf).await.expect("read failed") { diff --git a/websock-wasm-mux/Cargo.toml b/websock-wasm-mux/Cargo.toml index ff20d66..5377932 100644 --- a/websock-wasm-mux/Cargo.toml +++ b/websock-wasm-mux/Cargo.toml @@ -15,7 +15,7 @@ websock-proto = { workspace = true } websock-mux-proto = { workspace = true } websock-wasm = { workspace = true } bytes = { workspace = true } -futures-channel = "0.3" +futures-channel = { version = "0.3", features = ["sink"] } futures-io = "0.3" -futures-util = { version = "0.3", default-features = false, features = ["alloc"] } +futures-util = { version = "0.3", default-features = false, features = ["alloc", "sink"] } wasm-bindgen-futures = "0.4" diff --git a/websock-wasm-mux/src/client.rs b/websock-wasm-mux/src/client.rs index 02015d8..444252e 100644 --- a/websock-wasm-mux/src/client.rs +++ b/websock-wasm-mux/src/client.rs @@ -18,7 +18,8 @@ impl Client { /// Establish a browser WebSocket connection and create a mux [`Session`]. pub async fn connect(&self, url: &str) -> Result { + self.limits.validate()?; let conn = websock_wasm::connect(url, self.opts.clone()).await?; - Ok(Session::new(conn, self.limits.clone())) + Session::new(conn, self.limits.clone()) } } diff --git a/websock-wasm-mux/src/session.rs b/websock-wasm-mux/src/session.rs index 9be9f17..f363d8a 100644 --- a/websock-wasm-mux/src/session.rs +++ b/websock-wasm-mux/src/session.rs @@ -12,11 +12,11 @@ use futures_io::{AsyncRead as FuturesAsyncRead, AsyncWrite as FuturesAsyncWrite} use futures_util::lock::Mutex; use futures_util::stream::Stream; use futures_util::task::AtomicWaker; -use futures_util::{FutureExt, StreamExt}; +use futures_util::{FutureExt, SinkExt, StreamExt, future::poll_fn}; use wasm_bindgen_futures::spawn_local; use websock_proto::{Error, Message, Result}; -use websock_mux_proto::{Frame, StreamDir, StreamId}; +use websock_mux_proto::{Frame, StreamDir, StreamId, VarInt}; const MAX_WRITE_CHUNK: usize = 16 * 1024; @@ -50,13 +50,13 @@ pub struct Limits { impl Default for Limits { fn default() -> Self { Self { - max_ws_message_size: 1 * 1024 * 1024, + max_ws_message_size: 1024 * 1024, max_stream_data_per_frame: 256 * 1024, max_open_streams: 1024, recv_event_queue_len: 128, outbound_queue_len: 256, max_batch_frames: 64, - max_batch_bytes: 256 * 1024, + max_batch_bytes: 512 * 1024, initial_stream_window: 512 * 1024, stream_window_update_threshold: 256 * 1024, accept_uni_queue_len: 128, @@ -65,6 +65,59 @@ impl Default for Limits { } } +impl Limits { + /// Validate that the limits are non-zero and internally consistent. + pub fn validate(&self) -> Result<()> { + let non_zero = [ + ("max_ws_message_size", self.max_ws_message_size), + ("max_stream_data_per_frame", self.max_stream_data_per_frame), + ("max_open_streams", self.max_open_streams), + ("recv_event_queue_len", self.recv_event_queue_len), + ("outbound_queue_len", self.outbound_queue_len), + ("max_batch_frames", self.max_batch_frames), + ("max_batch_bytes", self.max_batch_bytes), + ("initial_stream_window", self.initial_stream_window), + ( + "stream_window_update_threshold", + self.stream_window_update_threshold, + ), + ("accept_uni_queue_len", self.accept_uni_queue_len), + ("accept_bi_queue_len", self.accept_bi_queue_len), + ]; + if let Some((name, _)) = non_zero.into_iter().find(|(_, value)| *value == 0) { + return Err(Error::Protocol(format!("{name} must be greater than zero"))); + } + if self.max_stream_data_per_frame > self.max_ws_message_size { + return Err(Error::Protocol( + "max_stream_data_per_frame must not exceed max_ws_message_size".into(), + )); + } + if self.max_batch_bytes > self.max_ws_message_size { + return Err(Error::Protocol( + "max_batch_bytes must not exceed max_ws_message_size".into(), + )); + } + if self.max_batch_bytes < self.max_stream_data_per_frame.saturating_add(33) { + return Err(Error::Protocol( + "max_batch_bytes must accommodate a maximum-size stream frame".into(), + )); + } + if self.stream_window_update_threshold > self.initial_stream_window { + return Err(Error::Protocol( + "stream_window_update_threshold must not exceed initial_stream_window".into(), + )); + } + if self.max_ws_message_size as u64 > VarInt::MAX.into_inner() + || self.initial_stream_window as u64 > VarInt::MAX.into_inner() + { + return Err(Error::Protocol( + "byte and flow-control limits must fit in a mux varint".into(), + )); + } + Ok(()) + } +} + #[derive(Clone)] pub struct Session { inner: Rc, @@ -73,7 +126,8 @@ pub struct Session { } impl Session { - pub(crate) fn new(conn: websock_wasm::Connection, limits: Limits) -> Self { + pub(crate) fn new(conn: websock_wasm::Connection, limits: Limits) -> Result { + limits.validate()?; let (outbound_tx, outbound_rx) = mpsc::channel::(limits.outbound_queue_len); let (accept_uni_tx, accept_uni_rx) = mpsc::channel::(limits.accept_uni_queue_len); @@ -94,29 +148,36 @@ impl Session { }; inner.spawn_task(conn, outbound_rx); - session + Ok(session) } pub fn open_uni(&self) -> Result { let id = self.inner.next_stream_id(StreamDir::Uni)?; - let flow = self - .inner - .register_send_flow(id, self.inner.limits.initial_stream_window as u64); - self.inner.send_frame(Frame::OpenUni { id })?; + let flow = self.inner.register_send_flow(id, 0)?; + if let Err(err) = self.inner.send_frame(Frame::OpenUni { id }) { + self.inner.remove_send_flow(id); + return Err(err); + } Ok(SendStream::new(id, self.inner.clone(), flow)) } pub fn open_bi(&self) -> Result<(SendStream, RecvStream)> { let id = self.inner.next_stream_id(StreamDir::Bi)?; - let flow = self - .inner - .register_send_flow(id, self.inner.limits.initial_stream_window as u64); + let flow = self.inner.register_send_flow(id, 0)?; let recv = self.inner.clone().register_recv_stream(id); - self.inner.send_frame(Frame::OpenBi { id })?; - self.inner.send_frame(Frame::MaxStreamData { + if let Err(err) = self.inner.send_frame(Frame::OpenBi { id }) { + self.inner.streams.borrow_mut().remove(&id); + self.inner.remove_send_flow(id); + return Err(err); + } + if let Err(err) = self.inner.send_frame(Frame::MaxStreamData { id, max: self.inner.limits.initial_stream_window as u64, - })?; + }) { + self.inner.streams.borrow_mut().remove(&id); + self.inner.remove_send_flow(id); + return Err(err); + } Ok((SendStream::new(id, self.inner.clone(), flow), recv)) } @@ -134,6 +195,7 @@ impl Session { struct SendFlowState { max_data: AtomicU64, sent_data: AtomicU64, + closed: AtomicBool, waker: AtomicWaker, } @@ -142,6 +204,7 @@ impl SendFlowState { Self { max_data: AtomicU64::new(initial_max), sent_data: AtomicU64::new(0), + closed: AtomicBool::new(false), waker: AtomicWaker::new(), } } @@ -177,6 +240,9 @@ impl SendFlowState { } fn update_max(&self, max: u64) { + if self.closed.load(Ordering::Acquire) { + return; + } let mut current = self.max_data.load(Ordering::Acquire); while max > current { match self @@ -191,6 +257,15 @@ impl SendFlowState { } } } + + fn close(&self) { + self.closed.store(true, Ordering::Release); + self.waker.wake(); + } + + fn is_closed(&self) -> bool { + self.closed.load(Ordering::Acquire) + } } pub struct SendStream { @@ -217,12 +292,15 @@ impl SendStream { } } - pub fn write(&self, data: &[u8]) -> Result<()> { - self.write_buf(Bytes::copy_from_slice(data)) + pub async fn write(&self, data: &[u8]) -> Result<()> { + self.write_buf(Bytes::copy_from_slice(data)).await } - pub fn write_buf(&self, data: Bytes) -> Result<()> { - if self.finished.load(Ordering::SeqCst) || self.session.closed.load(Ordering::SeqCst) { + pub async fn write_buf(&self, data: Bytes) -> Result<()> { + if self.finished.load(Ordering::SeqCst) + || self.flow.is_closed() + || self.session.closed.load(Ordering::SeqCst) + { return Err(Error::Closed); } let mut offset = 0usize; @@ -233,29 +311,53 @@ impl SendStream { if wanted == 0 { return Err(Error::Protocol("stream frame payload limit is zero".into())); } - let grant = self.flow.try_reserve(wanted); - if grant == 0 { - return Err(Error::Other("flow control blocked".into())); - } + let grant = poll_fn(|cx| { + self.flow.waker.register(cx.waker()); + let grant = self.flow.try_reserve(wanted); + if grant == 0 { + if self.flow.is_closed() + || self.finished.load(Ordering::SeqCst) + || self.session.closed.load(Ordering::SeqCst) + { + Poll::Ready(Err(Error::Closed)) + } else { + Poll::Pending + } + } else { + Poll::Ready(Ok(grant)) + } + }) + .await?; let chunk = data.slice(offset..offset + grant); - if let Err(err) = self.session.send_frame(Frame::Stream { - id: self.id, - data: chunk, - fin: false, - }) { + let mut outbound = self.outbound.clone(); + if outbound + .send(OutboundCmd::Frame(Frame::Stream { + id: self.id, + data: chunk, + fin: false, + })) + .await + .is_err() + { self.flow.release(grant); - return Err(err); + return Err(Error::Closed); } offset += grant; } Ok(()) } - pub fn write_all(&self, data: &[u8]) -> Result<()> { - self.write(data) + pub async fn write_all(&self, data: &[u8]) -> Result<()> { + self.write(data).await } - pub fn finish(&self) -> Result<()> { + pub async fn finish(&self) -> Result<()> { + if self.finished.load(Ordering::SeqCst) { + return Ok(()); + } + if self.flow.is_closed() || self.session.closed.load(Ordering::SeqCst) { + return Err(Error::Closed); + } if self .finished .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst) @@ -271,7 +373,9 @@ impl SendStream { Ok(()) } - pub fn reset(&self, code: u64) -> Result<()> { + pub async fn reset(&self, code: u64) -> Result<()> { + VarInt::from_u64(code) + .map_err(|_| Error::Protocol("reset code exceeds mux varint range".into()))?; self.finished.store(true, Ordering::SeqCst); self.session.remove_send_flow(self.id); self.session @@ -280,6 +384,8 @@ impl SendStream { pub fn closed(&self) -> bool { self.finished.load(Ordering::SeqCst) + || self.flow.is_closed() + || self.session.closed.load(Ordering::SeqCst) } } @@ -307,7 +413,10 @@ impl FuturesAsyncWrite for SendStream { if buf.is_empty() { return Poll::Ready(Ok(0)); } - if this.finished.load(Ordering::SeqCst) || this.session.closed.load(Ordering::SeqCst) { + if this.finished.load(Ordering::SeqCst) + || this.flow.is_closed() + || this.session.closed.load(Ordering::SeqCst) + { return Poll::Ready(Err(io_closed())); } @@ -322,7 +431,11 @@ impl FuturesAsyncWrite for SendStream { } let chunk_len = this.flow.try_reserve(wanted); if chunk_len == 0 { - return Poll::Pending; + return if this.flow.is_closed() || this.session.closed.load(Ordering::SeqCst) { + Poll::Ready(Err(io_closed())) + } else { + Poll::Pending + }; } match this.outbound.poll_ready(cx) { @@ -373,6 +486,11 @@ impl FuturesAsyncWrite for SendStream { fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { let this = self.get_mut(); + if !this.finished.load(Ordering::SeqCst) + && (this.flow.is_closed() || this.session.closed.load(Ordering::SeqCst)) + { + return Poll::Ready(Err(io_closed())); + } if !this.finished.load(Ordering::SeqCst) && !this.close_in_flight { match this.outbound.poll_ready(cx) { @@ -404,6 +522,9 @@ impl FuturesAsyncWrite for SendStream { impl Drop for SendStream { fn drop(&mut self) { + if Rc::strong_count(&self.finished) != 1 { + return; + } self.session.remove_send_flow(self.id); if !self.finished.load(Ordering::SeqCst) { let _ = self.session.try_send_frame(Frame::ResetStream { @@ -420,6 +541,12 @@ struct RecvEvent { fin: bool, } +struct RecvState { + sender: mpsc::Sender, + received: u64, + max_data: Rc, +} + pub struct RecvStream { id: StreamId, session: Rc, @@ -430,6 +557,8 @@ pub struct RecvStream { granted: u64, initial_window: u64, update_threshold: u64, + max_data: Rc, + stop_sent: AtomicBool, } impl RecvStream { @@ -439,6 +568,7 @@ impl RecvStream { receiver: mpsc::Receiver, initial_window: u64, update_threshold: u64, + max_data: Rc, ) -> Self { Self { id, @@ -450,6 +580,8 @@ impl RecvStream { granted: initial_window, initial_window, update_threshold, + max_data, + stop_sent: AtomicBool::new(false), } } @@ -458,7 +590,10 @@ impl RecvStream { return; } self.consumed = self.consumed.saturating_add(n as u64); - let target = self.consumed.saturating_add(self.initial_window); + let target = self + .consumed + .saturating_add(self.initial_window) + .min(VarInt::MAX.into_inner()); if target <= self.granted { return; } @@ -474,14 +609,15 @@ impl RecvStream { .is_ok() { self.granted = target; + self.max_data.store(target, Ordering::Release); } } pub async fn read(&mut self, buf: &mut [u8]) -> Result> { - if self.finished { - return Ok(None); - } if self.pending.is_empty() { + if self.finished { + return Ok(None); + } if let Some(chunk) = self.read_chunk_internal().await? { self.pending = chunk; } else { @@ -496,13 +632,13 @@ impl RecvStream { } pub async fn read_buf(&mut self, buf: &mut B) -> Result> { - if self.finished { - return Ok(None); - } if buf.remaining_mut() == 0 { return Ok(Some(0)); } if self.pending.is_empty() { + if self.finished { + return Ok(None); + } if let Some(chunk) = self.read_chunk_internal().await? { self.pending = chunk; } else { @@ -537,6 +673,20 @@ impl RecvStream { } pub async fn read_chunk(&mut self, max: usize) -> Result> { + if max == 0 { + return Err(Error::Protocol( + "read_chunk max must be greater than zero".into(), + )); + } + if !self.pending.is_empty() { + let amount = self.pending.len().min(max); + let chunk = self.pending.split_to(amount); + self.on_bytes_consumed(chunk.len()); + return Ok(Some(chunk)); + } + if self.finished { + return Ok(None); + } match self.receiver.next().await { Some(mut event) => { if event.fin { @@ -562,6 +712,12 @@ impl RecvStream { } pub fn stop(&self, code: u64) -> Result<()> { + VarInt::from_u64(code) + .map_err(|_| Error::Protocol("stop code exceeds mux varint range".into()))?; + if self.stop_sent.swap(true, Ordering::SeqCst) { + return Ok(()); + } + self.session.streams.borrow_mut().remove(&self.id); self.session .send_frame(Frame::StopSending { id: self.id, code }) } @@ -573,7 +729,8 @@ impl RecvStream { impl Drop for RecvStream { fn drop(&mut self) { - if !self.finished { + self.session.streams.borrow_mut().remove(&self.id); + if !self.finished && !self.stop_sent.swap(true, Ordering::SeqCst) { let _ = self.session.try_send_frame(Frame::StopSending { id: self.id, code: 0, @@ -646,10 +803,12 @@ struct SessionInner { outbound_tx: RefCell>, accept_uni_tx: Mutex>>, accept_bi_tx: Mutex>>, - streams: RefCell>>, + streams: RefCell>, send_flows: RefCell>>, next_uni: AtomicU64, next_bi: AtomicU64, + next_peer_uni: AtomicU64, + next_peer_bi: AtomicU64, closed: AtomicBool, } @@ -669,6 +828,8 @@ impl SessionInner { send_flows: RefCell::new(HashMap::new()), next_uni: AtomicU64::new(0), next_bi: AtomicU64::new(0), + next_peer_uni: AtomicU64::new(0), + next_peer_bi: AtomicU64::new(0), closed: AtomicBool::new(false), } } @@ -784,19 +945,29 @@ impl SessionInner { if !id.initiator_is_server() { return self.protocol_error(1, "OpenUni with wrong initiator").await; } - - let mut map = self.streams.borrow_mut(); - if map.len() >= self.limits.max_open_streams { - drop(map); - return self.protocol_error(3, "too many open streams").await; - } - if map.contains_key(&id) { - drop(map); - return self.protocol_error(1, "duplicate stream open").await; + if !self.validate_peer_stream_id(id) { + return self + .protocol_error(1, "OpenUni with non-monotonic StreamId") + .await; } - let recv = Self::register_recv_stream_locked(&self, &mut map, id); - drop(map); + let validation = { + let map = self.streams.borrow(); + if map.len() >= self.limits.max_open_streams { + Err((3, "too many open streams")) + } else if map.contains_key(&id) { + Err((1, "duplicate stream open")) + } else { + Ok(()) + } + }; + if let Err((code, reason)) = validation { + return self.protocol_error(code, reason).await; + } + let recv = { + let mut map = self.streams.borrow_mut(); + Self::register_recv_stream_locked(self, &mut map, id) + }; let _ = self.try_send_frame(Frame::MaxStreamData { id, max: self.limits.initial_stream_window as u64, @@ -827,25 +998,38 @@ impl SessionInner { if !id.initiator_is_server() { return self.protocol_error(1, "OpenBi with wrong initiator").await; } - - let mut map = self.streams.borrow_mut(); - if map.len() >= self.limits.max_open_streams { - drop(map); - return self.protocol_error(3, "too many open streams").await; - } - if map.contains_key(&id) { - drop(map); - return self.protocol_error(1, "duplicate stream open").await; + if !self.validate_peer_stream_id(id) { + return self + .protocol_error(1, "OpenBi with non-monotonic StreamId") + .await; } - let recv = Self::register_recv_stream_locked(&self, &mut map, id); - drop(map); + let validation = { + let map = self.streams.borrow(); + if map.len() >= self.limits.max_open_streams { + Err((3, "too many open streams")) + } else if map.contains_key(&id) { + Err((1, "duplicate stream open")) + } else { + Ok(()) + } + }; + if let Err((code, reason)) = validation { + return self.protocol_error(code, reason).await; + } + let recv = { + let mut map = self.streams.borrow_mut(); + Self::register_recv_stream_locked(self, &mut map, id) + }; let _ = self.try_send_frame(Frame::MaxStreamData { id, max: self.limits.initial_stream_window as u64, }); - let flow = self.register_send_flow(id, self.limits.initial_stream_window as u64); + let flow = match self.register_send_flow(id, 0) { + Ok(flow) => flow, + Err(err) => return self.protocol_error(3, &err.to_string()).await, + }; let send = SendStream::new(id, self.clone(), flow); let tx = self.accept_bi_tx.lock().await.clone(); if let Some(mut tx) = tx { @@ -870,40 +1054,49 @@ impl SessionInner { return self.protocol_error(2, "stream data too large").await; } - let mut map = self.streams.borrow_mut(); - let Some(tx) = map.get_mut(&id) else { - drop(map); - return self - .protocol_error(1, "Stream data on unknown stream") - .await; - }; - - match tx.try_send(RecvEvent { data, fin }) { - Ok(()) => {} - Err(e) => { - if e.is_full() { - map.remove(&id); - drop(map); - let _ = self.try_send_frame(Frame::ResetStream { id, code: 3 }); - return Ok(()); - } else { - map.remove(&id); - return Ok(()); - } + let result = { + let mut map = self.streams.borrow_mut(); + match map.get_mut(&id) { + None => Err("Stream data on unknown stream"), + Some(state) => match state.received.checked_add(data.len() as u64) { + None => Err("stream data overflow"), + Some(received) if received > state.max_data.load(Ordering::Acquire) => { + Err("stream flow-control limit exceeded") + } + Some(received) => { + state.received = received; + let (remove, reset) = + match state.sender.try_send(RecvEvent { data, fin }) { + Ok(()) => (fin, false), + Err(error) => (true, error.is_full()), + }; + if remove { + map.remove(&id); + }; + Ok(reset) + } + }, } + }; + let reset = match result { + Ok(reset) => reset, + Err(reason) => return self.protocol_error(2, reason).await, + }; + if reset { + let _ = self.try_send_frame(Frame::ResetStream { id, code: 3 }); } - - if fin { - self.remove_send_flow(id); - map.remove(&id); + Ok(()) + } + Frame::ResetStream { id, .. } => { + let removed = { self.streams.borrow_mut().remove(&id).is_some() }; + if !removed { + return self.protocol_error(1, "reset on unknown stream").await; } Ok(()) } - Frame::ResetStream { id, .. } | Frame::StopSending { id, .. } => { - let removed_recv = self.streams.borrow_mut().remove(&id).is_some(); - let removed_send = self.send_flows.borrow_mut().remove(&id).is_some(); - if !removed_recv && !removed_send { - return self.protocol_error(1, "reset/stop on unknown stream").await; + Frame::StopSending { id, .. } => { + if !self.remove_send_flow(id) { + return self.protocol_error(1, "stop on unknown stream").await; } Ok(()) } @@ -927,28 +1120,49 @@ impl SessionInner { fn register_recv_stream_locked( this: &Rc, - map: &mut HashMap>, + map: &mut HashMap, id: StreamId, ) -> RecvStream { let (tx, rx) = mpsc::channel(this.limits.recv_event_queue_len); - map.insert(id, tx); + let max_data = Rc::new(AtomicU64::new(this.limits.initial_stream_window as u64)); + map.insert( + id, + RecvState { + sender: tx, + received: 0, + max_data: max_data.clone(), + }, + ); RecvStream::new( id, this.clone(), rx, this.limits.initial_stream_window as u64, this.limits.stream_window_update_threshold as u64, + max_data, ) } - fn register_send_flow(&self, id: StreamId, initial_max: u64) -> Rc { + fn register_send_flow(&self, id: StreamId, initial_max: u64) -> Result> { let flow = Rc::new(SendFlowState::new(initial_max)); - self.send_flows.borrow_mut().insert(id, flow.clone()); - flow + let mut send_flows = self.send_flows.borrow_mut(); + if send_flows.len() >= self.limits.max_open_streams { + return Err(Error::Protocol("too many open send streams".into())); + } + if send_flows.contains_key(&id) { + return Err(Error::Protocol("duplicate send stream".into())); + } + send_flows.insert(id, flow.clone()); + Ok(flow) } - fn remove_send_flow(&self, id: StreamId) { - self.send_flows.borrow_mut().remove(&id); + fn remove_send_flow(&self, id: StreamId) -> bool { + if let Some(flow) = self.send_flows.borrow_mut().remove(&id) { + flow.close(); + true + } else { + false + } } fn try_send_frame(&self, frame: Frame) -> std::result::Result<(), Error> { @@ -981,23 +1195,51 @@ impl SessionInner { } self.streams.borrow_mut().clear(); - self.send_flows.borrow_mut().clear(); + { + let mut send_flows = self.send_flows.borrow_mut(); + for flow in send_flows.values() { + flow.close(); + } + send_flows.clear(); + } *self.accept_uni_tx.lock().await = None; *self.accept_bi_tx.lock().await = None; } + + fn validate_peer_stream_id(&self, id: StreamId) -> bool { + let next = match id.dir() { + StreamDir::Uni => &self.next_peer_uni, + StreamDir::Bi => &self.next_peer_bi, + }; + let mut current = next.load(Ordering::SeqCst); + loop { + if id.counter() < current { + return false; + } + match next.compare_exchange( + current, + id.counter().saturating_add(1), + Ordering::SeqCst, + Ordering::SeqCst, + ) { + Ok(_) => return true, + Err(actual) => current = actual, + } + } + } } impl websock_mux_proto::MuxSendStream for SendStream { fn write_buf<'a>(&'a self, data: Bytes) -> websock_proto::LocalBoxFuture<'a, Result<()>> { - Box::pin(async move { SendStream::write_buf(self, data) }) + Box::pin(async move { SendStream::write_buf(self, data).await }) } fn finish<'a>(&'a self) -> websock_proto::LocalBoxFuture<'a, Result<()>> { - Box::pin(async move { SendStream::finish(self) }) + Box::pin(async move { SendStream::finish(self).await }) } fn reset<'a>(&'a self, code: u64) -> websock_proto::LocalBoxFuture<'a, Result<()>> { - Box::pin(async move { SendStream::reset(self, code) }) + Box::pin(async move { SendStream::reset(self, code).await }) } fn closed(&self) -> bool { From 2d8aa628dc04f3629f7f469ec0ca763ac3b4d1b6 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 15:06:22 +0900 Subject: [PATCH 03/17] chore: make examples and TLS helpers clippy-clean --- websock-tungstenite/src/tls/cert.rs | 2 +- websock-tungstenite/src/tls/key.rs | 2 +- websock/examples/echo-server.rs | 15 +++++---------- 3 files changed, 7 insertions(+), 12 deletions(-) diff --git a/websock-tungstenite/src/tls/cert.rs b/websock-tungstenite/src/tls/cert.rs index a1cbf0f..f496a50 100644 --- a/websock-tungstenite/src/tls/cert.rs +++ b/websock-tungstenite/src/tls/cert.rs @@ -22,7 +22,7 @@ pub fn get_native_certs() -> Result { pub fn load_certs(cert_path: &Path) -> Result>> { let cert_bytes = fs::read(cert_path).map_err(|e| Error::Io(e.to_string()))?; - if cert_path.extension().map_or(false, |x| x == "der") { + if cert_path.extension().is_some_and(|x| x == "der") { return Ok(vec![CertificateDer::from(cert_bytes)]); } diff --git a/websock-tungstenite/src/tls/key.rs b/websock-tungstenite/src/tls/key.rs index 10c8987..0f45ad2 100644 --- a/websock-tungstenite/src/tls/key.rs +++ b/websock-tungstenite/src/tls/key.rs @@ -8,7 +8,7 @@ use websock_proto::{Error, Result}; pub fn load_key(key_path: &Path) -> Result> { let key = fs::read(key_path).map_err(|e| Error::Io(e.to_string()))?; - let key = if key_path.extension().map_or(false, |x| x == "der") { + let key = if key_path.extension().is_some_and(|x| x == "der") { // Treat raw DER as PKCS#8. PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(key)) } else { diff --git a/websock/examples/echo-server.rs b/websock/examples/echo-server.rs index a009bff..bdec9db 100644 --- a/websock/examples/echo-server.rs +++ b/websock/examples/echo-server.rs @@ -60,17 +60,12 @@ async fn main() -> anyhow::Result<()> { let mut conn = server.accept().await?; tracing::info!("accepted connection: {:?}", conn.peer_addr()); tokio::spawn(async move { - loop { - match conn.recv().await { - Ok(msg) => { - tracing::info!("received message: {:?}", msg); - if conn.send(msg).await.is_err() { - break; - } - tracing::info!("echoed message"); - } - Err(_) => break, + while let Ok(msg) = conn.recv().await { + tracing::info!("received message: {:?}", msg); + if conn.send(msg).await.is_err() { + break; } + tracing::info!("echoed message"); } let _ = conn.close().await; tracing::info!("connection closed"); From 7803cc0d532eb91c7b6aab0ed49729b31bbe7071 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 15:06:42 +0900 Subject: [PATCH 04/17] docs: document resource limits and production readiness --- README.md | 21 +++++++++++++++ TODO.md | 81 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 102 insertions(+) create mode 100644 TODO.md diff --git a/README.md b/README.md index bbdfd3a..6d40c04 100644 --- a/README.md +++ b/README.md @@ -38,6 +38,27 @@ See [examples][examples-url]. The `websock-wasm-demo` crate includes a small browser app that connects to an echo server. +## Resource limits + +WebSocket clients and servers use conservative message, frame, and write-buffer +limits by default. Customize them with `WebSocketLimits`: + +```rust +use websock::{ClientBuilder, WebSocketLimits}; + +let client = ClientBuilder::new() + .with_limits(WebSocketLimits { + max_message_size: 2 * 1024 * 1024, + max_frame_size: 512 * 1024, + max_write_buffer_size: 2 * 1024 * 1024, + }) + .build(); +``` + +The mux transports additionally expose `Limits` for stream counts, queue +capacities, batching, and per-stream flow-control windows. Invalid or +inconsistent limits return a protocol error rather than panicking. + ## Benchmarking Criterion benchmarks are available for `websock-mux-proto`. diff --git a/TODO.md b/TODO.md new file mode 100644 index 0000000..58140a2 --- /dev/null +++ b/TODO.md @@ -0,0 +1,81 @@ +# Production Readiness Review + +Reviewed on 2026-07-18. This file records concrete production-readiness work for +the workspace. Completed items describe fixes made as part of the review; open +items remain release blockers or follow-up work. + +## P0 — Safety and resource bounds + +- [x] Enforce mux receive-side flow control. A peer could send more than + the advertised `MaxStreamData` value, allowing each stream to exceed its + configured memory budget. + - Completion: track cumulative received bytes per stream, reject overflow and + credit violations as protocol errors, and test the boundary and violation. +- [x] Bound the browser WebSocket receive queue. `websock-wasm` used an + unbounded channel, so a fast peer can exhaust browser memory when the + application consumes messages slowly. + - Completion: use a bounded queue and close the socket on overflow. +- [x] Validate every mux `Limits` value before allocating channels or starting + tasks. Zero-capacity channels currently panic, and inconsistent byte/window + settings can deadlock or bypass configured bounds. + - Completion: invalid limits return `Error::Protocol` on both native and WASM + implementations, with regression tests. +- [x] Do not grant mux send credit before the peer advertises it. Using the local + receive-window configuration as peer credit breaks interoperability when + peers use different limits and weakens flow control. + - Completion: new send streams start with zero credit and wake only after a + valid `MaxStreamData` frame. + +## P1 — Protocol correctness and lifecycle + +- [x] Treat `ConnectionClose` as terminal and wake all blocked stream writers. + The previous implementation cleared maps but left the session open. +- [x] Reject non-monotonic peer stream IDs and cap locally opened streams. A + peer could reuse stale IDs, while local streams could grow the flow + map without respecting `max_open_streams`. +- [x] Remove stream state reliably. The native `try_lock` cleanup path could + silently retain entries when the mutex is contended. +- [x] Reset streams whose application receive queues are full instead of + blocking the entire native session behind one slow stream. +- [x] Preserve bidirectional send state after a receive FIN, separate reset and + stop-sending semantics, retain final-frame data across partial reads, and keep + a stream alive when one of several send-stream clones is dropped. +- [x] Reject stream-counter and reset/stop-code values that exceed the 62-bit + mux encoding range instead of wrapping IDs or panicking in background tasks. +- [x] Add native end-to-end tests for connect/subprotocol negotiation, + text/binary round trips, graceful close, oversized-message rejection, mux uni + streams, differing flow-control windows, and invalid-limit rejection. +- [ ] Extend native end-to-end tests with split operation, mux bidirectional + streams, malformed wire frames, reset/stop propagation, and TLS round trips. +- [ ] Add browser tests under `wasm-bindgen-test` for connect failure, queue + overflow, close, split ownership, and mux parity. +- [x] Ensure WebSocket-level frame/message limits are applied in Tungstenite + configuration before a complete oversized message is buffered. +- [x] Expose shared WebSocket message/frame/write-buffer limits and enforce + browser receive queue, message-size, and `bufferedAmount` bounds. +- [x] Advertise only HTTP/1.1 in the default ALPN list because the current + Tungstenite transport does not implement WebSocket over HTTP/2. +- [ ] Add explicit connection/session shutdown APIs with completion semantics + and task cancellation. Dropping the last mux session handle should not leave + detached tasks and a socket alive indefinitely. + +## P2 — API, observability, and release engineering + +- [ ] Preserve structured error sources and I/O error kinds instead of reducing + all transport errors to strings. This is a semver-sensitive public API change. +- [ ] Expose negotiated subprotocol and close-frame details consistently on + native and browser connections. +- [ ] Define and document the mux wire protocol, compatibility policy, error + codes, stream state machine, and flow-control rules. +- [ ] Expand crate-level and public-item documentation and enable + `#![warn(missing_docs)]` incrementally. +- [x] Add baseline CI for formatting, Clippy with warnings denied, tests, + documentation, Linux/macOS/Windows, and `wasm32-unknown-unknown`. +- [ ] Extend CI with an explicit MSRV job, minimal feature-set builds, dependency + policy checks, and security audits. +- [ ] Declare and test the MSRV, add changelog/release guidance, and verify + package contents with `cargo package --list` for every published crate. +- [ ] Review dependency features and duplicate native/mux TLS implementations to + reduce compile time, binary size, and maintenance drift. +- [ ] Add fuzz targets for varint/frame decoding and stateful mux frame + sequences, plus long-running concurrency and backpressure tests. From 66049e882194bc694785f1ac634919f70f3fc3b2 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 15:06:55 +0900 Subject: [PATCH 05/17] ci: add cross-platform Rust and WebAssembly checks --- .github/workflows/ci.yml | 58 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) create mode 100644 .github/workflows/ci.yml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..0694867 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,58 @@ +name: CI + +on: + push: + pull_request: + +permissions: + contents: read + +env: + CARGO_TERM_COLOR: always + RUSTFLAGS: "-Dwarnings" + +jobs: + quality: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 + - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 + with: + components: clippy, rustfmt + - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 + - run: cargo fmt --all -- --check + - run: cargo clippy --workspace --all-targets --all-features -- -D warnings + - run: cargo test --workspace --all-targets + - run: cargo doc --workspace --no-deps + env: + RUSTDOCFLAGS: "-Dwarnings" + + wasm: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 + - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 + with: + targets: wasm32-unknown-unknown + - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 + - name: Check browser crates + run: >- + cargo check + -p websock + -p websock-mux + -p websock-wasm + -p websock-wasm-mux + -p websock-wasm-demo + --target wasm32-unknown-unknown + + platforms: + strategy: + fail-fast: false + matrix: + os: [macos-latest, windows-latest] + runs-on: ${{ matrix.os }} + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 + - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 + - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 + - run: cargo test --workspace --all-targets From 4d6e968f6d90ff94dfb0db72c905837090cbe83e Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 21:23:47 +0900 Subject: [PATCH 06/17] feat: add deterministic session shutdown --- websock-mux-proto/Cargo.toml | 1 + websock-tungstenite-mux/src/session.rs | 143 +++++++++++- websock-tungstenite-mux/tests/session.rs | 271 +++++++++++++++++++++++ websock-wasm-mux/src/session.rs | 80 ++++++- 4 files changed, 483 insertions(+), 12 deletions(-) diff --git a/websock-mux-proto/Cargo.toml b/websock-mux-proto/Cargo.toml index 92b2cd8..32d1c60 100644 --- a/websock-mux-proto/Cargo.toml +++ b/websock-mux-proto/Cargo.toml @@ -2,6 +2,7 @@ name = "websock-mux-proto" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "Protocol for multiplexing WebSocket logical streams" repository = "https://github.com/foctal/websock" diff --git a/websock-tungstenite-mux/src/session.rs b/websock-tungstenite-mux/src/session.rs index cba71c8..9e4f043 100644 --- a/websock-tungstenite-mux/src/session.rs +++ b/websock-tungstenite-mux/src/session.rs @@ -3,7 +3,7 @@ use std::io; use std::pin::Pin; use std::sync::{ Arc, Mutex as StdMutex, MutexGuard as StdMutexGuard, - atomic::{AtomicBool, AtomicU64, Ordering}, + atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, }; use std::task::{Context, Poll}; @@ -12,9 +12,9 @@ use futures_util::future::poll_fn; use futures_util::task::AtomicWaker; use futures_util::{SinkExt, StreamExt}; use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; -use tokio::sync::{Mutex, mpsc, oneshot}; +use tokio::sync::{Mutex, Notify, mpsc, oneshot}; use tokio_tungstenite::tungstenite; -use tokio_util::sync::PollSender; +use tokio_util::sync::{CancellationToken, PollSender}; use websock_proto::{Error, Result}; use websock_mux_proto::VarInt; @@ -131,7 +131,6 @@ impl Limits { } } -#[derive(Clone)] pub struct Session { inner: Arc, accept_uni: Arc>>, @@ -206,6 +205,35 @@ impl Session { let mut rx = self.accept_bi.lock().await; rx.recv().await.ok_or(Error::Closed) } + + /// Gracefully close the WebSocket and wait for all session tasks to finish. + pub async fn shutdown(&self) -> Result<()> { + self.inner.shutdown().await + } + + /// Return whether the session has finished shutting down. + pub fn is_closed(&self) -> bool { + self.inner.is_closed() + } +} + +impl Clone for Session { + fn clone(&self) -> Self { + self.inner.session_handles.fetch_add(1, Ordering::Relaxed); + Self { + inner: self.inner.clone(), + accept_uni: self.accept_uni.clone(), + accept_bi: self.accept_bi.clone(), + } + } +} + +impl Drop for Session { + fn drop(&mut self) { + if self.inner.session_handles.fetch_sub(1, Ordering::AcqRel) == 1 { + self.inner.request_shutdown(); + } + } } struct SendFlowState { @@ -830,6 +858,7 @@ pub(crate) enum OutboundCmd { Frame(Frame), Ws(tungstenite::Message), Flush { ack: oneshot::Sender> }, + Shutdown { ack: oneshot::Sender> }, } pub(crate) struct SessionInner { @@ -845,6 +874,11 @@ pub(crate) struct SessionInner { next_peer_uni: AtomicU64, next_peer_bi: AtomicU64, closed: AtomicBool, + shutdown_started: AtomicBool, + session_handles: AtomicUsize, + active_tasks: AtomicUsize, + tasks_done: Notify, + cancel: CancellationToken, } impl SessionInner { @@ -868,6 +902,11 @@ impl SessionInner { next_peer_uni: AtomicU64::new(0), next_peer_bi: AtomicU64::new(0), closed: AtomicBool::new(false), + shutdown_started: AtomicBool::new(false), + session_handles: AtomicUsize::new(1), + active_tasks: AtomicUsize::new(2), + tasks_done: Notify::new(), + cancel: CancellationToken::new(), } } @@ -882,7 +921,14 @@ impl SessionInner { let inbound = self.clone(); tokio::spawn(async move { - while let Some(msg) = ws_stream.next().await { + loop { + let msg = tokio::select! { + _ = inbound.cancel.cancelled() => break, + msg = ws_stream.next() => msg, + }; + let Some(msg) = msg else { + break; + }; let msg = match msg { Ok(m) => m, Err(_) => break, @@ -924,7 +970,7 @@ impl SessionInner { } } - inbound.close_all().await; + inbound.task_finished().await; }); let outbound = self.clone(); @@ -949,7 +995,20 @@ impl SessionInner { let mut batch = BytesMut::new(); let mut batch_frames = 0usize; - while let Some(cmd) = outbound_rx.recv().await { + loop { + let cmd = tokio::select! { + _ = outbound.cancel.cancelled() => { + let _ = flush_batch(&mut ws_sink, &mut batch).await; + let _ = ws_sink.close().await; + break; + } + cmd = outbound_rx.recv() => cmd, + }; + let Some(cmd) = cmd else { + let _ = flush_batch(&mut ws_sink, &mut batch).await; + let _ = ws_sink.close().await; + break; + }; match cmd { OutboundCmd::Frame(frame) => { let encoded = frame.encode().freeze(); @@ -996,12 +1055,72 @@ impl SessionInner { let flush_res = ws_sink.flush().await.map_err(map_tungstenite_err); let _ = ack.send(flush_res); } + OutboundCmd::Shutdown { ack } => { + let result = if let Err(err) = flush_batch(&mut ws_sink, &mut batch).await { + Err(map_tungstenite_err(err)) + } else { + ws_sink.close().await.map_err(map_tungstenite_err) + }; + let _ = ack.send(result); + break; + } } } - outbound.close_all().await; + outbound.task_finished().await; }); } + fn request_shutdown(&self) { + self.shutdown_started.store(true, Ordering::Release); + self.cancel.cancel(); + } + + async fn shutdown(&self) -> Result<()> { + if self.closed.load(Ordering::Acquire) { + self.wait_for_tasks().await; + return Ok(()); + } + + let first = !self.shutdown_started.swap(true, Ordering::AcqRel); + let result = if first { + let (ack_tx, ack_rx) = oneshot::channel(); + match self + .outbound_tx + .send(OutboundCmd::Shutdown { ack: ack_tx }) + .await + { + Ok(()) => ack_rx.await.unwrap_or(Err(Error::Closed)), + Err(_) => Ok(()), + } + } else { + Ok(()) + }; + + if first { + self.cancel.cancel(); + } + self.wait_for_tasks().await; + result + } + + async fn task_finished(&self) { + self.close_all().await; + self.cancel.cancel(); + if self.active_tasks.fetch_sub(1, Ordering::AcqRel) == 1 { + self.tasks_done.notify_waiters(); + } + } + + async fn wait_for_tasks(&self) { + loop { + let notified = self.tasks_done.notified(); + if self.active_tasks.load(Ordering::Acquire) == 0 { + return; + } + notified.await; + } + } + pub(crate) async fn handle_frame(self: &Arc, frame: Frame) -> Result<()> { match frame { Frame::OpenUni { id } => { @@ -1446,6 +1565,14 @@ impl websock_mux_proto::MuxSession for Session { > { Box::pin(async move { Session::accept_bi(self).await }) } + + fn shutdown<'a>(&'a self) -> websock_proto::LocalBoxFuture<'a, Result<()>> { + Box::pin(async move { Session::shutdown(self).await }) + } + + fn closed(&self) -> bool { + Session::is_closed(self) + } } fn io_closed() -> io::Error { diff --git a/websock-tungstenite-mux/tests/session.rs b/websock-tungstenite-mux/tests/session.rs index 9f22582..ac42a4f 100644 --- a/websock-tungstenite-mux/tests/session.rs +++ b/websock-tungstenite-mux/tests/session.rs @@ -1,3 +1,8 @@ +use futures_util::SinkExt; +use std::time::Duration; +use tokio_tungstenite::tungstenite; +use tungstenite::client::IntoClientRequest; +use tungstenite::http::header::SEC_WEBSOCKET_PROTOCOL; use websock_tungstenite_mux::{ClientBuilder, Limits, ServerBuilder}; fn limits(window: usize) -> Limits { @@ -62,3 +67,269 @@ async fn invalid_limits_fail_before_connecting() { }; assert!(matches!(err, websock_proto::Error::Protocol(_))); } + +#[tokio::test] +async fn bidirectional_stream_round_trip_and_shutdown() { + let server = ServerBuilder::new() + .with_limits(limits(128)) + .build() + .await + .expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let session = server.accept().await.expect("accept session"); + let (send, mut recv) = session.accept_bi().await.expect("accept bi stream"); + let mut request = [0_u8; 4]; + assert_eq!( + recv.read(&mut request).await.expect("read request"), + Some(4) + ); + assert_eq!(&request, b"ping"); + send.write_all(b"pong").await.expect("write response"); + send.finish().await.expect("finish response"); + session.shutdown().await.expect("shutdown server session"); + assert!(session.is_closed()); + }); + + let client = ClientBuilder::new().with_limits(limits(128)).build(); + let session = client + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + let (send, mut recv) = session.open_bi().await.expect("open bi stream"); + send.write_all(b"ping").await.expect("write request"); + send.finish().await.expect("finish request"); + let mut response = [0_u8; 4]; + assert_eq!( + recv.read(&mut response).await.expect("read response"), + Some(4) + ); + assert_eq!(&response, b"pong"); + assert_eq!(recv.read(&mut response).await.expect("read FIN"), None); + session.shutdown().await.expect("shutdown client session"); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn reset_and_stop_sending_propagate_to_the_peer() { + let server = ServerBuilder::new() + .with_limits(limits(128)) + .build() + .await + .expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let session = server.accept().await.expect("accept session"); + + let mut reset_recv = session.accept_uni().await.expect("accept reset stream"); + let mut byte = [0_u8; 1]; + assert_eq!( + reset_recv.read(&mut byte).await.expect("observe reset"), + None + ); + + let stop_recv = session.accept_uni().await.expect("accept stopped stream"); + stop_recv.stop(7).await.expect("send stop"); + session.shutdown().await.expect("shutdown server session"); + }); + + let client = ClientBuilder::new().with_limits(limits(128)).build(); + let session = client + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + let reset_send = session.open_uni().await.expect("open reset stream"); + reset_send.reset(42).await.expect("reset stream"); + + let stopped_send = session.open_uni().await.expect("open stopped stream"); + let error = tokio::time::timeout(Duration::from_secs(2), async { + loop { + match stopped_send.write_all(b"x").await { + Ok(()) => tokio::task::yield_now().await, + Err(error) => break error, + } + } + }) + .await + .expect("stop propagation timed out"); + assert!(matches!(error, websock_proto::Error::Closed)); + session.shutdown().await.expect("shutdown client session"); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn malformed_wire_frame_closes_the_session() { + let server = ServerBuilder::new().build().await.expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let session = server.accept().await.expect("accept session"); + let result = tokio::time::timeout(Duration::from_secs(2), session.accept_uni()) + .await + .expect("session did not close"); + assert!(matches!(result, Err(websock_proto::Error::Closed))); + }); + + let mut request = format!("ws://{address}") + .into_client_request() + .expect("create request"); + request.headers_mut().insert( + SEC_WEBSOCKET_PROTOCOL, + websock_mux_proto::SUBPROTOCOL + .parse() + .expect("protocol header"), + ); + let (mut socket, _) = tokio_tungstenite::connect_async(request) + .await + .expect("connect raw WebSocket"); + socket + .send(tungstenite::Message::Binary(vec![0xff].into())) + .await + .expect("send malformed frame"); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn tls_mux_round_trip() { + let tls = + websock_tungstenite_mux::tls::TlsConfig::new_insecure_config().expect("create TLS config"); + let server = ServerBuilder::new() + .with_tls_config(tls.server_config) + .with_default_alpn() + .with_limits(limits(128)) + .build() + .await + .expect("bind TLS server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let session = server.accept().await.expect("accept TLS session"); + let mut recv = session.accept_uni().await.expect("accept TLS stream"); + let mut payload = [0_u8; 6]; + assert_eq!( + recv.read(&mut payload).await.expect("read TLS data"), + Some(6) + ); + assert_eq!(&payload, b"secure"); + session + .shutdown() + .await + .expect("shutdown TLS server session"); + }); + + let client = ClientBuilder::new() + .with_tls_config(tls.client_config) + .with_default_alpn() + .with_limits(limits(128)) + .build(); + let session = client + .connect(&format!("wss://localhost:{}", address.port())) + .await + .expect("connect TLS client"); + let send = session.open_uni().await.expect("open TLS stream"); + send.write_all(b"secure").await.expect("write TLS data"); + send.finish().await.expect("finish TLS stream"); + session + .shutdown() + .await + .expect("shutdown TLS client session"); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn dropping_last_session_handle_closes_the_peer() { + let server = ServerBuilder::new().build().await.expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let session = server.accept().await.expect("accept session"); + let result = tokio::time::timeout(Duration::from_secs(2), session.accept_uni()) + .await + .expect("peer session remained open"); + assert!(matches!(result, Err(websock_proto::Error::Closed))); + }); + + let session = ClientBuilder::new() + .build() + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + let clone = session.clone(); + drop(session); + tokio::task::yield_now().await; + drop(clone); + + server_task.await.expect("server task"); +} + +#[tokio::test] +#[ignore = "long-running concurrency and backpressure coverage"] +async fn many_concurrent_streams_remain_bounded_under_backpressure() { + const STREAMS: usize = 128; + const BYTES_PER_STREAM: usize = 64 * 1024; + + let soak_limits = Limits { + max_ws_message_size: 64 * 1024, + max_stream_data_per_frame: 1024, + max_open_streams: STREAMS * 2, + recv_event_queue_len: 4, + outbound_queue_len: 16, + max_batch_frames: 8, + max_batch_bytes: 8 * 1024, + initial_stream_window: 4 * 1024, + stream_window_update_threshold: 2 * 1024, + accept_uni_queue_len: STREAMS, + accept_bi_queue_len: 4, + }; + let server = ServerBuilder::new() + .with_limits(soak_limits.clone()) + .build() + .await + .expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let session = server.accept().await.expect("accept session"); + let mut readers = Vec::with_capacity(STREAMS); + for _ in 0..STREAMS { + let mut recv = session.accept_uni().await.expect("accept stream"); + readers.push(tokio::spawn(async move { + let mut total = 0; + while let Some(chunk) = recv.read_chunk(257).await.expect("read chunk") { + total += chunk.len(); + tokio::task::yield_now().await; + } + total + })); + } + for reader in readers { + assert_eq!(reader.await.expect("reader task"), BYTES_PER_STREAM); + } + session.shutdown().await.expect("shutdown server session"); + }); + + let session = ClientBuilder::new() + .with_limits(soak_limits) + .build() + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + let mut writers = Vec::with_capacity(STREAMS); + for _ in 0..STREAMS { + let session = session.clone(); + writers.push(tokio::spawn(async move { + let send = session.open_uni().await.expect("open stream"); + send.write_all(&vec![0x5a; BYTES_PER_STREAM]) + .await + .expect("write stream"); + send.finish().await.expect("finish stream"); + })); + } + for writer in writers { + writer.await.expect("writer task"); + } + session.shutdown().await.expect("shutdown client session"); + server_task.await.expect("server task"); +} diff --git a/websock-wasm-mux/src/session.rs b/websock-wasm-mux/src/session.rs index f363d8a..c6ff4fe 100644 --- a/websock-wasm-mux/src/session.rs +++ b/websock-wasm-mux/src/session.rs @@ -1,4 +1,4 @@ -use std::cell::RefCell; +use std::cell::{Cell, RefCell}; use std::collections::HashMap; use std::io; use std::pin::Pin; @@ -7,7 +7,7 @@ use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::task::{Context, Poll}; use bytes::{Buf, BufMut, Bytes, BytesMut}; -use futures_channel::mpsc; +use futures_channel::{mpsc, oneshot}; use futures_io::{AsyncRead as FuturesAsyncRead, AsyncWrite as FuturesAsyncWrite}; use futures_util::lock::Mutex; use futures_util::stream::Stream; @@ -118,7 +118,6 @@ impl Limits { } } -#[derive(Clone)] pub struct Session { inner: Rc, accept_uni: Rc>>, @@ -190,6 +189,40 @@ impl Session { let mut rx = self.accept_bi.lock().await; rx.next().await.ok_or(Error::Closed) } + + /// Close the WebSocket and wait for the session task to finish. + pub async fn shutdown(&self) -> Result<()> { + self.inner.request_shutdown(); + self.inner.wait_closed().await + } + + /// Return whether the session has finished shutting down. + pub fn is_closed(&self) -> bool { + self.inner.task_finished.load(Ordering::Acquire) + } +} + +impl Clone for Session { + fn clone(&self) -> Self { + self.inner + .session_handles + .set(self.inner.session_handles.get() + 1); + Self { + inner: self.inner.clone(), + accept_uni: self.accept_uni.clone(), + accept_bi: self.accept_bi.clone(), + } + } +} + +impl Drop for Session { + fn drop(&mut self) { + let remaining = self.inner.session_handles.get() - 1; + self.inner.session_handles.set(remaining); + if remaining == 0 { + self.inner.request_shutdown(); + } + } } struct SendFlowState { @@ -810,6 +843,10 @@ struct SessionInner { next_peer_uni: AtomicU64, next_peer_bi: AtomicU64, closed: AtomicBool, + shutdown_started: AtomicBool, + session_handles: Cell, + close_waiters: RefCell>>, + task_finished: AtomicBool, } impl SessionInner { @@ -831,7 +868,26 @@ impl SessionInner { next_peer_uni: AtomicU64::new(0), next_peer_bi: AtomicU64::new(0), closed: AtomicBool::new(false), + shutdown_started: AtomicBool::new(false), + session_handles: Cell::new(1), + close_waiters: RefCell::new(Vec::new()), + task_finished: AtomicBool::new(false), + } + } + + fn request_shutdown(&self) { + if !self.shutdown_started.swap(true, Ordering::AcqRel) { + self.outbound_tx.borrow_mut().close_channel(); + } + } + + async fn wait_closed(&self) -> Result<()> { + if self.task_finished.load(Ordering::Acquire) { + return Ok(()); } + let (tx, rx) = oneshot::channel(); + self.close_waiters.borrow_mut().push(tx); + rx.await.map_err(|_| Error::Closed) } fn spawn_task( @@ -917,8 +973,9 @@ impl SessionInner { } } - inner.close_all().await; let _ = conn.close().await; + inner.close_all().await; + inner.finish_task(); }); } @@ -1206,6 +1263,13 @@ impl SessionInner { *self.accept_bi_tx.lock().await = None; } + fn finish_task(&self) { + self.task_finished.store(true, Ordering::Release); + for waiter in self.close_waiters.borrow_mut().drain(..) { + let _ = waiter.send(()); + } + } + fn validate_peer_stream_id(&self, id: StreamId) -> bool { let next = match id.dir() { StreamDir::Uni => &self.next_peer_uni, @@ -1309,6 +1373,14 @@ impl websock_mux_proto::MuxSession for Session { > { Box::pin(async move { Session::accept_bi(self).await }) } + + fn shutdown<'a>(&'a self) -> websock_proto::LocalBoxFuture<'a, Result<()>> { + Box::pin(async move { Session::shutdown(self).await }) + } + + fn closed(&self) -> bool { + Session::is_closed(self) + } } fn io_closed() -> io::Error { From 7c97d5af0856b1da2519e5c8439dbebf1da44c95 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 21:26:01 +0900 Subject: [PATCH 07/17] feat: expose protocol and close metadata --- websock-tungstenite/src/connection.rs | 38 ++++++++++++++++---- websock-tungstenite/src/lib.rs | 2 +- websock-tungstenite/src/server.rs | 42 +++++++++++++++++++--- websock-wasm/src/connection.rs | 50 ++++++++++++++++++++++++--- websock/src/tungstenite.rs | 3 +- 5 files changed, 117 insertions(+), 18 deletions(-) diff --git a/websock-tungstenite/src/connection.rs b/websock-tungstenite/src/connection.rs index 867a7f6..fc90daa 100644 --- a/websock-tungstenite/src/connection.rs +++ b/websock-tungstenite/src/connection.rs @@ -6,9 +6,9 @@ use std::sync::Arc; use tokio::io::{AsyncRead, AsyncWrite}; use tokio_tungstenite::{Connector, WebSocketStream, tungstenite}; use tungstenite::client::IntoClientRequest; -use websock_proto::{ConnectOptions, Error, Message, Result}; +use websock_proto::{CloseFrame, ConnectOptions, Error, Message, Result}; -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone)] pub struct ConnectionInfo { /// Remote peer address for the connection. pub peer: std::net::SocketAddr, @@ -16,6 +16,8 @@ pub struct ConnectionInfo { pub local: std::net::SocketAddr, /// True when the connection is established over TLS. pub is_tls: bool, + /// WebSocket subprotocol selected during the opening handshake. + pub subprotocol: Option, } /// Establish a WebSocket connection using Tokio Tungstenite. @@ -62,7 +64,7 @@ pub async fn connect_with_tls( } let connector = tls.map(Connector::Rustls); - let (ws, _resp) = + let (ws, resp) = tokio_tungstenite::connect_async_tls_with_config(req, Some(config), false, connector) .await .map_err(map_tungstenite_err)?; @@ -79,15 +81,25 @@ pub async fn connect_with_tls( .local_addr() .map_err(|e| Error::Io(e.to_string()))?, is_tls: matches!(ws.get_ref(), tokio_tungstenite::MaybeTlsStream::Rustls(_)), + subprotocol: resp + .headers() + .get(tungstenite::http::header::SEC_WEBSOCKET_PROTOCOL) + .and_then(|value| value.to_str().ok()) + .map(str::to_owned), }; - Ok(Connection { ws, info }) + Ok(Connection { + ws, + info, + close_frame: None, + }) } /// WebSocket connection wrapper around a Tokio Tungstenite stream. pub struct Connection> { pub(crate) ws: WebSocketStream, pub(crate) info: ConnectionInfo, + pub(crate) close_frame: Option, } impl Connection @@ -126,7 +138,11 @@ where tungstenite::Message::Pong(_) => continue, tungstenite::Message::Text(s) => return Ok(Message::Text(s.to_string())), tungstenite::Message::Binary(b) => return Ok(Message::Binary(b)), - tungstenite::Message::Close(_) => { + tungstenite::Message::Close(frame) => { + self.close_frame = frame.map(|frame| CloseFrame { + code: frame.code.into(), + reason: frame.reason.to_string(), + }); let _ = self.ws.close(None).await; return Err(Error::Closed); } @@ -184,7 +200,17 @@ impl Connection { } /// Return the full connection metadata snapshot. pub fn info(&self) -> ConnectionInfo { - self.info + self.info.clone() + } + + /// Return the negotiated WebSocket subprotocol, if any. + pub fn negotiated_subprotocol(&self) -> Option<&str> { + self.info.subprotocol.as_deref() + } + + /// Return the most recently received close-frame metadata, if any. + pub fn close_frame(&self) -> Option { + self.close_frame.clone() } } diff --git a/websock-tungstenite/src/lib.rs b/websock-tungstenite/src/lib.rs index fb9bf03..43f57ff 100644 --- a/websock-tungstenite/src/lib.rs +++ b/websock-tungstenite/src/lib.rs @@ -7,6 +7,6 @@ mod server; pub mod tls; pub use builder::{Client, ClientBuilder, DangerousClientBuilder, ServerBuilder}; -pub use connection::{Connection, connect, connect_with_tls}; +pub use connection::{Connection, ConnectionInfo, connect, connect_with_tls}; pub use server::{Server, ServerStream, bind}; pub mod stream; diff --git a/websock-tungstenite/src/server.rs b/websock-tungstenite/src/server.rs index 0a379d2..4ce6be8 100644 --- a/websock-tungstenite/src/server.rs +++ b/websock-tungstenite/src/server.rs @@ -1,7 +1,7 @@ //! Server-side WebSocket acceptor for the Tokio Tungstenite transport. use std::collections::HashSet; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use tokio::io::{AsyncRead, AsyncWrite}; use tokio::net::{TcpListener, ToSocketAddrs}; use tokio_rustls::TlsAcceptor; @@ -66,6 +66,7 @@ pub struct Server { impl Server { /// Accept an incoming WebSocket connection. + #[allow(clippy::result_large_err)] pub async fn accept(&self) -> Result> { let (stream, _addr) = self .listener @@ -86,14 +87,17 @@ impl Server { (Box::new(stream), false) }; - let info = ConnectionInfo { + let mut info = ConnectionInfo { peer, local, is_tls, + subprotocol: None, }; let headers = Arc::clone(&self.headers); let protocols = Arc::clone(&self.protocols); + let selected_protocol = Arc::new(Mutex::new(None)); + let selected_protocol_callback = Arc::clone(&selected_protocol); let ws = accept_hdr_async_with_config( stream, @@ -106,6 +110,10 @@ impl Server { let value = HeaderValue::from_str(protocol).expect("protocol value validated on bind"); resp.headers_mut().insert(SEC_WEBSOCKET_PROTOCOL, value); + *selected_protocol_callback + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = + Some(protocol.to_owned()); } Ok(resp) @@ -115,10 +123,19 @@ impl Server { .await .map_err(map_tungstenite_err)?; - Ok(Connection { ws, info }) + info.subprotocol = selected_protocol + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone(); + Ok(Connection { + ws, + info, + close_frame: None, + }) } /// Accept an incoming WebSocket connection, returning the TLS stream type. + #[allow(clippy::result_large_err)] pub async fn accept_tls( &self, ) -> Result>> { @@ -139,7 +156,7 @@ impl Server { let headers = Arc::clone(&self.headers); let protocols = Arc::clone(&self.protocols); - let info = ConnectionInfo { + let mut info = ConnectionInfo { peer: tls_stream .get_ref() .0 @@ -151,7 +168,10 @@ impl Server { .local_addr() .map_err(|e| Error::Io(e.to_string()))?, is_tls: true, + subprotocol: None, }; + let selected_protocol = Arc::new(Mutex::new(None)); + let selected_protocol_callback = Arc::clone(&selected_protocol); let ws = accept_hdr_async_with_config( tls_stream, @@ -164,6 +184,10 @@ impl Server { let value = HeaderValue::from_str(protocol).expect("protocol value validated on bind"); resp.headers_mut().insert(SEC_WEBSOCKET_PROTOCOL, value); + *selected_protocol_callback + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = + Some(protocol.to_owned()); } Ok(resp) @@ -173,7 +197,15 @@ impl Server { .await .map_err(map_tungstenite_err)?; - Ok(Connection { ws, info }) + info.subprotocol = selected_protocol + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone(); + Ok(Connection { + ws, + info, + close_frame: None, + }) } /// Return the local address of the listener. diff --git a/websock-wasm/src/connection.rs b/websock-wasm/src/connection.rs index d31d564..88af1a8 100644 --- a/websock-wasm/src/connection.rs +++ b/websock-wasm/src/connection.rs @@ -3,7 +3,7 @@ use std::cell::RefCell; use std::rc::Rc; use websock_proto::Bytes; -use websock_proto::{ConnectOptions, Error, Message, Result}; +use websock_proto::{CloseFrame, ConnectOptions, Error, Message, Result}; use futures_channel::{mpsc, oneshot}; use futures_util::StreamExt; @@ -141,17 +141,26 @@ pub async fn connect(url: &str, opts: ConnectOptions) -> Result { }); ws.set_onerror(Some(onerror.as_ref().unchecked_ref())); + let close_frame = Rc::new(RefCell::new(None)); + let close_frame_handler = Rc::clone(&close_frame); let mut tx_close = tx; - let onclose = Closure::::new(move |_e: web_sys::CloseEvent| { + let onclose = Closure::::new(move |e: web_sys::CloseEvent| { + *close_frame_handler.borrow_mut() = Some(CloseFrame { + code: e.code(), + reason: e.reason(), + }); let _ = tx_close.try_send(Err(Error::Closed)); tx_close.close_channel(); }); ws.set_onclose(Some(onclose.as_ref().unchecked_ref())); + let negotiated_subprotocol = ws.protocol(); Ok(Connection { ws: Rc::new(ws), rx: Some(rx), max_write_buffer_size: opts.limits.max_write_buffer_size, + negotiated_subprotocol, + close_frame, _onmessage: Some(onmessage), _onerror: Some(onerror), _onclose: Some(onclose), @@ -163,6 +172,8 @@ pub struct Connection { pub(crate) ws: Rc, pub(crate) rx: Option>>, pub(crate) max_write_buffer_size: usize, + pub(crate) negotiated_subprotocol: String, + pub(crate) close_frame: Rc>>, pub(crate) _onmessage: Option>, pub(crate) _onerror: Option>, @@ -191,10 +202,39 @@ impl Connection { rx.next().await.ok_or(Error::Closed)? } - /// Close the WebSocket connection. + /// Close the WebSocket connection and wait for the browser close event. pub async fn close(&mut self) -> Result<()> { - self.ws.close().map_err(js_err)?; - Ok(()) + if self.ws.ready_state() == web_sys::WebSocket::CLOSED { + self.rx = None; + return Ok(()); + } + if self.ws.ready_state() != web_sys::WebSocket::CLOSING { + self.ws.close().map_err(js_err)?; + } + + let result = loop { + let next = match self.rx.as_mut() { + Some(rx) => rx.next().await, + None => return Ok(()), + }; + match next { + Some(Ok(_)) => continue, + Some(Err(Error::Closed)) | None => break Ok(()), + Some(Err(error)) => break Err(error), + } + }; + self.rx = None; + result + } + + /// Return the WebSocket subprotocol selected by the server, if any. + pub fn negotiated_subprotocol(&self) -> Option<&str> { + (!self.negotiated_subprotocol.is_empty()).then_some(self.negotiated_subprotocol.as_str()) + } + + /// Return the most recently received close-frame metadata, if any. + pub fn close_frame(&self) -> Option { + self.close_frame.borrow().clone() } } diff --git a/websock/src/tungstenite.rs b/websock/src/tungstenite.rs index da16d3c..4ffc31d 100644 --- a/websock/src/tungstenite.rs +++ b/websock/src/tungstenite.rs @@ -2,6 +2,7 @@ pub use websock_tungstenite::stream; pub use websock_tungstenite::{ - Client, ClientBuilder, Connection, DangerousClientBuilder, Server, ServerBuilder, + Client, ClientBuilder, Connection, ConnectionInfo, DangerousClientBuilder, Server, + ServerBuilder, }; pub use websock_tungstenite::{crypto, tls}; From d2ccd4a54190833ff70b2b2ef57985b771ea9ae6 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 21:29:01 +0900 Subject: [PATCH 08/17] fix: validate mux frames and document wire semantics --- README.md | 5 ++ docs/mux-protocol.md | 98 ++++++++++++++++++++++++++++++ websock-mux-proto/src/lib.rs | 4 ++ websock-mux-proto/src/stream.rs | 43 ++++++++++++- websock-mux-proto/src/transport.rs | 8 +++ websock-mux-proto/src/varint.rs | 6 +- websock-mux-proto/tests/frame.rs | 13 ++++ websock-proto/src/error.rs | 2 + websock-proto/src/lib.rs | 2 + websock-proto/src/options.rs | 1 + 10 files changed, 180 insertions(+), 2 deletions(-) create mode 100644 docs/mux-protocol.md diff --git a/README.md b/README.md index 6d40c04..968bc96 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,8 @@ A minimal WebSocket library for native and WebAssembly. +The minimum supported Rust version (MSRV) is 1.88. + ## Workspace crates - `websock`: top-level facade that selects native (`websock-tungstenite`) or browser (`websock-wasm`) transport. @@ -59,6 +61,9 @@ The mux transports additionally expose `Limits` for stream counts, queue capacities, batching, and per-stream flow-control windows. Invalid or inconsistent limits return a protocol error rather than panicking. +The multiplexing wire format, stream lifecycle, flow control, compatibility +policy, and error codes are specified in [docs/mux-protocol.md](docs/mux-protocol.md). + ## Benchmarking Criterion benchmarks are available for `websock-mux-proto`. diff --git a/docs/mux-protocol.md b/docs/mux-protocol.md new file mode 100644 index 0000000..95783f7 --- /dev/null +++ b/docs/mux-protocol.md @@ -0,0 +1,98 @@ +# WebSocket Multiplexing Protocol + +This document specifies version 1 of the `websock` multiplexing protocol. A +connection using this protocol MUST negotiate the WebSocket subprotocol +`websock-mux-1`. Peers MUST use binary WebSocket messages; text messages are a +protocol error. + +## Compatibility + +The subprotocol name carries the wire-protocol major version. Implementations +with incompatible framing or stream semantics MUST use a different +subprotocol. Additive behavior that old peers can safely ignore may remain +within version 1, but version 1 currently defines no ignorable frame types. +Unknown frame types are therefore a protocol error. + +One binary WebSocket message may contain one or more complete mux frames. +Frames MUST NOT span WebSocket messages. All integers use the QUIC +variable-length integer encoding and are limited to the inclusive range +`0..2^62-1`. + +## Stream identifiers + +A stream identifier encodes three fields: + +```text +stream_id = stream_counter * 4 | direction_bit * 2 | initiator_bit +``` + +- `initiator_bit`: client is `0`; server is `1`. +- `direction_bit`: bidirectional is `0`; unidirectional is `1`. +- `stream_counter`: a monotonically increasing counter maintained separately + for each direction. + +A peer MUST open its streams with monotonically increasing identifiers. An +identifier with the wrong initiator or direction is a protocol error. + +## Frames + +Each frame starts with a varint tag followed by the listed fields: + +| Tag | Frame | Fields | +| ---: | --- | --- | +| 0 | `OpenUni` | `stream_id` | +| 1 | `OpenBi` | `stream_id` | +| 2 | `Stream` | `stream_id`, `fin`, `length`, `length` data bytes | +| 3 | `ResetStream` | `stream_id`, application error code | +| 4 | `StopSending` | `stream_id`, application error code | +| 5 | `ConnectionClose` | connection error code, UTF-8 reason length, reason bytes | +| 6 | `MaxStreamData` | `stream_id`, cumulative byte limit | + +The `fin` field MUST be `0` or `1`. Length-prefixed data MUST be fully present +in the containing WebSocket message. A malformed field, invalid UTF-8 reason, +unknown tag, or truncated frame is a protocol error. + +## Stream state + +`OpenUni` creates a receive stream for the peer. `OpenBi` creates both a send +and a receive side. Stream data may be sent only after the corresponding open +frame and MUST NOT be sent after FIN or reset. + +A `Stream` frame with `fin = 1` closes only the sender's direction. Its data, +if any, remains readable before end-of-stream is reported. The other direction +of a bidirectional stream remains usable. + +`ResetStream` abruptly terminates the sender's direction. `StopSending` asks +the peer to stop its sender direction; the peer then treats that send side as +closed. Application error codes are opaque to the protocol and MUST fit in a +varint. + +## Flow control + +New send streams start with zero credit. A receiver grants credit with +`MaxStreamData`, whose `max` field is the cumulative number of stream-data +bytes the sender may transmit. Values MUST be monotonic. A sender MUST block +when it has exhausted the latest advertised limit. + +The receiver counts only `Stream` payload bytes, not frame overhead. Receiving +bytes beyond the advertised cumulative limit is a connection-level flow +control error. Implementations issue further cumulative credit as the +application consumes buffered bytes. + +## Connection closure and error codes + +`ConnectionClose` is terminal. After receiving it, a peer closes all streams, +wakes blocked operations, and closes the WebSocket session. + +The following connection error codes are defined: + +| Code | Meaning | +| ---: | --- | +| 0 | Graceful or unspecified closure | +| 1 | Malformed frame or stream state violation | +| 2 | Size or flow-control violation | +| 3 | Resource limit exceeded | + +Other codes are reserved for future versions. WebSocket close frames remain +part of the underlying transport; either a mux `ConnectionClose` or an +underlying WebSocket closure terminates the session. diff --git a/websock-mux-proto/src/lib.rs b/websock-mux-proto/src/lib.rs index b146394..3c2c9dd 100644 --- a/websock-mux-proto/src/lib.rs +++ b/websock-mux-proto/src/lib.rs @@ -1,6 +1,9 @@ //! Multiplexing protocol primitives and core trait contracts shared by //! native and WebAssembly mux transports. +#![warn(missing_docs)] + +/// Mux stream identifiers and wire frames. pub mod stream; pub mod transport; pub mod varint; @@ -9,4 +12,5 @@ pub use stream::{Frame, FrameDecodeError, StreamDir, StreamId}; pub use transport::{MuxRecvStream, MuxSendStream, MuxSession}; pub use varint::{VarInt, VarIntBoundsExceeded, VarIntUnexpectedEnd}; +/// WebSocket subprotocol identifier for mux wire protocol version 1. pub const SUBPROTOCOL: &str = "websock-mux-1"; diff --git a/websock-mux-proto/src/stream.rs b/websock-mux-proto/src/stream.rs index 152be0a..84d0be6 100644 --- a/websock-mux-proto/src/stream.rs +++ b/websock-mux-proto/src/stream.rs @@ -2,16 +2,21 @@ use bytes::{Buf, BufMut, Bytes, BytesMut}; use crate::varint::{VarInt, VarIntBoundsExceeded, VarIntUnexpectedEnd}; +/// Directionality of a multiplexed stream. #[derive(Debug, Copy, Clone, Eq, PartialEq)] pub enum StreamDir { + /// Bidirectional stream. Bi, + /// Unidirectional stream. Uni, } +/// Encoded stream identifier containing counter, initiator, and direction bits. #[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)] pub struct StreamId(pub u64); impl StreamId { + /// Construct a stream identifier from its component fields. pub fn new( counter: u64, is_server: bool, @@ -29,6 +34,7 @@ impl StreamId { Ok(Self(value)) } + /// Return the stream direction encoded in the identifier. pub fn dir(self) -> StreamDir { if (self.0 >> 1) & 1 == 1 { StreamDir::Uni @@ -37,6 +43,7 @@ impl StreamId { } } + /// Return whether the server initiated the stream. pub fn initiator_is_server(self) -> bool { self.0 & 1 == 1 } @@ -47,33 +54,54 @@ impl StreamId { } } +/// A single frame in the version 1 mux wire protocol. #[derive(Debug, Clone, PartialEq, Eq)] pub enum Frame { + /// Open a peer-receive-only stream. OpenUni { + /// Identifier of the new stream. id: StreamId, }, + /// Open a bidirectional stream. OpenBi { + /// Identifier of the new stream. id: StreamId, }, + /// Carry stream payload data and optional end-of-stream state. Stream { + /// Target stream identifier. id: StreamId, + /// Payload bytes. data: Bytes, + /// Whether this frame finishes the sender direction. fin: bool, }, + /// Abruptly terminate the sender direction. ResetStream { + /// Target stream identifier. id: StreamId, + /// Application-defined reset code. code: u64, }, + /// Ask the peer to stop its sender direction. StopSending { + /// Target stream identifier. id: StreamId, + /// Application-defined stop code. code: u64, }, + /// Increase the cumulative stream-data allowance. MaxStreamData { + /// Target stream identifier. id: StreamId, + /// New cumulative byte limit. max: u64, }, + /// Close the entire mux connection. ConnectionClose { + /// Connection-level error code. code: u64, + /// Human-readable UTF-8 reason. reason: String, }, } @@ -101,6 +129,7 @@ impl Frame { } } + /// Encode this frame into its canonical wire representation. pub fn encode(&self) -> BytesMut { let mut buf = BytesMut::with_capacity(self.encoded_len()); match self { @@ -144,6 +173,7 @@ impl Frame { buf } + /// Decode one complete frame from the front of a byte buffer. pub fn decode(buf: &mut B) -> Result { let tag = VarInt::decode(buf)?.into_inner(); match tag { @@ -155,7 +185,11 @@ impl Frame { }), 2 => { let id = StreamId(VarInt::decode(buf)?.into_inner()); - let fin = VarInt::decode(buf)?.into_inner() != 0; + let fin = match VarInt::decode(buf)?.into_inner() { + 0 => false, + 1 => true, + value => return Err(FrameDecodeError::InvalidFin(value)), + }; let len = VarInt::decode(buf)?.into_inner() as usize; if buf.remaining() < len { return Err(FrameDecodeError::UnexpectedEnd); @@ -192,14 +226,21 @@ impl Frame { } } +/// Errors produced while decoding a mux frame. #[derive(Debug, thiserror::Error)] pub enum FrameDecodeError { + /// The buffer ended before the frame was complete. #[error("unexpected end of buffer")] UnexpectedEnd, + /// The frame tag is not defined by this protocol version. #[error("unknown frame tag {0}")] UnknownTag(u64), + /// A connection-close reason was not valid UTF-8. #[error("invalid utf-8 in reason")] InvalidUtf8, + /// A stream FIN field contained a value other than zero or one. + #[error("invalid stream FIN value {0}")] + InvalidFin(u64), } impl From for FrameDecodeError { diff --git a/websock-mux-proto/src/transport.rs b/websock-mux-proto/src/transport.rs index 99ffa40..c039b6c 100644 --- a/websock-mux-proto/src/transport.rs +++ b/websock-mux-proto/src/transport.rs @@ -32,7 +32,9 @@ pub trait MuxRecvStream { /// Cross-platform mux session contract. pub trait MuxSession { + /// Concrete send-stream type. type SendStream: MuxSendStream; + /// Concrete receive-stream type. type RecvStream: MuxRecvStream; /// Open a unidirectional stream. @@ -46,4 +48,10 @@ pub trait MuxSession { /// Accept a peer-initiated bidirectional stream. fn accept_bi<'a>(&'a self) -> LocalBoxFuture<'a, Result<(Self::SendStream, Self::RecvStream)>>; + + /// Close the underlying WebSocket and wait for session tasks to finish. + fn shutdown<'a>(&'a self) -> LocalBoxFuture<'a, Result<()>>; + + /// Return true if the session has closed. + fn closed(&self) -> bool; } diff --git a/websock-mux-proto/src/varint.rs b/websock-mux-proto/src/varint.rs index 39b763c..a8f61ec 100644 --- a/websock-mux-proto/src/varint.rs +++ b/websock-mux-proto/src/varint.rs @@ -137,6 +137,7 @@ impl fmt::Display for VarInt { } impl VarInt { + /// Decode one QUIC variable-length integer from a byte buffer. pub fn decode(r: &mut B) -> Result { if !r.has_remaining() { return Err(VarIntUnexpectedEnd); @@ -179,7 +180,7 @@ impl VarInt { Ok(Self(x)) } - // Read a varint from an async stream. + /// Read one QUIC variable-length integer from an asynchronous stream. #[cfg(not(target_arch = "wasm32"))] pub async fn read(stream: &mut S) -> Result { // Eight bytes is the maximum encoded length. @@ -205,6 +206,7 @@ impl VarInt { Ok(v) } + /// Encode this value into a byte buffer. pub fn encode(&self, w: &mut B) { let x = self.0; if x <= Self::MAX_1BYTE { @@ -221,6 +223,7 @@ impl VarInt { } #[cfg(not(target_arch = "wasm32"))] + /// Write this value to an asynchronous stream. pub async fn write( &self, stream: &mut S, @@ -246,6 +249,7 @@ impl VarInt { #[error("value too large for varint encoding")] pub struct VarIntBoundsExceeded; +/// Error returned when a buffer or stream ends before a varint is complete. #[derive(Error, Debug, Copy, Clone, Eq, PartialEq)] #[error("unexpected end of buffer")] pub struct VarIntUnexpectedEnd; diff --git a/websock-mux-proto/tests/frame.rs b/websock-mux-proto/tests/frame.rs index 5597b46..a03ef37 100644 --- a/websock-mux-proto/tests/frame.rs +++ b/websock-mux-proto/tests/frame.rs @@ -128,3 +128,16 @@ fn frame_decode_invalid_utf8_reason() { other => panic!("unexpected error: {other:?}"), } } + +#[test] +fn frame_decode_rejects_non_boolean_fin() { + let id = StreamId::new(0, false, StreamDir::Bi).unwrap(); + let mut buf = BytesMut::new(); + VarInt::from_u32(2).encode(&mut buf); + VarInt::from_u64(id.0).unwrap().encode(&mut buf); + VarInt::from_u32(2).encode(&mut buf); + VarInt::from_u32(0).encode(&mut buf); + + let error = Frame::decode(&mut Cursor::new(buf.freeze())).unwrap_err(); + assert!(matches!(error, FrameDecodeError::InvalidFin(2))); +} diff --git a/websock-proto/src/error.rs b/websock-proto/src/error.rs index d54b9e6..f0572ed 100644 --- a/websock-proto/src/error.rs +++ b/websock-proto/src/error.rs @@ -30,9 +30,11 @@ pub enum Error { #[error("unsupported: {0}")] Unsupported(String), + /// A multiplexing frame could not be decoded. #[error("frame decode error: {0}")] FrameDecode(String), + /// A multiplexing stream identifier was invalid. #[error("stream id error: {0}")] StreamId(String), diff --git a/websock-proto/src/lib.rs b/websock-proto/src/lib.rs index b99a29b..de89c6f 100644 --- a/websock-proto/src/lib.rs +++ b/websock-proto/src/lib.rs @@ -4,6 +4,8 @@ //! API stays consistent across native and WebAssembly targets. //! It also defines the core trait contracts shared by native and WASM transports. +#![warn(missing_docs)] + mod error; mod message; mod options; diff --git a/websock-proto/src/options.rs b/websock-proto/src/options.rs index 0699912..4bc1a78 100644 --- a/websock-proto/src/options.rs +++ b/websock-proto/src/options.rs @@ -78,6 +78,7 @@ pub struct ServerOptions { pub limits: WebSocketLimits, } +/// Return the ALPN identifiers supported by the WebSocket transport. pub fn default_ws_alpn() -> Vec> { vec![ALPN_HTTP_1_1.to_vec()] } From b5c2378fe9ce05efae4c116d5bf667f69567196a Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 21:34:48 +0900 Subject: [PATCH 09/17] refactor: consolidate helpers and replace rustls-pemfile --- websock-mux/Cargo.toml | 2 +- websock-mux/examples/echo-server-mux.rs | 18 +- websock-tungstenite-mux/Cargo.toml | 8 +- websock-tungstenite-mux/src/tls/cert.rs | 96 --------- websock-tungstenite-mux/src/tls/key.rs | 22 -- websock-tungstenite-mux/src/tls/mod.rs | 260 +----------------------- websock/Cargo.toml | 2 +- websock/examples/echo-client.rs | 9 +- websock/examples/echo-server.rs | 14 +- websock/examples/split-client.rs | 9 +- 10 files changed, 21 insertions(+), 419 deletions(-) delete mode 100644 websock-tungstenite-mux/src/tls/cert.rs delete mode 100644 websock-tungstenite-mux/src/tls/key.rs diff --git a/websock-mux/Cargo.toml b/websock-mux/Cargo.toml index 898400f..177a57a 100644 --- a/websock-mux/Cargo.toml +++ b/websock-mux/Cargo.toml @@ -2,6 +2,7 @@ name = "websock-mux" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "WebSocket multiplexing layer for logical streams" repository = "https://github.com/foctal/websock" @@ -29,7 +30,6 @@ tracing-subscriber = "0.3" clap = { version = "4", features = ["derive"] } url = "2" rustls = { version = "0.23", default-features = false, features = ["ring", "std"] } -rustls-pemfile = "2.2" [[example]] name = "echo-client-mux" diff --git a/websock-mux/examples/echo-server-mux.rs b/websock-mux/examples/echo-server-mux.rs index 60eebc1..40ef443 100644 --- a/websock-mux/examples/echo-server-mux.rs +++ b/websock-mux/examples/echo-server-mux.rs @@ -3,7 +3,8 @@ //! This server accepts mux sessions and echoes bytes over each bidirectional stream. use clap::Parser; -use std::{fs, io, path}; +use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject}; +use std::path; use tracing::Level; use tracing_subscriber::FmtSubscriber; use websock_mux::{Server, ServerBuilder}; @@ -91,20 +92,11 @@ async fn main() -> anyhow::Result<()> { fn load_pem_cert_and_key( cert_path: &path::Path, key_path: &path::Path, -) -> anyhow::Result<( - Vec>, - rustls::pki_types::PrivateKeyDer<'static>, -)> { - let chain_file = fs::File::open(cert_path)?; - let mut chain_reader = io::BufReader::new(chain_file); - let chain: Vec> = - rustls_pemfile::certs(&mut chain_reader).collect::>()?; +) -> anyhow::Result<(Vec>, PrivateKeyDer<'static>)> { + let chain = CertificateDer::pem_file_iter(cert_path)?.collect::, _>>()?; anyhow::ensure!(!chain.is_empty(), "could not find certificate"); - let key_file = fs::File::open(key_path)?; - let mut key_reader = io::BufReader::new(key_file); - let key = rustls_pemfile::private_key(&mut key_reader)? - .ok_or_else(|| anyhow::anyhow!("missing private key"))?; + let key = PrivateKeyDer::from_pem_file(key_path)?; Ok((chain, key)) } diff --git a/websock-tungstenite-mux/Cargo.toml b/websock-tungstenite-mux/Cargo.toml index 6be528c..933b196 100644 --- a/websock-tungstenite-mux/Cargo.toml +++ b/websock-tungstenite-mux/Cargo.toml @@ -2,6 +2,7 @@ name = "websock-tungstenite-mux" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "Native WebSocket multiplexing layer for logical streams" repository = "https://github.com/foctal/websock" @@ -12,19 +13,14 @@ license = "MIT" [dependencies] bytes = { workspace = true } -thiserror = "2" websock-proto = { workspace = true } websock-mux-proto = { workspace = true } +websock-tungstenite = { workspace = true } tokio = { version = "1", features = ["rt", "sync", "macros"] } tokio-rustls = { version = "0.26", default-features = false, features = ["ring"]} rustls = { version = "0.23", default-features = false, features = ["ring", "std"] } -rustls-native-certs = "0.8" -rcgen = "0.14" -rustls-pemfile = "2.2" tokio-tungstenite = { version = "0.28", features = ["rustls-tls-webpki-roots"] } futures-util = { version = "0.3" } futures-core = { version = "0.3" } futures-sink = { version = "0.3" } tokio-util = { version = "0.7" } -#http = "1" -tracing = "0.1" diff --git a/websock-tungstenite-mux/src/tls/cert.rs b/websock-tungstenite-mux/src/tls/cert.rs deleted file mode 100644 index f496a50..0000000 --- a/websock-tungstenite-mux/src/tls/cert.rs +++ /dev/null @@ -1,96 +0,0 @@ -//! Certificate handling utilities. - -use rustls::client::danger::ServerCertVerifier; -use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; -use std::{fs, path::Path, sync::Arc}; -use websock_proto::{Error, Result}; - -/// Load native certificates from the host system. -pub fn get_native_certs() -> Result { - let mut root_store = rustls::RootCertStore::empty(); - - let cert_result = rustls_native_certs::load_native_certs(); - - for cert in cert_result.certs { - let _ = root_store.add(cert); - } - - Ok(root_store) -} - -/// Load a certificate chain from a file (.der or .pem). -pub fn load_certs(cert_path: &Path) -> Result>> { - let cert_bytes = fs::read(cert_path).map_err(|e| Error::Io(e.to_string()))?; - - if cert_path.extension().is_some_and(|x| x == "der") { - return Ok(vec![CertificateDer::from(cert_bytes)]); - } - - rustls_pemfile::certs(&mut &*cert_bytes) - .collect::, std::io::Error>>() - .map_err(|e| Error::Io(e.to_string())) -} - -/// Certificate verifier that unconditionally accepts certificates. -/// -/// # Warning -/// This is vulnerable to MITM attacks and must only be used for testing. -#[derive(Debug)] -pub struct SkipServerVerification(Arc); - -impl SkipServerVerification { - /// Create a verifier using the default ring provider. - pub fn new() -> Arc { - Self::with_provider(Arc::new(rustls::crypto::ring::default_provider())) - } - - /// Create a verifier with the provided crypto provider. - pub fn with_provider(provider: Arc) -> Arc { - Arc::new(Self(provider)) - } -} - -impl ServerCertVerifier for SkipServerVerification { - fn verify_server_cert( - &self, - _end_entity: &CertificateDer<'_>, - _intermediates: &[CertificateDer<'_>], - _server_name: &ServerName<'_>, - _ocsp: &[u8], - _now: UnixTime, - ) -> std::result::Result { - Ok(rustls::client::danger::ServerCertVerified::assertion()) - } - - fn verify_tls12_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> std::result::Result { - rustls::crypto::verify_tls12_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) - } - - fn verify_tls13_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> std::result::Result { - rustls::crypto::verify_tls13_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) - } - - fn supported_verify_schemes(&self) -> Vec { - self.0.signature_verification_algorithms.supported_schemes() - } -} diff --git a/websock-tungstenite-mux/src/tls/key.rs b/websock-tungstenite-mux/src/tls/key.rs deleted file mode 100644 index 0f45ad2..0000000 --- a/websock-tungstenite-mux/src/tls/key.rs +++ /dev/null @@ -1,22 +0,0 @@ -//! Private key handling utilities. - -use rustls::pki_types::{PrivateKeyDer, PrivatePkcs8KeyDer}; -use std::{fs, path::Path}; -use websock_proto::{Error, Result}; - -/// Load a private key from a file. -pub fn load_key(key_path: &Path) -> Result> { - let key = fs::read(key_path).map_err(|e| Error::Io(e.to_string()))?; - - let key = if key_path.extension().is_some_and(|x| x == "der") { - // Treat raw DER as PKCS#8. - PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(key)) - } else { - // Decode PEM. - rustls_pemfile::private_key(&mut &*key) - .map_err(|e| Error::Tls(e.to_string()))? - .ok_or_else(|| Error::Io("no keys found".into()))? - }; - - Ok(key) -} diff --git a/websock-tungstenite-mux/src/tls/mod.rs b/websock-tungstenite-mux/src/tls/mod.rs index 6079528..8c06568 100644 --- a/websock-tungstenite-mux/src/tls/mod.rs +++ b/websock-tungstenite-mux/src/tls/mod.rs @@ -1,259 +1,3 @@ -//! TLS configuration helpers for managing certificates and private keys. +//! TLS helpers shared with the native WebSocket transport. -pub mod cert; -pub mod key; - -use websock_proto::{Error, Result}; - -use cert::SkipServerVerification; -use rustls::client::{ClientConfig, WebPkiServerVerifier}; -use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer}; -use rustls::server::ServerConfig; -use std::path::Path; -use std::sync::Arc; - -/// Alias of [`rustls::server::ServerConfig`]. -pub type TlsServerConfig = rustls::server::ServerConfig; - -/// Alias of [`rustls::client::ClientConfig`]. -pub type TlsClientConfig = rustls::client::ClientConfig; - -/// Generate a self-signed certificate and private key (DER). -pub fn generate_self_signed_pair_der( - subject_alt_names: Vec, -) -> Result<(Vec>, PrivateKeyDer<'static>)> { - let cert = rcgen::generate_simple_self_signed(subject_alt_names) - .map_err(|e| Error::Tls(e.to_string()))?; - - let key = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der())); - let cert_chain = vec![CertificateDer::from(cert.cert)]; - Ok((cert_chain, key)) -} - -/// Generate a self-signed certificate and private key (PEM). -pub fn generate_self_signed_pair_pem( - subject_alt_names: Vec, -) -> Result<(Vec, String)> { - let cert = rcgen::generate_simple_self_signed(subject_alt_names) - .map_err(|e| Error::Tls(e.to_string()))?; - - let key = cert.signing_key.serialize_pem(); - let cert_chain = vec![cert.cert.pem()]; - Ok((cert_chain, key)) -} - -/// Load certificate chain and private key from files. -pub fn load_cert( - cert_path: &Path, - key_path: &Path, -) -> Result<(Vec>, PrivateKeyDer<'static>)> { - let cert_chain = cert::load_certs(cert_path)?; - let key = key::load_key(key_path)?; - Ok((cert_chain, key)) -} - -/// Bundled TLS configuration for both client and server use. -#[derive(Debug, Clone)] -pub struct TlsConfig { - /// Client-side rustls configuration. - pub client_config: ClientConfig, - /// Server-side rustls configuration. - pub server_config: ServerConfig, -} - -impl TlsConfig { - /// Create a new TLS configuration with the specified certificate and private key. - pub fn with_cert(cert_path: &Path, key_path: &Path) -> Result { - let client_config = TlsClientConfigBuilder::new_with_native_certs()? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - let server_config = TlsServerConfigBuilder::new_with_cert(cert_path, key_path)? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - Ok(Self { - client_config, - server_config, - }) - } - - /// Create a new TLS configuration with self-signed certificates (localhost). - pub fn with_self_signed_certs() -> Result { - let client_config = TlsClientConfigBuilder::new_with_native_certs()? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - let server_config = - TlsServerConfigBuilder::new_with_self_signed_certs(vec!["localhost".into()])? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - Ok(Self { - client_config, - server_config, - }) - } - - /// Create a new TLS configuration with system certificates (server side is self-signed localhost). - pub fn new_native_config() -> Result { - let client_config = TlsClientConfigBuilder::new_with_native_certs()? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - let server_config = - TlsServerConfigBuilder::new_with_self_signed_certs(vec!["localhost".into()])? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - Ok(Self { - client_config, - server_config, - }) - } - - /// Create a new TLS configuration with no certificate verification (testing only). - pub fn new_insecure_config() -> Result { - let client_config = TlsClientConfigBuilder::new_insecure()? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - let server_config = - TlsServerConfigBuilder::new_with_self_signed_certs(vec!["localhost".into()])? - .with_alpn_protocols(vec![b"h3".to_vec()]) - .build(); - - Ok(Self { - client_config, - server_config, - }) - } -} - -/// Server config builder (owned builder). -#[derive(Debug, Clone)] -pub struct TlsServerConfigBuilder { - inner: TlsServerConfig, -} - -impl TlsServerConfigBuilder { - /// Create an insecure server config with a self-signed certificate. - pub fn new_insecure(subject_alt_names: Vec) -> Result { - let (certs, key) = generate_self_signed_pair_der(subject_alt_names)?; - let inner = ServerConfig::builder() - .with_no_client_auth() - .with_single_cert(certs, key) - .map_err(|e| Error::Tls(e.to_string()))?; - Ok(Self { inner }) - } - - /// Create a server config using certificate and private key files. - pub fn new_with_cert(cert_path: &Path, key_path: &Path) -> Result { - let (certs, key) = load_cert(cert_path, key_path)?; - let inner = ServerConfig::builder() - .with_no_client_auth() - .with_single_cert(certs, key) - .map_err(|e| Error::Tls(e.to_string()))?; - Ok(Self { inner }) - } - - /// Create a server config using a self-signed certificate. - pub fn new_with_self_signed_certs(subject_alt_names: Vec) -> Result { - Self::new_insecure(subject_alt_names) - } - - /// Set ALPN protocol identifiers. - pub fn with_alpn_protocols(mut self, protocols: Vec>) -> Self { - self.inner.alpn_protocols = protocols; - self - } - - /// Finalize the builder and return the server config. - pub fn build(self) -> TlsServerConfig { - self.inner - } -} - -/// Client config builder (owned builder). -#[derive(Debug, Clone)] -pub struct TlsClientConfigBuilder { - inner: TlsClientConfig, -} - -impl TlsClientConfigBuilder { - /// Create an insecure client config that skips certificate verification. - pub fn new_insecure() -> Result { - let inner = ClientConfig::builder() - .dangerous() - .with_custom_certificate_verifier(SkipServerVerification::new()) - .with_no_client_auth(); - Ok(Self { inner }) - } - - /// Create a client config that uses the system root store. - pub fn new_with_native_certs() -> Result { - let native_certs = cert::get_native_certs()?; - let inner = ClientConfig::builder() - .with_root_certificates(native_certs) - .with_no_client_auth(); - Ok(Self { inner }) - } - - /// Create a client config with a custom WebPKI verifier. - pub fn new_with_webpki_verifier(verifier: Arc) -> Result { - let inner = ClientConfig::builder() - .with_webpki_verifier(verifier) - .with_no_client_auth(); - Ok(Self { inner }) - } - - /// Set ALPN protocol identifiers. - pub fn with_alpn_protocols(mut self, protocols: Vec>) -> Self { - self.inner.alpn_protocols = protocols; - self - } - - /// Finalize the builder and return the client config. - pub fn build(self) -> TlsClientConfig { - self.inner - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_generate_self_signed_pair_der() { - let (cert_chain, key) = generate_self_signed_pair_der(vec!["localhost".into()]).unwrap(); - let rustls_server_config = ServerConfig::builder() - .with_no_client_auth() - .with_single_cert(cert_chain, key); - - if let Err(e) = rustls_server_config { - panic!("Failed to create ServerConfig: {e}"); - } - } - - #[test] - fn test_generate_self_signed_pair_pem() { - let (cert_chain, key) = generate_self_signed_pair_pem(vec!["localhost".into()]).unwrap(); - - let cert_path = Path::new("cert.pem"); - let key_path = Path::new("key.pem"); - std::fs::write(cert_path, cert_chain.join("\n")).unwrap(); - std::fs::write(key_path, key).unwrap(); - - let (cert_chain, key) = load_cert(cert_path, key_path).unwrap(); - let rustls_server_config = ServerConfig::builder() - .with_no_client_auth() - .with_single_cert(cert_chain, key); - - if let Err(e) = rustls_server_config { - panic!("Failed to create ServerConfig: {e}"); - } - - std::fs::remove_file(cert_path).unwrap(); - std::fs::remove_file(key_path).unwrap(); - } -} +pub use websock_tungstenite::tls::*; diff --git a/websock/Cargo.toml b/websock/Cargo.toml index 80e3164..2bb896d 100644 --- a/websock/Cargo.toml +++ b/websock/Cargo.toml @@ -2,6 +2,7 @@ name = "websock" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "WebSocket library for native and WebAssembly" repository = "https://github.com/foctal/websock" @@ -24,7 +25,6 @@ tokio = { version = "1", features = ["rt", "rt-multi-thread", "signal", "macros" futures-util = { version = "0.3" } clap = { version = "4", features = ["derive"] } anyhow = "1" -rustls-pemfile = "2" tracing = "0.1" tracing-subscriber = "0.3" ring = "0.17" diff --git a/websock/examples/echo-client.rs b/websock/examples/echo-client.rs index a5a7ec0..8a1fae8 100644 --- a/websock/examples/echo-client.rs +++ b/websock/examples/echo-client.rs @@ -1,8 +1,8 @@ //! Echo client example for the websock native transport. use clap::Parser; -use rustls::pki_types::CertificateDer; -use std::{fs, io, path}; +use rustls::pki_types::{CertificateDer, pem::PemObject}; +use std::path; use tracing::Level; use tracing_subscriber::FmtSubscriber; use url::Url; @@ -98,9 +98,6 @@ fn is_localhost_url(url: &Url) -> bool { /// Load PEM-encoded certificates from disk into DER bytes. fn load_pem_certs(p: &path::Path) -> anyhow::Result>> { - let f = fs::File::open(p)?; - let mut r = io::BufReader::new(f); - let certs: Vec> = - rustls_pemfile::certs(&mut r).collect::>()?; + let certs = CertificateDer::pem_file_iter(p)?.collect::, _>>()?; Ok(certs.into_iter().map(|c| c.to_vec()).collect()) } diff --git a/websock/examples/echo-server.rs b/websock/examples/echo-server.rs index bdec9db..2c89215 100644 --- a/websock/examples/echo-server.rs +++ b/websock/examples/echo-server.rs @@ -1,8 +1,8 @@ //! Echo server example for the websock native transport. use clap::Parser; -use rustls::pki_types::{CertificateDer, PrivateKeyDer}; -use std::{fs, io, path}; +use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject}; +use std::path; use tracing::Level; use tracing_subscriber::FmtSubscriber; use websock::{Server, ServerBuilder}; @@ -78,16 +78,10 @@ fn load_pem_cert_and_key( cert_path: &path::Path, key_path: &path::Path, ) -> anyhow::Result<(Vec>, PrivateKeyDer<'static>)> { - let chain_file = fs::File::open(cert_path)?; - let mut chain_reader = io::BufReader::new(chain_file); - let chain: Vec> = - rustls_pemfile::certs(&mut chain_reader).collect::>()?; + let chain = CertificateDer::pem_file_iter(cert_path)?.collect::, _>>()?; anyhow::ensure!(!chain.is_empty(), "could not find certificate"); - let key_file = fs::File::open(key_path)?; - let mut key_reader = io::BufReader::new(key_file); - let key = rustls_pemfile::private_key(&mut key_reader)? - .ok_or_else(|| anyhow::anyhow!("missing private key"))?; + let key = PrivateKeyDer::from_pem_file(key_path)?; Ok((chain, key)) } diff --git a/websock/examples/split-client.rs b/websock/examples/split-client.rs index 911915f..91a4f4c 100644 --- a/websock/examples/split-client.rs +++ b/websock/examples/split-client.rs @@ -2,8 +2,8 @@ use clap::Parser; use futures_util::{SinkExt, StreamExt}; -use rustls::pki_types::CertificateDer; -use std::{fs, io, path}; +use rustls::pki_types::{CertificateDer, pem::PemObject}; +use std::path; use tracing::Level; use tracing_subscriber::FmtSubscriber; use url::Url; @@ -104,9 +104,6 @@ fn is_localhost_url(url: &Url) -> bool { /// Load PEM-encoded certificates from disk into DER bytes. fn load_pem_certs(p: &path::Path) -> anyhow::Result>> { - let f = fs::File::open(p)?; - let mut r = io::BufReader::new(f); - let certs: Vec> = - rustls_pemfile::certs(&mut r).collect::>()?; + let certs = CertificateDer::pem_file_iter(p)?.collect::, _>>()?; Ok(certs.into_iter().map(|c| c.to_vec()).collect()) } From 50d9c97fe5bab0e220ed83c3a0a20c3a6330bc45 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 21:38:37 +0900 Subject: [PATCH 10/17] test: expand native and browser integration coverage --- .github/workflows/ci.yml | 57 +++++++++++ websock-browser-test-server/Cargo.toml | 14 +++ websock-browser-test-server/src/main.rs | 122 ++++++++++++++++++++++++ websock-proto/Cargo.toml | 1 + websock-tungstenite-mux/src/server.rs | 1 + websock-tungstenite/Cargo.toml | 2 +- websock-tungstenite/src/tls/cert.rs | 6 +- websock-tungstenite/src/tls/key.rs | 6 +- websock-tungstenite/src/tls/mod.rs | 16 ++-- websock-tungstenite/tests/connection.rs | 121 +++++++++++++++++++++++ websock-wasm-demo/Cargo.toml | 6 +- websock-wasm-mux/Cargo.toml | 6 +- websock-wasm-mux/tests/browser.rs | 23 +++++ websock-wasm/Cargo.toml | 9 +- websock-wasm/tests/browser.rs | 73 ++++++++++++++ 15 files changed, 442 insertions(+), 21 deletions(-) create mode 100644 websock-browser-test-server/Cargo.toml create mode 100644 websock-browser-test-server/src/main.rs create mode 100644 websock-wasm-mux/tests/browser.rs create mode 100644 websock-wasm/tests/browser.rs diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 0694867..a3206e8 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -44,6 +44,63 @@ jobs: -p websock-wasm-mux -p websock-wasm-demo --target wasm32-unknown-unknown + - name: Install wasm-bindgen test runner + run: cargo install wasm-bindgen-cli --version 0.2.126 --locked + - name: Build browser test server + run: cargo build -p websock-browser-test-server + - name: Run browser integration tests + shell: bash + run: | + target/debug/websock-browser-test-server & + server_pid=$! + trap 'kill "$server_pid"' EXIT + ready=false + for attempt in {1..30}; do + if (echo > /dev/tcp/127.0.0.1/32123) 2>/dev/null && + (echo > /dev/tcp/127.0.0.1/32124) 2>/dev/null; then + ready=true + break + fi + sleep 1 + done + if [ "$ready" != true ]; then + echo "browser test server failed to start" >&2 + exit 1 + fi + CARGO_TARGET_WASM32_UNKNOWN_UNKNOWN_RUNNER=wasm-bindgen-test-runner \ + cargo test -p websock-wasm -p websock-wasm-mux \ + --tests --target wasm32-unknown-unknown + + msrv: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 + - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 + with: + toolchain: "1.88.0" + - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 + - run: cargo check --workspace --all-targets + + minimal-features: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 + - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 + - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 + - run: cargo check --workspace --all-targets --no-default-features + + dependency-policy: + runs-on: ubuntu-latest + strategy: + matrix: + checks: + - advisories + - bans licenses sources + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 + - uses: EmbarkStudios/cargo-deny-action@v2 + with: + command: check ${{ matrix.checks }} platforms: strategy: diff --git a/websock-browser-test-server/Cargo.toml b/websock-browser-test-server/Cargo.toml new file mode 100644 index 0000000..9443a40 --- /dev/null +++ b/websock-browser-test-server/Cargo.toml @@ -0,0 +1,14 @@ +[package] +name = "websock-browser-test-server" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +authors.workspace = true +publish = false +license = "MIT" + +[dependencies] +futures-util = "0.3" +tokio = { version = "1", features = ["macros", "net", "rt-multi-thread", "time"] } +tokio-tungstenite = "0.28" +websock-tungstenite-mux = { workspace = true } diff --git a/websock-browser-test-server/src/main.rs b/websock-browser-test-server/src/main.rs new file mode 100644 index 0000000..34da850 --- /dev/null +++ b/websock-browser-test-server/src/main.rs @@ -0,0 +1,122 @@ +//! Local WebSocket endpoints used by the browser integration tests. + +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use futures_util::{SinkExt, StreamExt}; +use tokio::net::{TcpListener, TcpStream}; +use tokio_tungstenite::tungstenite; +use websock_tungstenite_mux::ServerBuilder; + +const WEBSOCKET_ADDRESS: &str = "127.0.0.1:32123"; +const MUX_ADDRESS: &str = "127.0.0.1:32124"; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let websocket = serve_websocket_endpoints(); + let mux = serve_mux_endpoint(); + tokio::try_join!(websocket, mux)?; + Ok(()) +} + +async fn serve_websocket_endpoints() -> std::io::Result<()> { + let listener = TcpListener::bind(WEBSOCKET_ADDRESS).await?; + loop { + let (stream, _) = listener.accept().await?; + tokio::spawn(async move { + let _ = handle_websocket(stream).await; + }); + } +} + +#[allow(clippy::result_large_err)] +async fn handle_websocket( + stream: TcpStream, +) -> Result<(), Box> { + let path = Arc::new(Mutex::new(String::new())); + let callback_path = Arc::clone(&path); + let mut socket = tokio_tungstenite::accept_hdr_async( + stream, + move |request: &tungstenite::handshake::server::Request, response| { + *callback_path + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = + request.uri().path().to_owned(); + Ok(response) + }, + ) + .await?; + let path = path + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone(); + + match path.as_str() { + "/overflow" => { + for value in 0..256_u16 { + socket + .send(tungstenite::Message::Binary( + value.to_be_bytes().to_vec().into(), + )) + .await?; + } + tokio::time::sleep(Duration::from_millis(250)).await; + socket.close(None).await?; + } + "/close" => { + socket + .send(tungstenite::Message::Close(Some( + tungstenite::protocol::CloseFrame { + code: tungstenite::protocol::frame::coding::CloseCode::Policy, + reason: "browser-test".into(), + }, + ))) + .await?; + } + "/delayed" => { + tokio::time::sleep(Duration::from_millis(100)).await; + socket + .send(tungstenite::Message::Text("still-open".into())) + .await?; + let _ = socket.next().await; + } + _ => { + while let Some(message) = socket.next().await { + let message = message?; + if message.is_close() { + break; + } + socket.send(message).await?; + } + } + } + Ok(()) +} + +async fn serve_mux_endpoint() -> std::io::Result<()> { + let address: std::net::SocketAddr = MUX_ADDRESS.parse().expect("valid mux test address"); + let server = ServerBuilder::new() + .with_addr(address) + .build() + .await + .map_err(std::io::Error::other)?; + + loop { + let session = match server.accept().await { + Ok(session) => session, + Err(_) => continue, + }; + tokio::spawn(async move { + while let Ok((send, mut recv)) = session.accept_bi().await { + while let Ok(Some(chunk)) = recv.read_chunk(16 * 1024).await { + if send.write_buf(chunk).await.is_err() { + return; + } + } + if send.finish().await.is_err() { + return; + } + } + }); + } +} diff --git a/websock-proto/Cargo.toml b/websock-proto/Cargo.toml index 0abbb26..981c079 100644 --- a/websock-proto/Cargo.toml +++ b/websock-proto/Cargo.toml @@ -2,6 +2,7 @@ name = "websock-proto" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "Protocol-level primitives shared across websock transports." repository = "https://github.com/foctal/websock" diff --git a/websock-tungstenite-mux/src/server.rs b/websock-tungstenite-mux/src/server.rs index 6bcfee9..06b6ff4 100644 --- a/websock-tungstenite-mux/src/server.rs +++ b/websock-tungstenite-mux/src/server.rs @@ -103,6 +103,7 @@ pub struct Server { } impl Server { + #[allow(clippy::result_large_err)] pub async fn accept(&self) -> Result { let (stream, _) = self .listener diff --git a/websock-tungstenite/Cargo.toml b/websock-tungstenite/Cargo.toml index 2c3455d..2b64715 100644 --- a/websock-tungstenite/Cargo.toml +++ b/websock-tungstenite/Cargo.toml @@ -2,6 +2,7 @@ name = "websock-tungstenite" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "Native transport implementation based on tokio-tungstenite." repository = "https://github.com/foctal/websock" @@ -17,7 +18,6 @@ tokio-rustls = { version = "0.26", default-features = false, features = ["ring"] rustls = { version = "0.23", default-features = false, features = ["ring", "std"] } rustls-native-certs = "0.8" rcgen = "0.14" -rustls-pemfile = "2.2" tokio-tungstenite = { version = "0.28", features = ["rustls-tls-webpki-roots"] } futures-util = { version = "0.3" } futures-core = { version = "0.3" } diff --git a/websock-tungstenite/src/tls/cert.rs b/websock-tungstenite/src/tls/cert.rs index f496a50..2337efe 100644 --- a/websock-tungstenite/src/tls/cert.rs +++ b/websock-tungstenite/src/tls/cert.rs @@ -1,7 +1,7 @@ //! Certificate handling utilities. use rustls::client::danger::ServerCertVerifier; -use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; +use rustls::pki_types::{CertificateDer, ServerName, UnixTime, pem::PemObject}; use std::{fs, path::Path, sync::Arc}; use websock_proto::{Error, Result}; @@ -26,8 +26,8 @@ pub fn load_certs(cert_path: &Path) -> Result>> { return Ok(vec![CertificateDer::from(cert_bytes)]); } - rustls_pemfile::certs(&mut &*cert_bytes) - .collect::, std::io::Error>>() + CertificateDer::pem_slice_iter(&cert_bytes) + .collect::, _>>() .map_err(|e| Error::Io(e.to_string())) } diff --git a/websock-tungstenite/src/tls/key.rs b/websock-tungstenite/src/tls/key.rs index 0f45ad2..4e264c4 100644 --- a/websock-tungstenite/src/tls/key.rs +++ b/websock-tungstenite/src/tls/key.rs @@ -1,6 +1,6 @@ //! Private key handling utilities. -use rustls::pki_types::{PrivateKeyDer, PrivatePkcs8KeyDer}; +use rustls::pki_types::{PrivateKeyDer, PrivatePkcs8KeyDer, pem::PemObject}; use std::{fs, path::Path}; use websock_proto::{Error, Result}; @@ -13,9 +13,7 @@ pub fn load_key(key_path: &Path) -> Result> { PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(key)) } else { // Decode PEM. - rustls_pemfile::private_key(&mut &*key) - .map_err(|e| Error::Tls(e.to_string()))? - .ok_or_else(|| Error::Io("no keys found".into()))? + PrivateKeyDer::from_pem_slice(&key).map_err(|e| Error::Tls(e.to_string()))? }; Ok(key) diff --git a/websock-tungstenite/src/tls/mod.rs b/websock-tungstenite/src/tls/mod.rs index 6079528..0e40273 100644 --- a/websock-tungstenite/src/tls/mod.rs +++ b/websock-tungstenite/src/tls/mod.rs @@ -65,11 +65,11 @@ impl TlsConfig { /// Create a new TLS configuration with the specified certificate and private key. pub fn with_cert(cert_path: &Path, key_path: &Path) -> Result { let client_config = TlsClientConfigBuilder::new_with_native_certs()? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); let server_config = TlsServerConfigBuilder::new_with_cert(cert_path, key_path)? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); Ok(Self { @@ -81,12 +81,12 @@ impl TlsConfig { /// Create a new TLS configuration with self-signed certificates (localhost). pub fn with_self_signed_certs() -> Result { let client_config = TlsClientConfigBuilder::new_with_native_certs()? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); let server_config = TlsServerConfigBuilder::new_with_self_signed_certs(vec!["localhost".into()])? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); Ok(Self { @@ -98,12 +98,12 @@ impl TlsConfig { /// Create a new TLS configuration with system certificates (server side is self-signed localhost). pub fn new_native_config() -> Result { let client_config = TlsClientConfigBuilder::new_with_native_certs()? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); let server_config = TlsServerConfigBuilder::new_with_self_signed_certs(vec!["localhost".into()])? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); Ok(Self { @@ -115,12 +115,12 @@ impl TlsConfig { /// Create a new TLS configuration with no certificate verification (testing only). pub fn new_insecure_config() -> Result { let client_config = TlsClientConfigBuilder::new_insecure()? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); let server_config = TlsServerConfigBuilder::new_with_self_signed_certs(vec!["localhost".into()])? - .with_alpn_protocols(vec![b"h3".to_vec()]) + .with_alpn_protocols(websock_proto::default_ws_alpn()) .build(); Ok(Self { diff --git a/websock-tungstenite/tests/connection.rs b/websock-tungstenite/tests/connection.rs index a06822e..6a32d0d 100644 --- a/websock-tungstenite/tests/connection.rs +++ b/websock-tungstenite/tests/connection.rs @@ -1,3 +1,4 @@ +use futures_util::{SinkExt, StreamExt}; use websock_proto::{Error, Message, WebSocketLimits}; use websock_tungstenite::{ClientBuilder, ServerBuilder}; @@ -12,6 +13,7 @@ async fn client_server_round_trip_text_and_binary() { let server_task = tokio::spawn(async move { let mut connection = server.accept().await.expect("accept connection"); + assert_eq!(connection.negotiated_subprotocol(), Some("test.v1")); assert_eq!( connection.recv().await.expect("receive text"), Message::Text("hello".into()) @@ -28,6 +30,7 @@ async fn client_server_round_trip_text_and_binary() { .connect(&format!("ws://{address}")) .await .expect("connect client"); + assert_eq!(connection.negotiated_subprotocol(), Some("test.v1")); connection .send(Message::Text("hello".into())) .await @@ -83,3 +86,121 @@ async fn invalid_server_limits_return_an_error() { }; assert!(matches!(err, Error::Protocol(_))); } + +#[tokio::test] +async fn split_connection_sends_and_receives_independently() { + let server = ServerBuilder::new().build().await.expect("bind server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let mut connection = server.accept().await.expect("accept connection"); + let message = connection.recv().await.expect("receive message"); + connection.send(message).await.expect("echo message"); + assert!(matches!(connection.recv().await, Err(Error::Closed))); + }); + + let connection = ClientBuilder::new() + .build() + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + let (mut sink, mut stream) = websock_tungstenite::stream::split(connection); + sink.send(Message::Text("split".into())) + .await + .expect("send through split sink"); + assert_eq!( + stream.next().await.expect("stream item").expect("message"), + Message::Text("split".into()) + ); + sink.close().await.expect("close split sink"); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn tls_client_server_round_trip() { + let (certificates, key) = + websock_tungstenite::tls::generate_self_signed_pair_der(vec!["localhost".into()]) + .expect("generate certificate"); + let certificate_bytes = certificates + .iter() + .map(|certificate| certificate.as_ref().to_vec()) + .collect::>(); + let server = ServerBuilder::new() + .with_certificate(certificates, key) + .expect("configure certificate") + .with_default_alpn() + .build() + .await + .expect("bind TLS server"); + let address = server.local_addr().expect("server address"); + + let server_task = tokio::spawn(async move { + let mut connection = server.accept().await.expect("accept TLS connection"); + assert_eq!( + connection.recv().await.expect("receive TLS message"), + Message::Text("secure".into()) + ); + connection + .send(Message::Binary(b"reply".as_slice().into())) + .await + .expect("send TLS reply"); + connection.close().await.expect("close TLS connection"); + }); + + let client = ClientBuilder::new() + .with_default_alpn() + .with_server_certificates(certificate_bytes) + .expect("configure client roots"); + let mut connection = client + .connect(&format!("wss://localhost:{}", address.port())) + .await + .expect("connect TLS client"); + connection + .send(Message::Text("secure".into())) + .await + .expect("send TLS message"); + assert_eq!( + connection.recv().await.expect("receive TLS reply"), + Message::Binary(b"reply".as_slice().into()) + ); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn received_close_frame_details_are_preserved() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind raw server"); + let address = listener.local_addr().expect("server address"); + let server_task = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept TCP connection"); + let mut socket = tokio_tungstenite::accept_async(stream) + .await + .expect("accept WebSocket"); + socket + .send(tokio_tungstenite::tungstenite::Message::Close(Some( + tokio_tungstenite::tungstenite::protocol::CloseFrame { + code: + tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode::Policy, + reason: "policy".into(), + }, + ))) + .await + .expect("send close frame"); + }); + + let mut connection = ClientBuilder::new() + .build() + .connect(&format!("ws://{address}")) + .await + .expect("connect client"); + assert!(matches!(connection.recv().await, Err(Error::Closed))); + assert_eq!( + connection.close_frame(), + Some(websock_proto::CloseFrame { + code: 1008, + reason: "policy".into(), + }) + ); + server_task.await.expect("server task"); +} diff --git a/websock-wasm-demo/Cargo.toml b/websock-wasm-demo/Cargo.toml index b8cfc26..dcba82b 100644 --- a/websock-wasm-demo/Cargo.toml +++ b/websock-wasm-demo/Cargo.toml @@ -2,8 +2,10 @@ name = "websock-wasm-demo" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true publish = false +license = "MIT" [lib] crate-type = ["cdylib"] @@ -11,8 +13,8 @@ crate-type = ["cdylib"] [dependencies] websock = { path = "../websock" } websock-mux = { path = "../websock-mux" } -wasm-bindgen = "0.2" -wasm-bindgen-futures = "0.4" +wasm-bindgen = { workspace = true } +wasm-bindgen-futures = { workspace = true } console_error_panic_hook = "0.1" web-sys = { version = "0.3", features = ["console", "Document", "HtmlInputElement", "HtmlButtonElement", "Window"] } futures-util = { version = "0.3", features = ["sink"] } diff --git a/websock-wasm-mux/Cargo.toml b/websock-wasm-mux/Cargo.toml index 5377932..754eba2 100644 --- a/websock-wasm-mux/Cargo.toml +++ b/websock-wasm-mux/Cargo.toml @@ -2,6 +2,7 @@ name = "websock-wasm-mux" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "WebAssembly WebSocket multiplexing transport (WebTransport-like streams over WebSocket)." repository = "https://github.com/foctal/websock" @@ -18,4 +19,7 @@ bytes = { workspace = true } futures-channel = { version = "0.3", features = ["sink"] } futures-io = "0.3" futures-util = { version = "0.3", default-features = false, features = ["alloc", "sink"] } -wasm-bindgen-futures = "0.4" +wasm-bindgen-futures = { workspace = true } + +[dev-dependencies] +wasm-bindgen-test = { workspace = true } diff --git a/websock-wasm-mux/tests/browser.rs b/websock-wasm-mux/tests/browser.rs new file mode 100644 index 0000000..82da36a --- /dev/null +++ b/websock-wasm-mux/tests/browser.rs @@ -0,0 +1,23 @@ +use wasm_bindgen_test::*; + +wasm_bindgen_test_configure!(run_in_browser); + +#[wasm_bindgen_test(async)] +async fn bidirectional_stream_matches_native_behavior() { + let client = websock_wasm_mux::ClientBuilder::new().build(); + let session = client + .connect("ws://127.0.0.1:32124") + .await + .expect("connect mux endpoint"); + let (send, mut recv) = session.open_bi().expect("open bi stream"); + + send.write_all(b"browser-mux").await.expect("write request"); + send.finish().await.expect("finish request"); + + let mut response = Vec::new(); + while let Some(chunk) = recv.read_chunk(1024).await.expect("read response") { + response.extend_from_slice(&chunk); + } + assert_eq!(response, b"browser-mux"); + session.shutdown().await.expect("shutdown mux session"); +} diff --git a/websock-wasm/Cargo.toml b/websock-wasm/Cargo.toml index 4a0553c..7ae04a6 100644 --- a/websock-wasm/Cargo.toml +++ b/websock-wasm/Cargo.toml @@ -2,6 +2,7 @@ name = "websock-wasm" version.workspace = true edition.workspace = true +rust-version.workspace = true authors.workspace = true description = "WebAssembly transport implementation for browser-based WebSockets." repository = "https://github.com/foctal/websock" @@ -12,8 +13,8 @@ license = "MIT" [dependencies] websock-proto = { workspace = true } -wasm-bindgen = "0.2" -wasm-bindgen-futures = "0.4" +wasm-bindgen = { workspace = true } +wasm-bindgen-futures = { workspace = true } web-sys = { version = "0.3", features = [ "WebSocket", "MessageEvent", @@ -26,3 +27,7 @@ futures-util = { version = "0.3" } futures-channel = "0.3" futures-core = { version = "0.3" } futures-sink = { version = "0.3" } + +[dev-dependencies] +gloo-timers = { version = "0.3", features = ["futures"] } +wasm-bindgen-test = { workspace = true } diff --git a/websock-wasm/tests/browser.rs b/websock-wasm/tests/browser.rs new file mode 100644 index 0000000..e3ad9ed --- /dev/null +++ b/websock-wasm/tests/browser.rs @@ -0,0 +1,73 @@ +use futures_util::StreamExt; +use gloo_timers::future::TimeoutFuture; +use wasm_bindgen_test::*; +use websock_proto::{Error, Message}; + +wasm_bindgen_test_configure!(run_in_browser); + +#[wasm_bindgen_test(async)] +async fn connect_failure_is_reported() { + let result = websock_wasm::connect( + "not-a-websocket-url", + websock_proto::ConnectOptions::default(), + ) + .await; + assert!(result.is_err()); +} + +#[wasm_bindgen_test(async)] +async fn receive_queue_overflow_closes_the_connection() { + let mut connection = websock_wasm::connect( + "ws://127.0.0.1:32123/overflow", + websock_proto::ConnectOptions::default(), + ) + .await + .expect("connect overflow endpoint"); + + TimeoutFuture::new(250).await; + let mut received = 0; + let terminal = loop { + match connection.recv().await { + Ok(_) => received += 1, + Err(error) => break error, + } + }; + assert!(received < 256); + assert!(matches!(terminal, Error::Closed | Error::Protocol(_))); +} + +#[wasm_bindgen_test(async)] +async fn close_frame_details_are_preserved() { + let mut connection = websock_wasm::connect( + "ws://127.0.0.1:32123/close", + websock_proto::ConnectOptions::default(), + ) + .await + .expect("connect close endpoint"); + + assert!(matches!(connection.recv().await, Err(Error::Closed))); + assert_eq!( + connection.close_frame(), + Some(websock_proto::CloseFrame { + code: 1008, + reason: "browser-test".into(), + }) + ); +} + +#[wasm_bindgen_test(async)] +async fn dropping_one_split_half_keeps_the_other_half_alive() { + let connection = websock_wasm::connect( + "ws://127.0.0.1:32123/delayed", + websock_proto::ConnectOptions::default(), + ) + .await + .expect("connect delayed endpoint"); + let (sink, mut stream) = websock_wasm::stream::split(connection); + drop(sink); + + assert_eq!( + stream.next().await.expect("stream item").expect("message"), + Message::Text("still-open".into()) + ); +} From 39e04216bc4351b8b303dfee8291f18342d37286 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 21:39:28 +0900 Subject: [PATCH 11/17] build: add fuzzing and dependency policy checks --- .github/workflows/rust.yml | 22 -------------- Cargo.toml | 6 ++++ deny.toml | 24 +++++++++++++++ fuzz/Cargo.toml | 34 +++++++++++++++++++++ fuzz/fuzz_targets/frame_decode.rs | 14 +++++++++ fuzz/fuzz_targets/mux_sequence.rs | 47 ++++++++++++++++++++++++++++++ fuzz/fuzz_targets/varint_decode.rs | 14 +++++++++ 7 files changed, 139 insertions(+), 22 deletions(-) delete mode 100644 .github/workflows/rust.yml create mode 100644 deny.toml create mode 100644 fuzz/Cargo.toml create mode 100644 fuzz/fuzz_targets/frame_decode.rs create mode 100644 fuzz/fuzz_targets/mux_sequence.rs create mode 100644 fuzz/fuzz_targets/varint_decode.rs diff --git a/.github/workflows/rust.yml b/.github/workflows/rust.yml deleted file mode 100644 index 9fd45e0..0000000 --- a/.github/workflows/rust.yml +++ /dev/null @@ -1,22 +0,0 @@ -name: Rust - -on: - push: - branches: [ "main" ] - pull_request: - branches: [ "main" ] - -env: - CARGO_TERM_COLOR: always - -jobs: - build: - - runs-on: ubuntu-latest - - steps: - - uses: actions/checkout@v4 - - name: Build - run: cargo build --verbose - - name: Run tests - run: cargo test --verbose diff --git a/Cargo.toml b/Cargo.toml index 01c3ead..17dbef4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -6,15 +6,18 @@ members = [ "websock-tungstenite", "websock-wasm", "websock-wasm-demo", + "websock-browser-test-server", "websock-mux", "websock-tungstenite-mux", "websock-wasm-mux", "websock-mux-proto" ] +exclude = ["fuzz"] [workspace.package] version = "0.4.0" edition = "2024" +rust-version = "1.88" authors = ["shellrow "] [workspace.dependencies] @@ -25,3 +28,6 @@ websock-mux-proto = { path = "websock-mux-proto", version = "0.4.0" } websock-tungstenite-mux = { path = "websock-tungstenite-mux", version = "0.4.0" } websock-wasm-mux = { path = "websock-wasm-mux", version = "0.4.0" } bytes = "1" +wasm-bindgen = "=0.2.126" +wasm-bindgen-futures = "=0.4.76" +wasm-bindgen-test = "=0.3.76" diff --git a/deny.toml b/deny.toml new file mode 100644 index 0000000..64ea477 --- /dev/null +++ b/deny.toml @@ -0,0 +1,24 @@ +[advisories] +yanked = "deny" + +[licenses] +confidence-threshold = 0.8 +allow = [ + "Apache-2.0", + "Apache-2.0 WITH LLVM-exception", + "BSD-2-Clause", + "BSD-3-Clause", + "CDLA-Permissive-2.0", + "ISC", + "MIT", + "Unicode-3.0", +] + +[bans] +multiple-versions = "warn" +wildcards = "allow" + +[sources] +unknown-registry = "deny" +unknown-git = "deny" +allow-registry = ["https://github.com/rust-lang/crates.io-index"] diff --git a/fuzz/Cargo.toml b/fuzz/Cargo.toml new file mode 100644 index 0000000..e42805a --- /dev/null +++ b/fuzz/Cargo.toml @@ -0,0 +1,34 @@ +[package] +name = "websock-fuzz" +version = "0.0.0" +edition = "2024" +publish = false + +[package.metadata] +cargo-fuzz = true + +[dependencies] +bytes = "1" +libfuzzer-sys = "0.4" +websock-mux-proto = { path = "../websock-mux-proto" } + +[[bin]] +name = "varint_decode" +path = "fuzz_targets/varint_decode.rs" +test = false +doc = false +bench = false + +[[bin]] +name = "frame_decode" +path = "fuzz_targets/frame_decode.rs" +test = false +doc = false +bench = false + +[[bin]] +name = "mux_sequence" +path = "fuzz_targets/mux_sequence.rs" +test = false +doc = false +bench = false diff --git a/fuzz/fuzz_targets/frame_decode.rs b/fuzz/fuzz_targets/frame_decode.rs new file mode 100644 index 0000000..d09cc6c --- /dev/null +++ b/fuzz/fuzz_targets/frame_decode.rs @@ -0,0 +1,14 @@ +#![no_main] + +use bytes::{Buf, Bytes}; +use libfuzzer_sys::fuzz_target; +use websock_mux_proto::Frame; + +fuzz_target!(|data: &[u8]| { + let mut input = Bytes::copy_from_slice(data); + while input.has_remaining() { + if Frame::decode(&mut input).is_err() { + break; + } + } +}); diff --git a/fuzz/fuzz_targets/mux_sequence.rs b/fuzz/fuzz_targets/mux_sequence.rs new file mode 100644 index 0000000..c393fca --- /dev/null +++ b/fuzz/fuzz_targets/mux_sequence.rs @@ -0,0 +1,47 @@ +#![no_main] + +use std::collections::HashSet; + +use bytes::{Buf, Bytes}; +use libfuzzer_sys::fuzz_target; +use websock_mux_proto::{Frame, StreamId}; + +fuzz_target!(|data: &[u8]| { + let mut input = Bytes::copy_from_slice(data); + let mut receive_streams = HashSet::::new(); + let mut send_streams = HashSet::::new(); + + while input.has_remaining() { + let Ok(frame) = Frame::decode(&mut input) else { + break; + }; + match frame { + Frame::OpenUni { id } => { + receive_streams.insert(id); + } + Frame::OpenBi { id } => { + receive_streams.insert(id); + send_streams.insert(id); + } + Frame::Stream { id, fin, .. } => { + if fin { + receive_streams.remove(&id); + } + } + Frame::ResetStream { id, .. } => { + receive_streams.remove(&id); + } + Frame::StopSending { id, .. } => { + send_streams.remove(&id); + } + Frame::MaxStreamData { id, .. } => { + let _ = send_streams.contains(&id); + } + Frame::ConnectionClose { .. } => { + receive_streams.clear(); + send_streams.clear(); + break; + } + } + } +}); diff --git a/fuzz/fuzz_targets/varint_decode.rs b/fuzz/fuzz_targets/varint_decode.rs new file mode 100644 index 0000000..ce9960f --- /dev/null +++ b/fuzz/fuzz_targets/varint_decode.rs @@ -0,0 +1,14 @@ +#![no_main] + +use bytes::{Buf, Bytes}; +use libfuzzer_sys::fuzz_target; +use websock_mux_proto::VarInt; + +fuzz_target!(|data: &[u8]| { + let mut input = Bytes::copy_from_slice(data); + while input.has_remaining() { + if VarInt::decode(&mut input).is_err() { + break; + } + } +}); From e7043369a85bb8f8360270d50d2bfc271bb2ce87 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 18 Jul 2026 21:40:51 +0900 Subject: [PATCH 12/17] docs: update docs --- CHANGELOG.md | 29 +++++++++++++++++++++++++++++ RELEASING.md | 26 ++++++++++++++++++++++++++ TODO.md | 23 ++++++++++++----------- 3 files changed, 67 insertions(+), 11 deletions(-) create mode 100644 CHANGELOG.md create mode 100644 RELEASING.md diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..fa2c2fb --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,29 @@ +# Changelog + +All notable changes to this project will be documented in this file. + +The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), +and this project uses [Semantic Versioning](https://semver.org/spec/v2.0.0.html). + +## Unreleased + +### Added + +- Explicit mux session shutdown with task-completion semantics. +- Negotiated subprotocol and received close-frame metadata accessors. +- Native end-to-end coverage for split connections, bidirectional mux streams, + malformed frames, reset/stop propagation, TLS, and session-handle cleanup. +- Browser integration coverage for connection failures, receive-queue overflow, + close metadata, split ownership, and mux parity. +- A versioned mux wire-protocol specification. + +### Changed + +- The mux TLS helpers now reuse the native transport implementation. +- PEM parsing now uses `rustls-pki-types` directly instead of the unmaintained + `rustls-pemfile` crate. +- Default TLS ALPN configuration advertises only HTTP/1.1. + +### Fixed + +- Non-boolean mux stream FIN fields are rejected as malformed. diff --git a/RELEASING.md b/RELEASING.md new file mode 100644 index 0000000..7fe1132 --- /dev/null +++ b/RELEASING.md @@ -0,0 +1,26 @@ +# Release Guide + +Release and publishing operations are intentionally manual. + +1. Confirm that `CHANGELOG.md` describes all user-visible changes and move the + Unreleased entries under the new version and date. +2. Run `cargo fmt --all -- --check`. +3. Run `cargo clippy --workspace --all-targets --all-features -- -D warnings`. +4. Run `cargo test --workspace --all-targets`. +5. Run `cargo doc --workspace --no-deps` with `RUSTDOCFLAGS=-Dwarnings`. +6. Check the WASM packages for `wasm32-unknown-unknown`. +7. Run `cargo package --list -p ` for every publishable crate and inspect + each file list. Then run `cargo package -p ` in dependency order. +8. Update versions and internal dependency requirements together. +9. Create the release commit and tag manually. +10. Publish crates in dependency order, then create the GitHub Release manually. + +The publishable dependency order is: + +1. `websock-proto` +2. `websock-mux-proto` +3. `websock-tungstenite` and `websock-wasm` +4. `websock-tungstenite-mux` and `websock-wasm-mux` +5. `websock` and `websock-mux` + +`websock-wasm-demo` is not published. diff --git a/TODO.md b/TODO.md index 58140a2..d88e0db 100644 --- a/TODO.md +++ b/TODO.md @@ -45,9 +45,9 @@ items remain release blockers or follow-up work. - [x] Add native end-to-end tests for connect/subprotocol negotiation, text/binary round trips, graceful close, oversized-message rejection, mux uni streams, differing flow-control windows, and invalid-limit rejection. -- [ ] Extend native end-to-end tests with split operation, mux bidirectional +- [x] Extend native end-to-end tests with split operation, mux bidirectional streams, malformed wire frames, reset/stop propagation, and TLS round trips. -- [ ] Add browser tests under `wasm-bindgen-test` for connect failure, queue +- [x] Add browser tests under `wasm-bindgen-test` for connect failure, queue overflow, close, split ownership, and mux parity. - [x] Ensure WebSocket-level frame/message limits are applied in Tungstenite configuration before a complete oversized message is buffered. @@ -55,7 +55,7 @@ items remain release blockers or follow-up work. browser receive queue, message-size, and `bufferedAmount` bounds. - [x] Advertise only HTTP/1.1 in the default ALPN list because the current Tungstenite transport does not implement WebSocket over HTTP/2. -- [ ] Add explicit connection/session shutdown APIs with completion semantics +- [x] Add explicit connection/session shutdown APIs with completion semantics and task cancellation. Dropping the last mux session handle should not leave detached tasks and a socket alive indefinitely. @@ -63,19 +63,20 @@ items remain release blockers or follow-up work. - [ ] Preserve structured error sources and I/O error kinds instead of reducing all transport errors to strings. This is a semver-sensitive public API change. -- [ ] Expose negotiated subprotocol and close-frame details consistently on +- [x] Expose negotiated subprotocol and close-frame details consistently on native and browser connections. -- [ ] Define and document the mux wire protocol, compatibility policy, error +- [x] Define and document the mux wire protocol, compatibility policy, error codes, stream state machine, and flow-control rules. -- [ ] Expand crate-level and public-item documentation and enable +- [x] Expand crate-level and public-item documentation and enable `#![warn(missing_docs)]` incrementally. - [x] Add baseline CI for formatting, Clippy with warnings denied, tests, documentation, Linux/macOS/Windows, and `wasm32-unknown-unknown`. -- [ ] Extend CI with an explicit MSRV job, minimal feature-set builds, dependency +- [x] Extend CI with an explicit MSRV job, minimal feature-set builds, dependency policy checks, and security audits. -- [ ] Declare and test the MSRV, add changelog/release guidance, and verify - package contents with `cargo package --list` for every published crate. -- [ ] Review dependency features and duplicate native/mux TLS implementations to +- [x] Declare and test the MSRV, and add changelog and release guidance. +- [ ] Manually verify package contents with `cargo package --list` for every + published crate as part of the release process. +- [x] Review dependency features and duplicate native/mux TLS implementations to reduce compile time, binary size, and maintenance drift. -- [ ] Add fuzz targets for varint/frame decoding and stateful mux frame +- [x] Add fuzz targets for varint/frame decoding and stateful mux frame sequences, plus long-running concurrency and backpressure tests. From e8b42b1e5b1e472725af53cefd6f5b0c37531879 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sun, 19 Jul 2026 17:36:35 +0900 Subject: [PATCH 13/17] feat: preserve structured transport error sources --- websock-proto/src/error.rs | 78 ++++++++++++++++++++++++- websock-proto/src/lib.rs | 2 +- websock-tungstenite-mux/src/client.rs | 4 +- websock-tungstenite-mux/src/server.rs | 19 ++---- websock-tungstenite-mux/src/session.rs | 25 ++++---- websock-tungstenite/src/builder.rs | 6 +- websock-tungstenite/src/connection.rs | 60 ++++++++++++------- websock-tungstenite/src/server.rs | 45 ++++---------- websock-tungstenite/src/tls/cert.rs | 4 +- websock-tungstenite/src/tls/key.rs | 4 +- websock-tungstenite/src/tls/mod.rs | 10 ++-- websock-tungstenite/tests/connection.rs | 7 ++- 12 files changed, 163 insertions(+), 101 deletions(-) diff --git a/websock-proto/src/error.rs b/websock-proto/src/error.rs index f0572ed..0ba3503 100644 --- a/websock-proto/src/error.rs +++ b/websock-proto/src/error.rs @@ -1,8 +1,14 @@ +use std::error::Error as StdError; +use std::io; + use thiserror::Error; /// Result type for protocol and transport operations. pub type Result = std::result::Result; +/// A thread-safe, type-erased error retained as the source of a transport error. +pub type BoxError = Box; + /// Errors that can occur when sending or receiving WebSocket messages. #[derive(Debug, Error)] pub enum Error { @@ -20,11 +26,15 @@ pub enum Error { /// An IO failure occurred while reading or writing. #[error("io error: {0}")] - Io(String), + Io(#[from] io::Error), /// A TLS handshake or validation error occurred. #[error("tls error: {0}")] - Tls(String), + Tls(#[source] BoxError), + + /// An error reported by the underlying WebSocket transport. + #[error("transport error: {0}")] + Transport(#[source] BoxError), /// The operation is not supported on the current platform or configuration. #[error("unsupported: {0}")] @@ -44,8 +54,72 @@ pub enum Error { } impl Error { + /// Create an `Error::Tls` while retaining the concrete source error. + pub fn tls(error: E) -> Self + where + E: StdError + Send + Sync + 'static, + { + Self::Tls(Box::new(error)) + } + + /// Create an `Error::Transport` while retaining the concrete source error. + pub fn transport(error: E) -> Self + where + E: StdError + Send + Sync + 'static, + { + Self::Transport(Box::new(error)) + } + /// Create an `Error::Other` from any displayable error value. pub fn other(e: E) -> Self { Self::Other(e.to_string()) } + + /// Return the underlying I/O error kind, when this is an I/O error. + pub fn io_kind(&self) -> Option { + match self { + Self::Io(error) => Some(error.kind()), + _ => None, + } + } +} + +#[cfg(test)] +mod tests { + use std::error::Error as _; + + use super::Error; + + #[test] + fn io_error_preserves_kind_and_source() { + let error = Error::from(std::io::Error::new( + std::io::ErrorKind::ConnectionReset, + "peer reset", + )); + + assert_eq!(error.io_kind(), Some(std::io::ErrorKind::ConnectionReset)); + assert_eq!( + error + .source() + .and_then(|source| source.downcast_ref::()) + .map(std::io::Error::kind), + Some(std::io::ErrorKind::ConnectionReset) + ); + } + + #[test] + fn transport_error_preserves_concrete_source() { + let error = Error::transport(std::io::Error::new( + std::io::ErrorKind::TimedOut, + "transport timeout", + )); + + assert_eq!( + error + .source() + .and_then(|source| source.downcast_ref::()) + .map(std::io::Error::kind), + Some(std::io::ErrorKind::TimedOut) + ); + } } diff --git a/websock-proto/src/lib.rs b/websock-proto/src/lib.rs index de89c6f..233b581 100644 --- a/websock-proto/src/lib.rs +++ b/websock-proto/src/lib.rs @@ -12,7 +12,7 @@ mod options; mod transport; pub use bytes::Bytes; -pub use error::{Error, Result}; +pub use error::{BoxError, Error, Result}; pub use message::{CloseFrame, Message}; pub use options::{ConnectOptions, ServerOptions, WebSocketLimits, default_ws_alpn}; pub use transport::{LocalBoxFuture, WebSocketConnection}; diff --git a/websock-tungstenite-mux/src/client.rs b/websock-tungstenite-mux/src/client.rs index bc395e7..f537b74 100644 --- a/websock-tungstenite-mux/src/client.rs +++ b/websock-tungstenite-mux/src/client.rs @@ -67,9 +67,7 @@ impl Client { self.limits.validate()?; validate_client_protocols(&self.opts)?; - let mut request = url - .into_client_request() - .map_err(|e| websock_proto::Error::InvalidUrl(e.to_string()))?; + let mut request = url.into_client_request().map_err(Error::transport)?; let headers = request.headers_mut(); for (k, v) in self.opts.headers.iter() { diff --git a/websock-tungstenite-mux/src/server.rs b/websock-tungstenite-mux/src/server.rs index 06b6ff4..bdca0b0 100644 --- a/websock-tungstenite-mux/src/server.rs +++ b/websock-tungstenite-mux/src/server.rs @@ -70,9 +70,7 @@ where A: ToSocketAddrs, { limits.validate()?; - let listener = TcpListener::bind(addr) - .await - .map_err(|e| Error::Io(e.to_string()))?; + let listener = TcpListener::bind(addr).await.map_err(Error::Io)?; let headers = prepare_headers(&opts)?; validate_protocols(&opts)?; @@ -105,17 +103,10 @@ pub struct Server { impl Server { #[allow(clippy::result_large_err)] pub async fn accept(&self) -> Result { - let (stream, _) = self - .listener - .accept() - .await - .map_err(|e| Error::Io(e.to_string()))?; + let (stream, _) = self.listener.accept().await.map_err(Error::Io)?; let (stream, _is_tls): (ServerStream, bool) = if let Some(acceptor) = &self.acceptor { - let tls_stream = acceptor - .accept(stream) - .await - .map_err(|e| Error::Tls(e.to_string()))?; + let tls_stream = acceptor.accept(stream).await.map_err(Error::tls)?; (Box::new(tls_stream), true) } else { (Box::new(stream), false) @@ -164,8 +155,6 @@ impl Server { } pub fn local_addr(&self) -> Result { - self.listener - .local_addr() - .map_err(|e| Error::Io(e.to_string())) + self.listener.local_addr().map_err(Error::Io) } } diff --git a/websock-tungstenite-mux/src/session.rs b/websock-tungstenite-mux/src/session.rs index 9e4f043..35e7d0e 100644 --- a/websock-tungstenite-mux/src/session.rs +++ b/websock-tungstenite-mux/src/session.rs @@ -1586,10 +1586,10 @@ fn io_invalid_input(message: &'static str) -> io::Error { fn io_from_error(err: Error) -> io::Error { match err { Error::Closed => io_closed(), - Error::Io(message) => io::Error::other(message), + Error::Io(error) => error, + Error::Tls(error) | Error::Transport(error) => io::Error::other(error), Error::Protocol(message) | Error::InvalidUrl(message) - | Error::Tls(message) | Error::StreamId(message) | Error::Unsupported(message) | Error::FrameDecode(message) @@ -1602,14 +1602,9 @@ pub(crate) fn map_tungstenite_err(e: tungstenite::Error) -> Error { use tungstenite::Error as E; match e { E::ConnectionClosed | E::AlreadyClosed => Error::Closed, - E::Io(io) => Error::Io(io.to_string()), - E::Tls(tls) => Error::Tls(tls.to_string()), - E::Url(url) => Error::InvalidUrl(url.to_string()), - E::Protocol(err) => Error::Protocol(err.to_string()), - E::Utf8(err) => Error::Protocol(err), - E::Capacity(err) => Error::Protocol(err.to_string()), - E::HttpFormat(err) => Error::Protocol(err.to_string()), - other => Error::Other(other.to_string()), + E::Io(io) => Error::Io(io), + E::Tls(tls) => Error::tls(tls), + other => Error::transport(other), } } @@ -1905,4 +1900,14 @@ mod tests { assert!(outbound_rx.try_recv().is_err()); send.write(b"x").await.expect("remaining clone is usable"); } + + #[test] + fn stream_io_conversion_preserves_io_error_kind() { + let error = io_from_error(Error::Io(io::Error::new( + io::ErrorKind::ConnectionReset, + "peer reset", + ))); + + assert_eq!(error.kind(), io::ErrorKind::ConnectionReset); + } } diff --git a/websock-tungstenite/src/builder.rs b/websock-tungstenite/src/builder.rs index 27e3fdf..567fc8b 100644 --- a/websock-tungstenite/src/builder.rs +++ b/websock-tungstenite/src/builder.rs @@ -110,9 +110,7 @@ impl ClientBuilder { { let mut roots = RootCertStore::empty(); for cert in chain { - roots - .add(CertificateDer::from(cert)) - .map_err(|e| Error::Tls(e.to_string()))?; + roots.add(CertificateDer::from(cert)).map_err(Error::tls)?; } let config = ClientConfig::builder() .with_root_certificates(roots) @@ -265,7 +263,7 @@ impl ServerBuilder { let config = ServerConfig::builder() .with_no_client_auth() .with_single_cert(chain, key) - .map_err(|e| Error::Tls(e.to_string()))?; + .map_err(Error::tls)?; self.tls = Some(config); Ok(self) } diff --git a/websock-tungstenite/src/connection.rs b/websock-tungstenite/src/connection.rs index fc90daa..d0f40ef 100644 --- a/websock-tungstenite/src/connection.rs +++ b/websock-tungstenite/src/connection.rs @@ -39,9 +39,7 @@ pub async fn connect_with_tls( .max_write_buffer_size(opts.limits.max_write_buffer_size) .max_message_size(Some(opts.limits.max_message_size)) .max_frame_size(Some(opts.limits.max_frame_size)); - let mut req = url - .into_client_request() - .map_err(|e| Error::InvalidUrl(e.to_string()))?; + let mut req = url.into_client_request().map_err(Error::transport)?; // Apply configured headers and subprotocols. { @@ -70,16 +68,8 @@ pub async fn connect_with_tls( .map_err(map_tungstenite_err)?; let info = ConnectionInfo { - peer: ws - .get_ref() - .get_ref() - .peer_addr() - .map_err(|e| Error::Io(e.to_string()))?, - local: ws - .get_ref() - .get_ref() - .local_addr() - .map_err(|e| Error::Io(e.to_string()))?, + peer: ws.get_ref().get_ref().peer_addr().map_err(Error::Io)?, + local: ws.get_ref().get_ref().local_addr().map_err(Error::Io)?, is_tls: matches!(ws.get_ref(), tokio_tungstenite::MaybeTlsStream::Rustls(_)), subprotocol: resp .headers() @@ -219,13 +209,41 @@ pub(crate) fn map_tungstenite_err(e: tungstenite::Error) -> Error { use tungstenite::Error as E; match e { E::ConnectionClosed | E::AlreadyClosed => Error::Closed, - E::Io(io) => Error::Io(io.to_string()), - E::Tls(tls) => Error::Tls(tls.to_string()), - E::Url(url) => Error::InvalidUrl(url.to_string()), - E::Protocol(err) => Error::Protocol(err.to_string()), - E::Utf8(err) => Error::Protocol(err), - E::Capacity(err) => Error::Protocol(err.to_string()), - E::HttpFormat(err) => Error::Protocol(err.to_string()), - other => Error::Other(other.to_string()), + E::Io(io) => Error::Io(io), + E::Tls(tls) => Error::tls(tls), + other => Error::transport(other), + } +} + +#[cfg(test)] +mod tests { + use std::error::Error as _; + use std::io; + + use super::map_tungstenite_err; + use websock_proto::Error; + + #[test] + fn tungstenite_io_error_preserves_kind() { + let error = map_tungstenite_err(tokio_tungstenite::tungstenite::Error::Io(io::Error::new( + io::ErrorKind::ConnectionAborted, + "connection aborted", + ))); + + assert_eq!(error.io_kind(), Some(io::ErrorKind::ConnectionAborted)); + } + + #[test] + fn tungstenite_protocol_error_is_retained_as_source() { + let error = map_tungstenite_err(tokio_tungstenite::tungstenite::Error::Protocol( + tokio_tungstenite::tungstenite::error::ProtocolError::ResetWithoutClosingHandshake, + )); + + assert!(matches!(error, Error::Transport(_))); + assert!( + error + .source() + .is_some_and(|source| source.is::()) + ); } } diff --git a/websock-tungstenite/src/server.rs b/websock-tungstenite/src/server.rs index 4ce6be8..c2c1866 100644 --- a/websock-tungstenite/src/server.rs +++ b/websock-tungstenite/src/server.rs @@ -23,9 +23,7 @@ where A: ToSocketAddrs, { opts.limits.validate()?; - let listener = TcpListener::bind(addr) - .await - .map_err(|e| Error::Io(e.to_string()))?; + let listener = TcpListener::bind(addr).await.map_err(Error::Io)?; let headers = prepare_headers(&opts)?; validate_protocols(&opts)?; @@ -68,20 +66,13 @@ impl Server { /// Accept an incoming WebSocket connection. #[allow(clippy::result_large_err)] pub async fn accept(&self) -> Result> { - let (stream, _addr) = self - .listener - .accept() - .await - .map_err(|e| Error::Io(e.to_string()))?; + let (stream, _addr) = self.listener.accept().await.map_err(Error::Io)?; - let peer = stream.peer_addr().map_err(|e| Error::Io(e.to_string()))?; - let local = stream.local_addr().map_err(|e| Error::Io(e.to_string()))?; + let peer = stream.peer_addr().map_err(Error::Io)?; + let local = stream.local_addr().map_err(Error::Io)?; let (stream, is_tls): (ServerStream, bool) = if let Some(acceptor) = &self.acceptor { - let tls_stream = acceptor - .accept(stream) - .await - .map_err(|e| Error::Tls(e.to_string()))?; + let tls_stream = acceptor.accept(stream).await.map_err(Error::tls)?; (Box::new(tls_stream), true) } else { (Box::new(stream), false) @@ -139,34 +130,22 @@ impl Server { pub async fn accept_tls( &self, ) -> Result>> { - let (stream, _addr) = self - .listener - .accept() - .await - .map_err(|e| Error::Io(e.to_string()))?; + let (stream, _addr) = self.listener.accept().await.map_err(Error::Io)?; let tls_stream = self .acceptor .as_ref() - .ok_or_else(|| Error::Tls("missing tls acceptor".into()))? + .ok_or_else(|| Error::Protocol("missing tls acceptor".into()))? .accept(stream) .await - .map_err(|e| Error::Tls(e.to_string()))?; + .map_err(Error::tls)?; let headers = Arc::clone(&self.headers); let protocols = Arc::clone(&self.protocols); let mut info = ConnectionInfo { - peer: tls_stream - .get_ref() - .0 - .peer_addr() - .map_err(|e| Error::Io(e.to_string()))?, - local: tls_stream - .get_ref() - .0 - .local_addr() - .map_err(|e| Error::Io(e.to_string()))?, + peer: tls_stream.get_ref().0.peer_addr().map_err(Error::Io)?, + local: tls_stream.get_ref().0.local_addr().map_err(Error::Io)?, is_tls: true, subprotocol: None, }; @@ -210,9 +189,7 @@ impl Server { /// Return the local address of the listener. pub fn local_addr(&self) -> Result { - self.listener - .local_addr() - .map_err(|e| Error::Io(e.to_string())) + self.listener.local_addr().map_err(Error::Io) } } diff --git a/websock-tungstenite/src/tls/cert.rs b/websock-tungstenite/src/tls/cert.rs index 2337efe..101e1d1 100644 --- a/websock-tungstenite/src/tls/cert.rs +++ b/websock-tungstenite/src/tls/cert.rs @@ -20,7 +20,7 @@ pub fn get_native_certs() -> Result { /// Load a certificate chain from a file (.der or .pem). pub fn load_certs(cert_path: &Path) -> Result>> { - let cert_bytes = fs::read(cert_path).map_err(|e| Error::Io(e.to_string()))?; + let cert_bytes = fs::read(cert_path).map_err(Error::Io)?; if cert_path.extension().is_some_and(|x| x == "der") { return Ok(vec![CertificateDer::from(cert_bytes)]); @@ -28,7 +28,7 @@ pub fn load_certs(cert_path: &Path) -> Result>> { CertificateDer::pem_slice_iter(&cert_bytes) .collect::, _>>() - .map_err(|e| Error::Io(e.to_string())) + .map_err(Error::tls) } /// Certificate verifier that unconditionally accepts certificates. diff --git a/websock-tungstenite/src/tls/key.rs b/websock-tungstenite/src/tls/key.rs index 4e264c4..a89bd76 100644 --- a/websock-tungstenite/src/tls/key.rs +++ b/websock-tungstenite/src/tls/key.rs @@ -6,14 +6,14 @@ use websock_proto::{Error, Result}; /// Load a private key from a file. pub fn load_key(key_path: &Path) -> Result> { - let key = fs::read(key_path).map_err(|e| Error::Io(e.to_string()))?; + let key = fs::read(key_path).map_err(Error::Io)?; let key = if key_path.extension().is_some_and(|x| x == "der") { // Treat raw DER as PKCS#8. PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(key)) } else { // Decode PEM. - PrivateKeyDer::from_pem_slice(&key).map_err(|e| Error::Tls(e.to_string()))? + PrivateKeyDer::from_pem_slice(&key).map_err(Error::tls)? }; Ok(key) diff --git a/websock-tungstenite/src/tls/mod.rs b/websock-tungstenite/src/tls/mod.rs index 0e40273..d6af4c1 100644 --- a/websock-tungstenite/src/tls/mod.rs +++ b/websock-tungstenite/src/tls/mod.rs @@ -22,8 +22,7 @@ pub type TlsClientConfig = rustls::client::ClientConfig; pub fn generate_self_signed_pair_der( subject_alt_names: Vec, ) -> Result<(Vec>, PrivateKeyDer<'static>)> { - let cert = rcgen::generate_simple_self_signed(subject_alt_names) - .map_err(|e| Error::Tls(e.to_string()))?; + let cert = rcgen::generate_simple_self_signed(subject_alt_names).map_err(Error::tls)?; let key = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der())); let cert_chain = vec![CertificateDer::from(cert.cert)]; @@ -34,8 +33,7 @@ pub fn generate_self_signed_pair_der( pub fn generate_self_signed_pair_pem( subject_alt_names: Vec, ) -> Result<(Vec, String)> { - let cert = rcgen::generate_simple_self_signed(subject_alt_names) - .map_err(|e| Error::Tls(e.to_string()))?; + let cert = rcgen::generate_simple_self_signed(subject_alt_names).map_err(Error::tls)?; let key = cert.signing_key.serialize_pem(); let cert_chain = vec![cert.cert.pem()]; @@ -143,7 +141,7 @@ impl TlsServerConfigBuilder { let inner = ServerConfig::builder() .with_no_client_auth() .with_single_cert(certs, key) - .map_err(|e| Error::Tls(e.to_string()))?; + .map_err(Error::tls)?; Ok(Self { inner }) } @@ -153,7 +151,7 @@ impl TlsServerConfigBuilder { let inner = ServerConfig::builder() .with_no_client_auth() .with_single_cert(certs, key) - .map_err(|e| Error::Tls(e.to_string()))?; + .map_err(Error::tls)?; Ok(Self { inner }) } diff --git a/websock-tungstenite/tests/connection.rs b/websock-tungstenite/tests/connection.rs index 6a32d0d..5efc698 100644 --- a/websock-tungstenite/tests/connection.rs +++ b/websock-tungstenite/tests/connection.rs @@ -1,4 +1,5 @@ use futures_util::{SinkExt, StreamExt}; +use std::error::Error as _; use websock_proto::{Error, Message, WebSocketLimits}; use websock_tungstenite::{ClientBuilder, ServerBuilder}; @@ -62,7 +63,11 @@ async fn server_rejects_message_above_configured_limit() { .recv() .await .expect_err("oversized message must fail"); - assert!(matches!(err, Error::Protocol(_))); + assert!(matches!(err, Error::Transport(_))); + assert!( + err.source() + .is_some_and(|source| source.is::()) + ); }); let client = ClientBuilder::new().build(); From b8929b5606f3fa2c655e8ed3d98eafa6a0d7d61e Mon Sep 17 00:00:00 2001 From: shellrow Date: Sun, 19 Jul 2026 17:36:57 +0900 Subject: [PATCH 14/17] docs: update docs --- CHANGELOG.md | 4 ++++ README.md | 8 ++++++++ RELEASING.md | 13 +++++++++++++ TODO.md | 2 +- 4 files changed, 26 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index fa2c2fb..9f0bf54 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -19,6 +19,10 @@ and this project uses [Semantic Versioning](https://semver.org/spec/v2.0.0.html) ### Changed +- The shared error API now retains native `std::io::Error` values and structured + TLS and WebSocket transport sources instead of storing their display strings. + This changes the public payload of `Error::Io` and `Error::Tls` and adds + `Error::Transport`. - The mux TLS helpers now reuse the native transport implementation. - PEM parsing now uses `rustls-pki-types` directly instead of the unmaintained `rustls-pemfile` crate. diff --git a/README.md b/README.md index 968bc96..cd23a65 100644 --- a/README.md +++ b/README.md @@ -64,6 +64,14 @@ inconsistent limits return a protocol error rather than panicking. The multiplexing wire format, stream lifecycle, flow control, compatibility policy, and error codes are specified in [docs/mux-protocol.md](docs/mux-protocol.md). +## Error handling + +Native I/O failures retain their original `std::io::Error`, including the +`ErrorKind`, in `websock_proto::Error::Io`. TLS and underlying WebSocket +transport failures retain their concrete errors as standard error sources. +Call `std::error::Error::source` to inspect or downcast those sources, and use +`Error::io_kind` when only the I/O classification is needed. + ## Benchmarking Criterion benchmarks are available for `websock-mux-proto`. diff --git a/RELEASING.md b/RELEASING.md index 7fe1132..33787b4 100644 --- a/RELEASING.md +++ b/RELEASING.md @@ -24,3 +24,16 @@ The publishable dependency order is: 5. `websock` and `websock-mux` `websock-wasm-demo` is not published. + +For the manual package-content inspection in step 7, run: + +```text +cargo package --list -p websock-proto +cargo package --list -p websock-mux-proto +cargo package --list -p websock-tungstenite +cargo package --list -p websock-wasm +cargo package --list -p websock-tungstenite-mux +cargo package --list -p websock-wasm-mux +cargo package --list -p websock +cargo package --list -p websock-mux +``` diff --git a/TODO.md b/TODO.md index d88e0db..f12423c 100644 --- a/TODO.md +++ b/TODO.md @@ -61,7 +61,7 @@ items remain release blockers or follow-up work. ## P2 — API, observability, and release engineering -- [ ] Preserve structured error sources and I/O error kinds instead of reducing +- [x] Preserve structured error sources and I/O error kinds instead of reducing all transport errors to strings. This is a semver-sensitive public API change. - [x] Expose negotiated subprotocol and close-frame details consistently on native and browser connections. From 9109a8dfec4e5e00a91ebcac900f072c3f15fad7 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 25 Jul 2026 21:13:29 +0900 Subject: [PATCH 15/17] ci: streamline Rust checks --- .github/workflows/ci.yml | 82 ++++------------------------------------ 1 file changed, 8 insertions(+), 74 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a3206e8..28597fd 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -2,11 +2,17 @@ name: CI on: push: + branches: + - main pull_request: permissions: contents: read +concurrency: + group: ci-${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + env: CARGO_TERM_COLOR: always RUSTFLAGS: "-Dwarnings" @@ -21,11 +27,8 @@ jobs: components: clippy, rustfmt - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 - run: cargo fmt --all -- --check - - run: cargo clippy --workspace --all-targets --all-features -- -D warnings - - run: cargo test --workspace --all-targets - - run: cargo doc --workspace --no-deps - env: - RUSTDOCFLAGS: "-Dwarnings" + - run: cargo clippy --workspace --all-targets + - run: cargo test --workspace wasm: runs-on: ubuntu-latest @@ -44,72 +47,3 @@ jobs: -p websock-wasm-mux -p websock-wasm-demo --target wasm32-unknown-unknown - - name: Install wasm-bindgen test runner - run: cargo install wasm-bindgen-cli --version 0.2.126 --locked - - name: Build browser test server - run: cargo build -p websock-browser-test-server - - name: Run browser integration tests - shell: bash - run: | - target/debug/websock-browser-test-server & - server_pid=$! - trap 'kill "$server_pid"' EXIT - ready=false - for attempt in {1..30}; do - if (echo > /dev/tcp/127.0.0.1/32123) 2>/dev/null && - (echo > /dev/tcp/127.0.0.1/32124) 2>/dev/null; then - ready=true - break - fi - sleep 1 - done - if [ "$ready" != true ]; then - echo "browser test server failed to start" >&2 - exit 1 - fi - CARGO_TARGET_WASM32_UNKNOWN_UNKNOWN_RUNNER=wasm-bindgen-test-runner \ - cargo test -p websock-wasm -p websock-wasm-mux \ - --tests --target wasm32-unknown-unknown - - msrv: - runs-on: ubuntu-latest - steps: - - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 - - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 - with: - toolchain: "1.88.0" - - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 - - run: cargo check --workspace --all-targets - - minimal-features: - runs-on: ubuntu-latest - steps: - - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 - - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 - - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 - - run: cargo check --workspace --all-targets --no-default-features - - dependency-policy: - runs-on: ubuntu-latest - strategy: - matrix: - checks: - - advisories - - bans licenses sources - steps: - - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 - - uses: EmbarkStudios/cargo-deny-action@v2 - with: - command: check ${{ matrix.checks }} - - platforms: - strategy: - fail-fast: false - matrix: - os: [macos-latest, windows-latest] - runs-on: ${{ matrix.os }} - steps: - - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 - - uses: dtolnay/rust-toolchain@4cda84d5c5c54efe2404f9d843567869ab1699d4 - - uses: Swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae # v2 - - run: cargo test --workspace --all-targets From 1a9ef52980cb656eeaaf01ebde5fb9c0fac5a6c2 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 25 Jul 2026 21:31:13 +0900 Subject: [PATCH 16/17] chore: remove unused docs --- CHANGELOG.md | 33 --------------------- RELEASING.md | 39 ------------------------- TODO.md | 82 ---------------------------------------------------- 3 files changed, 154 deletions(-) delete mode 100644 CHANGELOG.md delete mode 100644 RELEASING.md delete mode 100644 TODO.md diff --git a/CHANGELOG.md b/CHANGELOG.md deleted file mode 100644 index 9f0bf54..0000000 --- a/CHANGELOG.md +++ /dev/null @@ -1,33 +0,0 @@ -# Changelog - -All notable changes to this project will be documented in this file. - -The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), -and this project uses [Semantic Versioning](https://semver.org/spec/v2.0.0.html). - -## Unreleased - -### Added - -- Explicit mux session shutdown with task-completion semantics. -- Negotiated subprotocol and received close-frame metadata accessors. -- Native end-to-end coverage for split connections, bidirectional mux streams, - malformed frames, reset/stop propagation, TLS, and session-handle cleanup. -- Browser integration coverage for connection failures, receive-queue overflow, - close metadata, split ownership, and mux parity. -- A versioned mux wire-protocol specification. - -### Changed - -- The shared error API now retains native `std::io::Error` values and structured - TLS and WebSocket transport sources instead of storing their display strings. - This changes the public payload of `Error::Io` and `Error::Tls` and adds - `Error::Transport`. -- The mux TLS helpers now reuse the native transport implementation. -- PEM parsing now uses `rustls-pki-types` directly instead of the unmaintained - `rustls-pemfile` crate. -- Default TLS ALPN configuration advertises only HTTP/1.1. - -### Fixed - -- Non-boolean mux stream FIN fields are rejected as malformed. diff --git a/RELEASING.md b/RELEASING.md deleted file mode 100644 index 33787b4..0000000 --- a/RELEASING.md +++ /dev/null @@ -1,39 +0,0 @@ -# Release Guide - -Release and publishing operations are intentionally manual. - -1. Confirm that `CHANGELOG.md` describes all user-visible changes and move the - Unreleased entries under the new version and date. -2. Run `cargo fmt --all -- --check`. -3. Run `cargo clippy --workspace --all-targets --all-features -- -D warnings`. -4. Run `cargo test --workspace --all-targets`. -5. Run `cargo doc --workspace --no-deps` with `RUSTDOCFLAGS=-Dwarnings`. -6. Check the WASM packages for `wasm32-unknown-unknown`. -7. Run `cargo package --list -p ` for every publishable crate and inspect - each file list. Then run `cargo package -p ` in dependency order. -8. Update versions and internal dependency requirements together. -9. Create the release commit and tag manually. -10. Publish crates in dependency order, then create the GitHub Release manually. - -The publishable dependency order is: - -1. `websock-proto` -2. `websock-mux-proto` -3. `websock-tungstenite` and `websock-wasm` -4. `websock-tungstenite-mux` and `websock-wasm-mux` -5. `websock` and `websock-mux` - -`websock-wasm-demo` is not published. - -For the manual package-content inspection in step 7, run: - -```text -cargo package --list -p websock-proto -cargo package --list -p websock-mux-proto -cargo package --list -p websock-tungstenite -cargo package --list -p websock-wasm -cargo package --list -p websock-tungstenite-mux -cargo package --list -p websock-wasm-mux -cargo package --list -p websock -cargo package --list -p websock-mux -``` diff --git a/TODO.md b/TODO.md deleted file mode 100644 index f12423c..0000000 --- a/TODO.md +++ /dev/null @@ -1,82 +0,0 @@ -# Production Readiness Review - -Reviewed on 2026-07-18. This file records concrete production-readiness work for -the workspace. Completed items describe fixes made as part of the review; open -items remain release blockers or follow-up work. - -## P0 — Safety and resource bounds - -- [x] Enforce mux receive-side flow control. A peer could send more than - the advertised `MaxStreamData` value, allowing each stream to exceed its - configured memory budget. - - Completion: track cumulative received bytes per stream, reject overflow and - credit violations as protocol errors, and test the boundary and violation. -- [x] Bound the browser WebSocket receive queue. `websock-wasm` used an - unbounded channel, so a fast peer can exhaust browser memory when the - application consumes messages slowly. - - Completion: use a bounded queue and close the socket on overflow. -- [x] Validate every mux `Limits` value before allocating channels or starting - tasks. Zero-capacity channels currently panic, and inconsistent byte/window - settings can deadlock or bypass configured bounds. - - Completion: invalid limits return `Error::Protocol` on both native and WASM - implementations, with regression tests. -- [x] Do not grant mux send credit before the peer advertises it. Using the local - receive-window configuration as peer credit breaks interoperability when - peers use different limits and weakens flow control. - - Completion: new send streams start with zero credit and wake only after a - valid `MaxStreamData` frame. - -## P1 — Protocol correctness and lifecycle - -- [x] Treat `ConnectionClose` as terminal and wake all blocked stream writers. - The previous implementation cleared maps but left the session open. -- [x] Reject non-monotonic peer stream IDs and cap locally opened streams. A - peer could reuse stale IDs, while local streams could grow the flow - map without respecting `max_open_streams`. -- [x] Remove stream state reliably. The native `try_lock` cleanup path could - silently retain entries when the mutex is contended. -- [x] Reset streams whose application receive queues are full instead of - blocking the entire native session behind one slow stream. -- [x] Preserve bidirectional send state after a receive FIN, separate reset and - stop-sending semantics, retain final-frame data across partial reads, and keep - a stream alive when one of several send-stream clones is dropped. -- [x] Reject stream-counter and reset/stop-code values that exceed the 62-bit - mux encoding range instead of wrapping IDs or panicking in background tasks. -- [x] Add native end-to-end tests for connect/subprotocol negotiation, - text/binary round trips, graceful close, oversized-message rejection, mux uni - streams, differing flow-control windows, and invalid-limit rejection. -- [x] Extend native end-to-end tests with split operation, mux bidirectional - streams, malformed wire frames, reset/stop propagation, and TLS round trips. -- [x] Add browser tests under `wasm-bindgen-test` for connect failure, queue - overflow, close, split ownership, and mux parity. -- [x] Ensure WebSocket-level frame/message limits are applied in Tungstenite - configuration before a complete oversized message is buffered. -- [x] Expose shared WebSocket message/frame/write-buffer limits and enforce - browser receive queue, message-size, and `bufferedAmount` bounds. -- [x] Advertise only HTTP/1.1 in the default ALPN list because the current - Tungstenite transport does not implement WebSocket over HTTP/2. -- [x] Add explicit connection/session shutdown APIs with completion semantics - and task cancellation. Dropping the last mux session handle should not leave - detached tasks and a socket alive indefinitely. - -## P2 — API, observability, and release engineering - -- [x] Preserve structured error sources and I/O error kinds instead of reducing - all transport errors to strings. This is a semver-sensitive public API change. -- [x] Expose negotiated subprotocol and close-frame details consistently on - native and browser connections. -- [x] Define and document the mux wire protocol, compatibility policy, error - codes, stream state machine, and flow-control rules. -- [x] Expand crate-level and public-item documentation and enable - `#![warn(missing_docs)]` incrementally. -- [x] Add baseline CI for formatting, Clippy with warnings denied, tests, - documentation, Linux/macOS/Windows, and `wasm32-unknown-unknown`. -- [x] Extend CI with an explicit MSRV job, minimal feature-set builds, dependency - policy checks, and security audits. -- [x] Declare and test the MSRV, and add changelog and release guidance. -- [ ] Manually verify package contents with `cargo package --list` for every - published crate as part of the release process. -- [x] Review dependency features and duplicate native/mux TLS implementations to - reduce compile time, binary size, and maintenance drift. -- [x] Add fuzz targets for varint/frame decoding and stateful mux frame - sequences, plus long-running concurrency and backpressure tests. From ff7e02218a831d0ef911eb08340427f1e6ca0484 Mon Sep 17 00:00:00 2001 From: shellrow Date: Sat, 25 Jul 2026 21:34:48 +0900 Subject: [PATCH 17/17] chore: bump version to 0.5.0 --- Cargo.toml | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 17dbef4..4505937 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -15,18 +15,18 @@ members = [ exclude = ["fuzz"] [workspace.package] -version = "0.4.0" +version = "0.5.0" edition = "2024" rust-version = "1.88" authors = ["shellrow "] [workspace.dependencies] -websock-proto = { path = "websock-proto", version = "0.4.0" } -websock-tungstenite = { path = "websock-tungstenite", version = "0.4.0" } -websock-wasm = { path = "websock-wasm", version = "0.4.0" } -websock-mux-proto = { path = "websock-mux-proto", version = "0.4.0" } -websock-tungstenite-mux = { path = "websock-tungstenite-mux", version = "0.4.0" } -websock-wasm-mux = { path = "websock-wasm-mux", version = "0.4.0" } +websock-proto = { path = "websock-proto", version = "0.5.0" } +websock-tungstenite = { path = "websock-tungstenite", version = "0.5.0" } +websock-wasm = { path = "websock-wasm", version = "0.5.0" } +websock-mux-proto = { path = "websock-mux-proto", version = "0.5.0" } +websock-tungstenite-mux = { path = "websock-tungstenite-mux", version = "0.5.0" } +websock-wasm-mux = { path = "websock-wasm-mux", version = "0.5.0" } bytes = "1" wasm-bindgen = "=0.2.126" wasm-bindgen-futures = "=0.4.76"