diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..28597fd --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,49 @@ +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" + +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 + - run: cargo test --workspace + + 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 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..4505937 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -6,22 +6,28 @@ 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" +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" +wasm-bindgen-test = "=0.3.76" diff --git a/README.md b/README.md index bbdfd3a..cd23a65 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. @@ -38,6 +40,38 @@ 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. + +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/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/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/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; + } + } +}); 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-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-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 33ebb20..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, @@ -22,13 +27,14 @@ 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)?; 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,38 +43,65 @@ impl StreamId { } } + /// Return whether the server initiated the stream. 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 + } } +/// 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, }, } @@ -96,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 { @@ -139,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 { @@ -150,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); @@ -187,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 d03e9a2..a03ef37 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(); @@ -122,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-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-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-proto/src/error.rs b/websock-proto/src/error.rs index d54b9e6..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,19 +26,25 @@ 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}")] 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), @@ -42,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 fd9770e..233b581 100644 --- a/websock-proto/src/lib.rs +++ b/websock-proto/src/lib.rs @@ -4,13 +4,15 @@ //! 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; 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, 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..4bc1a78 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,22 @@ pub struct ServerOptions { /// Additional response headers (native only). pub headers: Vec<(String, String)>, + + /// Resource limits for accepted WebSocket connections. + 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(), 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-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/client.rs b/websock-tungstenite-mux/src/client.rs index 7688cd9..f537b74 100644 --- a/websock-tungstenite-mux/src/client.rs +++ b/websock-tungstenite-mux/src/client.rs @@ -64,11 +64,10 @@ impl Client { url: &str, tls: Option>, ) -> Result { + 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() { @@ -88,10 +87,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..bdca0b0 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,9 +69,8 @@ pub async fn bind( where A: ToSocketAddrs, { - let listener = TcpListener::bind(addr) - .await - .map_err(|e| Error::Io(e.to_string()))?; + limits.validate()?; + let listener = TcpListener::bind(addr).await.map_err(Error::Io)?; let headers = prepare_headers(&opts)?; validate_protocols(&opts)?; @@ -102,18 +101,12 @@ 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) @@ -122,7 +115,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 +146,7 @@ impl Server { Ok(resp) }, + Some(self.limits.websocket_config()), ) .await .map_err(map_tungstenite_err)?; @@ -161,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 5be3723..35e7d0e 100644 --- a/websock-tungstenite-mux/src/session.rs +++ b/websock-tungstenite-mux/src/session.rs @@ -2,8 +2,8 @@ use std::collections::HashMap; use std::io; use std::pin::Pin; use std::sync::{ - Arc, - atomic::{AtomicBool, AtomicU64, Ordering}, + Arc, Mutex as StdMutex, MutexGuard as StdMutexGuard, + atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, }; use std::task::{Context, Poll}; @@ -12,11 +12,12 @@ 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; 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,7 +68,69 @@ impl Default for Limits { } } -#[derive(Clone)] +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)) + } +} + pub struct Session { inner: Arc, accept_uni: Arc>>, @@ -83,6 +146,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 +171,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)) } @@ -136,11 +205,41 @@ 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 { max_data: AtomicU64, sent_data: AtomicU64, + closed: AtomicBool, waker: AtomicWaker, } @@ -149,6 +248,7 @@ impl SendFlowState { Self { max_data: AtomicU64::new(initial_max), sent_data: AtomicU64::new(0), + closed: AtomicBool::new(false), waker: AtomicWaker::new(), } } @@ -185,6 +285,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 +302,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 +360,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 +397,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 +421,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 +431,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 +460,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 +475,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 +550,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 +582,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 +601,12 @@ struct RecvEvent { fin: bool, } +struct RecvState { + sender: mpsc::Sender, + received: u64, + max_data: Arc, +} + pub struct RecvStream { id: StreamId, session: Arc, @@ -475,6 +617,8 @@ pub struct RecvStream { granted: u64, initial_window: u64, update_threshold: u64, + max_data: Arc, + stop_sent: AtomicBool, } impl RecvStream { @@ -484,6 +628,7 @@ impl RecvStream { receiver: mpsc::Receiver, initial_window: u64, update_threshold: u64, + max_data: Arc, ) -> Self { Self { id, @@ -495,6 +640,8 @@ impl RecvStream { granted: initial_window, initial_window, update_threshold, + max_data, + stop_sent: AtomicBool::new(false), } } @@ -503,7 +650,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 +669,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 +692,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 +733,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 +772,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 +790,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, @@ -686,6 +858,7 @@ pub(crate) enum OutboundCmd { Frame(Frame), Ws(tungstenite::Message), Flush { ack: oneshot::Sender> }, + Shutdown { ack: oneshot::Sender> }, } pub(crate) struct SessionInner { @@ -694,11 +867,18 @@ 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, + shutdown_started: AtomicBool, + session_handles: AtomicUsize, + active_tasks: AtomicUsize, + tasks_done: Notify, + cancel: CancellationToken, } impl SessionInner { @@ -715,11 +895,18 @@ 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), + shutdown_started: AtomicBool::new(false), + session_handles: AtomicUsize::new(1), + active_tasks: AtomicUsize::new(2), + tasks_done: Notify::new(), + cancel: CancellationToken::new(), } } @@ -734,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, @@ -776,7 +970,7 @@ impl SessionInner { } } - inbound.close_all().await; + inbound.task_finished().await; }); let outbound = self.clone(); @@ -801,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(); @@ -848,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 } => { @@ -865,6 +1132,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 +1150,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 +1168,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 +1189,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 +1205,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 +1272,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 +1298,7 @@ impl SessionInner { rx, initial_window, self.limits.stream_window_update_threshold as u64, + max_data, ); self.send_frame(Frame::MaxStreamData { id, @@ -1007,8 +1311,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 +1329,7 @@ impl SessionInner { rx, initial_window, self.limits.stream_window_update_threshold as u64, + max_data, ) } @@ -1036,19 +1349,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 +1417,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 +1449,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 { @@ -1194,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 { @@ -1207,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) @@ -1223,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), } } @@ -1258,7 +1632,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 +1732,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 +1757,157 @@ 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"); + } + + #[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-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-tungstenite-mux/tests/session.rs b/websock-tungstenite-mux/tests/session.rs new file mode 100644 index 0000000..ac42a4f --- /dev/null +++ b/websock-tungstenite-mux/tests/session.rs @@ -0,0 +1,335 @@ +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 { + 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(_))); +} + +#[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-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/builder.rs b/websock-tungstenite/src/builder.rs index ec71c4c..567fc8b 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()); @@ -104,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) @@ -238,6 +242,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()); @@ -253,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 757fce8..d0f40ef 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. @@ -29,9 +31,15 @@ pub async fn connect_with_tls( opts: ConnectOptions, tls: Option>, ) -> Result { - let mut req = url - .into_client_request() - .map_err(|e| Error::InvalidUrl(e.to_string()))?; + 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(Error::transport)?; // Apply configured headers and subprotocols. { @@ -54,31 +62,34 @@ 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 - .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() + .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 @@ -117,7 +128,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); } @@ -175,7 +190,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() } } @@ -184,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/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 0d04dbe..c2c1866 100644 --- a/websock-tungstenite/src/server.rs +++ b/websock-tungstenite/src/server.rs @@ -1,11 +1,11 @@ //! 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; -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,19 +22,26 @@ pub async fn bind( where A: ToSocketAddrs, { - let listener = TcpListener::bind(addr) - .await - .map_err(|e| Error::Io(e.to_string()))?; + opts.limits.validate()?; + let listener = TcpListener::bind(addr).await.map_err(Error::Io)?; let headers = prepare_headers(&opts)?; 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,40 +59,38 @@ pub struct Server { protocols: Arc>, headers: Arc>, acceptor: Option, + config: tungstenite::protocol::WebSocketConfig, } 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) }; - 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 = 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() { @@ -96,53 +101,58 @@ 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) }, + Some(self.config), ) .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>> { - 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 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()))?, + let mut info = ConnectionInfo { + 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, }; + let selected_protocol = Arc::new(Mutex::new(None)); + let selected_protocol_callback = Arc::clone(&selected_protocol); - 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() { @@ -153,22 +163,33 @@ 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) }, + Some(self.config), ) .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. 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) } } @@ -201,10 +222,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/src/tls/cert.rs b/websock-tungstenite/src/tls/cert.rs index a1cbf0f..101e1d1 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}; @@ -20,15 +20,15 @@ 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().map_or(false, |x| x == "der") { + 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())) + CertificateDer::pem_slice_iter(&cert_bytes) + .collect::, _>>() + .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 10c8987..a89bd76 100644 --- a/websock-tungstenite/src/tls/key.rs +++ b/websock-tungstenite/src/tls/key.rs @@ -1,21 +1,19 @@ //! 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}; /// 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().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 { // 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(Error::tls)? }; Ok(key) diff --git a/websock-tungstenite/src/tls/mod.rs b/websock-tungstenite/src/tls/mod.rs index 6079528..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()]; @@ -65,11 +63,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 +79,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 +96,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 +113,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 { @@ -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 new file mode 100644 index 0000000..5efc698 --- /dev/null +++ b/websock-tungstenite/tests/connection.rs @@ -0,0 +1,211 @@ +use futures_util::{SinkExt, StreamExt}; +use std::error::Error as _; +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.negotiated_subprotocol(), Some("test.v1")); + 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"); + assert_eq!(connection.negotiated_subprotocol(), Some("test.v1")); + 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::Transport(_))); + assert!( + err.source() + .is_some_and(|source| source.is::()) + ); + }); + + 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(_))); +} + +#[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-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..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" @@ -15,7 +16,10 @@ 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"] } -wasm-bindgen-futures = "0.4" +futures-util = { version = "0.3", default-features = false, features = ["alloc", "sink"] } +wasm-bindgen-futures = { workspace = true } + +[dev-dependencies] +wasm-bindgen-test = { workspace = true } 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..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,16 +7,16 @@ 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; 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,7 +65,59 @@ impl Default for Limits { } } -#[derive(Clone)] +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 struct Session { inner: Rc, accept_uni: Rc>>, @@ -73,7 +125,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 +147,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)) } @@ -129,11 +189,46 @@ 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 { max_data: AtomicU64, sent_data: AtomicU64, + closed: AtomicBool, waker: AtomicWaker, } @@ -142,6 +237,7 @@ impl SendFlowState { Self { max_data: AtomicU64::new(initial_max), sent_data: AtomicU64::new(0), + closed: AtomicBool::new(false), waker: AtomicWaker::new(), } } @@ -177,6 +273,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 +290,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 +325,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 +344,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 +406,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 +417,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 +446,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 +464,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 +519,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 +555,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 +574,12 @@ struct RecvEvent { fin: bool, } +struct RecvState { + sender: mpsc::Sender, + received: u64, + max_data: Rc, +} + pub struct RecvStream { id: StreamId, session: Rc, @@ -430,6 +590,8 @@ pub struct RecvStream { granted: u64, initial_window: u64, update_threshold: u64, + max_data: Rc, + stop_sent: AtomicBool, } impl RecvStream { @@ -439,6 +601,7 @@ impl RecvStream { receiver: mpsc::Receiver, initial_window: u64, update_threshold: u64, + max_data: Rc, ) -> Self { Self { id, @@ -450,6 +613,8 @@ impl RecvStream { granted: initial_window, initial_window, update_threshold, + max_data, + stop_sent: AtomicBool::new(false), } } @@ -458,7 +623,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 +642,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 +665,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 +706,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 +745,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 +762,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,11 +836,17 @@ 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, + shutdown_started: AtomicBool, + session_handles: Cell, + close_waiters: RefCell>>, + task_finished: AtomicBool, } impl SessionInner { @@ -669,10 +865,31 @@ 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), + 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( self: Rc, mut conn: websock_wasm::Connection, @@ -756,8 +973,9 @@ impl SessionInner { } } - inner.close_all().await; let _ = conn.close().await; + inner.close_all().await; + inner.finish_task(); }); } @@ -784,19 +1002,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 +1055,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 +1111,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 +1177,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 +1252,58 @@ 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 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, + 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 { @@ -1067,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 { 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/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..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; @@ -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,88 @@ 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 onclose = Closure::::new(move |_e: web_sys::CloseEvent| { - let _ = tx_close.unbounded_send(Err(Error::Closed)); + 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| { + *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), @@ -124,7 +170,10 @@ 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) negotiated_subprotocol: String, + pub(crate) close_frame: Rc>>, pub(crate) _onmessage: Option>, pub(crate) _onerror: Option>, @@ -135,8 +184,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,14 +199,42 @@ 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. + /// 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() } } @@ -195,3 +278,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), ) } 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()) + ); +} 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 a009bff..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}; @@ -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"); @@ -83,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()) } 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};