diff --git a/.markdownlint-cli2.jsonc b/.markdownlint-cli2.jsonc index a822f004..3c6d268f 100644 --- a/.markdownlint-cli2.jsonc +++ b/.markdownlint-cli2.jsonc @@ -1,12 +1,14 @@ { - // Ignore MD013 (line length) inside tables since reflowing - // Markdown tables often breaks formatting and readability. "config": { + "MD004": { "style": "dash" }, + "MD010": { "code_blocks": false }, "MD013": { "line_length": 80, "code_block_line_length": 120, - "tables": false + "tables": false, + "headings": false }, - "MD040": false - } + "MD029": { "style": "ordered" } + }, + "ignores": ["**/.venv/**", ".node_modules/**", "**/node_modules/**", "**/target/**"] } diff --git a/AGENTS.md b/AGENTS.md index 5c69367e..df1c6860 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -2,7 +2,7 @@ ## Code Style and Structure -- **Code is for humans.** Write your code with clarity and empathy—assume a +- **Code is for humans.** Write code with clarity and empathy—assume a tired teammate will need to debug it at 3 a.m. - **Comment *why*, not *what*.** Explain assumptions, edge cases, trade-offs, or complexity. Don't echo the obvious. @@ -43,8 +43,10 @@ relevant file(s) in the `docs/` directory to reflect the latest state. **Ensure the documentation remains accurate and current.** - Documentation must use en-GB-oxendict ("-ize" / "-yse" / "-our") spelling - and grammar. (EXCEPTION: the naming of the "LICENSE" file, which is to be - left unchanged for community consistency.) + and grammar. (EXCEPTION: the filename `LICENSE` is left unchanged for + community consistency.) +- A documentation style guide is provided at + `docs/documentation-style-guide.md`. ## Change Quality & Committing @@ -68,7 +70,7 @@ - **Imperative Mood:** Use the imperative mood in the subject line (e.g., "Fix bug", "Add feature" instead of "Fixed bug", "Added feature"). - **Subject Line:** The first line should be a concise summary of the change - (ideally 50 characters or less). + (ideally 50 characters or fewer). - **Body:** Separate the subject from the body with a blank line. Subsequent lines should explain the *what* and *why* of the change in more detail, including rationale, goals, and scope. Wrap the body at 72 characters. @@ -79,7 +81,7 @@ ## Refactoring Heuristics & Workflow - **Recognizing Refactoring Needs:** Regularly assess the codebase for potential - refactoring opportunities. Consider refactoring when you observe: + refactoring opportunities. Perform refactoring when you observe: - **Long Methods/Functions:** Functions or methods that are excessively long or try to do too many things. - **Duplicated Code:** Identical or very similar code blocks appearing in @@ -103,7 +105,7 @@ - **Separate Atomic Refactors:** If refactoring is deemed necessary: - Perform the refactoring as a **separate, atomic commit** *after* the functional change commit. - - Ensure the refactoring adheres to the testing guidelines (behavioral tests + - Ensure refactoring adheres to the testing guidelines (behavioural tests pass before and after, unit tests added for new units). - Ensure the refactoring commit itself passes all quality gates. @@ -114,18 +116,18 @@ management. Contributors should follow these best practices when working on the project: - Run `make check-fmt`, `make lint`, and `make test` before committing. These - targets wrap the following commands so contributors understand the exact + targets wrap the following commands, so contributors understand the exact behaviour and policy enforced: - `make check-fmt` executes: - ``` + ```sh cargo fmt --workspace -- --check ``` validating formatting across the entire workspace without modifying files. - `make lint` executes: - ``` + ```sh cargo clippy --workspace --all-targets --all-features -- -D warnings ``` @@ -133,7 +135,7 @@ project: warnings. - `make test` executes: - ``` + ```sh cargo test --workspace ``` @@ -147,8 +149,8 @@ project: adhering to separation of concerns and CQRS. - Where a function has too many parameters, group related parameters in meaningfully named structs. -- Where a function is returning a large error consider using `Arc` to reduce the - amount of data returned. +- Where a function is returning a large error, consider using `Arc` to reduce + the amount of data returned. - Write unit and behavioural tests for new functionality. Run both before and after making any change. - Every module **must** begin with a module level (`//!`) comment explaining the @@ -156,31 +158,25 @@ project: - Document public APIs using Rustdoc comments (`///`) so documentation can be generated with cargo doc. - Prefer immutable data and avoid unnecessary `mut` bindings. -- Handle errors with the `Result` type instead of panicking where feasible. - Use explicit version ranges in `Cargo.toml` and keep dependencies up-to-date. -- Avoid `unsafe` code unless absolutely necessary and document any usage - clearly. +- Avoid `unsafe` code unless absolutely necessary, and document any usage + clearly with a "SAFETY" comment. - Place function attributes **after** doc comments. - Do not use `return` in single-line functions. - Use predicate functions for conditional criteria with more than two branches. - Lints must not be silenced except as a **last resort**. - Lint rule suppressions must be tightly scoped and include a clear reason. -- Prefer `expect` over `allow`. -- Use `rstest` fixtures for shared setup. -- Replace duplicated tests with `#[rstest(...)]` parameterised cases. -- Prefer `mockall` for mocks/stubs. -- Prefer `.expect()` over `.unwrap()`. - Use `concat!()` to combine long string literals rather than escaping newlines with a backslash. -- Prefer single line versions of functions where appropriate. I.e., +- Prefer single line versions of functions where appropriate. i.e., - ``` + ```rust pub fn new(id: u64) -> Self { Self(id) } ``` Instead of: - ``` + ```rust pub fn new(id: u64) -> Self { Self(id) } @@ -190,16 +186,30 @@ project: `newt-hype` when introducing many homogeneous wrappers that share behaviour; add small shims such as `From<&str>` and `AsRef` for string-backed wrappers. For path-centric wrappers implement `AsRef` alongside - `into_inner()` and `to_path_buf()`; avoid attempting + `into_inner()` and `to_path_buf()`, avoid attempting `impl From for PathBuf` because of the orphan rule. Prefer explicit tuple structs whenever bespoke validation or tailored trait surfaces are - required, customising `Deref`, `AsRef`, and `TryFrom` per type. Use + required, customizing `Deref`, `AsRef`, and `TryFrom` per type. Use `the-newtype` when defining traits and needing blanket implementations that apply across wrappers satisfying `Newtype + AsRef/AsMut`, or when establishing a coherent internal convention that keeps trait forwarding consistent without per-type boilerplate. Combine approaches: lean on `newt-hype` for the common case, tuple structs for outliers, and - `the-newtype` to unify behaviour when you own the trait definitions. + `the-newtype` to unify behaviour when owning the trait definitions. +- Use `cap_std` and `cap_std::fs_utf8` / `camino` in place of `std::fs` and + `std::path` for enhanced cross-platform support and capability-oriented + filesystem access. + +### Testing + +- Use `rstest` fixtures for shared setup. +- Replace duplicated tests with `#[rstest(...)]` parameterized cases. +- Prefer `mockall` for ad hoc mocks/stubs. +- For testing of functionality depending upon environment variables, dependency + injection and the `mockable` crate are the preferred option. +- If mockable cannot be used, env mutations in tests MUST be wrapped in shared + guards and mutexes placed in a shared `test_utils` or `test_helpers` crate. + Direct environment mutation is FORBIDDEN in tests. ### Dependency Management @@ -225,6 +235,18 @@ project: - **Never export the opaque type from a library**. Convert to domain enums at API boundaries, and to `eyre` only in the main `main()` entrypoint or top-level async task. +- In tests, prefer `.expect(...)` over `.unwrap()` to surface clearer failure + diagnostics. +- In production code and shared fixtures, avoid `.expect()` entirely: return + `Result` and use `?` to propagate errors instead of panicking. +- Keep `expect_used` **strict**; do not suppress the lint. +- Recognize that `allow-expect-in-tests = true` **doesn’t cover** helpers + outside `#[cfg(test)]` or `#[test]`; avoid `expect` in such fixtures. +- Use `anyhow`/`eyre` with `.context(...)` to **preserve backtraces** and + provide clear, typed failure paths. +- Update helpers (e.g., `set_dir`) to **return errors** rather than panicking. +- Consume fallible fixtures in `rstest` by **making the test return `Result`** + and applying `?` to the fixture. ## Markdown Guidance @@ -243,39 +265,39 @@ project: The following tooling is available in this environment: -- `mbake` – A Makefile validator. Run using `mbake validate Makefile`. -- `strace` – Traces system calls and signals made by a process; useful for +- `mbake` — A Makefile validator. Run using `mbake validate Makefile`. +- `strace` — Traces system calls and signals made by a process; useful for debugging runtime behaviour and syscalls. -- `gdb` – The GNU Debugger, for inspecting and controlling programs as they +- `gdb` — The GNU Debugger, for inspecting and controlling programs as they execute (or post-mortem via core dumps). -- `ripgrep` – Fast, recursive text search tool (`grep` alternative) that +- `ripgrep` — Fast, recursive text search tool (`grep` alternative) that respects `.gitignore` files. -- `ltrace` – Traces calls to dynamic library functions made by a process. -- `valgrind` – Suite for detecting memory leaks, profiling, and debugging +- `ltrace` — Traces calls to dynamic library functions made by a process. +- `valgrind` — Suite for detecting memory leaks, profiling, and debugging low-level memory errors. -- `bpftrace` – High-level tracing tool for eBPF, using a custom scripting +- `bpftrace` — High-level tracing tool for eBPF, using a custom scripting language for kernel and application tracing. -- `lsof` – Lists open files and the processes using them. -- `htop` – Interactive process viewer (visual upgrade to `top`). -- `iotop` – Displays and monitors I/O usage by processes. -- `ncdu` – NCurses-based disk usage viewer for finding large files/folders. -- `tree` – Displays directory structure as a tree. -- `bat` – `cat` clone with syntax highlighting, Git integration, and paging. -- `delta` – Syntax-highlighted pager for Git and diff output. -- `tcpdump` – Captures and analyses network traffic at the packet level. -- `nmap` – Network scanner for host discovery, port scanning, and service +- `lsof` — Lists open files and the processes using them. +- `htop` — Interactive process viewer (visual upgrade to `top`). +- `iotop` — Displays and monitors I/O usage by processes. +- `ncdu` — NCurses-based disk usage viewer for finding large files/folders. +- `tree` — Displays directory structure as a tree. +- `bat` — `cat` clone with syntax highlighting, Git integration, and paging. +- `delta` — Syntax-highlighted pager for Git and diff output. +- `tcpdump` — Captures and analyses network traffic at the packet level. +- `nmap` — Network scanner for host discovery, port scanning, and service identification. -- `lldb` – LLVM debugger, alternative to `gdb`. -- `eza` – Modern `ls` replacement with more features and better defaults. -- `fzf` – Interactive fuzzy finder for selecting files, commands, etc. -- `hyperfine` – Command-line benchmarking tool with statistical output. -- `shellcheck` – Linter for shell scripts, identifying errors and bad practices. -- `fd` – Fast, user-friendly `find` alternative with sensible defaults. -- `checkmake` – Linter for `Makefile`s, ensuring they follow best practices and +- `lldb` — LLVM debugger, alternative to `gdb`. +- `eza` — Modern `ls` replacement with more features and better defaults. +- `fzf` — Interactive fuzzy finder for selecting files, commands, etc. +- `hyperfine` — Command-line benchmarking tool with statistical output. +- `shellcheck` — Linter for shell scripts, identifying errors and bad practices. +- `fd` — Fast, user-friendly `find` alternative with sensible defaults. +- `checkmake` — Linter for `Makefile`s, ensuring they follow best practices and conventions. -- `srgn` – [Structural grep](https://github.com/alexpovel/srgn), searches code +- `srgn` — [Structural grep](https://github.com/alexpovel/srgn), searches code and enables editing by syntax tree patterns. -- `difft` **(Difftastic)** – Semantic diff tool that compares code structure +- `difft` **(Difftastic)** — Semantic diff tool that compares code structure rather than just text differences. ## Key Takeaway diff --git a/Cargo.toml b/Cargo.toml index 08ff909f..038b5ce1 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -39,10 +39,10 @@ log = "0.4.28" dashmap = "6.1.0" leaky-bucket = "1.1.2" tracing = { version = "0.1.41", features = ["log", "log-always"] } -tracing-subscriber = "0.3" +tracing-subscriber = "0.3.18" metrics = { version = "0.24.2", optional = true } thiserror = "2.0.16" -static_assertions = "1" +static_assertions = "1.1.0" derive_more = { version = "2.0.1", features = ["display", "from"] } [dev-dependencies] @@ -85,11 +85,45 @@ cucumber-tests = [] test-support = [] [lints.clippy] -pedantic = "warn" +pedantic = { level = "warn", priority = -1 } + +# 1. hygiene +allow_attributes = "deny" +allow_attributes_without_reason = "deny" +blanket_clippy_restriction_lints = "deny" +cognitive_complexity = "deny" +needless_pass_by_value = "deny" +implicit_hasher = "deny" + +# 2. debugging leftovers +dbg_macro = "deny" +print_stdout = "deny" +print_stderr = "deny" + +# 3. panic-prone operations +unwrap_used = "deny" +expect_used = "deny" +indexing_slicing = "deny" +string_slice = "deny" +integer_division = "deny" +integer_division_remainder_used = "deny" +panic_in_result_fn = "deny" +unreachable = "deny" [lints.rust] +unknown_lints = "deny" +renamed_and_removed_lints = "deny" unexpected_cfgs = { level = "warn", check-cfg = ['cfg(loom)'] } +[lints.rustdoc] +missing_crate_level_docs = "deny" +broken_intra_doc_links = "deny" +private_intra_doc_links = "deny" +bare_urls = "deny" +invalid_html_tags = "deny" +invalid_codeblock_attributes = "deny" +unescaped_backticks = "deny" + [[example]] name = "echo" path = "examples/echo.rs" diff --git a/Makefile b/Makefile index c1ee70ae..925b92bd 100644 --- a/Makefile +++ b/Makefile @@ -4,6 +4,7 @@ CRATE ?= wireframe CARGO ?= cargo BUILD_JOBS ?= CLIPPY_FLAGS ?= --all-targets --all-features -- -D warnings +RUSTDOC_FLAGS ?= --cfg docsrs -D warnings MDLINT ?= markdownlint NIXIE ?= nixie @@ -29,6 +30,7 @@ target/%/lib$(CRATE).rlib: ## Build library in debug or release $@ lint: ## Run Clippy with warnings denied + RUSTDOCFLAGS="$(RUSTDOC_FLAGS)" $(CARGO) doc --no-deps $(CARGO) clippy $(CLIPPY_FLAGS) fmt: ## Format Rust and Markdown sources diff --git a/clippy.toml b/clippy.toml new file mode 100644 index 00000000..effc6040 --- /dev/null +++ b/clippy.toml @@ -0,0 +1,7 @@ +# Align with CodeScene’s ceiling +cognitive-complexity-threshold = 9 # default is 25 +too-many-arguments-threshold = 4 # default is 7 +too-many-lines-threshold = 70 # default is 100 +excessive-nesting-threshold = 4 # default is off + +allow-expect-in-tests = true diff --git a/examples/async_stream.rs b/examples/async_stream.rs index ae1fa78e..981cbfe3 100644 --- a/examples/async_stream.rs +++ b/examples/async_stream.rs @@ -6,6 +6,7 @@ use async_stream::try_stream; use futures::StreamExt; +use tracing::info; use wireframe::response::Response; #[derive(bincode::Encode, bincode::BorrowDecode, Debug, PartialEq)] @@ -22,10 +23,12 @@ fn stream_response() -> Response { #[tokio::main] async fn main() { + tracing_subscriber::fmt::init(); + let Response::Stream(mut stream) = stream_response() else { return; }; while let Some(Ok(frame)) = stream.next().await { - println!("received frame: {frame:?}"); + info!(?frame, "received frame"); } } diff --git a/examples/echo.rs b/examples/echo.rs index 91ae5f1d..c46a1bc4 100644 --- a/examples/echo.rs +++ b/examples/echo.rs @@ -10,27 +10,56 @@ use wireframe::{ }; type App = wireframe::app::WireframeApp; +type EchoHandler = + Arc Pin + Send>> + Send + Sync>; + +use std::{io, net::SocketAddr, pin::Pin, sync::Arc}; + +use tokio::signal; +use tracing::{error, info}; + +fn echo_handler() -> Pin + Send>> { + Box::pin(async { + info!("echo request received"); + // `WireframeApp` automatically echoes the envelope back. + }) +} + +fn build_app(handler: EchoHandler) -> wireframe::app::Result { App::new()?.route(1, handler) } #[tokio::main] async fn main() -> Result<(), ServerError> { - let factory = || { - App::new() - .expect("failed to create WireframeApp") - .route( - 1, - std::sync::Arc::new(|_: &Envelope| { - Box::pin(async move { - println!("echo request received"); - // `WireframeApp` automatically echoes the envelope back. - }) - }), - ) - .expect("failed to register route 1") + tracing_subscriber::fmt::init(); + + let handler: EchoHandler = Arc::new(|_: &Envelope| echo_handler()); + build_app(handler.clone()).map_err(|err| { + error!("failed to build echo app: {err}"); + ServerError::Bind(io::Error::other(err)) + })?; + + let factory = { + let handler = Arc::clone(&handler); + move || match build_app(Arc::clone(&handler)) { + Ok(app) => app, + Err(err) => { + error!("failed to rebuild echo app: {err}"); + App::default() + } + } }; - WireframeServer::new(factory) - .bind("127.0.0.1:7878".parse().expect("invalid socket address"))? - .run() + let addr: SocketAddr = "127.0.0.1:7878".parse().map_err(|err| { + ServerError::Bind(std::io::Error::new(std::io::ErrorKind::InvalidInput, err)) + })?; + let server = WireframeServer::new(factory).bind(addr)?; + + server + .run_with_shutdown(async { + match signal::ctrl_c().await { + Ok(()) => info!("shutdown signal received, stopping echo server"), + Err(err) => error!("failed to wait for shutdown signal: {err}"), + } + }) .await?; Ok(()) } diff --git a/examples/metadata_routing.rs b/examples/metadata_routing.rs index 4670d042..b8125123 100644 --- a/examples/metadata_routing.rs +++ b/examples/metadata_routing.rs @@ -41,16 +41,25 @@ impl FrameMetadata for HeaderSerializer { type Error = io::Error; fn parse(&self, src: &[u8]) -> Result<(Envelope, usize), io::Error> { - if src.len() < 3 { - return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "header")); - } - let id = u32::from(u16::from_be_bytes([src[0], src[1]])); + let id_bytes: [u8; 2] = src + .get(..2) + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "header"))? + .try_into() + .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "header id width"))?; + // The third byte carries message flags. This example intentionally // ignores the flags, but a real protocol might parse and act on these - // bits. - let _ = src[2]; + // bits. We still validate its presence to avoid panics. + let _flags = src + .get(2) + .copied() + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "header flags"))?; + // Only extract metadata here; defer payload handling to the serializer. - Ok((Envelope::new(id, None, Vec::new()), 3)) + Ok(( + Envelope::new(u32::from(u16::from_be_bytes(id_bytes)), None, Vec::new()), + 3, + )) } } @@ -60,7 +69,7 @@ struct Ping; #[tokio::main] async fn main() -> io::Result<()> { let app = App::with_serializer(HeaderSerializer) - .expect("failed to create app") + .map_err(io::Error::other)? .buffer_capacity(MAX_FRAME) .route( 1, @@ -70,7 +79,7 @@ async fn main() -> io::Result<()> { }) }), ) - .expect("failed to add ping route") + .map_err(io::Error::other)? .route( 2, Arc::new(|_env: &Envelope| { @@ -79,27 +88,23 @@ async fn main() -> io::Result<()> { }) }), ) - .expect("failed to add pong route"); + .map_err(io::Error::other)?; let mut codec = app.length_codec(); let (mut client, server) = duplex(1024); - let server_task = tokio::spawn(async move { - app.handle_connection(server).await; - }); + let server_task = tokio::spawn(async move { app.handle_connection_result(server).await }); - let payload = Ping.to_bytes().expect("failed to serialize Ping message"); + let payload = Ping.to_bytes().map_err(io::Error::other)?; let mut frame = Vec::new(); frame.extend_from_slice(&1u16.to_be_bytes()); frame.push(0); frame.extend_from_slice(&payload); let mut bytes = BytesMut::with_capacity(frame.len() + 4); // +4 for the length prefix - codec - .encode(frame.into(), &mut bytes) - .expect("failed to encode frame"); + codec.encode(frame.into(), &mut bytes)?; client.write_all(&bytes).await?; client.shutdown().await?; - server_task.await.expect("server task failed"); + server_task.await.map_err(io::Error::other)??; Ok(()) } diff --git a/examples/multi_packet.rs b/examples/multi_packet.rs index 1af1cb88..d0ffae46 100644 --- a/examples/multi_packet.rs +++ b/examples/multi_packet.rs @@ -9,7 +9,8 @@ use std::time::Duration; use futures::TryStreamExt; use tokio::time::sleep; -use wireframe::Response; +use tracing::info; +use wireframe::{Response, WireframeError}; const TRANSCRIPT: &[&str] = &[ "Client: HELLO", @@ -58,7 +59,7 @@ fn multi_packet_response() -> Response { for (index, line) in TRANSCRIPT.iter().enumerate() { let frame = Frame::chunk(index, line); if sender.send(frame).await.is_err() { - // The connection dropped; stop work early. + tracing::trace!("connection dropped, stopping chunk task early"); break; } @@ -72,7 +73,7 @@ fn multi_packet_response() -> Response { let _ = chunk_task.await; let summary = Frame::summary(TRANSCRIPT.len()); if summary_sender.send(summary).await.is_err() { - // The connection dropped; stop work early. + tracing::trace!("connection dropped, summary not sent"); } }); @@ -80,24 +81,24 @@ fn multi_packet_response() -> Response { } #[tokio::main] -async fn main() { +async fn main() -> Result<(), WireframeError<()>> { + tracing_subscriber::fmt::init(); + let response = multi_packet_response(); let mut stream = response.into_stream(); - while let Some(frame) = stream - .try_next() - .await - .expect("multi-packet stream should not fail") - { + while let Some(frame) = stream.try_next().await? { match frame { Frame { kind: FrameKind::Chunk(index), data, - } => println!("Chunk {index}: {data}"), + } => info!("Chunk {index}: {data}"), Frame { kind: FrameKind::Summary, data, - } => println!("Summary: {data}"), + } => info!("Summary: {data}"), } } + + Ok(()) } diff --git a/examples/packet_enum.rs b/examples/packet_enum.rs index 544fe5b4..7b7ff207 100644 --- a/examples/packet_enum.rs +++ b/examples/packet_enum.rs @@ -3,20 +3,22 @@ //! The application defines an enum representing different packet variants and //! shows how to dispatch handlers based on the variant received. -use std::{collections::HashMap, future::Future, pin::Pin}; +use std::{collections::HashMap, future::Future, net::SocketAddr, pin::Pin, sync::Arc}; use async_trait::async_trait; -use tracing::{info, warn}; +use tokio::{net::TcpListener, signal}; +use tracing::{error, info, warn}; use wireframe::{ app::Envelope, message::Message, middleware::{HandlerService, Service, ServiceRequest, ServiceResponse, Transform}, serializer::BincodeSerializer, - server::{ServerError, WireframeServer}, }; type App = wireframe::app::WireframeApp; +const DEFAULT_ADDR: &str = "127.0.0.1:7879"; + #[derive(bincode::Encode, bincode::BorrowDecode, Debug)] enum ExamplePacket { Ping, @@ -78,22 +80,51 @@ fn handle_packet(_env: &Envelope) -> Pin + Send>> { }) } +fn build_app() -> wireframe::app::Result { + App::new()? + .wrap(DecodeMiddleware)? + .route(1, Arc::new(handle_packet)) +} + #[tokio::main] -async fn main() -> Result<(), ServerError> { - let factory = || { - App::new() - .expect("Failed to create WireframeApp") - .wrap(DecodeMiddleware) - .expect("Failed to wrap middleware") - .route(1, std::sync::Arc::new(handle_packet)) - .expect("Failed to add route") - }; - - let addr = std::env::var("SERVER_ADDR").unwrap_or_else(|_| "127.0.0.1:7879".to_string()); - - WireframeServer::new(factory) - .bind(addr.parse().expect("Invalid server address"))? - .run() - .await?; +#[expect( + clippy::integer_division_remainder_used, + reason = "tokio::select! macro expansion performs modulo internally" +)] +async fn main() -> std::io::Result<()> { + tracing_subscriber::fmt::init(); + + let app = Arc::new(build_app().map_err(std::io::Error::other)?); + + let addr_str = std::env::var("SERVER_ADDR").unwrap_or_else(|_| DEFAULT_ADDR.to_string()); + let addr: SocketAddr = addr_str.parse().map_err(|e| { + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("SERVER_ADDR must be a valid socket address: {e}"), + ) + })?; + + let listener = TcpListener::bind(addr).await?; + loop { + tokio::select! { + res = listener.accept() => { + let (stream, _) = res?; + let app = Arc::clone(&app); + tokio::spawn(async move { + if let Err(e) = app.handle_connection_result(stream).await { + error!("connection handling failed: {e}"); + } + }); + } + ctrl_c = signal::ctrl_c() => { + match ctrl_c { + Ok(()) => info!("packet_enum server received shutdown signal"), + Err(e) => error!("failed waiting for shutdown signal: {e}"), + } + break; + } + } + } + Ok(()) } diff --git a/examples/ping_pong.rs b/examples/ping_pong.rs index 7410d31b..d27a078f 100644 --- a/examples/ping_pong.rs +++ b/examples/ping_pong.rs @@ -6,12 +6,13 @@ use std::{net::SocketAddr, sync::Arc}; use async_trait::async_trait; +use tokio::{net::TcpListener, signal}; +use tracing::{error, info}; use wireframe::{ app::{Envelope, Packet, Result as AppResult}, message::Message, middleware::{HandlerService, Service, ServiceRequest, ServiceResponse, Transform}, serializer::BincodeSerializer, - server::{ServerError, WireframeServer}, }; type App = wireframe::app::WireframeApp; @@ -30,7 +31,7 @@ fn encode_error(msg: impl Into) -> Vec { match err.to_bytes() { Ok(bytes) => bytes, Err(e) => { - eprintln!("failed to encode error: {e:?}"); + error!(error = ?e, "failed to encode error"); Vec::new() } } @@ -66,7 +67,7 @@ where let (ping_req, _) = match Ping::from_bytes(req.frame()) { Ok(val) => val, Err(e) => { - eprintln!("failed to decode ping: {e:?}"); + error!(error = ?e, "failed to decode ping"); return Ok(ServiceResponse::new( encode_error(format!("decode error: {e:?}")), cid, @@ -77,13 +78,13 @@ where let pong_resp = if let Some(v) = ping_req.0.checked_add(1) { Pong(v) } else { - eprintln!("ping overflowed at {}", ping_req.0); + error!(value = ping_req.0, "ping overflowed"); return Ok(ServiceResponse::new(encode_error("overflow"), cid)); }; match pong_resp.to_bytes() { Ok(bytes) => *response.frame_mut() = bytes, Err(e) => { - eprintln!("failed to encode pong: {e:?}"); + error!(error = ?e, "failed to encode pong"); return Ok(ServiceResponse::new( encode_error(format!("encode error: {e:?}")), cid, @@ -118,9 +119,9 @@ where type Error = std::convert::Infallible; async fn call(&self, req: ServiceRequest) -> Result { - println!("request: {:?}", req.frame()); + info!(frame = ?req.frame(), "request"); let resp = self.inner.call(req).await?; - println!("response: {:?}", resp.frame()); + info!(frame = ?resp.frame(), "response"); Ok(resp) } } @@ -144,14 +145,40 @@ fn build_app() -> AppResult { } #[tokio::main] -async fn main() -> Result<(), ServerError> { - let factory = || build_app().expect("app build failed"); +#[expect( + clippy::integer_division_remainder_used, + reason = "tokio::select! macro expansion performs modulo internally" +)] +async fn main() -> std::io::Result<()> { + tracing_subscriber::fmt::init(); let default_addr = "127.0.0.1:7878"; let addr_str = std::env::args() .nth(1) .unwrap_or_else(|| default_addr.into()); - let addr: SocketAddr = addr_str.parse().expect("invalid address"); - WireframeServer::new(factory).bind(addr)?.run().await?; + + let app = Arc::new(build_app().map_err(std::io::Error::other)?); + let addr: SocketAddr = addr_str.parse().map_err(std::io::Error::other)?; + let listener = TcpListener::bind(addr).await?; + loop { + tokio::select! { + res = listener.accept() => { + let (stream, _) = res?; + let app = Arc::clone(&app); + tokio::spawn(async move { + if let Err(err) = app.handle_connection_result(stream).await { + error!("connection handling failed: {err}"); + } + }); + } + ctrl_c = signal::ctrl_c() => { + match ctrl_c { + Ok(()) => info!("ping-pong server received shutdown signal"), + Err(e) => error!("failed waiting for shutdown signal: {e}"), + } + break; + } + } + } Ok(()) } diff --git a/src/app/connection.rs b/src/app/connection.rs index 08509839..d84a5765 100644 --- a/src/app/connection.rs +++ b/src/app/connection.rs @@ -26,13 +26,25 @@ use crate::{ serializer::Serializer, }; +fn purge_expired(fragmentation: &mut Option) { + if let Some(frag) = fragmentation.as_mut() { + frag.purge_expired(); + } +} + /// Maximum consecutive deserialization failures before closing a connection. const MAX_DESER_FAILURES: u32 = 10; -#[derive(Debug)] -enum EnvelopeDecodeError { - Parse(E), - Deserialize(Box), +/// Per-frame processing state bundled for `handle_frame`. +struct FrameHandlingContext<'a, E, W> +where + E: Packet, + W: AsyncRead + AsyncWrite + Unpin, +{ + framed: &'a mut Framed, + deser_failures: &'a mut u32, + routes: &'a HashMap>, + fragmentation: &'a mut Option, } impl WireframeApp @@ -115,23 +127,19 @@ where fn parse_envelope( &self, frame: &[u8], - ) -> std::result::Result<(Envelope, usize), EnvelopeDecodeError> { + ) -> std::result::Result<(Envelope, usize), Box> { self.serializer .parse(frame) - .map_err(EnvelopeDecodeError::Parse) - .or_else(|_| { - self.serializer - .deserialize::(frame) - .map_err(EnvelopeDecodeError::Deserialize) - }) + .map_err(Box::::from) + .or_else(|_| self.serializer.deserialize::(frame)) } - /// Handle an accepted connection end-to-end. + /// Handle an accepted connection end-to-end, returning any processing error. /// - /// Runs optional connection setup to produce per-connection state, - /// initializes (and caches) route chains, processes the framed stream - /// with per-frame timeouts, and finally runs optional teardown. - pub async fn handle_connection(&self, stream: W) + /// # Errors + /// + /// Returns an [`io::Error`] if stream processing or handler execution fails. + pub async fn handle_connection_result(&self, stream: W) -> io::Result<()> where W: AsyncRead + AsyncWrite + Send + Unpin + 'static, { @@ -152,11 +160,27 @@ where "connection terminated with error: correlation_id={:?}, error={e:?}", None:: ); + return Err(e); } if let (Some(teardown), Some(state)) = (&self.on_disconnect, state) { teardown(state).await; } + + Ok(()) + } + + /// Handle an accepted connection end-to-end, logging errors and swallowing the result. + pub async fn handle_connection(&self, stream: W) + where + W: AsyncRead + AsyncWrite + Send + Unpin + 'static, + { + if let Err(e) = self.handle_connection_result(stream).await { + warn!( + "connection handling completed with error: correlation_id={:?}, error={e:?}", + None:: + ); + } } async fn build_chains(&self) -> HashMap> { @@ -190,11 +214,13 @@ where match timeout(timeout_dur, framed.next()).await { Ok(Some(Ok(buf))) => { self.handle_frame( - &mut framed, buf.as_ref(), - &mut deser_failures, - routes, - &mut fragmentation, + FrameHandlingContext { + framed: &mut framed, + deser_failures: &mut deser_failures, + routes, + fragmentation: &mut fragmentation, + }, ) .await?; } @@ -202,9 +228,7 @@ where Ok(None) => break, Err(_) => { debug!("read timeout elapsed; continuing to wait for next frame"); - if let Some(state) = fragmentation.as_mut() { - state.purge_expired(); - } + purge_expired(&mut fragmentation); } } } @@ -214,15 +238,19 @@ where async fn handle_frame( &self, - framed: &mut Framed, frame: &[u8], - deser_failures: &mut u32, - routes: &HashMap>, - fragmentation: &mut Option, + ctx: FrameHandlingContext<'_, E, W>, ) -> io::Result<()> where W: AsyncRead + AsyncWrite + Unpin, { + let FrameHandlingContext { + framed, + deser_failures, + routes, + fragmentation, + } = ctx; + crate::metrics::inc_frames(crate::metrics::Direction::Inbound); let Some(env) = self.decode_envelope(frame, deser_failures)? else { return Ok(()); @@ -238,8 +266,16 @@ where }; if let Some(service) = routes.get(&env.id) { - frame_handling::forward_response(&self.serializer, env, service, framed, fragmentation) - .await?; + frame_handling::forward_response( + env, + service, + frame_handling::ResponseContext { + serializer: &self.serializer, + framed, + fragmentation, + }, + ) + .await?; } else { warn!( "no handler for message id: id={}, correlation_id={:?}", @@ -250,6 +286,25 @@ where Ok(()) } + /// Increment deserialization failures and close the connection if the threshold is exceeded. + fn handle_decode_failure( + deser_failures: &mut u32, + context: &str, + err: impl std::fmt::Debug, + ) -> Result, io::Error> { + *deser_failures += 1; + warn!("{context}: correlation_id={:?}, error={err:?}", None::); + crate::metrics::inc_deser_errors(); + if *deser_failures >= MAX_DESER_FAILURES { + warn!("closing connection after {deser_failures} deserialization failures: {context}"); + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "too many deserialization failures", + )); + } + Ok(None) + } + fn decode_envelope( &self, frame: &[u8], @@ -260,35 +315,8 @@ where *deser_failures = 0; Ok(Some(env)) } - Err(EnvelopeDecodeError::Parse(e)) => { - *deser_failures += 1; - warn!( - "failed to parse message: correlation_id={:?}, error={e:?}", - None:: - ); - crate::metrics::inc_deser_errors(); - if *deser_failures >= MAX_DESER_FAILURES { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - "too many deserialization failures", - )); - } - Ok(None) - } - Err(EnvelopeDecodeError::Deserialize(e)) => { - *deser_failures += 1; - warn!( - "failed to deserialize message: correlation_id={:?}, error={e:?}", - None:: - ); - crate::metrics::inc_deser_errors(); - if *deser_failures >= MAX_DESER_FAILURES { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - "too many deserialization failures", - )); - } - Ok(None) + Err(err) => { + Self::handle_decode_failure(deser_failures, "failed to decode message", err) } } } diff --git a/src/app/envelope.rs b/src/app/envelope.rs index d229fd58..02ebb36e 100644 --- a/src/app/envelope.rs +++ b/src/app/envelope.rs @@ -1,16 +1,17 @@ //! Packet abstraction and envelope types. //! //! These types decouple serialisation from routing by wrapping raw payloads in -//! identifiers understood by [`crate::app::WireframeApp`]. This allows the -//! builder (`crate::app::WireframeApp`) to route frames before full -//! deserialisation. See [`crate::app::builder::WireframeApp`] for how envelopes -//! are used when registering routes. +//! identifiers understood by [`crate::app::builder::WireframeApp`]. This +//! allows the builder to route frames before full deserialisation. See +//! [`crate::app::builder::WireframeApp`] for how envelopes are used when +//! registering routes. use crate::{correlation::CorrelatableFrame, message::Message}; /// Envelope-like type used to wrap incoming and outgoing messages. /// -/// Custom envelope types must implement this trait so [`WireframeApp`] can +/// Custom envelope types must implement this trait so +/// [`crate::app::builder::WireframeApp`] can /// route messages and construct responses. /// /// # Example @@ -66,7 +67,8 @@ pub struct PacketParts { payload: Vec, } -/// Basic envelope type used by [`WireframeApp::handle_connection`]. +/// Basic envelope type used by +/// [`crate::app::builder::WireframeApp::handle_connection`]. /// /// Incoming frames are deserialised into an `Envelope` containing the /// message identifier and raw payload bytes. diff --git a/src/app/frame_handling.rs b/src/app/frame_handling.rs index 0eefd97e..d916e88c 100644 --- a/src/app/frame_handling.rs +++ b/src/app/frame_handling.rs @@ -20,51 +20,65 @@ use crate::{ serializer::Serializer, }; -/// Attempt to reassemble a potentially fragmented envelope. -pub(crate) fn reassemble_if_needed( - fragmentation: &mut Option, - deser_failures: &mut u32, - env: Envelope, - max_deser_failures: u32, -) -> io::Result> { - fn handle_fragment_error( - deser_failures: &mut u32, - max_deser_failures: u32, +struct DeserFailureTracker<'a> { + count: &'a mut u32, + limit: u32, +} + +impl<'a> DeserFailureTracker<'a> { + fn new(count: &'a mut u32, limit: u32) -> Self { Self { count, limit } } + + fn record( + &mut self, correlation_id: Option, context: &str, err: impl std::fmt::Debug, - ) -> io::Result> { - *deser_failures += 1; + ) -> io::Result<()> { + *self.count += 1; warn!("{context}: correlation_id={correlation_id:?}, error={err:?}"); crate::metrics::inc_deser_errors(); - if *deser_failures >= max_deser_failures { + if *self.count >= self.limit { return Err(io::Error::new( io::ErrorKind::InvalidData, "too many deserialization failures", )); } - Ok(None) + Ok(()) } +} + +pub(crate) struct ResponseContext<'a, S, W> +where + S: Serializer + Send + Sync, + W: AsyncRead + AsyncWrite + Unpin, +{ + pub(crate) serializer: &'a S, + pub(crate) framed: &'a mut Framed, + pub(crate) fragmentation: &'a mut Option, +} + +/// Attempt to reassemble a potentially fragmented envelope. +pub(crate) fn reassemble_if_needed( + fragmentation: &mut Option, + deser_failures: &mut u32, + env: Envelope, + max_deser_failures: u32, +) -> io::Result> { + let mut failures = DeserFailureTracker::new(deser_failures, max_deser_failures); if let Some(state) = fragmentation.as_mut() { let correlation_id = env.correlation_id; match state.reassemble(env) { Ok(Some(env)) => Ok(Some(env)), Ok(None) => Ok(None), - Err(FragmentProcessError::Decode(err)) => handle_fragment_error( - deser_failures, - max_deser_failures, - correlation_id, - "failed to decode fragment header", - err, - ), - Err(FragmentProcessError::Reassembly(err)) => handle_fragment_error( - deser_failures, - max_deser_failures, - correlation_id, - "fragment reassembly failed", - err, - ), + Err(FragmentProcessError::Decode(err)) => { + failures.record(correlation_id, "failed to decode fragment header", err)?; + Ok(None) + } + Err(FragmentProcessError::Reassembly(err)) => { + failures.record(correlation_id, "fragment reassembly failed", err)?; + Ok(None) + } } } else { Ok(Some(env)) @@ -73,11 +87,9 @@ pub(crate) fn reassemble_if_needed( /// Forward a handler response, fragmenting if required, and write to the framed stream. pub(crate) async fn forward_response( - serializer: &S, env: Envelope, service: &HandlerService, - framed: &mut Framed, - fragmentation: &mut Option, + ctx: ResponseContext<'_, S, W>, ) -> io::Result<()> where S: Serializer + Send + Sync, @@ -85,59 +97,97 @@ where W: AsyncRead + AsyncWrite + Unpin, { let request = ServiceRequest::new(env.payload, env.correlation_id); - match service.call(request).await { - Ok(resp) => { - let parts = PacketParts::new(env.id, resp.correlation_id(), resp.into_inner()) - .inherit_correlation(env.correlation_id); - let correlation_id = parts.correlation_id(); - let responses = if let Some(state) = fragmentation.as_mut() { - match state.fragment(Envelope::from_parts(parts)) { - Ok(fragmented) => fragmented, - Err(err) => { - warn!( - "failed to fragment response: id={}, correlation_id={:?}, \ - error={err:?}", - env.id, correlation_id - ); - crate::metrics::inc_handler_errors(); - return Ok(()); - } - } - } else { - vec![Envelope::from_parts(parts)] - }; - - for response in responses { - match serializer.serialize(&response) { - Ok(bytes) => { - if let Err(e) = framed.send(bytes.into()).await { - warn!( - "failed to send response: id={}, correlation_id={:?}, error={e:?}", - env.id, correlation_id - ); - crate::metrics::inc_handler_errors(); - break; - } - } - Err(e) => { - warn!( - "failed to serialize response: id={}, correlation_id={:?}, error={e:?}", - env.id, correlation_id - ); - crate::metrics::inc_handler_errors(); - break; - } - } - } - } + let resp = match service.call(request).await { + Ok(resp) => resp, Err(e) => { warn!( "handler error: id={}, correlation_id={:?}, error={e:?}", env.id, env.correlation_id ); crate::metrics::inc_handler_errors(); + return Ok(()); + } + }; + + let parts = PacketParts::new(env.id, resp.correlation_id(), resp.into_inner()) + .inherit_correlation(env.correlation_id); + let correlation_id = parts.correlation_id(); + let Ok(responses) = fragment_responses(ctx.fragmentation, parts, env.id, correlation_id) else { + return Ok(()); // already logged + }; + + for response in responses { + let Ok(bytes) = serialize_response(ctx.serializer, &response, env.id, correlation_id) + else { + break; // already logged + }; + + if send_response_bytes(ctx.framed, bytes, env.id, correlation_id) + .await + .is_err() + { + break; } } Ok(()) } + +fn fragment_responses( + fragmentation: &mut Option, + parts: PacketParts, + id: u32, + correlation_id: Option, +) -> io::Result> { + let envelope = Envelope::from_parts(parts); + match fragmentation.as_mut() { + Some(state) => match state.fragment(envelope) { + Ok(fragmented) => Ok(fragmented), + Err(err) => { + warn!( + "failed to fragment response: id={id}, correlation_id={correlation_id:?}, \ + error={err:?}" + ); + crate::metrics::inc_handler_errors(); + Err(io::Error::other("fragmentation failed")) + } + }, + None => Ok(vec![envelope]), + } +} + +fn serialize_response( + serializer: &S, + response: &Envelope, + id: u32, + correlation_id: Option, +) -> io::Result> { + match serializer.serialize(response) { + Ok(bytes) => Ok(bytes), + Err(e) => { + warn!( + "failed to serialize response: id={id}, correlation_id={correlation_id:?}, \ + error={e:?}" + ); + crate::metrics::inc_handler_errors(); + Err(io::Error::other("serialization failed")) + } + } +} + +async fn send_response_bytes( + framed: &mut Framed, + bytes: Vec, + id: u32, + correlation_id: Option, +) -> io::Result<()> +where + W: AsyncRead + AsyncWrite + Unpin, +{ + if let Err(e) = framed.send(bytes.into()).await { + warn!("failed to send response: id={id}, correlation_id={correlation_id:?}, error={e:?}"); + crate::metrics::inc_handler_errors(); + return Err(io::Error::other("send failed")); + } + Ok(()) +} diff --git a/src/connection.rs b/src/connection.rs index f2c15db6..c8dc9bc9 100644 --- a/src/connection.rs +++ b/src/connection.rs @@ -94,6 +94,18 @@ pub struct FairnessConfig { pub time_slice: Option, } +/// Bundles push queues with their shared handle for actor construction. +pub struct ConnectionChannels { + pub queues: PushQueues, + pub handle: PushHandle, +} + +impl ConnectionChannels { + /// Create a new bundle of push queues and their associated handle. + #[must_use] + pub fn new(queues: PushQueues, handle: PushHandle) -> Self { Self { queues, handle } } +} + impl Default for FairnessConfig { fn default() -> Self { Self { @@ -267,8 +279,7 @@ where shutdown: CancellationToken, ) -> Self { Self::with_hooks( - queues, - handle, + ConnectionChannels::new(queues, handle), response, shutdown, ProtocolHooks::::default(), @@ -278,12 +289,12 @@ where /// Create a new `ConnectionActor` with custom protocol hooks. #[must_use] pub fn with_hooks( - queues: PushQueues, - handle: PushHandle, + channels: ConnectionChannels, response: Option>, shutdown: CancellationToken, hooks: ProtocolHooks, ) -> Self { + let ConnectionChannels { queues, handle } = channels; let ctx = ConnectionContext; let counter = ActiveConnection::new(); let mut actor = Self { @@ -415,7 +426,7 @@ where ); } MultiPacketStamp::Disabled => { - unreachable!("multi-packet correlation invoked without configuration"); + // No channel is active, so there is nothing to stamp. } } } @@ -477,6 +488,10 @@ where /// /// The `strict_priority_order` and `shutdown_signal_precedence` tests /// assert that this ordering is preserved across refactors. + #[expect( + clippy::integer_division_remainder_used, + reason = "tokio::select! expands to modulus operations internally" + )] async fn next_event(&mut self, state: &ActorState) -> Event { let high_available = self.high_rx.is_some(); let low_available = self.low_rx.is_some(); @@ -646,13 +661,12 @@ where where F: Packet, { - if let Some(fragmenter) = &self.fragmenter { - match fragment_packet(fragmenter, frame) { - Ok(frames) => { - for frame in frames { - self.push_frame(frame, out); - } - } + if let Some(fragmenter) = self.fragmenter.as_deref() { + let fragmented = fragment_packet(fragmenter, frame); + match fragmented { + Ok(frames) => frames + .into_iter() + .for_each(|frame| self.push_frame(frame, out)), Err(err) => { warn!( "failed to fragment frame: connection_id={:?}, peer={:?}, error={err:?}", @@ -754,10 +768,14 @@ where fn try_opportunistic_drain(&mut self, kind: QueueKind, ctx: DrainContext<'_, F>) -> bool { let DrainContext { out, state } = ctx; match kind { - QueueKind::High => unreachable!(concat!( - "try_opportunistic_drain(High) is unsupported; ", - "High is handled by biased polling", - )), + QueueKind::High => { + debug_assert!( + false, + "try_opportunistic_drain(High) is unsupported; High is handled by biased \ + polling" + ); + false + } QueueKind::Low => { let res = match self.low_rx.as_mut() { Some(receiver) => receiver.try_recv(), diff --git a/src/connection/test_support.rs b/src/connection/test_support.rs index c630eced..deaa9bbb 100644 --- a/src/connection/test_support.rs +++ b/src/connection/test_support.rs @@ -9,6 +9,7 @@ use tokio_util::sync::CancellationToken; use super::{ ActorState, ConnectionActor, + ConnectionChannels, DrainContext, MultiPacketTerminationReason, ProtocolHooks, @@ -54,8 +55,7 @@ pub fn create_test_actor_with_hooks( .low_capacity(4) .build()?; Ok(ConnectionActor::with_hooks( - queues, - handle, + ConnectionChannels::new(queues, handle), None, CancellationToken::new(), hooks, @@ -69,10 +69,6 @@ pub struct ActorHarness { pub out: Vec, } -impl Default for ActorHarness { - fn default() -> Self { Self::new().expect("failed to build ActorHarness") } -} - impl ActorHarness { /// Create a harness with custom hooks and state flags. /// diff --git a/src/extractor.rs b/src/extractor.rs index bfbd1a14..0a15f583 100644 --- a/src/extractor.rs +++ b/src/extractor.rs @@ -144,7 +144,7 @@ impl Payload<'_> { /// ``` pub fn advance(&mut self, count: usize) { let n = count.min(self.data.len()); - self.data = &self.data[n..]; + self.data = self.data.get(n..).unwrap_or_default(); } /// Returns the number of bytes remaining. diff --git a/src/fairness.rs b/src/fairness.rs index 73b46303..bb9b4ab1 100644 --- a/src/fairness.rs +++ b/src/fairness.rs @@ -145,14 +145,16 @@ mod tests { } } + #[expect(clippy::expect_used, reason = "poisoned lock should fail tests loudly")] fn advance(&self, dur: Duration) { - let mut now = self.now.lock().expect("lock poisoned"); + let mut now = self.now.lock().expect("MockClock mutex poisoned"); *now += dur; } } impl Clock for MockClock { - fn now(&self) -> Instant { *self.now.lock().expect("lock poisoned") } + #[expect(clippy::expect_used, reason = "poisoned lock should fail tests loudly")] + fn now(&self) -> Instant { *self.now.lock().expect("MockClock mutex poisoned") } } #[rstest] diff --git a/src/fragment/error.rs b/src/fragment/error.rs index 78c4d855..9e3c04bb 100644 --- a/src/fragment/error.rs +++ b/src/fragment/error.rs @@ -52,6 +52,16 @@ pub enum FragmentationError { /// The fragment index cannot advance because it would overflow `u32`. #[error("fragment index overflow after {last}")] IndexOverflow { last: FragmentIndex }, + /// Calculated fragment slice exceeded payload bounds. + #[error("fragment slice out of bounds: offset={offset}, end={end}, total={total}")] + SliceBounds { + /// Start offset attempted. + offset: usize, + /// Exclusive end offset attempted. + end: usize, + /// Total payload length. + total: usize, + }, } /// Errors produced while re-assembling inbound fragments. diff --git a/src/fragment/fragmenter.rs b/src/fragment/fragmenter.rs index eb7f153b..de7b90b3 100644 --- a/src/fragment/fragmenter.rs +++ b/src/fragment/fragmenter.rs @@ -20,6 +20,16 @@ pub struct Fragmenter { next_message_id: AtomicU64, } +#[derive(Debug, Clone, Copy)] +pub(crate) struct FragmentCursor { + offset: usize, + index: FragmentIndex, +} + +impl FragmentCursor { + pub(crate) const fn new(offset: usize, index: FragmentIndex) -> Self { Self { offset, index } } +} + impl Fragmenter { /// Create a new fragmenter that caps fragment payloads at `max_fragment_size` bytes. #[must_use] @@ -64,7 +74,8 @@ impl Fragmenter { /// /// Returns [`FragmentationError::Encode`] if serialization fails, or /// [`FragmentationError::IndexOverflow`] if the fragment index would - /// overflow `u32`. + /// overflow `u32`. Slice calculations that exceed payload bounds return + /// [`FragmentationError::SliceBounds`]. pub fn fragment_message( &self, message: &M, @@ -78,7 +89,9 @@ impl Fragmenter { /// # Errors /// /// Returns [`FragmentationError::IndexOverflow`] if more than - /// `u32::MAX + 1` fragments are required. + /// `u32::MAX + 1` fragments are required, or + /// [`FragmentationError::SliceBounds`] if slice calculation exceeds payload + /// bounds. pub fn fragment_bytes( &self, payload: impl AsRef<[u8]>, @@ -92,7 +105,9 @@ impl Fragmenter { /// # Errors /// /// Returns [`FragmentationError::IndexOverflow`] if more than - /// `u32::MAX + 1` fragments are required. + /// `u32::MAX + 1` fragments are required, or + /// [`FragmentationError::SliceBounds`] if slice calculation exceeds payload + /// bounds. pub fn fragment_with_id( &self, message_id: MessageId, @@ -106,6 +121,19 @@ impl Fragmenter { &self, message_id: MessageId, payload: &[u8], + ) -> Result, FragmentationError> { + self.build_fragments_from( + message_id, + payload, + FragmentCursor::new(0, FragmentIndex::zero()), + ) + } + + fn build_fragments_from( + &self, + message_id: MessageId, + payload: &[u8], + mut cursor: FragmentCursor, ) -> Result, FragmentationError> { let max = self.max_fragment_size.get(); if payload.is_empty() { @@ -114,32 +142,59 @@ impl Fragmenter { } let total = payload.len(); + if cursor.offset > total { + return Err(FragmentationError::SliceBounds { + offset: cursor.offset, + end: cursor.offset, + total, + }); + } let mut fragments = Vec::with_capacity(div_ceil(total, max)); - let mut index = FragmentIndex::zero(); - let mut offset = 0usize; - while offset < total { - let end = (offset + max).min(total); + while cursor.offset < total { + let end = (cursor.offset + max).min(total); let is_last = end == total; + let chunk = if let Some(slice) = payload.get(cursor.offset..end) { + slice.to_vec() + } else { + return Err(FragmentationError::SliceBounds { + offset: cursor.offset, + end, + total, + }); + }; fragments.push(FragmentFrame::new( - FragmentHeader::new(message_id, index, is_last), - payload[offset..end].to_vec(), + FragmentHeader::new(message_id, cursor.index, is_last), + chunk, )); if is_last { break; } - offset = end; - index = index + cursor.offset = end; + cursor.index = cursor + .index .checked_increment() - .ok_or(FragmentationError::IndexOverflow { last: index })?; + .ok_or(FragmentationError::IndexOverflow { last: cursor.index })?; } Ok(fragments) } } +#[cfg(test)] +impl Fragmenter { + pub(crate) fn build_fragments_from_for_tests( + &self, + message_id: MessageId, + payload: &[u8], + cursor: FragmentCursor, + ) -> Result, FragmentationError> { + self.build_fragments_from(message_id, payload, cursor) + } +} + /// Metadata and payload for a single outbound fragment. #[derive(Clone, Debug, PartialEq, Eq)] pub struct FragmentFrame { diff --git a/src/fragment/payload.rs b/src/fragment/payload.rs index 3821e3da..56560027 100644 --- a/src/fragment/payload.rs +++ b/src/fragment/payload.rs @@ -8,7 +8,12 @@ use std::num::NonZeroUsize; -use bincode::{borrow_decode_from_slice, config, encode_to_vec, error::DecodeError}; +use bincode::{ + borrow_decode_from_slice, + config, + encode_to_vec, + error::{DecodeError, EncodeError}, +}; use super::{FragmentHeader, FragmentIndex, MessageId}; @@ -27,11 +32,14 @@ pub fn fragment_overhead() -> NonZeroUsize { // header size is stable for the fixed-width fields used here and must // remain well below `u16::MAX` to satisfy the framing format. let header = FragmentHeader::new(MessageId::new(0), FragmentIndex::zero(), false); - let header_bytes = encode_to_vec(header, config::standard()) - .expect("fragment header encoding must be infallible for constants"); + let header_bytes = encode_to_vec(header, config::standard()).unwrap_or_else(|err| { + panic!("fragment header encoding must be infallible for constants: {err}") + }); // Magic + length prefix (u16 big-endian) + encoded header. let overhead = FRAGMENT_MAGIC.len() + std::mem::size_of::() + header_bytes.len(); - NonZeroUsize::new(overhead).expect("fragment overhead must be non-zero") + NonZeroUsize::new(overhead).unwrap_or_else(|| { + panic!("fragment overhead must be non-zero (computed {overhead})"); + }) } /// Encode a fragment for transport by prefixing marker and header bytes. @@ -42,20 +50,13 @@ pub fn fragment_overhead() -> NonZeroUsize { /// # Errors /// /// Returns a [`bincode::error::EncodeError`] if the header cannot be encoded. -/// -/// # Panics -/// -/// Panics if the encoded header exceeds `u16::MAX` bytes, which should be -/// impossible for the fixed-size `FragmentHeader`. pub fn encode_fragment_payload( header: FragmentHeader, payload: &[u8], ) -> Result, bincode::error::EncodeError> { let header_bytes = encode_to_vec(header, config::standard())?; - let header_len: u16 = header_bytes - .len() - .try_into() - .expect("fragment header length must fit in u16"); + let header_len = u16::try_from(header_bytes.len()) + .map_err(|_| EncodeError::Other("fragment header length must fit within u16::MAX"))?; let mut buf = Vec::with_capacity( FRAGMENT_MAGIC.len() + std::mem::size_of::() + header_bytes.len() + payload.len(), @@ -80,37 +81,48 @@ pub fn encode_fragment_payload( pub fn decode_fragment_payload( payload: &[u8], ) -> Result, DecodeError> { - if payload.len() < FRAGMENT_MAGIC.len() + std::mem::size_of::() { + let minimum_len = FRAGMENT_MAGIC.len() + std::mem::size_of::(); + if payload.len() < minimum_len { return Ok(None); } - if &payload[..FRAGMENT_MAGIC.len()] != FRAGMENT_MAGIC { + let Some(prefix) = payload.get(..FRAGMENT_MAGIC.len()) else { + return Ok(None); + }; + if prefix != FRAGMENT_MAGIC { return Ok(None); } let header_len_offset = FRAGMENT_MAGIC.len(); - let len_bytes = [payload[header_len_offset], payload[header_len_offset + 1]]; + let len_hi = payload + .get(header_len_offset) + .copied() + .ok_or(DecodeError::UnexpectedEnd { additional: 0 })?; + let len_lo = payload + .get(header_len_offset + 1) + .copied() + .ok_or(DecodeError::UnexpectedEnd { additional: 0 })?; + let len_bytes = [len_hi, len_lo]; let header_len = u16::from_be_bytes(len_bytes) as usize; let header_start = header_len_offset + std::mem::size_of::(); let header_end = header_start + header_len; - if payload.len() < header_end { + let Some(header_bytes) = payload.get(header_start..header_end) else { return Err(DecodeError::UnexpectedEnd { - additional: header_end - payload.len(), + additional: header_end.saturating_sub(payload.len()), }); - } + }; - let (header, consumed) = borrow_decode_from_slice::( - &payload[header_start..header_end], - config::standard(), - )?; + let (header, consumed) = + borrow_decode_from_slice::(header_bytes, config::standard())?; if consumed != header_len { return Err(DecodeError::OtherString( "fragment header length mismatch".to_string(), )); } - Ok(Some((header, &payload[header_end..]))) + let remainder = payload.get(header_end..).unwrap_or_default(); + Ok(Some((header, remainder))) } #[cfg(test)] @@ -140,6 +152,16 @@ mod tests { ); } + #[test] + fn decode_returns_none_when_shorter_than_prefix_and_length() { + let payload = [b'F', b'R', b'A', b'G', 0]; + assert!( + decode_fragment_payload(&payload) + .expect("decode ok") + .is_none() + ); + } + #[test] fn fragment_overhead_matches_encoded_header() { let header = FragmentHeader::new(MessageId::new(1), FragmentIndex::zero(), true); @@ -149,48 +171,78 @@ mod tests { assert!(encoded.len() < u16::MAX as usize, "header must fit in u16"); } - #[test] - fn decode_fragment_payload_rejects_truncated_header() { - let header = FragmentHeader::new(MessageId::new(2), FragmentIndex::new(1), false); + /// Helper to test fragment decode errors with custom manipulation and assertions. + fn assert_fragment_decode_error(header: FragmentHeader, manipulate: F, assert_error: E) + where + F: FnOnce(Vec) -> (u16, Vec), // (advertised_len, header_bytes) + E: FnOnce(DecodeError), + { let encoded = encode_to_vec(header, config::standard()).expect("encode header"); + let (advertised_len, header_bytes) = manipulate(encoded); - // Advertise a longer header than provided to force `UnexpectedEnd`. - let advertised_len: u16 = (encoded.len() + 4) - .try_into() - .expect("encoded header length must stay within u16"); let mut payload = Vec::new(); payload.extend_from_slice(FRAGMENT_MAGIC); payload.extend_from_slice(&advertised_len.to_be_bytes()); - payload.extend_from_slice(&encoded); + payload.extend_from_slice(&header_bytes); let err = decode_fragment_payload(&payload).expect_err("expected decode failure"); - match err { - DecodeError::UnexpectedEnd { .. } => {} - other => panic!("expected UnexpectedEnd, got {other:?}"), - } + assert_error(err); } #[test] - fn decode_fragment_payload_rejects_length_mismatch() { - let header = FragmentHeader::new(MessageId::new(3), FragmentIndex::new(5), true); - let mut encoded = encode_to_vec(header, config::standard()).expect("encode header"); - encoded.extend_from_slice(&[0_u8, 1]); // pad so the advertised length exceeds consumed. - let advertised_len: u16 = encoded - .len() - .try_into() - .expect("padded header length must fit in u16"); + fn decode_fragment_payload_rejects_truncated_header() { + let header = FragmentHeader::new(MessageId::new(2), FragmentIndex::new(1), false); + assert_fragment_decode_error( + header, + |encoded| { + // Advertise a longer header than provided to force `UnexpectedEnd`. + let advertised_len: u16 = (encoded.len() + 4) + .try_into() + .expect("encoded header length must stay within u16"); + (advertised_len, encoded) + }, + |err| match err { + DecodeError::UnexpectedEnd { .. } => {} + other => panic!("expected UnexpectedEnd, got {other:?}"), + }, + ); + } + #[test] + fn decode_fragment_payload_rejects_missing_header_bytes() { + let advertised_len: u16 = 4; let mut payload = Vec::new(); payload.extend_from_slice(FRAGMENT_MAGIC); payload.extend_from_slice(&advertised_len.to_be_bytes()); - payload.extend_from_slice(&encoded); + // No header bytes provided. let err = decode_fragment_payload(&payload).expect_err("expected decode failure"); match err { - DecodeError::OtherString(msg) => { - assert_eq!(msg, "fragment header length mismatch"); - } - other => panic!("expected length mismatch error, got {other:?}"), + DecodeError::UnexpectedEnd { additional } => assert_eq!(additional, 4), + other => panic!("expected UnexpectedEnd, got {other:?}"), } } + + #[test] + fn decode_fragment_payload_rejects_length_mismatch() { + let header = FragmentHeader::new(MessageId::new(3), FragmentIndex::new(5), true); + assert_fragment_decode_error( + header, + |mut encoded| { + // Pad so the advertised length exceeds consumed. + encoded.extend_from_slice(&[0_u8, 1]); + let advertised_len: u16 = encoded + .len() + .try_into() + .expect("padded header length must fit in u16"); + (advertised_len, encoded) + }, + |err| match err { + DecodeError::OtherString(msg) => { + assert_eq!(msg, "fragment header length mismatch"); + } + other => panic!("expected length mismatch error, got {other:?}"), + }, + ); + } } diff --git a/src/fragment/reassembler.rs b/src/fragment/reassembler.rs index f5a5b326..1162114b 100644 --- a/src/fragment/reassembler.rs +++ b/src/fragment/reassembler.rs @@ -1,8 +1,8 @@ //! Inbound helper that stitches fragments back into complete messages. //! //! [`Reassembler`] mirrors the outbound [`Fragmenter`](crate::fragment::Fragmenter) by -//! collecting fragment payloads keyed by [`MessageId`](crate::fragment::MessageId). -//! It enforces ordering via [`FragmentSeries`](crate::fragment::FragmentSeries), guards +//! collecting fragment payloads keyed by [`MessageId`]. +//! It enforces ordering via [`FragmentSeries`], guards //! against unbounded allocation with a configurable cap, and purges stale partial //! assemblies after a fixed timeout. The helper is transport-agnostic so codecs and //! behavioural tests can reuse it without depending on socket types. @@ -154,14 +154,12 @@ impl Reassembler { Ok(FragmentStatus::Incomplete) => Self::append_and_maybe_complete( self.max_message_size, occupied, - header.message_id(), payload, false, ), Ok(FragmentStatus::Complete) => Self::append_and_maybe_complete( self.max_message_size, occupied, - header.message_id(), payload, true, ), @@ -237,10 +235,10 @@ impl Reassembler { fn append_and_maybe_complete( limit: NonZeroUsize, mut occupied: OccupiedEntry<'_, MessageId, PartialMessage>, - message_id: MessageId, payload: &[u8], completes: bool, ) -> Result, ReassemblyError> { + let message_id = *occupied.key(); let Some(attempted) = occupied.get().len().checked_add(payload.len()) else { occupied.remove(); return Err(ReassemblyError::MessageTooLarge { diff --git a/src/fragment/tests.rs b/src/fragment/tests.rs index c4e4b61a..2d20bbe3 100644 --- a/src/fragment/tests.rs +++ b/src/fragment/tests.rs @@ -7,6 +7,7 @@ use bincode::{BorrowDecode, Encode}; use rstest::rstest; use super::*; +use crate::fragment::fragmenter::FragmentCursor; fn setup_reassembler_with_first_fragment( message_id: u64, @@ -114,13 +115,9 @@ fn fragmenter_splits_payload_into_multiple_frames() { assert!(batch.is_fragmented()); assert_eq!(batch.message_id(), MessageId::new(0)); - let fragments = batch.fragments(); - assert_eq!(fragments[0].payload(), &[0, 1, 2]); - assert!(!fragments[0].header().is_last_fragment()); - assert_eq!(fragments[1].payload(), &[3, 4, 5]); - assert!(!fragments[1].header().is_last_fragment()); - assert_eq!(fragments[2].payload(), &[6, 7]); - assert!(fragments[2].header().is_last_fragment()); + assert_fragment(&batch, 0, &[0, 1, 2], false); + assert_fragment(&batch, 1, &[3, 4, 5], false); + assert_fragment(&batch, 2, &[6, 7], true); } #[test] @@ -130,7 +127,10 @@ fn fragmenter_handles_empty_payload() { assert_eq!(batch.len(), 1); assert!(!batch.is_fragmented()); - let fragment = &batch.fragments()[0]; + let fragment = batch + .fragments() + .first() + .expect("batch should contain at least one fragment"); assert_eq!(fragment.payload(), &[]); assert!(fragment.header().is_last_fragment()); assert_eq!(fragment.header().fragment_index(), FragmentIndex::zero()); @@ -139,6 +139,15 @@ fn fragmenter_handles_empty_payload() { #[derive(Debug, Encode, BorrowDecode)] struct DummyMessage(Vec); +fn assert_fragment(batch: &FragmentBatch, index: usize, payload: &[u8], is_last: bool) { + let fragment = batch + .fragments() + .get(index) + .expect("fragment missing at requested index"); + assert_eq!(fragment.payload(), payload); + assert_eq!(fragment.header().is_last_fragment(), is_last); +} + #[test] fn fragmenter_fragments_messages_and_increments_ids() { let fragmenter = @@ -190,6 +199,28 @@ fn fragmenter_respects_explicit_message_ids() { assert_eq!(next.message_id(), MessageId::new(10)); } +#[test] +fn fragmenter_returns_error_for_out_of_bounds_slice() { + let fragmenter = Fragmenter::new(NonZeroUsize::new(4).expect("non-zero")); + let payload = [1_u8, 2, 3, 4]; + + let err = fragmenter + .build_fragments_from_for_tests( + MessageId::new(1), + &payload, + FragmentCursor::new(payload.len() + 1, FragmentIndex::zero()), + ) + .expect_err("invalid slice should produce an error"); + match err { + FragmentationError::SliceBounds { offset, end, total } => { + assert_eq!(offset, payload.len() + 1); + assert_eq!(end, payload.len() + 1); + assert_eq!(total, payload.len()); + } + other => panic!("expected SliceBounds, got {other:?}"), + } +} + #[test] fn reassembler_allows_single_fragment_at_max_message_size() { let max_message_size = NonZeroUsize::new(16).expect("non-zero"); diff --git a/src/frame/conversion.rs b/src/frame/conversion.rs index a3aab575..50ce5ed5 100644 --- a/src/frame/conversion.rs +++ b/src/frame/conversion.rs @@ -37,9 +37,33 @@ pub fn bytes_to_u64(bytes: &[u8], size: usize, endianness: Endianness) -> io::Re } let mut buf = [0u8; 8]; + // NOTE: size is validated above; this is a defensive fallback. + let prefix = bytes + .get(..size) + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, ERR_INCOMPLETE_PREFIX))?; match endianness { - Endianness::Big => buf[8 - size..].copy_from_slice(&bytes[..size]), - Endianness::Little => buf[..size].copy_from_slice(&bytes[..size]), + Endianness::Big => { + if let Some(dst) = buf.get_mut(8 - size..) { + dst.copy_from_slice(prefix); + } else { + debug_assert!(false, "validated size should fit into prefix buffer"); + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + ERR_UNSUPPORTED_PREFIX, + )); + } + } + Endianness::Little => { + if let Some(dst) = buf.get_mut(..size) { + dst.copy_from_slice(prefix); + } else { + debug_assert!(false, "validated size should fit into prefix buffer"); + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + ERR_UNSUPPORTED_PREFIX, + )); + } + } } let val = match endianness { @@ -49,18 +73,53 @@ pub fn bytes_to_u64(bytes: &[u8], size: usize, endianness: Endianness) -> io::Re Ok(val) } +/// Convert a length value into a u64 based on the prefix size. +/// +/// Callers are expected to validate `size` against the supported set +/// `{1, 2, 4, 8}` before invoking this helper. +fn convert_len_to_value(len: usize, size: usize) -> io::Result { + let value = match size { + 1 => u64::from(checked_prefix_cast::(len)?), + 2 => u64::from(checked_prefix_cast::(len)?), + 4 => u64::from(checked_prefix_cast::(len)?), + 8 => checked_prefix_cast(len)?, + _ => { + debug_assert!(false, "size should be validated upstream"); + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + ERR_UNSUPPORTED_PREFIX, + )); + } + }; + Ok(value) +} + +/// Write a u64 into `prefix` according to the specified endianness. +fn write_bytes_with_endianness(value: u64, endianness: Endianness, prefix: &mut [u8]) { + let size = prefix.len(); + match endianness { + Endianness::Big => { + for (i, byte) in prefix.iter_mut().enumerate() { + let shift = 8 * (size - 1 - i); + *byte = ((value >> shift) & 0xff) as u8; + } + } + Endianness::Little => { + for (i, byte) in prefix.iter_mut().enumerate() { + let shift = 8 * i; + *byte = ((value >> shift) & 0xff) as u8; + } + } + } +} + /// Encodes an integer directly into `out` according to `size` and `endianness`. /// /// The function supports prefix sizes of `1`, `2`, `4`, or `8` bytes. /// /// # Errors -/// Returns [`io::ErrorKind::InvalidInput`] if the size is unsupported or if -/// `len` does not fit into the prefix. -/// -/// # Panics -/// Panics if the bit-shifting within the `write_bytes` closure leaves bits of -/// `value` outside the `u8` range. This cannot occur for valid prefix sizes and -/// checked values. +/// Returns [`io::ErrorKind::InvalidInput`] when the prefix size is unsupported +/// or when `len` does not fit into the requested prefix. #[must_use = "length prefix byte count must be used"] pub fn u64_to_bytes( len: usize, @@ -75,42 +134,19 @@ pub fn u64_to_bytes( )); } - let write_bytes = |value: u64, e: Endianness, size: usize, out: &mut [u8]| match e { - Endianness::Big => { - for (i, b) in out.iter_mut().enumerate().take(size) { - let shift = 8 * (size - 1 - i); - *b = u8::try_from((value >> shift) & 0xff).expect("masked < 256"); - } - } - Endianness::Little => { - for (i, b) in out.iter_mut().enumerate().take(size) { - let shift = 8 * i; - *b = u8::try_from((value >> shift) & 0xff).expect("masked < 256"); - } - } - }; + let value = convert_len_to_value(len, size)?; - match size { - 1 => { - let v: u8 = checked_prefix_cast(len)?; - write_bytes(u64::from(v), endianness, 1, &mut out[..1]); - } - 2 => { - let v: u16 = checked_prefix_cast(len)?; - write_bytes(u64::from(v), endianness, 2, &mut out[..2]); - } - 4 => { - let v: u32 = checked_prefix_cast(len)?; - write_bytes(u64::from(v), endianness, 4, &mut out[..4]); - } - 8 => { - let v: u64 = checked_prefix_cast(len)?; - write_bytes(v, endianness, 8, &mut out[..8]); - } - _ => unreachable!(), - } + #[expect( + clippy::indexing_slicing, + reason = "size validated to be within the 8-byte prefix buffer" + )] + let prefix = &mut out[..size]; - out[size..].fill(0); + write_bytes_with_endianness(value, endianness, prefix); + + if let Some(tail) = out.get_mut(size..) { + tail.fill(0); + } Ok(size) } diff --git a/src/frame/format.rs b/src/frame/format.rs index bc1467c2..d556b314 100644 --- a/src/frame/format.rs +++ b/src/frame/format.rs @@ -96,7 +96,14 @@ impl LengthFormat { pub fn write_len(&self, len: usize, dst: &mut BytesMut) -> io::Result<()> { let mut buf = [0u8; 8]; let written = u64_to_bytes(len, self.bytes, self.endianness, &mut buf)?; - dst.extend_from_slice(&buf[..written]); + let prefix = buf.get(..written).ok_or_else(|| { + debug_assert!(false, "written prefix length must never exceed buffer"); + io::Error::new( + io::ErrorKind::InvalidInput, + "internal: prefix slice exceeds buffer", + ) + })?; + dst.extend_from_slice(prefix); Ok(()) } } diff --git a/src/frame/tests.rs b/src/frame/tests.rs index 39aea809..c67b6678 100644 --- a/src/frame/tests.rs +++ b/src/frame/tests.rs @@ -47,14 +47,25 @@ fn u64_to_bytes_ok( let mut buf = [0u8; 8]; let written = u64_to_bytes(value, size, endianness, &mut buf).expect("failed to encode u64"); assert_eq!(written, size); - assert_eq!(&buf[..written], expected.as_slice()); + assert_eq!( + buf.get(..written) + .expect("written value must be within buffer bounds"), + expected.as_slice() + ); } #[rstest] #[case(vec![0x01], 2, Endianness::Big)] #[case(vec![0x02, 0x03], 4, Endianness::Little)] fn bytes_to_u64_short(#[case] bytes: Vec, #[case] size: usize, #[case] endianness: Endianness) { - let err = bytes_to_u64(&bytes, size, endianness).unwrap_err(); + let err = bytes_to_u64(&bytes, size, endianness) + .expect_err("short input must fail with UnexpectedEof"); + assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); +} + +#[test] +fn bytes_to_u64_rejects_empty_input() { + let err = bytes_to_u64(&[], 2, Endianness::Big).expect_err("empty slice must error"); assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof); } @@ -66,14 +77,16 @@ fn bytes_to_u64_unsupported( #[case] size: usize, #[case] endianness: Endianness, ) { - let err = bytes_to_u64(&bytes, size, endianness).unwrap_err(); + let err = bytes_to_u64(&bytes, size, endianness) + .expect_err("unsupported size must fail with InvalidInput"); assert_eq!(err.kind(), io::ErrorKind::InvalidInput); } #[rstest] fn u64_to_bytes_large() { let mut buf = [0u8; 8]; - let err = u64_to_bytes(300, 1, Endianness::Big, &mut buf).unwrap_err(); + let err = u64_to_bytes(300, 1, Endianness::Big, &mut buf) + .expect_err("value 300 must fail for 1-byte width (max 255)"); assert_eq!(err.kind(), io::ErrorKind::InvalidInput); } @@ -85,6 +98,14 @@ fn u64_to_bytes_zero_length() { assert_eq!(err.kind(), io::ErrorKind::InvalidInput); } +#[test] +fn u64_to_bytes_rejects_oversized_prefix() { + let mut buf = [0u8; 8]; + let err = u64_to_bytes(1, 9, Endianness::Big, &mut buf) + .expect_err("unsupported prefix size must fail"); + assert_eq!(err.kind(), io::ErrorKind::InvalidInput); +} + #[rstest] #[case(1usize, 3, Endianness::Big)] #[case(1usize, 3, Endianness::Little)] @@ -94,7 +115,8 @@ fn u64_to_bytes_unsupported( #[case] endianness: Endianness, ) { let mut buf = [0u8; 8]; - let err = u64_to_bytes(value, size, endianness, &mut buf).unwrap_err(); + let err = u64_to_bytes(value, size, endianness, &mut buf) + .expect_err("unsupported size must fail with InvalidInput"); assert_eq!(err.kind(), io::ErrorKind::InvalidInput); } @@ -107,6 +129,12 @@ fn u64_to_bytes_zeroes_remainder( #[case] endianness: Endianness, ) { let mut buf = [0xaau8; 8]; - u64_to_bytes(value, size, endianness, &mut buf).unwrap(); - assert!(buf[size..].iter().all(|&b| b == 0)); + u64_to_bytes(value, size, endianness, &mut buf) + .expect("conversion should succeed for valid size"); + assert!( + buf.get(size..) + .expect("size must be within buffer bounds") + .iter() + .all(|&b| b == 0) + ); } diff --git a/src/main.rs b/src/main.rs deleted file mode 100644 index c73477d4..00000000 --- a/src/main.rs +++ /dev/null @@ -1,10 +0,0 @@ -//! Minimal binary demonstrating `wireframe` usage. -//! -//! Currently prints a greeting and exits. - -fn main() { - // Enable structured logging for examples and integration tests. - // Applications embedding the library should install their own subscriber. - tracing_subscriber::fmt::init(); - println!("Hello from Wireframe!"); -} diff --git a/src/preamble.rs b/src/preamble.rs index 7600e462..66fbaaf9 100644 --- a/src/preamble.rs +++ b/src/preamble.rs @@ -11,9 +11,10 @@ const MAX_PREAMBLE_LEN: usize = 1024; /// Trait bound for types accepted as connection preambles. /// /// The bound allows decoding borrowed data for any lifetime without -/// requiring an external decoding context. -pub trait Preamble: for<'de> BorrowDecode<'de, ()> + Send + 'static {} -impl Preamble for T where for<'de> T: BorrowDecode<'de, ()> + Send + 'static {} +/// requiring an external decoding context. `Sync` is required because +/// preamble values are shared by reference with asynchronous handlers. +pub trait Preamble: for<'de> BorrowDecode<'de, ()> + Send + Sync + 'static {} +impl Preamble for T where for<'de> T: BorrowDecode<'de, ()> + Send + Sync + 'static {} async fn read_more( reader: &mut R, @@ -30,10 +31,12 @@ where buf.resize(start + additional, 0); let mut read = 0; while read < additional { - match reader - .read(&mut buf[start + read..start + additional]) - .await - { + let range_start = start + read; + let range_end = start + additional; + let chunk = buf + .get_mut(range_start..range_end) + .ok_or(DecodeError::Other("preamble buffer range invalid"))?; + match reader.read(chunk).await { Ok(0) => { return Err(DecodeError::Io { inner: io::Error::from(io::ErrorKind::UnexpectedEof), diff --git a/src/push/queues/builder.rs b/src/push/queues/builder.rs index 3694ea70..f8d1d2bf 100644 --- a/src/push/queues/builder.rs +++ b/src/push/queues/builder.rs @@ -18,7 +18,7 @@ use super::{ /// Allows configuration of queue capacities, rate limiting and an optional /// dead-letter queue before constructing [`PushQueues`] and its paired /// [`PushHandle`]. Defaults mirror the previous constructors: both queues have -/// a capacity of one and pushes are limited to [`DEFAULT_PUSH_RATE`] per +/// a capacity of one and pushes are limited to the default push rate per /// second unless overridden. Construct via [`PushQueues::builder`] or /// [`Default::default`]. /// diff --git a/src/push/queues/handle.rs b/src/push/queues/handle.rs index dfa398a5..9412759b 100644 --- a/src/push/queues/handle.rs +++ b/src/push/queues/handle.rs @@ -105,19 +105,7 @@ impl PushHandle { // the next refill window, tokens remain available to actively polled // tasks. if let Some(ref limiter) = self.0.limiter { - loop { - // Prefer a non-blocking acquisition. If not available, back - // off briefly before trying again. We intentionally do not - // poll the limiter's async acquire future to avoid enqueuing - // this task as a waiter and reserving a token prematurely. - if limiter.try_acquire(1) { - break; - } - // The limiter is configured with a 1s refill interval; a - // short sleep yields to the scheduler and advances virtual - // time in tests (tokio::time::pause/advance). - sleep(Duration::from_millis(10)).await; - } + self.wait_for_permit(limiter).await; } // Then send the frame, awaiting capacity if the queue is currently @@ -197,31 +185,59 @@ impl PushHandle { /// Send a frame to the configured dead letter queue if available. fn route_to_dlq(&self, frame: F) + where + F: std::fmt::Debug, + { + if let Some(dlq) = &self.0.dlq_tx + && let Err(mpsc::error::TrySendError::Full(f) | mpsc::error::TrySendError::Closed(f)) = + dlq.try_send(frame) + { + let dropped = self.0.dlq_drops.fetch_add(1, Ordering::Relaxed) + 1; + let mut last = match self.0.dlq_last_log.lock() { + Ok(guard) => guard, + Err(poisoned) => { + warn!("DLQ last-log mutex poisoned; continuing with stale state"); + poisoned.into_inner() + } + }; + self.log_dlq_drop(&f, dropped, &mut last); + } + } + + /// Interval between attempts to acquire a rate-limit permit. + /// + /// Kept short so tests that use `tokio::time::pause/advance` progress + /// quickly while remaining negligible relative to the 1s refill window. + const PERMIT_POLL_INTERVAL: Duration = Duration::from_millis(10); + + async fn wait_for_permit(&self, limiter: &RateLimiter) { + loop { + if limiter.try_acquire(1) { + break; + } + // The limiter is configured with a 1s refill interval; a short + // sleep yields to the scheduler and advances virtual time in + // tests (tokio::time::pause/advance). + sleep(Self::PERMIT_POLL_INTERVAL).await; + } + } + + fn log_dlq_drop(&self, frame: &F, dropped: usize, last_log: &mut Instant) where F: std::fmt::Debug, { let log_every_n = self.0.dlq_log_every_n; let log_interval = self.0.dlq_log_interval; + let should_log = (log_every_n != 0 && dropped.is_multiple_of(log_every_n)) + || last_log.elapsed() > log_interval; - if let Some(dlq) = &self.0.dlq_tx { - match dlq.try_send(frame) { - Ok(()) => {} - Err(mpsc::error::TrySendError::Full(f) | mpsc::error::TrySendError::Closed(f)) => { - let dropped = self.0.dlq_drops.fetch_add(1, Ordering::Relaxed) + 1; - let mut last = self.0.dlq_last_log.lock().expect("lock poisoned"); - let now = Instant::now(); - if (log_every_n != 0 && dropped.is_multiple_of(log_every_n)) - || now.duration_since(*last) > log_interval - { - warn!( - "DLQ dropped frames (full or closed): frame={f:?}, dropped={dropped}, \ - log_every_n={log_every_n}, log_interval={log_interval:?}" - ); - *last = now; - self.0.dlq_drops.store(0, Ordering::Relaxed); - } - } - } + if should_log { + warn!( + "DLQ dropped frames (full or closed): frame={frame:?}, dropped={dropped}, \ + log_every_n={log_every_n}, log_interval={log_interval:?}" + ); + *last_log = Instant::now(); + self.0.dlq_drops.store(0, Ordering::Relaxed); } } diff --git a/src/push/queues/mod.rs b/src/push/queues/mod.rs index 9636b200..eafff53a 100644 --- a/src/push/queues/mod.rs +++ b/src/push/queues/mod.rs @@ -114,9 +114,12 @@ impl PushQueues { #[must_use] pub fn builder() -> PushQueuesBuilder { PushQueuesBuilder::default() } - /// Validates whether the provided rate is invalid (zero or exceeds the maximum). - fn is_invalid_rate(rate: Option) -> bool { - matches!(rate, Some(r) if r == 0 || r > MAX_PUSH_RATE) + /// Returns the invalid rate if it is zero or exceeds the maximum. + fn invalid_rate(rate: Option) -> Option { + match rate { + Some(r) if r == 0 || r > MAX_PUSH_RATE => Some(r), + _ => None, + } } pub(super) fn build_with_config( @@ -130,11 +133,10 @@ impl PushQueues { dlq_log_every_n, dlq_log_interval, } = config; - if Self::is_invalid_rate(rate) { + if let Some(invalid) = Self::invalid_rate(rate) { // Reject unsupported rates early to avoid building queues that cannot // be used. The bounds prevent runaway resource consumption. - let r = rate.unwrap(); - return Err(PushConfigError::InvalidRate(r)); + return Err(PushConfigError::InvalidRate(invalid)); } if high_capacity == 0 || low_capacity == 0 { return Err(PushConfigError::InvalidCapacity { @@ -192,31 +194,30 @@ impl PushQueues { /// Create a new set of queues with the specified bounds for each priority /// and return them along with a [`PushHandle`] for producers. /// - /// # Panics + /// # Errors /// - /// Panics if either queue capacity is zero. Prefer `PushQueues::builder()` - /// to receive a [`Result`] instead. + /// Returns [`PushConfigError::InvalidCapacity`] if either queue capacity is + /// zero or [`PushConfigError::InvalidRate`] if the default rate is invalid. #[deprecated(since = "0.1.0", note = "Use `PushQueues::builder` instead")] - #[must_use] - pub fn bounded(high_capacity: usize, low_capacity: usize) -> (Self, PushHandle) { + pub fn bounded( + high_capacity: usize, + low_capacity: usize, + ) -> Result<(Self, PushHandle), PushConfigError> { Self::build_via_builder(high_capacity, low_capacity, Some(DEFAULT_PUSH_RATE), None) - .expect("invalid capacities or rate in deprecated bounded()") } /// Create queues with no rate limiting. /// - /// # Panics + /// # Errors /// - /// Panics if either queue capacity is zero. Prefer `PushQueues::builder()` - /// to receive a [`Result`] instead. + /// Returns [`PushConfigError::InvalidCapacity`] if either queue capacity is + /// zero. #[deprecated(since = "0.1.0", note = "Use `PushQueues::builder` instead")] - #[must_use] pub fn bounded_no_rate_limit( high_capacity: usize, low_capacity: usize, - ) -> (Self, PushHandle) { + ) -> Result<(Self, PushHandle), PushConfigError> { Self::build_via_builder(high_capacity, low_capacity, None, None) - .expect("invalid capacities in deprecated bounded_no_rate_limit()") } /// Create queues with a custom rate limit in pushes per second. @@ -276,6 +277,10 @@ impl PushQueues { /// assert_eq!(frame, 2); /// } /// ``` + #[expect( + clippy::integer_division_remainder_used, + reason = "tokio::select! expands to modulus internally" + )] pub async fn recv(&mut self) -> Option<(PushPriority, F)> { let mut high_closed = false; let mut low_closed = false; diff --git a/src/rewind_stream.rs b/src/rewind_stream.rs index d40f4791..1831b7db 100644 --- a/src/rewind_stream.rs +++ b/src/rewind_stream.rs @@ -31,18 +31,32 @@ impl RewindStream { } } +#[cfg(test)] +impl RewindStream { + pub(crate) fn set_pos_for_tests(&mut self, pos: usize) { self.pos = pos; } + + pub(crate) fn leftover_len_for_tests(&self) -> usize { self.leftover.len() } +} + impl AsyncRead for RewindStream { fn poll_read( mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>, ) -> Poll> { - if self.pos < self.leftover.len() { - let remaining = self.leftover.len() - self.pos; + if self.pos != self.leftover.len() { + let remaining = self.leftover.len().saturating_sub(self.pos); let to_copy = remaining.min(buf.remaining()); let start = self.pos; let end = start + to_copy; - buf.put_slice(&self.leftover[start..end]); + if let Some(slice) = self.leftover.get(start..end) { + buf.put_slice(slice); + } else { + return Poll::Ready(Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "rewind buffer slice out of bounds", + ))); + } self.pos += to_copy; if self.pos < self.leftover.len() || to_copy > 0 { return Poll::Ready(Ok(())); @@ -77,3 +91,32 @@ impl AsyncWrite for RewindStream { } impl Unpin for RewindStream {} + +#[cfg(test)] +mod tests { + use std::{pin::Pin, task::Context}; + + use futures::task::noop_waker_ref; + use tokio::io::{AsyncRead, ReadBuf}; + + use super::*; + + #[test] + fn poll_read_returns_error_for_invalid_leftover_slice_bounds() { + let mut stream = RewindStream::new(vec![1_u8, 2, 3], tokio::io::empty()); + stream.set_pos_for_tests(stream.leftover_len_for_tests() + 1); + + let waker = noop_waker_ref(); + let mut cx = Context::from_waker(waker); + let mut buffer = [0_u8; 2]; + let mut read_buf = ReadBuf::new(&mut buffer); + + let mut pinned = Pin::new(&mut stream); + let result = RewindStream::poll_read(Pin::as_mut(&mut pinned), &mut cx, &mut read_buf); + + match result { + Poll::Ready(Err(err)) => assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof), + other => panic!("expected UnexpectedEof, got {other:?}"), + } + } +} diff --git a/src/server/config/binding.rs b/src/server/config/binding.rs index 2231a7f1..38c07cb2 100644 --- a/src/server/config/binding.rs +++ b/src/server/config/binding.rs @@ -23,7 +23,10 @@ trait WireframePreamble: Preamble {} impl WireframePreamble for T where T: Preamble {} /// Blanket impl uses private trait aliases; suppress visibility lint -#[allow(private_bounds)] +#[expect( + private_bounds, + reason = "helper trait aliases are module-private by design" +)] impl WireframeServer where F: WireframeFactory, @@ -68,7 +71,10 @@ where } /// Blanket impl uses private trait aliases; suppress visibility lint -#[allow(private_bounds)] +#[expect( + private_bounds, + reason = "helper trait aliases are module-private by design" +)] impl WireframeServer where F: WireframeFactory, @@ -140,7 +146,10 @@ where } /// Blanket impl uses private trait aliases; suppress visibility lint -#[allow(private_bounds)] +#[expect( + private_bounds, + reason = "helper trait aliases are module-private by design" +)] impl WireframeServer where F: WireframeFactory, diff --git a/src/server/connection.rs b/src/server/connection.rs index b7b0a279..7ea4fc19 100644 --- a/src/server/connection.rs +++ b/src/server/connection.rs @@ -7,12 +7,11 @@ use log::{error, warn}; use tokio::{net::TcpStream, time::timeout}; use tokio_util::task::TaskTracker; -use super::{PreambleFailure, PreambleHandler}; use crate::{ app::WireframeApp, preamble::{Preamble, read_preamble}, rewind_stream::RewindStream, - server::runtime::PreambleHooks, + server::{PreambleFailure, PreambleHandler, runtime::PreambleHooks}, }; /// Spawn a task to process a single TCP connection, logging and discarding any panics. @@ -33,20 +32,8 @@ pub(super) fn spawn_connection_task( } }; tracker.spawn(async move { - let PreambleHooks { - on_success, - on_failure, - timeout: preamble_timeout, - } = hooks; - let fut = std::panic::AssertUnwindSafe(process_stream( - stream, - peer_addr, - factory, - on_success, - on_failure, - preamble_timeout, - )) - .catch_unwind(); + let fut = std::panic::AssertUnwindSafe(process_stream(stream, peer_addr, factory, hooks)) + .catch_unwind(); if let Err(panic) = fut.await { crate::metrics::inc_connection_panics(); @@ -62,48 +49,27 @@ async fn process_stream( mut stream: TcpStream, peer_addr: Option, factory: F, - on_success: Option>, - on_failure: Option, - preamble_timeout: Option, + hooks: PreambleHooks, ) where F: Fn() -> WireframeApp + Send + Sync + 'static, T: Preamble, { - let preamble_result = match preamble_timeout { - Some(limit) => match timeout(limit, read_preamble::<_, T>(&mut stream)).await { - Ok(result) => result, - Err(_) => Err(timeout_error()), - }, - None => read_preamble::<_, T>(&mut stream).await, - }; + let PreambleHooks { + on_success, + on_failure, + timeout: preamble_timeout, + } = hooks; - match preamble_result { + match read_preamble_with_timeout::(&mut stream, preamble_timeout).await { Ok((preamble, leftover)) => { - if let Some(handler) = on_success.as_ref() - && let Err(e) = handler(&preamble, &mut stream).await - { - error!( - "preamble handler error: error={e}, error_debug={e:?}, peer_addr={peer_addr:?}" - ); - } + run_preamble_success(on_success.as_ref(), &preamble, &mut stream, peer_addr).await; let stream = RewindStream::new(leftover, stream); - let app = (factory)(); - app.handle_connection(stream).await; + if let Err(e) = (factory)().handle_connection_result(stream).await { + warn!("connection task error: {e:?}"); + } } Err(err) => { - if let Some(handler) = on_failure.as_ref() { - if let Err(e) = handler(&err, &mut stream).await { - error!( - "preamble failure handler error: error={e}, error_debug={e:?}, \ - peer_addr={peer_addr:?}" - ); - } - } else { - error!( - "preamble decode failed and no failure handler set: error={err:?}, \ - peer_addr={peer_addr:?}" - ); - } + run_preamble_failure(on_failure.as_ref(), err, &mut stream, peer_addr).await; } } } @@ -115,6 +81,53 @@ fn timeout_error() -> bincode::error::DecodeError { } } +async fn read_preamble_with_timeout( + stream: &mut TcpStream, + preamble_timeout: Option, +) -> Result<(T, Vec), bincode::error::DecodeError> { + match preamble_timeout { + Some(limit) => match timeout(limit, read_preamble::<_, T>(stream)).await { + Ok(result) => result, + Err(_) => Err(timeout_error()), + }, + None => read_preamble::<_, T>(stream).await, + } +} + +async fn run_preamble_success( + handler: Option<&PreambleHandler>, + preamble: &T, + stream: &mut TcpStream, + peer_addr: Option, +) { + if let Some(handler) = handler + && let Err(e) = handler(preamble, stream).await + { + error!("preamble handler error: error={e}, error_debug={e:?}, peer_addr={peer_addr:?}"); + } +} + +async fn run_preamble_failure( + handler: Option<&PreambleFailure>, + err: bincode::error::DecodeError, + stream: &mut TcpStream, + peer_addr: Option, +) { + if let Some(handler) = handler { + if let Err(e) = handler(&err, stream).await { + error!( + "preamble failure handler error: error={e}, error_debug={e:?}, \ + peer_addr={peer_addr:?}" + ); + } + } else { + error!( + "preamble decode failed and no failure handler set: error={err:?}, \ + peer_addr={peer_addr:?}" + ); + } +} + #[cfg(test)] mod tests { use rstest::rstest; @@ -144,7 +157,7 @@ mod tests { let app_factory = move || { factory() .on_connection_setup(|| async { panic!("boom") }) - .unwrap() + .expect("failed to install panic setup callback") }; let tracker = TaskTracker::new(); let listener = TcpListener::bind("127.0.0.1:0") diff --git a/src/server/runtime.rs b/src/server/runtime.rs index 9d815425..5d613ca0 100644 --- a/src/server/runtime.rs +++ b/src/server/runtime.rs @@ -1,6 +1,6 @@ //! Runtime control for [`WireframeServer`]. -use std::{io, net::SocketAddr, sync::Arc}; +use std::{fmt, io, net::SocketAddr, sync::Arc}; use async_trait::async_trait; use futures::Future; @@ -82,6 +82,21 @@ impl BackoffConfig { } } +#[derive(Debug)] +pub(super) struct AcceptLoopOptions { + pub preamble: PreambleHooks, + pub shutdown: CancellationToken, + pub tracker: TaskTracker, + pub backoff: BackoffConfig, +} + +struct AcceptHandles<'a, T> { + preamble: &'a PreambleHooks, + shutdown: &'a CancellationToken, + tracker: &'a TaskTracker, + backoff: &'a BackoffConfig, +} + #[derive(Default)] pub(super) struct PreambleHooks { pub on_success: Option>, @@ -99,6 +114,22 @@ impl Clone for PreambleHooks { } } +impl fmt::Debug for PreambleHooks { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("PreambleHooks") + .field( + "on_success", + &self.on_success.as_ref().map(|_| "Some()"), + ) + .field( + "on_failure", + &self.on_failure.as_ref().map(|_| "Some()"), + ) + .field("timeout", &self.timeout) + .finish() + } +} + impl WireframeServer where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, @@ -200,6 +231,10 @@ where /// Returns an [`io::Error`] if the server was not bound to a listener. /// Accept failures are retried with exponential back-off and do not /// surface as errors. + #[expect( + clippy::integer_division_remainder_used, + reason = "tokio::select! expands to modulus internally" + )] pub async fn run_with_shutdown(self, shutdown: S) -> Result<(), ServerError> where S: Future + Send, @@ -232,10 +267,12 @@ where tracker.spawn(accept_loop( listener, factory, - preamble_hooks, - token, - t, - backoff_config, + AcceptLoopOptions { + preamble: preamble_hooks, + shutdown: token, + tracker: t, + backoff: backoff_config, + }, )); } @@ -293,10 +330,12 @@ where /// accept_loop::<_, (), _>( /// listener, /// || WireframeApp::default(), -/// PreambleHooks::default(), -/// token, -/// tracker, -/// BackoffConfig::default(), +/// AcceptLoopOptions { +/// preamble: PreambleHooks::default(), +/// shutdown: token, +/// tracker, +/// backoff: BackoffConfig::default(), +/// }, /// ) /// .await; /// } @@ -304,49 +343,74 @@ where pub(super) async fn accept_loop( listener: Arc, factory: F, - preamble: PreambleHooks, - shutdown: CancellationToken, - tracker: TaskTracker, - backoff_config: BackoffConfig, + options: AcceptLoopOptions, ) where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, T: Preamble, L: AcceptListener + Send + Sync + 'static, { + let AcceptLoopOptions { + preamble, + shutdown, + tracker, + backoff, + } = options; debug_assert!( - backoff_config.initial_delay <= backoff_config.max_delay, + backoff.initial_delay <= backoff.max_delay, "BackoffConfig invariant violated: initial_delay > max_delay" ); debug_assert!( - backoff_config.initial_delay >= Duration::from_millis(1), + backoff.initial_delay >= Duration::from_millis(1), "BackoffConfig invariant violated: initial_delay < 1ms" ); - let mut delay = backoff_config.initial_delay; - loop { - select! { - biased; - - () = shutdown.cancelled() => break, - - res = listener.accept() => match res { - Ok((stream, _)) => { - let hooks = preamble.clone(); - spawn_connection_task( - stream, - factory.clone(), - hooks, - &tracker, - ); - delay = backoff_config.initial_delay; - } - Err(e) => { - let local_addr = listener.local_addr().ok(); - warn!("accept error: error={e:?}, local_addr={local_addr:?}"); - sleep(delay).await; - delay = (delay * 2).min(backoff_config.max_delay); - } - }, - } + let mut delay = backoff.initial_delay; + let handles = AcceptHandles { + preamble: &preamble, + shutdown: &shutdown, + tracker: &tracker, + backoff: &backoff, + }; + while let Some(next_delay) = accept_iteration(&listener, &factory, &handles, delay).await { + delay = next_delay; + } +} + +#[expect( + clippy::integer_division_remainder_used, + reason = "tokio::select! expands to modulus internally" +)] +async fn accept_iteration( + listener: &Arc, + factory: &F, + handles: &AcceptHandles<'_, T>, + delay: Duration, +) -> Option +where + F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, + T: Preamble, + L: AcceptListener + Send + Sync + 'static, +{ + select! { + biased; + + () = handles.shutdown.cancelled() => None, + res = listener.accept() => Some(match res { + Ok((stream, _)) => { + spawn_connection_task( + stream, + (*factory).clone(), + handles.preamble.clone(), + handles.tracker, + ); + handles.backoff.initial_delay + } + Err(e) => { + let local_addr = listener.local_addr().ok(); + warn!("accept error: error={e:?}, local_addr={local_addr:?}"); + sleep(delay).await; + (delay * 2).min(handles.backoff.max_delay) + } + }), } } @@ -444,10 +508,12 @@ mod tests { tracker.spawn(accept_loop::<_, (), _>( listener, factory, - PreambleHooks::default(), - token.clone(), - tracker.clone(), - BackoffConfig::default(), + AcceptLoopOptions { + preamble: PreambleHooks::default(), + shutdown: token.clone(), + tracker: tracker.clone(), + backoff: BackoffConfig::default(), + }, )); token.cancel(); @@ -457,14 +523,13 @@ mod tests { assert!(result.is_ok()); } - #[rstest] - #[tokio::test(start_paused = true)] - async fn test_accept_loop_exponential_backoff_async( - factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, - ) { - let calls = Arc::new(Mutex::new(Vec::new())); + /// Creates a mock listener that fails with exponential backoff tracking. + fn setup_backoff_mock_listener( + calls: &Arc>>, + num_calls: usize, + ) -> MockAcceptListener { let mut listener = MockAcceptListener::new(); - let call_log = calls.clone(); + let call_log = Arc::clone(calls); listener .expect_accept() .returning(move || { @@ -474,12 +539,46 @@ mod tests { Err(io::Error::other("mock error")) }) }) - .times(4); + .times(num_calls); listener .expect_local_addr() .returning(|| Ok("127.0.0.1:0".parse().expect("addr parse"))) - .times(4); - let listener = Arc::new(listener); + .times(num_calls); + listener + } + + /// Validates that recorded call intervals match expected backoff delays. + fn assert_backoff_intervals(calls: &[Instant], expected: &[Duration]) { + let intervals: Vec<_> = calls + .windows(2) + .map(|w| { + let a = w.first().expect("window has first element"); + let b = w.get(1).expect("window has second element"); + b.checked_duration_since(*a) + .expect("instants should be monotonically increasing") + }) + .collect(); + + assert_eq!( + intervals.len(), + expected.len(), + "interval count mismatch: got {}, expected {}", + intervals.len(), + expected.len() + ); + + for (interval, expected) in intervals.into_iter().zip(expected.iter()) { + assert_eq!(interval, *expected); + } + } + + #[rstest] + #[tokio::test(start_paused = true)] + async fn test_accept_loop_exponential_backoff_async( + factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, + ) { + let calls = Arc::new(Mutex::new(Vec::new())); + let listener = Arc::new(setup_backoff_mock_listener(&calls, 4)); let token = CancellationToken::new(); let tracker = TaskTracker::new(); let backoff = BackoffConfig { @@ -490,10 +589,12 @@ mod tests { tracker.spawn(accept_loop::<_, (), _>( listener, factory, - PreambleHooks::default(), - token.clone(), - tracker.clone(), - backoff, + AcceptLoopOptions { + preamble: PreambleHooks::default(), + shutdown: token.clone(), + tracker: tracker.clone(), + backoff, + }, )); yield_now().await; @@ -501,7 +602,7 @@ mod tests { let first_call = { let calls = calls.lock().expect("lock"); assert_eq!(calls.len(), 1); - calls[0] + calls.first().copied().expect("call record missing") }; for ms in [5, 10, 20] { @@ -517,15 +618,13 @@ mod tests { let calls = calls.lock().expect("lock"); assert_eq!(calls.len(), 4); - assert_eq!(calls[0], first_call); - let intervals: Vec<_> = calls.windows(2).map(|w| w[1] - w[0]).collect(); + let first = calls.first().copied().expect("at least one call logged"); + assert_eq!(first, first_call); let expected = [ Duration::from_millis(5), Duration::from_millis(10), Duration::from_millis(20), ]; - for (interval, expected) in intervals.into_iter().zip(expected) { - assert_eq!(interval, expected); - } + assert_backoff_intervals(&calls, &expected); } } diff --git a/src/server/test_util.rs b/src/server/test_util.rs index 98d1ea60..eb7e34fa 100644 --- a/src/server/test_util.rs +++ b/src/server/test_util.rs @@ -8,11 +8,6 @@ use rstest::fixture; use super::{Bound, WireframeServer}; use crate::app::WireframeApp; -#[cfg_attr( - not(test), - expect(dead_code, reason = "Used in builder tests via fixtures") -)] -#[cfg_attr(test, allow(dead_code, reason = "Used in builder tests via fixtures"))] #[derive(Debug, Clone, PartialEq, Encode, Decode)] pub struct TestPreamble { pub id: u32, diff --git a/tests/app_data.rs b/tests/app_data.rs index bc7b6811..62300eef 100644 --- a/tests/app_data.rs +++ b/tests/app_data.rs @@ -12,14 +12,30 @@ use wireframe::extractor::{ SharedState, }; +#[expect( + clippy::allow_attributes, + reason = "rstest single-line fixtures need allow to avoid unfulfilled lint expectations" +)] #[allow( + unfulfilled_lint_expectations, + reason = "rstest occasionally misses the expected lint for single-line fixtures on stable" +)] +#[expect( unused_braces, reason = "rustc false positive for single line rstest fixtures" )] #[fixture] fn request() -> MessageRequest { MessageRequest::default() } +#[expect( + clippy::allow_attributes, + reason = "rstest single-line fixtures need allow to avoid unfulfilled lint expectations" +)] #[allow( + unfulfilled_lint_expectations, + reason = "rstest occasionally misses the expected lint for single-line fixtures on stable" +)] +#[expect( unused_braces, reason = "rustc false positive for single line rstest fixtures" )] diff --git a/tests/common/mod.rs b/tests/common/mod.rs index 16c83285..2ad0fa0f 100644 --- a/tests/common/mod.rs +++ b/tests/common/mod.rs @@ -7,8 +7,10 @@ use std::net::{Ipv4Addr, SocketAddr, TcpListener as StdTcpListener}; /// Create a TCP listener bound to a free local port. -#[expect(dead_code, reason = "Used by tests that bind to random ports")] -#[allow(unfulfilled_lint_expectations)] +#[expect( + clippy::expect_used, + reason = "binding to an ephemeral localhost port must abort the test immediately" +)] pub fn unused_listener() -> StdTcpListener { let addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0); StdTcpListener::bind(addr).expect("failed to bind port") @@ -18,13 +20,16 @@ use rstest::fixture; use wireframe::{app::Envelope, serializer::BincodeSerializer}; pub type TestApp = wireframe::app::WireframeApp; +pub type TestResult = Result>; #[fixture] -#[expect( - unused_braces, - reason = "rustc false positive for single line rstest fixtures" -)] -#[allow(unfulfilled_lint_expectations)] pub fn factory() -> impl Fn() -> TestApp + Send + Sync + Clone + 'static { - || TestApp::new().expect("TestApp::new failed") + fn build() -> TestApp { TestApp::default() } + build +} + +#[cfg(test)] +mod tests { + #[test] + fn unused_listener_is_callable() { let _ = super::unused_listener(); } } diff --git a/tests/connection.rs b/tests/connection.rs index b15792e3..4426acaa 100644 --- a/tests/connection.rs +++ b/tests/connection.rs @@ -17,6 +17,9 @@ use wireframe::{ }; use wireframe_testing::{LoggerHandle, logger}; +mod common; +use common::TestResult; + #[derive(Clone, Copy, Debug, PartialEq, Eq)] struct HookCounts { before: usize, @@ -140,7 +143,10 @@ struct HarnessFactory { impl HarnessFactory { /// Build a connection actor harness with shared hook counters. - fn create(&self, config: HarnessConfig) -> ActorHarness { + fn create( + &self, + config: HarnessConfig, + ) -> Result { let HarnessConfig { has_response, has_multi_packet, @@ -153,7 +159,6 @@ impl HarnessFactory { .counters .build_hooks_with_increment(increment, stream_end_fn); ActorHarness::new_with_state(hooks, has_response, has_multi_packet) - .expect("failed to create harness") } /// Read the accumulated hook counters for the most recent harness. @@ -233,8 +238,8 @@ fn assert_reason_logged( } #[rstest] -fn process_multi_packet_forwards_frame(harness_factory: HarnessFactory) { - let mut harness = harness_factory.create(HarnessConfig::new()); +fn process_multi_packet_forwards_frame(harness_factory: HarnessFactory) -> TestResult { + let mut harness = harness_factory.create(HarnessConfig::new())?; harness.process_multi_packet(Some(5)); assert_multi_packet_processing_result( @@ -243,16 +248,17 @@ fn process_multi_packet_forwards_frame(harness_factory: HarnessFactory) { &[6], HookCounts { before: 1, end: 0 }, ); + Ok(()) } #[rstest] -fn process_multi_packet_none_emits_end_frame(harness_factory: HarnessFactory) { +fn process_multi_packet_none_emits_end_frame(harness_factory: HarnessFactory) -> TestResult { let mut harness = harness_factory.create( HarnessConfig::new() .with_multi_packet() .with_increment(2) .with_stream_end(|_| Some(9)), - ); + )?; let (_tx, rx) = mpsc::channel(1); harness.set_multi_queue(Some(rx)); @@ -264,6 +270,7 @@ fn process_multi_packet_none_emits_end_frame(harness_factory: HarnessFactory) { &[11], HookCounts { before: 1, end: 1 }, ); + Ok(()) } #[rstest( @@ -278,23 +285,30 @@ fn handle_multi_packet_closed_behaviour( terminator: Option, expected_output: Vec, expected_before: usize, -) { +) -> TestResult { let mut harness = harness_factory.create( HarnessConfig::new() .with_multi_packet() .with_stream_end(move |_| terminator), - ); + )?; let (_tx, rx) = mpsc::channel(1); harness.set_multi_queue(Some(rx)); harness.handle_multi_packet_closed(); let snapshot = harness.snapshot(); - assert!(snapshot.is_active && !snapshot.is_shutting_down && !snapshot.is_done); - assert!( - !harness.has_multi_queue(), - "multi-packet channel should be cleared", - ); + if !snapshot.is_active { + return Err("connection should be active".into()); + } + if snapshot.is_shutting_down { + return Err("connection should not be shutting down".into()); + } + if snapshot.is_done { + return Err("connection should not be done".into()); + } + if harness.has_multi_queue() { + return Err("multi-packet channel should be cleared".into()); + } assert_frame_processed( &harness.out, &expected_output, @@ -304,26 +318,32 @@ fn handle_multi_packet_closed_behaviour( }, harness_factory.counts(), ); + Ok(()) } #[rstest] -fn try_opportunistic_drain_forwards_frame(harness_factory: HarnessFactory) { - let mut harness = harness_factory.create(HarnessConfig::new()); +fn try_opportunistic_drain_forwards_frame(harness_factory: HarnessFactory) -> TestResult { + let mut harness = harness_factory.create(HarnessConfig::new())?; let (tx, rx) = mpsc::channel(1); - tx.try_send(9).expect("send frame"); + tx.try_send(9)?; drop(tx); harness.set_low_queue(Some(rx)); let drained = harness.try_drain_low(); - assert!(drained, "queue should report a drained frame"); - assert!(harness.has_low_queue(), "queue remains available"); + if !drained { + return Err("queue should report a drained frame".into()); + } + if !harness.has_low_queue() { + return Err("queue remains available".into()); + } assert_frame_processed( &harness.out, &[10], HookCounts { before: 1, end: 0 }, harness_factory.counts(), ); + Ok(()) } #[rstest] @@ -331,13 +351,13 @@ fn try_opportunistic_drain_forwards_frame(harness_factory: HarnessFactory) { fn handle_multi_packet_closed_logs_reason( harness_factory: HarnessFactory, mut logger: LoggerHandle, -) { +) -> TestResult { logger.clear(); let mut harness = harness_factory.create( HarnessConfig::new() .with_multi_packet() .with_stream_end(|_| Some(5)), - ); + )?; let (_tx, rx) = mpsc::channel(1); harness .actor_mut() @@ -345,6 +365,7 @@ fn handle_multi_packet_closed_logs_reason( logger.clear(); harness.handle_multi_packet_closed(); assert_reason_logged(&mut logger, Level::Info, "drained", Some(11)); + Ok(()) } #[rstest] @@ -352,13 +373,13 @@ fn handle_multi_packet_closed_logs_reason( fn try_opportunistic_drain_multi_disconnect_logs_reason( harness_factory: HarnessFactory, mut logger: LoggerHandle, -) { +) -> TestResult { logger.clear(); let mut harness = harness_factory.create( HarnessConfig::new() .with_multi_packet() .with_stream_end(|_| Some(5)), - ); + )?; let (tx, rx) = mpsc::channel(1); harness .actor_mut() @@ -366,19 +387,25 @@ fn try_opportunistic_drain_multi_disconnect_logs_reason( drop(tx); logger.clear(); let drained = harness.try_drain_multi(); - assert!(!drained, "disconnect should not report a drained frame"); + if drained { + return Err("disconnect should not report a drained frame".into()); + } assert_reason_logged(&mut logger, Level::Warn, "disconnected", Some(12)); + Ok(()) } #[rstest] #[serial(connection_logs)] -fn start_shutdown_logs_reason(harness_factory: HarnessFactory, mut logger: LoggerHandle) { +fn start_shutdown_logs_reason( + harness_factory: HarnessFactory, + mut logger: LoggerHandle, +) -> TestResult { logger.clear(); let mut harness = harness_factory.create( HarnessConfig::new() .with_multi_packet() .with_stream_end(|_| Some(5)), - ); + )?; let (_tx, rx) = mpsc::channel(1); harness .actor_mut() @@ -386,67 +413,88 @@ fn start_shutdown_logs_reason(harness_factory: HarnessFactory, mut logger: Logge logger.clear(); harness.start_shutdown(); assert_reason_logged(&mut logger, Level::Info, "shutdown", Some(13)); - assert!( - !harness.has_multi_queue(), - "multi-packet queue should be cleared after shutdown", - ); + if harness.has_multi_queue() { + return Err("multi-packet queue should be cleared after shutdown".into()); + } + Ok(()) } #[rstest] -fn try_opportunistic_drain_multi_disconnect_emits_terminator(harness_factory: HarnessFactory) { +fn try_opportunistic_drain_multi_disconnect_emits_terminator( + harness_factory: HarnessFactory, +) -> TestResult { let mut harness = harness_factory.create( HarnessConfig::new() .with_multi_packet() .with_stream_end(|_| Some(5)), - ); + )?; let (tx, rx) = mpsc::channel(1); harness.set_multi_queue(Some(rx)); drop(tx); let drained = harness.try_drain_multi(); - assert!(!drained, "disconnect should not report a drained frame",); - assert!( - !harness.has_multi_queue(), - "multi-packet queue should be cleared after disconnect", - ); + if drained { + return Err("disconnect should not report a drained frame".into()); + } + if harness.has_multi_queue() { + return Err("multi-packet queue should be cleared after disconnect".into()); + } assert_frame_processed( &harness.out, &[6], HookCounts { before: 1, end: 1 }, harness_factory.counts(), ); + Ok(()) } #[test] -fn try_opportunistic_drain_returns_false_when_empty() { - let mut harness = ActorHarness::new().expect("failed to create harness"); +fn try_opportunistic_drain_returns_false_when_empty() -> TestResult { + let mut harness = ActorHarness::new()?; let (_tx, rx) = mpsc::channel(1); harness.set_low_queue(Some(rx)); let drained = harness.try_drain_low(); - assert!(!drained, "no frame should be drained"); - assert!(harness.has_low_queue(), "queue should remain available"); - assert!(harness.out.is_empty(), "no frames should be emitted"); + if drained { + return Err("no frame should be drained".into()); + } + if !harness.has_low_queue() { + return Err("queue should remain available".into()); + } + if !harness.out.is_empty() { + return Err("no frames should be emitted".into()); + } + Ok(()) } #[test] -fn try_opportunistic_drain_handles_disconnect() { - let mut harness = ActorHarness::new().expect("failed to create harness"); +fn try_opportunistic_drain_handles_disconnect() -> TestResult { + let mut harness = ActorHarness::new()?; let (tx, rx) = mpsc::channel(1); harness.set_low_queue(Some(rx)); drop(tx); let drained = harness.try_drain_low(); - assert!(!drained, "disconnect should not produce a frame"); - assert!( - !harness.has_low_queue(), - "queue should be cleared after disconnect", - ); + if drained { + return Err("disconnect should not produce a frame".into()); + } + if harness.has_low_queue() { + return Err("queue should be cleared after disconnect".into()); + } let snapshot = harness.snapshot(); - assert!(snapshot.is_active && !snapshot.is_shutting_down && !snapshot.is_done); + if !snapshot.is_active { + return Err("connection should be active".into()); + } + if snapshot.is_shutting_down { + return Err("connection should not be shutting down".into()); + } + if snapshot.is_done { + return Err("connection should not be done".into()); + } + Ok(()) } #[tokio::test] diff --git a/tests/connection_actor_errors.rs b/tests/connection_actor_errors.rs index 95d356c4..b9c44393 100644 --- a/tests/connection_actor_errors.rs +++ b/tests/connection_actor_errors.rs @@ -1,9 +1,12 @@ #![cfg(not(loom))] //! Error propagation and protocol hook tests for `ConnectionActor`. -use std::sync::{ - Arc, - atomic::{AtomicUsize, Ordering}, +use std::{ + io, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, }; use futures::stream; @@ -13,46 +16,62 @@ use tokio_util::sync::CancellationToken; use wireframe::{ ConnectionContext, ProtocolHooks, - connection::ConnectionActor, + connection::{ConnectionActor, ConnectionChannels}, push::PushQueues, response::WireframeError, }; use wireframe_testing::{LoggerHandle, logger, push_expect}; -#[fixture] +mod common; +use common::TestResult; + +#[expect( + clippy::allow_attributes, + reason = "rstest single-line fixtures need allow to avoid unfulfilled lint expectations" +)] +#[allow( + unfulfilled_lint_expectations, + reason = "rstest occasionally misses the expected lint for single-line fixtures on stable" +)] #[expect( unused_braces, reason = "rustc false positive for single line rstest fixtures" )] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] -fn queues() -> (PushQueues, wireframe::push::PushHandle) { +#[fixture] +fn queues() +-> Result<(PushQueues, wireframe::push::PushHandle), wireframe::push::PushConfigError> { PushQueues::::builder() .high_capacity(8) .low_capacity(8) .build() - .expect("failed to build PushQueues") } -#[fixture] +#[expect( + clippy::allow_attributes, + reason = "rstest single-line fixtures need allow to avoid unfulfilled lint expectations" +)] +#[allow( + unfulfilled_lint_expectations, + reason = "rstest occasionally misses the expected lint for single-line fixtures on stable" +)] #[expect( unused_braces, reason = "rustc false positive for single line rstest fixtures" )] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] +#[fixture] fn shutdown_token() -> CancellationToken { CancellationToken::new() } #[rstest] #[tokio::test] #[serial] async fn before_send_hook_modifies_frames( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; push_expect!(handle.push_high_priority(1), "push high-priority"); let stream = stream::iter(vec![Ok(2u8)]); @@ -62,25 +81,31 @@ async fn before_send_hook_modifies_frames( }; let mut actor: ConnectionActor<_, ()> = ConnectionActor::with_hooks( - queues, - handle, + ConnectionChannels::new(queues, handle), Some(Box::pin(stream)), shutdown_token, hooks, ); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| io::Error::other(format!("actor run failed: {e:?}")))?; assert_eq!(out, vec![2, 3]); + Ok(()) } #[rstest] #[tokio::test] #[serial] async fn on_command_end_hook_runs( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; let stream = stream::iter(vec![Ok(1u8)]); let counter = Arc::new(AtomicUsize::new(0)); @@ -93,15 +118,18 @@ async fn on_command_end_hook_runs( }; let mut actor: ConnectionActor<_, ()> = ConnectionActor::with_hooks( - queues, - handle, + ConnectionChannels::new(queues, handle), Some(Box::pin(stream)), shutdown_token, hooks, ); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| io::Error::other(format!("actor run failed: {e:?}")))?; assert_eq!(counter.load(Ordering::SeqCst), 1); + Ok(()) } #[derive(Debug)] @@ -113,10 +141,13 @@ enum TestError { #[tokio::test] #[serial] async fn error_propagation_from_stream( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; let stream = stream::iter(vec![ Ok(1u8), Ok(2u8), @@ -133,32 +164,41 @@ async fn error_propagation_from_stream( ..ProtocolHooks::::default() }; let mut actor: ConnectionActor<_, TestError> = ConnectionActor::with_hooks( - queues, - handle, + ConnectionChannels::new(queues, handle), Some(Box::pin(stream)), shutdown_token, hooks, ); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| io::Error::other(format!("actor run failed: {e:?}")))?; assert_eq!(called.load(Ordering::SeqCst), 1); assert_eq!(out, vec![1, 2]); + Ok(()) } #[rstest] #[tokio::test] #[serial] async fn protocol_error_logs_warning( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, mut logger: LoggerHandle, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; let stream = stream::iter(vec![Err(WireframeError::Protocol(TestError::Kaboom))]); let mut actor: ConnectionActor<_, TestError> = ConnectionActor::new(queues, handle, Some(Box::pin(stream)), shutdown_token); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| io::Error::other(format!("actor run failed: {e:?}")))?; assert!(out.is_empty()); let mut found = false; while let Some(record) = logger.pop() { @@ -168,16 +208,20 @@ async fn protocol_error_logs_warning( } } assert!(found, "warning log not found"); + Ok(()) } #[rstest] #[tokio::test] #[serial] async fn io_error_terminates_connection( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; let stream = stream::iter(vec![ Ok(1u8), Err(WireframeError::Io(std::io::Error::other("fail"))), @@ -188,4 +232,5 @@ async fn io_error_terminates_connection( let result = actor.run(&mut out).await; assert!(matches!(result, Err(WireframeError::Io(_)))); assert_eq!(out, vec![1]); + Ok(()) } diff --git a/tests/connection_actor_fairness.rs b/tests/connection_actor_fairness.rs index 476f7830..79e0eb73 100644 --- a/tests/connection_actor_fairness.rs +++ b/tests/connection_actor_fairness.rs @@ -1,5 +1,5 @@ -#![cfg(not(loom))] //! Fairness and priority tests for `ConnectionActor`. +#![cfg(not(loom))] use futures::stream; use rstest::{fixture, rstest}; @@ -15,40 +15,32 @@ use wireframe::{ }; use wireframe_testing::push_expect; +mod common; +use common::TestResult; + #[fixture] -#[expect( - unused_braces, - reason = "rustc false positive for single line rstest fixtures" -)] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] -fn queues() -> (PushQueues, wireframe::push::PushHandle) { +fn queues() -> TestResult<(PushQueues, wireframe::push::PushHandle)> { PushQueues::::builder() .high_capacity(8) .low_capacity(8) .build() - .expect("failed to build PushQueues") + .map_err(Into::into) } #[fixture] -#[expect( - unused_braces, - reason = "rustc false positive for single line rstest fixtures" -)] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] -fn shutdown_token() -> CancellationToken { CancellationToken::new() } +fn shutdown_token() -> CancellationToken { + // Provide a fresh cancellation token for each rstest. + CancellationToken::new() +} #[rstest] #[tokio::test] #[serial] async fn strict_priority_order( - queues: (PushQueues, wireframe::push::PushHandle), + queues: TestResult<(PushQueues, wireframe::push::PushHandle)>, shutdown_token: CancellationToken, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; push_expect!(handle.push_low_priority(2), "push low-priority"); push_expect!(handle.push_high_priority(1), "push high-priority"); @@ -56,18 +48,22 @@ async fn strict_priority_order( let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, Some(Box::pin(stream)), shutdown_token); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); - assert_eq!(out, vec![1, 2, 3]); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("actor run failed: {e:?}")))?; + assert_eq!(out, vec![1, 2, 3], "unexpected frame ordering"); + Ok(()) } #[rstest] #[tokio::test] #[serial] async fn fairness_yields_low_after_burst( - queues: (PushQueues, wireframe::push::PushHandle), + queues: TestResult<(PushQueues, wireframe::push::PushHandle)>, shutdown_token: CancellationToken, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; let fairness = FairnessConfig { max_high_before_low: 2, time_slice: None, @@ -82,8 +78,16 @@ async fn fairness_yields_low_after_burst( ConnectionActor::new(queues, handle, None, shutdown_token); actor.set_fairness(fairness); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); - assert_eq!(out, vec![1, 2, 99, 3, 4, 5]); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("actor run failed: {e:?}")))?; + assert_eq!( + out, + vec![1, 2, 99, 3, 4, 5], + "unexpected frame order under fairness" + ); + Ok(()) } #[derive(Debug, Clone, Copy)] @@ -98,9 +102,9 @@ async fn queue_frames( order: &[Priority], handle: &wireframe::push::PushHandle, high_count: usize, -) -> Vec { +) -> TestResult> { let mut next_high = 1u8; - let mut next_low = u8::try_from(high_count).expect("too many high frames") + 1; + let mut next_low = u8::try_from(high_count).map_err(|_| "high_count exceeds u8 range")? + 1; let mut highs = Vec::new(); let mut lows = Vec::new(); @@ -126,18 +130,24 @@ async fn queue_frames( } } - highs.into_iter().chain(lows.into_iter()).collect() + Ok(highs.into_iter().chain(lows.into_iter()).collect()) } // Ensure the helper correctly handles edge cases without queued frames. #[rstest] #[tokio::test] #[serial] -async fn queue_frames_empty_input(queues: (PushQueues, wireframe::push::PushHandle)) { - let (_, handle) = queues; +async fn queue_frames_empty_input( + queues: TestResult<(PushQueues, wireframe::push::PushHandle)>, +) -> TestResult { + let (_, handle) = queues?; let priorities: &[Priority] = &[]; - let result = queue_frames(priorities, &handle, 0).await; - assert!(result.is_empty(), "Expected empty output for empty input"); + let result = queue_frames(priorities, &handle, 0).await?; + assert!( + result.is_empty(), + "expected empty output for empty input but got {result:?}" + ); + Ok(()) } #[rstest] @@ -160,36 +170,43 @@ async fn queue_frames_empty_input(queues: (PushQueues, wireframe::push::Push #[serial] async fn processes_all_priorities_in_order( #[case] order: Vec, - queues: (PushQueues, wireframe::push::PushHandle), + queues: TestResult<(PushQueues, wireframe::push::PushHandle)>, shutdown_token: CancellationToken, -) { - let (queues, handle) = queues; +) -> TestResult { + let (queues, handle) = queues?; let fairness = FairnessConfig { max_high_before_low: 0, time_slice: None, }; let high_count = order.iter().filter(|p| matches!(p, Priority::High)).count(); - let expected = queue_frames(&order, &handle, high_count).await; + let expected = queue_frames(&order, &handle, high_count).await?; let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, None, shutdown_token); actor.set_fairness(fairness); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); - assert_eq!(out, expected); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("actor run failed: {e:?}")))?; + assert_eq!( + out, expected, + "unexpected frame ordering with fairness disabled" + ); + Ok(()) } #[rstest] #[tokio::test] #[serial] async fn fairness_yields_low_with_time_slice( - queues: (PushQueues, wireframe::push::PushHandle), + queues: TestResult<(PushQueues, wireframe::push::PushHandle)>, shutdown_token: CancellationToken, -) { +) -> TestResult { // Use Tokio's virtual clock so timing-dependent fairness is deterministic. time::pause(); - let (queues, handle) = queues; + let (queues, handle) = queues?; let fairness = FairnessConfig { max_high_before_low: 0, time_slice: Some(Duration::from_millis(10)), @@ -216,14 +233,15 @@ async fn fairness_yields_low_with_time_slice( } drop(handle); - let out = rx.await.expect("actor output missing"); - assert!(out.contains(&42), "Low-priority item was not yielded"); + let out = rx.await.map_err(|_| "actor output missing")?; + assert!(out.contains(&42), "low-priority item was not yielded"); let pos = out .iter() .position(|x| *x == 42) - .expect("value 42 should be present"); + .ok_or("value 42 should be present")?; assert!( pos > 0 && pos < out.len() - 1, - "Low-priority item should be yielded in the middle", + "low-priority item should be yielded in the middle: pos={pos}, out={out:?}" ); + Ok(()) } diff --git a/tests/connection_actor_shutdown.rs b/tests/connection_actor_shutdown.rs index 70c5e35b..80bde3d2 100644 --- a/tests/connection_actor_shutdown.rs +++ b/tests/connection_actor_shutdown.rs @@ -9,40 +9,58 @@ use tokio_util::{sync::CancellationToken, task::TaskTracker}; use wireframe::{connection::ConnectionActor, push::PushQueues}; use wireframe_testing::push_expect; -#[fixture] -#[expect( - unused_braces, - reason = "rustc false positive for single line rstest fixtures" -)] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] -fn queues() -> (PushQueues, wireframe::push::PushHandle) { - PushQueues::::builder() - .high_capacity(8) - .low_capacity(8) - .build() - .expect("failed to build PushQueues") +mod common; +use common::TestResult; + +// Apply expected lint suppressions for single-line rstest fixtures. +// Context: https://github.com/la10736/rstest/issues/222 +macro_rules! single_line_fixture { + ($item:item) => { + #[expect( + clippy::allow_attributes, + reason = "rstest single-line fixtures need allow to avoid unfulfilled lint \ + expectations" + )] + #[allow( + unfulfilled_lint_expectations, + reason = "rstest occasionally misses the expected lint for single-line fixtures on \ + stable" + )] + #[expect( + unused_braces, + reason = "rustc false positive for single line rstest fixtures" + )] + $item + }; } -#[fixture] -#[expect( - unused_braces, - reason = "rustc false positive for single line rstest fixtures" -)] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] -fn shutdown_token() -> CancellationToken { CancellationToken::new() } +single_line_fixture! { + #[fixture] + fn queues() + -> Result<(PushQueues, wireframe::push::PushHandle), wireframe::push::PushConfigError> { + PushQueues::::builder() + .high_capacity(8) + .low_capacity(8) + .build() + } +} + +single_line_fixture! { + #[fixture] + fn shutdown_token() -> CancellationToken { CancellationToken::new() } +} #[rstest] #[tokio::test] #[serial] async fn shutdown_signal_precedence( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, ) { - let (queues, handle) = queues; + let (queues, handle) = queues.expect("fixture should build queues"); shutdown_token.cancel(); let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, None, shutdown_token); @@ -56,10 +74,13 @@ async fn shutdown_signal_precedence( #[tokio::test] #[serial] async fn complete_draining_of_sources( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, ) { - let (queues, handle) = queues; + let (queues, handle) = queues.expect("fixture should build queues"); push_expect!(handle.push_high_priority(1), "push high-priority"); let stream = stream::iter(vec![Ok(2u8), Ok(3u8)]); @@ -75,10 +96,13 @@ async fn complete_draining_of_sources( #[tokio::test] #[serial] async fn interleaved_shutdown_during_stream( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, ) { - let (queues, handle) = queues; + let (queues, handle) = queues.expect("fixture should build queues"); let token = shutdown_token.clone(); tokio::spawn(async move { sleep(Duration::from_millis(50)).await; @@ -121,7 +145,7 @@ async fn push_queue_exhaustion_backpressure() { #[rstest] #[tokio::test] #[serial] -async fn graceful_shutdown_waits_for_tasks() { +async fn graceful_shutdown_waits_for_tasks() -> TestResult { let tracker = TaskTracker::new(); let token = CancellationToken::new(); @@ -130,8 +154,7 @@ async fn graceful_shutdown_waits_for_tasks() { let (queues, handle) = PushQueues::::builder() .high_capacity(1) .low_capacity(1) - .build() - .expect("failed to build PushQueues"); + .build()?; let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle.clone(), None, token.clone()); handles.push(handle); @@ -149,15 +172,19 @@ async fn graceful_shutdown_waits_for_tasks() { .await .is_ok(), ); + Ok(()) } #[rstest] #[tokio::test] #[serial] async fn connection_count_decrements_on_abort( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, ) { - let (queues, handle) = queues; + let (queues, handle) = queues.expect("fixture should build queues"); let token = CancellationToken::new(); token.cancel(); @@ -176,10 +203,13 @@ async fn connection_count_decrements_on_abort( #[tokio::test] #[serial] async fn connection_count_decrements_on_close( - queues: (PushQueues, wireframe::push::PushHandle), + queues: Result< + (PushQueues, wireframe::push::PushHandle), + wireframe::push::PushConfigError, + >, shutdown_token: CancellationToken, ) { - let (queues, handle) = queues; + let (queues, handle) = queues.expect("fixture should build queues"); let before = wireframe::connection::active_connection_count(); let stream = stream::iter(vec![Ok(1u8)]); let mut actor: ConnectionActor<_, ()> = diff --git a/tests/connection_fragmentation.rs b/tests/connection_fragmentation.rs index 102abf1e..e42d1ba7 100644 --- a/tests/connection_fragmentation.rs +++ b/tests/connection_fragmentation.rs @@ -4,7 +4,7 @@ //! multiple fragments and that small frames pass through unfragmented. #![cfg(not(loom))] -use std::{num::NonZeroUsize, time::Duration}; +use std::{io, num::NonZeroUsize, time::Duration}; use tokio_util::sync::CancellationToken; use wireframe::{ @@ -16,42 +16,48 @@ use wireframe::{ const ROUTE_ID: u32 = 7; -fn setup_fragmented_actor() -> ( +mod common; +use common::TestResult; + +fn setup_fragmented_actor() -> TestResult<( ConnectionActor, PushHandle, FragmentationConfig, -) { +)> { let (queues, handle) = PushQueues::::builder() .high_capacity(4) .low_capacity(4) - .build() - .expect("build queues"); + .build()?; let shutdown = CancellationToken::new(); let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle.clone(), None, shutdown); - let cfg = FragmentationConfig::for_frame_budget( - 96, - NonZeroUsize::new(256).expect("non-zero message cap"), - Duration::from_secs(5), - ) - .expect("frame budget must exceed overhead"); + let message_cap = NonZeroUsize::new(256).ok_or("message cap must be non-zero")?; + let cfg = FragmentationConfig::for_frame_budget(96, message_cap, Duration::from_secs(5)) + .ok_or("frame budget must exceed overhead")?; actor.enable_fragmentation(cfg); - (actor, handle, cfg) + Ok((actor, handle, cfg)) } #[tokio::test] -async fn connection_actor_fragments_outbound_frames() { - let (mut actor, handle, cfg) = setup_fragmented_actor(); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn connection_actor_fragments_outbound_frames() -> TestResult { + let (mut actor, handle, cfg) = setup_fragmented_actor()?; let cap = cfg.fragment_payload_cap.get(); let payload = vec![1_u8; cap.saturating_add(16)]; let frame = Envelope::new(ROUTE_ID, Some(9), payload.clone()); - handle.push_low_priority(frame).await.expect("push frame"); + handle.push_low_priority(frame).await?; drop(handle); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|err| io::Error::other(format!("actor run failed: {err:?}")))?; assert!( out.len() > 1, @@ -63,37 +69,51 @@ async fn connection_actor_fragments_outbound_frames() { let mut assembled: Option> = None; for env in out { let payload = env.into_parts().payload(); - if let Some((header, frag)) = decode_fragment_payload(&payload).expect("decode payload") { - if let Some(message) = reassembler.push(header, frag).expect("reassemble fragment") { - assembled = Some(message.into_payload()); - } - } else { + let Some((header, frag)) = decode_fragment_payload(&payload)? else { assembled = Some(payload); + continue; + }; + + if let Some(message) = reassembler.push(header, frag)? { + assembled = Some(message.into_payload()); } } - assert_eq!(assembled.expect("assembled payload"), payload); + let assembled = assembled.ok_or("missing reassembled payload")?; + assert_eq!(assembled, payload, "reassembled payload mismatch"); + Ok(()) } #[tokio::test] -async fn connection_actor_passes_through_small_outbound_frames_unfragmented() { - let (mut actor, handle, cfg) = setup_fragmented_actor(); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn connection_actor_passes_through_small_outbound_frames_unfragmented() -> TestResult { + let (mut actor, handle, cfg) = setup_fragmented_actor()?; let payload_cap = cfg.fragment_payload_cap.get(); let payload = vec![5_u8; payload_cap.saturating_sub(1)]; let frame = Envelope::new(ROUTE_ID, Some(1), payload.clone()); - handle.push_low_priority(frame).await.expect("push frame"); + handle.push_low_priority(frame).await?; drop(handle); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|err| io::Error::other(format!("actor run failed: {err:?}")))?; assert_eq!(out.len(), 1, "expected unfragmented single frame"); - let only = out.into_iter().next().expect("frame present"); + let only = out + .into_iter() + .next() + .ok_or("expected single frame but none found")?; let payload_out = only.into_parts().payload(); - match decode_fragment_payload(&payload_out) { - Ok(None) => {} - other => panic!("expected unfragmented payload, got {other:?}"), + match decode_fragment_payload(&payload_out)? { + None => {} + Some(_) => return Err("expected unfragmented payload".into()), } - assert_eq!(payload_out, payload); + assert_eq!(payload_out, payload, "payload mutated during round trip"); + Ok(()) } diff --git a/tests/correlation_id.rs b/tests/correlation_id.rs index 63b39313..9fbd8cbd 100644 --- a/tests/correlation_id.rs +++ b/tests/correlation_id.rs @@ -1,5 +1,7 @@ #![cfg(not(loom))] //! Tests for `correlation_id` propagation in streaming responses. +use std::io; + use async_stream::try_stream; use rstest::rstest; use tokio::sync::mpsc; @@ -7,14 +9,21 @@ use tokio_util::sync::CancellationToken; use wireframe::{ CorrelatableFrame, app::Envelope, - connection::ConnectionActor, + connection::{ConnectionActor, ConnectionChannels}, hooks::{ConnectionContext, ProtocolHooks}, push::PushQueues, response::FrameStream, }; +mod common; +use common::TestResult; + #[tokio::test] -async fn stream_frames_carry_request_correlation_id() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn stream_frames_carry_request_correlation_id() -> TestResult { let cid = 42u64; let stream: FrameStream = Box::pin(try_stream! { yield Envelope::new(1, Some(cid), vec![1]); @@ -24,28 +33,32 @@ async fn stream_frames_carry_request_correlation_id() { .high_capacity(1) .low_capacity(1) .unlimited() - .build() - .expect("failed to build PushQueues"); + .build()?; let shutdown = CancellationToken::new(); let mut actor = ConnectionActor::new(queues, handle, Some(stream), shutdown); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); - assert!(out.iter().all(|e| e.correlation_id() == Some(cid))); + actor + .run(&mut out) + .await + .map_err(|e| io::Error::other(format!("actor run failed: {e:?}")))?; + assert!( + out.iter().all(|e| e.correlation_id() == Some(cid)), + "frames lost correlation id" + ); + Ok(()) } async fn run_multi_packet_channel( request_correlation: Option, frame_correlations: &[Option], hooks: ProtocolHooks, -) -> Vec { +) -> TestResult> { let capacity = frame_correlations.len().max(1); let (tx, rx) = mpsc::channel(capacity); for (idx, correlation) in frame_correlations.iter().enumerate() { let marker = (idx + 1) as u64; let payload = marker.to_le_bytes().to_vec(); - tx.send(Envelope::new(1, *correlation, payload)) - .await - .expect("send frame"); + tx.send(Envelope::new(1, *correlation, payload)).await?; } drop(tx); @@ -53,16 +66,22 @@ async fn run_multi_packet_channel( .high_capacity(2) .low_capacity(2) .unlimited() - .build() - .expect("failed to build PushQueues"); + .build()?; let shutdown = CancellationToken::new(); - let mut actor: ConnectionActor = - ConnectionActor::with_hooks(queues, handle, None, shutdown, hooks); + let mut actor: ConnectionActor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + None, + shutdown, + hooks, + ); actor.set_multi_packet_with_correlation(Some(rx), request_correlation); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); - out + actor + .run(&mut out) + .await + .map_err(|e| io::Error::other(format!("actor run failed: {e:?}")))?; + Ok(out) } #[rstest] @@ -74,13 +93,17 @@ async fn multi_packet_frames_apply_expected_correlation( #[case] request: Option, #[case] initial: Vec>, #[case] expected: Vec>, -) { - let frames = run_multi_packet_channel(request, &initial, ProtocolHooks::default()).await; +) -> TestResult { + let frames = run_multi_packet_channel(request, &initial, ProtocolHooks::default()).await?; let correlations: Vec> = frames .iter() .map(CorrelatableFrame::correlation_id) .collect(); - assert_eq!(correlations, expected); + assert_eq!( + correlations, expected, + "unexpected correlation ids: {correlations:?}, expected {expected:?}" + ); + Ok(()) } #[rstest] @@ -90,7 +113,7 @@ async fn multi_packet_frames_apply_expected_correlation( async fn multi_packet_terminator_applies_correlation( #[case] request: Option, #[case] expected: Option, -) { +) -> TestResult { let hooks = ProtocolHooks { stream_end: Some(Box::new(|_ctx: &mut ConnectionContext| { Some(Envelope::new(255, None, vec![])) @@ -98,8 +121,14 @@ async fn multi_packet_terminator_applies_correlation( ..ProtocolHooks::default() }; - let frames = run_multi_packet_channel(request, &[], hooks).await; - assert_eq!(frames.len(), 1, "terminator frame missing"); - let terminator = frames.last().expect("terminator frame missing"); - assert_eq!(terminator.correlation_id(), expected); + let frames = run_multi_packet_channel(request, &[], hooks).await?; + let [terminator] = frames.as_slice() else { + return Err(io::Error::other("expected exactly one terminator frame").into()); + }; + assert_eq!( + terminator.correlation_id(), + expected, + "unexpected terminator correlation" + ); + Ok(()) } diff --git a/tests/extractor.rs b/tests/extractor.rs index 720b4bdb..677cab9e 100644 --- a/tests/extractor.rs +++ b/tests/extractor.rs @@ -11,19 +11,17 @@ use wireframe::{ message::Message as MessageTrait, }; -#[allow( - unused_braces, - reason = "rustc false positive for single line rstest fixtures" -)] #[fixture] -fn request() -> MessageRequest { MessageRequest::default() } +fn request() -> MessageRequest { + // default request used across extractor tests + MessageRequest::default() +} -#[allow( - unused_braces, - reason = "rustc false positive for single line rstest fixtures" -)] #[fixture] -fn empty_payload() -> Payload<'static> { Payload::default() } +fn empty_payload() -> Payload<'static> { + // simple empty payload ensures extractors handle zero-length bodies + Payload::default() +} #[derive(bincode::Encode, bincode::BorrowDecode, PartialEq, Debug)] struct TestMsg(u8); diff --git a/tests/fragment_transport.rs b/tests/fragment_transport.rs index 6f3578fb..74419d65 100644 --- a/tests/fragment_transport.rs +++ b/tests/fragment_transport.rs @@ -1,10 +1,11 @@ #![cfg(not(loom))] //! Integration tests for transport-level fragmentation and reassembly. -use std::{num::NonZeroUsize, time::Duration}; +use std::{io, num::NonZeroUsize, time::Duration}; use futures::{SinkExt, StreamExt}; use rstest::rstest; +use thiserror::Error; use tokio::{ io::AsyncWriteExt, sync::mpsc, @@ -19,85 +20,131 @@ use wireframe::{ FragmentationConfig, Fragmenter, Reassembler, + ReassemblyError, decode_fragment_payload, encode_fragment_payload, }, serializer::BincodeSerializer, }; +mod common; +use common::TestResult; + +#[derive(Debug, Error)] +enum TestError { + #[error("test setup failed: {0}")] + Setup(&'static str), + #[error("fragmentation failed: {0}")] + Fragmentation(#[from] wireframe::fragment::FragmentationError), + #[error("encoding failed: {0}")] + Encode(#[from] bincode::error::EncodeError), + #[error("decoding failed: {0}")] + Decode(#[from] bincode::error::DecodeError), + #[error("reassembly failed: {0}")] + Reassembly(#[from] ReassemblyError), + #[error("send failed: {0}")] + Send(String), + #[error("application error: {0}")] + App(#[from] wireframe::app::WireframeError), + #[error(transparent)] + Other(#[from] Box), + #[error("assertion failed: {0}")] + Assertion(String), + #[error("io failed: {0}")] + Io(#[from] std::io::Error), + #[error("timeout: {0}")] + Timeout(#[from] tokio::time::error::Elapsed), + #[error("task join failed: {0}")] + Join(#[from] tokio::task::JoinError), +} + +impl From> for TestError { + fn from(err: mpsc::error::SendError) -> Self { TestError::Send(err.to_string()) } +} + const ROUTE_ID: u32 = 42; const CORRELATION: Option = Some(7); -fn fragmentation_config(capacity: usize) -> FragmentationConfig { - FragmentationConfig::for_frame_budget( - capacity, - NonZeroUsize::new(capacity * 16).expect("non-zero message limit"), - Duration::from_millis(30), - ) - .expect("frame budget must exceed fragment overhead") +fn fragmentation_config(capacity: usize) -> TestResult { + let message_limit = NonZeroUsize::new(capacity.saturating_mul(16)) + .ok_or(TestError::Setup("non-zero message limit"))?; + + let config = + FragmentationConfig::for_frame_budget(capacity, message_limit, Duration::from_millis(30)) + .ok_or(TestError::Setup( + "frame budget must exceed fragment overhead", + ))?; + + Ok(config) +} + +fn fragmentation_config_with_timeout( + capacity: usize, + timeout_ms: u64, +) -> TestResult { + let mut config = fragmentation_config(capacity)?; + config.reassembly_timeout = Duration::from_millis(timeout_ms); + Ok(config) } -fn fragment_envelope(env: &Envelope, fragmenter: &Fragmenter) -> Vec { +fn fragment_envelope(env: &Envelope, fragmenter: &Fragmenter) -> TestResult> { let parts = env.clone().into_parts(); let id = parts.id(); let correlation = parts.correlation_id(); let payload = parts.payload(); if payload.len() <= fragmenter.max_fragment_size().get() { - return vec![Envelope::new(id, correlation, payload)]; + return Ok(vec![Envelope::new(id, correlation, payload)]); } - fragmenter - .fragment_bytes(payload) - .expect("fragment payload") + let envelopes = fragmenter + .fragment_bytes(payload)? .into_iter() .map(|fragment| { let (header, payload) = fragment.into_parts(); - let encoded = encode_fragment_payload(header, &payload).expect("encode fragment"); - Envelope::new(id, correlation, encoded) + encode_fragment_payload(header, &payload) + .map(|encoded| Envelope::new(id, correlation, encoded)) + .map_err(TestError::from) }) - .collect() + .collect::, TestError>>()?; + + Ok(envelopes) } async fn send_envelopes( client: &mut Framed, envelopes: &[Envelope], -) { +) -> TestResult { let serializer = BincodeSerializer; for env in envelopes { - let bytes = serializer.serialize(env).expect("serialize envelope"); - client.send(bytes.into()).await.expect("send frame"); + let bytes = serializer.serialize(env)?; + client.send(bytes.into()).await?; } + Ok(()) } async fn read_reassembled_response( client: &mut Framed, cfg: &FragmentationConfig, -) -> Vec { +) -> TestResult> { let serializer = BincodeSerializer; let mut reassembler = Reassembler::new(cfg.max_message_size, cfg.reassembly_timeout); while let Some(frame) = client.next().await { - let bytes = frame.expect("read frame"); - let (env, _) = serializer - .deserialize::(&bytes) - .expect("decode envelope"); + let bytes = frame?; + let (env, _) = serializer.deserialize::(&bytes)?; let payload = env.into_parts().payload(); - if let Some((header, fragment)) = - decode_fragment_payload(&payload).expect("decode fragment payload") - { - if let Some(message) = reassembler - .push(header, fragment) - .expect("reassemble fragment") - { - return message.into_payload(); + match decode_fragment_payload(&payload)? { + Some((header, fragment)) => { + if let Some(message) = reassembler.push(header, fragment)? { + return Ok(message.into_payload()); + } } - } else { - return payload; + None => return Ok(payload), } } - panic!("response stream ended before reassembly completed"); + Err(TestError::Setup("response stream ended before reassembly completed").into()) } fn make_handler(sender: &mpsc::UnboundedSender>) -> Handler { @@ -106,7 +153,10 @@ fn make_handler(sender: &mpsc::UnboundedSender>) -> Handler { let tx = tx.clone(); let payload = env.clone().into_parts().payload(); Box::pin(async move { - tx.send(payload).expect("record payload"); + assert!( + tx.send(payload).is_ok(), + "handler channel send must succeed in tests" + ); }) }) } @@ -115,87 +165,126 @@ fn make_app( capacity: usize, config: FragmentationConfig, sender: &mpsc::UnboundedSender>, -) -> WireframeApp { - WireframeApp::new() - .expect("build app") +) -> TestResult { + Ok(WireframeApp::new()? .buffer_capacity(capacity) .fragmentation(Some(config)) - .route(ROUTE_ID, make_handler(sender)) - .expect("register route") + .route(ROUTE_ID, make_handler(sender))?) } fn spawn_app( app: WireframeApp, ) -> ( Framed, - tokio::task::JoinHandle<()>, + tokio::task::JoinHandle>, ) { let codec = app.length_codec(); let (client_stream, server_stream) = tokio::io::duplex(256); let client = Framed::new(client_stream, codec.clone()); - let server = tokio::spawn(async move { app.handle_connection(server_stream).await }); + let server = tokio::spawn(async move { app.handle_connection_result(server_stream).await }); (client, server) } -#[tokio::test] -async fn fragmented_request_and_response_round_trip() { - let buffer_capacity = 512; - let config = fragmentation_config(buffer_capacity); +fn build_envelopes( + request: Envelope, + config: &FragmentationConfig, + should_fragment: bool, +) -> TestResult> { + if should_fragment { + let fragmenter = Fragmenter::new(config.fragment_payload_cap); + fragment_envelope(&request, &fragmenter) + } else { + Ok(vec![request]) + } +} + +async fn assert_handler_observed( + rx: &mut mpsc::UnboundedReceiver>, + expected: &[u8], +) -> TestResult<()> { + let observed = timeout(Duration::from_secs(1), rx.recv()) + .await? + .ok_or(TestError::Setup("handler payload missing"))?; + assert_eq!( + observed, expected, + "observed payload mismatch: expected {expected:?}, got {observed:?}" + ); + Ok(()) +} + +async fn read_response_payload( + client: &mut Framed, + config: &FragmentationConfig, +) -> TestResult> { + let response = timeout( + Duration::from_secs(1), + read_reassembled_response(client, config), + ) + .await??; + Ok(response) +} + +/// Common helper for round-trip fragmentation tests. +/// Returns the response payload for additional test-specific assertions. +async fn run_round_trip_test( + buffer_capacity: usize, + payload: Vec, + should_fragment: bool, +) -> TestResult> { + let config = fragmentation_config(buffer_capacity)?; let (tx, mut rx) = mpsc::unbounded_channel(); - let app = make_app(buffer_capacity, config, &tx); + let app = make_app(buffer_capacity, config, &tx)?; let (mut client, server) = spawn_app(app); - let payload = vec![b'Z'; 1_200]; let request = Envelope::new(ROUTE_ID, CORRELATION, payload.clone()); - let fragmenter = Fragmenter::new(config.fragment_payload_cap); - let fragments = fragment_envelope(&request, &fragmenter); - send_envelopes(&mut client, &fragments).await; - client.flush().await.expect("flush client"); + let envelopes = build_envelopes(request, &config, should_fragment)?; - let observed = rx.recv().await.expect("handler payload"); - assert_eq!(observed, payload); + send_envelopes(&mut client, &envelopes).await?; + client.flush().await?; - client.get_mut().shutdown().await.expect("shutdown write"); - let response = read_reassembled_response(&mut client, &config).await; + assert_handler_observed(&mut rx, &payload).await?; + client.get_mut().shutdown().await?; + let response = read_response_payload(&mut client, &config).await?; assert_eq!(response, payload); - server.await.expect("server task"); + server.await??; + + Ok(response) } #[tokio::test] -async fn unfragmented_request_and_response_round_trip() { +async fn fragmented_request_and_response_round_trip() -> TestResult { let buffer_capacity = 512; - let config = fragmentation_config(buffer_capacity); - let (tx, mut rx) = mpsc::unbounded_channel(); - let app = make_app(buffer_capacity, config, &tx); - let (mut client, server) = spawn_app(app); + let payload = vec![b'Z'; 1_200]; + run_round_trip_test(buffer_capacity, payload, true).await?; + Ok(()) +} +#[tokio::test] +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn unfragmented_request_and_response_round_trip() -> TestResult { + let buffer_capacity = 512; + let config = fragmentation_config(buffer_capacity)?; let cap = config.fragment_payload_cap.get(); let payload_len = cap.saturating_sub(8).max(1); let payload = vec![b's'; payload_len]; - let request = Envelope::new(ROUTE_ID, CORRELATION, payload.clone()); - - send_envelopes(&mut client, &[request]).await; - client.flush().await.expect("flush client"); - - let observed = rx.recv().await.expect("handler payload"); - assert_eq!(observed, payload); - client.get_mut().shutdown().await.expect("shutdown write"); - let response = read_reassembled_response(&mut client, &config).await; - assert_eq!(response, payload); + let response = run_round_trip_test(buffer_capacity, payload, false).await?; assert!( - matches!(decode_fragment_payload(&response), Ok(None)), + decode_fragment_payload(&response)?.is_none(), "small payload should pass through unfragmented" ); - server.await.expect("server task"); + Ok(()) } struct FragmentRejectionSetup { client: Framed, - server: tokio::task::JoinHandle<()>, + server: tokio::task::JoinHandle>, fragments: Vec, rx: mpsc::UnboundedReceiver>, } @@ -204,68 +293,83 @@ impl FragmentRejectionSetup { fn new( capacity: usize, config: FragmentationConfig, - fragment_mutator: impl FnOnce(Vec) -> Vec, - ) -> Self { + fragment_mutator: impl FnOnce(Vec) -> TestResult>, + ) -> TestResult { let (tx, rx) = mpsc::unbounded_channel(); - let app = make_app(capacity, config, &tx); + let app = make_app(capacity, config, &tx)?; let (client, server) = spawn_app(app); let fragmenter = Fragmenter::new(config.fragment_payload_cap); let payload = vec![1_u8; 800]; let request = Envelope::new(ROUTE_ID, CORRELATION, payload); - let fragments = fragment_mutator(fragment_envelope(&request, &fragmenter)); + let fragments = fragment_mutator(fragment_envelope(&request, &fragmenter)?)?; - Self { + Ok(Self { client, server, fragments, rx, - } + }) } } -async fn test_fragment_rejection(fragment_mutator: F, rejection_message: &str) +async fn test_fragment_rejection(fragment_mutator: F, rejection_message: &str) -> TestResult where - F: FnOnce(Vec) -> Vec, + F: FnOnce(Vec) -> TestResult>, { let buffer_capacity = 512; - let config = fragmentation_config(buffer_capacity); + let config = fragmentation_config(buffer_capacity)?; let FragmentRejectionSetup { mut client, server, fragments, mut rx, - } = FragmentRejectionSetup::new(buffer_capacity, config, fragment_mutator); + } = FragmentRejectionSetup::new(buffer_capacity, config, fragment_mutator)?; - send_envelopes(&mut client, &fragments).await; - client.get_mut().shutdown().await.expect("shutdown write"); + send_envelopes(&mut client, &fragments).await?; + client.get_mut().shutdown().await?; if let Ok(Some(_)) = timeout(Duration::from_millis(200), rx.recv()).await { - panic!("{rejection_message}"); + return Err(TestError::Assertion(rejection_message.to_string()).into()); } drop(client); - server.await.expect("server task"); + server.await??; + + Ok(()) } -type FragmentMutator = fn(Vec) -> Vec; +type FragmentMutator = fn(Vec) -> TestResult>; + +fn mutate_out_of_order(mut fragments: Vec) -> TestResult> { + if fragments.len() < 2 { + return Err(TestError::Setup("expected at least two fragments").into()); + } -fn mutate_out_of_order(mut fragments: Vec) -> Vec { fragments.swap(0, 1); - fragments + Ok(fragments) } -fn mutate_duplicate(mut fragments: Vec) -> Vec { - let duplicate = fragments[0].clone(); +fn mutate_duplicate(mut fragments: Vec) -> TestResult> { + let duplicate = fragments + .first() + .cloned() + .ok_or(TestError::Setup("fragmenter produced no fragments"))?; fragments.insert(1, duplicate); - fragments + Ok(fragments) } -fn mutate_malformed_header(mut fragments: Vec) -> Vec { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +fn mutate_malformed_header(mut fragments: Vec) -> TestResult> { let parts = fragments .first() .cloned() - .expect("fragmenter must produce at least one fragment") + .ok_or(TestError::Setup( + "fragmenter must produce at least one fragment", + ))? .into_parts(); let mut payload = parts.clone().payload(); assert!( @@ -280,13 +384,17 @@ fn mutate_malformed_header(mut fragments: Vec) -> Vec { payload.push(0); } } - fragments[0] = Envelope::from_parts(PacketParts::new( - parts.id(), - parts.correlation_id(), - payload, - )); + if let Some(first) = fragments.get_mut(0) { + *first = Envelope::from_parts(PacketParts::new( + parts.id(), + parts.correlation_id(), + payload, + )); + } else { + return Err(TestError::Setup("fragment list unexpectedly empty").into()); + } fragments.truncate(1); - fragments + Ok(fragments) } #[rstest] @@ -303,22 +411,21 @@ fn mutate_malformed_header(mut fragments: Vec) -> Vec { async fn fragment_rejection_cases( #[case] mutator: FragmentMutator, #[case] rejection_message: &str, -) { - test_fragment_rejection(mutator, rejection_message).await; +) -> TestResult { + test_fragment_rejection(mutator, rejection_message).await } #[tokio::test] -async fn expired_fragments_are_evicted() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn expired_fragments_are_evicted() -> TestResult { let buffer_capacity = 512; let timeout_ms = 10; - let config = FragmentationConfig::for_frame_budget( - buffer_capacity, - NonZeroUsize::new(buffer_capacity * 2).expect("non-zero message limit"), - Duration::from_millis(timeout_ms), - ) - .expect("frame budget must exceed fragment overhead"); + let config = fragmentation_config_with_timeout(buffer_capacity, timeout_ms)?; let (tx, mut rx) = mpsc::unbounded_channel(); - let app = make_app(buffer_capacity, config, &tx); + let app = make_app(buffer_capacity, config, &tx)?; let codec = app.length_codec(); let (client_stream, server_stream) = tokio::io::duplex(256); let mut client = Framed::new(client_stream, codec.clone()); @@ -326,53 +433,81 @@ async fn expired_fragments_are_evicted() { let payload = vec![3_u8; 800]; let request = Envelope::new(ROUTE_ID, CORRELATION, payload); - let fragments = fragment_envelope(&request, &fragmenter); + let fragments = fragment_envelope(&request, &fragmenter)?; - let server = tokio::spawn(async move { app.handle_connection(server_stream).await }); + let server = tokio::spawn(async move { app.handle_connection_result(server_stream).await }); // Send the first fragment then pause long enough for eviction. - send_envelopes(&mut client, &fragments[..1]).await; + let first_fragment = fragments + .get(..1) + .ok_or(TestError::Setup("fragmenter produced no fragments"))?; + send_envelopes(&mut client, first_fragment).await?; sleep(Duration::from_millis(timeout_ms * 2)).await; - send_envelopes(&mut client, &fragments[1..]).await; - client.get_mut().shutdown().await.expect("shutdown write"); + if let Some(rest) = fragments.get(1..) { + send_envelopes(&mut client, rest).await?; + } + client.get_mut().shutdown().await?; + let recv_result = timeout(Duration::from_millis(200), rx.recv()).await; assert!( - timeout(Duration::from_millis(200), rx.recv()) - .await - .is_err(), + recv_result.is_err(), "handler should not receive after timeout eviction" ); drop(client); - server.await.expect("server task"); + server.await??; + + Ok(()) } #[tokio::test] -async fn fragmentation_can_be_disabled_via_public_api() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn fragmentation_can_be_disabled_via_public_api() -> TestResult { let capacity = 1024; let (tx, mut rx) = mpsc::unbounded_channel(); + let config = fragmentation_config(capacity)?; let handler = make_handler(&tx); - let app: WireframeApp = WireframeApp::new() - .expect("build app") + let app: WireframeApp = WireframeApp::new()? .buffer_capacity(capacity) .fragmentation(None) - .route(ROUTE_ID, handler) - .expect("register route"); + .route(ROUTE_ID, handler)?; let (mut client, server) = spawn_app(app); - let payload = vec![b'X'; capacity / 2]; + let half_capacity = capacity + .checked_div(2) + .ok_or(TestError::Setup("capacity must be at least two"))?; + let payload = vec![b'X'; half_capacity]; let request = Envelope::new(ROUTE_ID, CORRELATION, payload.clone()); let serializer = BincodeSerializer; - let bytes = serializer.serialize(&request).expect("serialize envelope"); - client.send(bytes.into()).await.expect("send frame"); - client.get_mut().shutdown().await.expect("shutdown write"); - drop(client); + let bytes = serializer.serialize(&request)?; + client.send(bytes.into()).await?; + + let observed = timeout(Duration::from_secs(1), rx.recv()) + .await? + .ok_or(TestError::Setup("handler payload missing"))?; + assert_eq!( + observed, payload, + "observed payload mismatch: expected {payload:?}, got {observed:?}" + ); + + client.get_mut().shutdown().await?; + let response = timeout( + Duration::from_secs(1), + read_reassembled_response(&mut client, &config), + ) + .await??; + assert!( + decode_fragment_payload(&response)?.is_none(), + "expected no fragmentation when fragmentation is disabled" + ); - let observed = rx.recv().await.expect("handler payload"); - assert_eq!(observed, payload); + server.await??; - server.await.expect("server task"); + Ok(()) } diff --git a/tests/lifecycle.rs b/tests/lifecycle.rs index a4d54339..8fd2a76c 100644 --- a/tests/lifecycle.rs +++ b/tests/lifecycle.rs @@ -26,6 +26,9 @@ use wireframe_testing::{ run_with_duplex_server, }; +mod common; +use common::TestResult; + type App = wireframe::app::WireframeApp; type BasicApp = wireframe::app::WireframeApp; @@ -52,61 +55,89 @@ fn wireframe_app_with_lifecycle_callbacks( setup: &Arc, teardown: &Arc, state: u32, -) -> App +) -> wireframe::app::Result> where E: Packet, { let setup_cb = call_counting_callback(setup, state); let teardown_cb = call_counting_callback(teardown, ()); - App::::new() - .expect("failed to create app") - .on_connection_setup(move || setup_cb(())) - .expect("setup callback") - .on_connection_teardown(teardown_cb) - .expect("teardown callback") + let app = App::::new()? + .on_connection_setup(move || setup_cb(()))? + .on_connection_teardown(teardown_cb)?; + + Ok(app) } #[tokio::test] -async fn setup_and_teardown_callbacks_run() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn setup_and_teardown_callbacks_run() -> TestResult<()> { let setup_count = Arc::new(AtomicUsize::new(0)); let teardown_count = Arc::new(AtomicUsize::new(0)); - let app = wireframe_app_with_lifecycle_callbacks::(&setup_count, &teardown_count, 42); + let app = + wireframe_app_with_lifecycle_callbacks::(&setup_count, &teardown_count, 42)?; run_with_duplex_server(app).await; - assert_eq!(setup_count.load(Ordering::SeqCst), 1); - assert_eq!(teardown_count.load(Ordering::SeqCst), 1); + assert_eq!( + setup_count.load(Ordering::SeqCst), + 1, + "setup callback did not run exactly once" + ); + assert_eq!( + teardown_count.load(Ordering::SeqCst), + 1, + "teardown callback did not run exactly once" + ); + + Ok(()) } #[tokio::test] -async fn setup_without_teardown_runs() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn setup_without_teardown_runs() -> TestResult<()> { let setup_count = Arc::new(AtomicUsize::new(0)); let cb = call_counting_callback(&setup_count, ()); - let app = BasicApp::new() - .expect("failed to create app") - .on_connection_setup(move || cb(())) - .expect("setup callback"); + let app = BasicApp::new()?.on_connection_setup(move || cb(()))?; run_with_duplex_server(app).await; - assert_eq!(setup_count.load(Ordering::SeqCst), 1); + assert_eq!( + setup_count.load(Ordering::SeqCst), + 1, + "setup callback did not run" + ); + + Ok(()) } #[tokio::test] -async fn teardown_without_setup_does_not_run() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn teardown_without_setup_does_not_run() -> TestResult<()> { let teardown_count = Arc::new(AtomicUsize::new(0)); let cb = call_counting_callback(&teardown_count, ()); - let app = BasicApp::new() - .expect("failed to create app") - .on_connection_teardown(cb) - .expect("teardown callback"); + let app = BasicApp::new()?.on_connection_teardown(cb)?; run_with_duplex_server(app).await; - assert_eq!(teardown_count.load(Ordering::SeqCst), 0); + assert_eq!( + teardown_count.load(Ordering::SeqCst), + 0, + "teardown callback should not run" + ); + + Ok(()) } #[derive(bincode::Encode, bincode::BorrowDecode, PartialEq, Debug)] @@ -138,40 +169,47 @@ impl Packet for StateEnvelope { } #[tokio::test] -async fn helpers_preserve_correlation_id_and_run_callbacks() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn helpers_preserve_correlation_id_and_run_callbacks() -> TestResult<()> { let setup = Arc::new(AtomicUsize::new(0)); let teardown = Arc::new(AtomicUsize::new(0)); - let app = wireframe_app_with_lifecycle_callbacks::(&setup, &teardown, 7) - .route(1, Arc::new(|_: &StateEnvelope| Box::pin(async {}))) - .expect("route registration failed"); + let app = wireframe_app_with_lifecycle_callbacks::(&setup, &teardown, 7)? + .route(1, Arc::new(|_: &StateEnvelope| Box::pin(async {})))?; let env = StateEnvelope { id: 1, correlation_id: Some(0), payload: vec![1], }; - let bytes = BincodeSerializer - .serialize(&env) - .expect("failed to serialise envelope"); + let bytes = BincodeSerializer.serialize(&env)?; let mut frame = BytesMut::with_capacity(bytes.len() + 4); let mut codec = new_test_codec(TEST_MAX_FRAME); - codec - .encode(bytes.into(), &mut frame) - .expect("encode should succeed"); + codec.encode(bytes.into(), &mut frame)?; - let out = run_app(app, vec![frame.to_vec()], None) - .await - .expect("app run failed"); - assert!(!out.is_empty()); + let out = run_app(app, vec![frame.to_vec()], None).await?; + assert!(!out.is_empty(), "expected response frames"); let frames = decode_frames(out); - assert_eq!(frames.len(), 1, "expected a single response frame"); - let (resp, _) = BincodeSerializer - .deserialize::(&frames[0]) - .expect("deserialize failed"); - assert_eq!(resp.correlation_id, Some(0)); - - assert_eq!(setup.load(Ordering::SeqCst), 1); - assert_eq!(teardown.load(Ordering::SeqCst), 1); + let [first] = frames.as_slice() else { + panic!("expected a single response frame"); + }; + let (resp, _) = BincodeSerializer.deserialize::(first)?; + assert_eq!(resp.correlation_id, Some(0), "correlation id not preserved"); + + assert_eq!( + setup.load(Ordering::SeqCst), + 1, + "setup callback did not run exactly once" + ); + assert_eq!( + teardown.load(Ordering::SeqCst), + 1, + "teardown callback did not run exactly once" + ); + + Ok(()) } diff --git a/tests/metadata.rs b/tests/metadata.rs index 2619eead..03e2464e 100644 --- a/tests/metadata.rs +++ b/tests/metadata.rs @@ -1,7 +1,7 @@ -#![cfg(not(loom))] //! Tests for frame metadata parsing using custom serializers. //! //! They ensure parse callbacks run before deserialization and errors fall back correctly. +#![cfg(not(loom))] use std::sync::{ Arc, @@ -15,16 +15,19 @@ use wireframe::{ }; use wireframe_testing::{TestSerializer, drive_with_bincode}; +mod common; +use common::TestResult; + type TestApp = wireframe::app::WireframeApp; -fn mock_wireframe_app_with_serializer(serializer: S) -> TestApp +fn mock_wireframe_app_with_serializer( + serializer: S, +) -> Result, wireframe::app::WireframeError> where S: TestSerializer + Default, { - wireframe::app::WireframeApp::::with_serializer(serializer) - .expect("failed to create app") + wireframe::app::WireframeApp::::with_serializer(serializer)? .route(1, Arc::new(|_| Box::pin(async {}))) - .expect("route registration failed") } #[derive(Default)] @@ -42,7 +45,7 @@ impl Serializer for CountingSerializer { &self, _bytes: &[u8], ) -> Result<(M, usize), Box> { - panic!("unexpected deserialize call") + Err("unexpected deserialize call".into()) } } @@ -57,18 +60,21 @@ impl FrameMetadata for CountingSerializer { } #[tokio::test] -async fn metadata_parser_invoked_before_deserialize() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn metadata_parser_invoked_before_deserialize() -> TestResult<()> { let counter = Arc::new(AtomicUsize::new(0)); let serializer = CountingSerializer(counter.clone()); - let app = mock_wireframe_app_with_serializer(serializer); + let app = mock_wireframe_app_with_serializer(serializer)?; let env = Envelope::new(1, Some(0), vec![42]); - let out = drive_with_bincode(app, env) - .await - .expect("drive_with_bincode failed"); - assert!(!out.is_empty()); - assert_eq!(counter.load(Ordering::Relaxed), 1); + let out = drive_with_bincode(app, env).await?; + assert!(!out.is_empty(), "no frames emitted"); + assert_eq!(counter.load(Ordering::Relaxed), 1, "expected 1 parse call"); + Ok(()) } #[derive(Default)] @@ -102,18 +108,29 @@ impl FrameMetadata for FallbackSerializer { } #[tokio::test] -async fn falls_back_to_deserialize_after_parse_error() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn falls_back_to_deserialize_after_parse_error() -> TestResult<()> { let parse_calls = Arc::new(AtomicUsize::new(0)); let deser_calls = Arc::new(AtomicUsize::new(0)); let serializer = FallbackSerializer(parse_calls.clone(), deser_calls.clone()); - let app = mock_wireframe_app_with_serializer(serializer); + let app = mock_wireframe_app_with_serializer(serializer)?; let env = Envelope::new(1, Some(0), vec![7]); - let out = drive_with_bincode(app, env) - .await - .expect("drive_with_bincode failed"); - assert!(!out.is_empty()); - assert_eq!(parse_calls.load(Ordering::Relaxed), 1); - assert_eq!(deser_calls.load(Ordering::Relaxed), 1); + let out = drive_with_bincode(app, env).await?; + assert!(!out.is_empty(), "no frames emitted"); + assert_eq!( + parse_calls.load(Ordering::Relaxed), + 1, + "expected 1 parse call" + ); + assert_eq!( + deser_calls.load(Ordering::Relaxed), + 1, + "expected 1 deserialize call" + ); + Ok(()) } diff --git a/tests/middleware_order.rs b/tests/middleware_order.rs index f57d29e0..d6b47021 100644 --- a/tests/middleware_order.rs +++ b/tests/middleware_order.rs @@ -12,6 +12,9 @@ use wireframe::{ }; use wireframe_testing::{decode_frames, encode_frame}; +mod common; +use common::TestResult; + type TestApp = wireframe::app::WireframeApp; struct TagMiddleware(u8); @@ -53,7 +56,11 @@ impl Transform> for TagMiddleware { } #[tokio::test] -async fn middleware_applied_in_reverse_order() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn middleware_applied_in_reverse_order() -> TestResult<()> { let handler: Handler = std::sync::Arc::new(|_env: &Envelope| Box::pin(async {})); let app = TestApp::new() .expect("failed to create app") @@ -68,26 +75,31 @@ async fn middleware_applied_in_reverse_order() { let env = Envelope::new(1, Some(7), vec![b'X']); let serializer = BincodeSerializer; - let bytes = serializer.serialize(&env).expect("serialization failed"); + let bytes = serializer.serialize(&env)?; let mut codec = app.length_codec(); let frame = encode_frame(&mut codec, bytes); - client.write_all(&frame).await.expect("write failed"); - client.shutdown().await.expect("shutdown failed"); + client.write_all(&frame).await?; + client.shutdown().await?; - let handle = tokio::spawn(async move { app.handle_connection(server).await }); + let handle = tokio::spawn(async move { app.handle_connection_result(server).await }); let mut out = Vec::new(); - client.read_to_end(&mut out).await.expect("read failed"); - handle.await.expect("join failed"); + client.read_to_end(&mut out).await?; + handle.await??; let frames = decode_frames(out); - assert_eq!(frames.len(), 1, "expected a single response frame"); - let (resp, _) = serializer - .deserialize::(&frames[0]) - .expect("deserialize failed"); + let [first] = frames.as_slice() else { + return Err("expected a single response frame".into()); + }; + let (resp, _) = serializer.deserialize::(first)?; let parts = wireframe::app::Packet::into_parts(resp); let correlation_id = parts.correlation_id(); let payload = parts.payload(); - assert_eq!(payload, vec![b'X', b'A', b'B', b'B', b'A']); - assert_eq!(correlation_id, Some(7)); + assert_eq!( + payload, + [b'X', b'A', b'B', b'B', b'A'], + "unexpected payload" + ); + assert_eq!(correlation_id, Some(7), "unexpected correlation id"); + Ok(()) } diff --git a/tests/multi_packet.rs b/tests/multi_packet.rs index 391b0436..383d8fa6 100644 --- a/tests/multi_packet.rs +++ b/tests/multi_packet.rs @@ -1,7 +1,7 @@ //! Tests for multi-packet responses using channels. #![cfg(not(loom))] -use std::time::Duration; +use std::{error::Error, time::Duration}; use futures::TryStreamExt; use rstest::{fixture, rstest}; @@ -13,6 +13,16 @@ use wireframe::{ push::{PushHandle, PushQueues}, }; +mod common; +use common::TestResult; + +fn boxed_err( + context: &str, + err: E, +) -> Box { + format!("{context}: {err:?}").into() +} + #[derive(PartialEq, Debug)] struct TestMsg(u8); @@ -20,18 +30,23 @@ const CAPACITY: usize = 2; /// Provide push queues, handle, and shutdown token for connection actor tests. #[fixture] -fn actor_components() -> (PushQueues, PushHandle, CancellationToken) { +fn actor_components() -> TestResult<(PushQueues, PushHandle, CancellationToken)> { let (queues, handle) = PushQueues::::builder() .high_capacity(4) .low_capacity(4) - .build() - .expect("failed to build PushQueues"); - (queues, handle, CancellationToken::new()) + .build()?; + + Ok((queues, handle, CancellationToken::new())) } /// Drain all messages from a `FrameStream` for non-channel response variants. -async fn drain_all(stream: wireframe::FrameStream) -> Vec { - stream.try_collect::>().await.expect("stream error") +async fn drain_all( + stream: wireframe::FrameStream, +) -> TestResult> { + stream + .try_collect::>() + .await + .map_err(|err| boxed_err("stream error", err)) } /// Multi-packet responses drain every frame regardless of channel state. @@ -40,24 +55,27 @@ async fn drain_all(stream: wireframe::FrameStream) /// channel's capacity. #[rstest(count, case(0), case(1), case(2), case(CAPACITY + 1))] #[tokio::test] -async fn multi_packet_drains_all_messages(count: usize) { +async fn multi_packet_drains_all_messages(count: usize) -> TestResult { let (tx, rx) = mpsc::channel(CAPACITY); let send_task = tokio::spawn(async move { for i in 0..count { - tx.send(TestMsg(u8::try_from(i).expect("<= u8::MAX"))) - .await - .expect("send"); + tx.send(TestMsg(u8::try_from(i)?)).await?; } + Ok::<_, Box>(()) }); let resp: Response = Response::MultiPacket(rx); - let received = drain_all(resp.into_stream()).await; - send_task.await.expect("sender join"); - assert_eq!( - received, - (0..count) - .map(|i| TestMsg(u8::try_from(i).expect("<= u8::MAX"))) - .collect::>() - ); + let received = drain_all(resp.into_stream()).await?; + send_task + .await + .map_err(|e| boxed_err("send task join", e))??; + let expected = (0..count) + .map(u8::try_from) + .collect::, _>>()? + .into_iter() + .map(TestMsg) + .collect::>(); + assert_eq!(received, expected); + Ok(()) } /// Drains frames from a multi-packet channel via the connection actor. @@ -70,54 +88,46 @@ async fn multi_packet_drains_all_messages(count: usize) { #[tokio::test] async fn connection_actor_drains_multi_packet_channel( frames: Vec, - actor_components: (PushQueues, PushHandle, CancellationToken), -) { + actor_components: TestResult<(PushQueues, PushHandle, CancellationToken)>, +) -> TestResult { let capacity = frames.len().max(1); let (tx, rx) = mpsc::channel(capacity); for &value in &frames { - tx.send(value).await.expect("send frame"); + tx.send(value).await?; } drop(tx); - let (queues, handle, shutdown) = actor_components; + let (queues, handle, shutdown) = actor_components?; let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, None, shutdown); actor.set_multi_packet(Some(rx)); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| boxed_err("connection actor error", e))?; assert_eq!(out, frames); + Ok(()) } #[rstest] #[tokio::test] async fn connection_actor_interleaves_multi_packet_and_priority_frames( - actor_components: (PushQueues, PushHandle, CancellationToken), -) { - let (queues, handle, shutdown) = actor_components; + actor_components: TestResult<(PushQueues, PushHandle, CancellationToken)>, +) -> TestResult { + let (queues, handle, shutdown) = actor_components?; let multi_frames = [1_u8, 2, 3]; let (multi_tx, multi_rx) = mpsc::channel(multi_frames.len()); for &frame in &multi_frames { - multi_tx.send(frame).await.expect("send multi-packet frame"); + multi_tx.send(frame).await?; } drop(multi_tx); - handle - .push_high_priority(10) - .await - .expect("push high-priority frame"); - handle - .push_high_priority(11) - .await - .expect("push high-priority frame"); - handle - .push_low_priority(100) - .await - .expect("push low-priority frame"); - handle - .push_low_priority(101) - .await - .expect("push low-priority frame"); + handle.push_high_priority(10).await?; + handle.push_high_priority(11).await?; + handle.push_low_priority(100).await?; + handle.push_low_priority(101).await?; let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, None, shutdown); actor.set_fairness(FairnessConfig { @@ -127,17 +137,21 @@ async fn connection_actor_interleaves_multi_packet_and_priority_frames( actor.set_multi_packet(Some(multi_rx)); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| boxed_err("connection actor error", e))?; assert_eq!(out, vec![10, 100, 11, 101, 1, 2, 3]); + Ok(()) } #[rstest] #[tokio::test] async fn shutdown_completes_multi_packet_channel( - actor_components: (PushQueues, PushHandle, CancellationToken), -) { - let (queues, handle, shutdown) = actor_components; + actor_components: TestResult<(PushQueues, PushHandle, CancellationToken)>, +) -> TestResult { + let (queues, handle, shutdown) = actor_components?; let (tx, rx) = mpsc::channel(1); let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, None, shutdown); actor.set_multi_packet(Some(rx)); @@ -146,28 +160,32 @@ async fn shutdown_completes_multi_packet_channel( let join = tokio::spawn(async move { let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); - out + actor + .run(&mut out) + .await + .map_err(|e| boxed_err("connection actor error", e))?; + Ok::<_, Box>(out) }); yield_now().await; cancel.cancel(); - let out = timeout(Duration::from_millis(1000), join) + let join_result = timeout(Duration::from_millis(1000), join) .await - .expect("actor shutdown timed out") - .expect("actor task panicked"); + .map_err(|e| boxed_err("connection actor shutdown timeout", e))??; + let out = join_result?; assert!(out.is_empty()); drop(tx); + Ok(()) } #[rstest] #[tokio::test] async fn shutdown_during_active_multi_packet_send( - actor_components: (PushQueues, PushHandle, CancellationToken), -) { - let (queues, handle, shutdown) = actor_components; + actor_components: TestResult<(PushQueues, PushHandle, CancellationToken)>, +) -> TestResult { + let (queues, handle, shutdown) = actor_components?; let (tx, rx) = mpsc::channel(4); let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, None, shutdown); actor.set_multi_packet(Some(rx)); @@ -176,37 +194,57 @@ async fn shutdown_during_active_multi_packet_send( let join = tokio::spawn(async move { let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); - out + actor + .run(&mut out) + .await + .map_err(|e| boxed_err("connection actor error", e))?; + Ok::<_, Box>(out) }); - tx.send(1).await.expect("send frame"); - tx.send(2).await.expect("send frame"); + tx.send(1).await?; + tx.send(2).await?; yield_now().await; cancel.cancel(); let _ = tx.send(3).await; - let out = timeout(Duration::from_millis(1000), join) + let join_result = timeout(Duration::from_millis(1000), join) .await - .expect("actor shutdown timed out") - .expect("actor task panicked"); + .map_err(|e| boxed_err("connection actor shutdown timeout", e))??; + let out = join_result?; assert!(out.is_empty() || out == vec![1, 2], "actor output: {out:?}"); drop(tx); + Ok(()) } /// Returns an empty stream for an empty vector response. #[tokio::test] -async fn vec_empty_returns_empty_stream() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn vec_empty_returns_empty_stream() -> TestResult { let resp: Response = Response::Vec(Vec::new()); - let received = drain_all(resp.into_stream()).await; - assert!(received.is_empty()); + let received = drain_all(resp.into_stream()).await?; + assert!( + received.is_empty(), + "expected empty stream, got {received:?}" + ); + Ok(()) } /// `Response::Empty` yields no frames. #[tokio::test] -async fn empty_returns_empty_stream() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn empty_returns_empty_stream() -> TestResult { let resp: Response = Response::Empty; - let received = drain_all(resp.into_stream()).await; - assert!(received.is_empty()); + let received = drain_all(resp.into_stream()).await?; + assert!( + received.is_empty(), + "expected empty stream, got {received:?}" + ); + Ok(()) } diff --git a/tests/multi_packet_streaming.rs b/tests/multi_packet_streaming.rs index f8b7d68c..07249de5 100644 --- a/tests/multi_packet_streaming.rs +++ b/tests/multi_packet_streaming.rs @@ -7,7 +7,11 @@ //! responses to ensure correlation identifiers allow clients to demultiplex //! concurrent activity. -use std::sync::{Arc, OnceLock}; +use std::sync::{ + Arc, + OnceLock, + atomic::{AtomicBool, Ordering}, +}; use log::Level as LogLevel; use logtest as flexi_logger; @@ -16,12 +20,15 @@ use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; use wireframe::{ app::{Envelope, Packet, PacketParts}, - connection::{ConnectionActor, FairnessConfig}, + connection::{ConnectionActor, ConnectionChannels, FairnessConfig}, hooks::{ConnectionContext, ProtocolHooks}, push::{PushHandle, PushQueues}, }; use wireframe_testing::{LoggerHandle, logger}; +mod common; +use common::TestResult; + const STREAM_ID: u32 = 7; const TERMINATOR_ID: u32 = 255; @@ -37,20 +44,21 @@ struct ActorHarness { } impl ActorHarness { - fn new() -> Self { + fn new() -> TestResult { let (queues, handle) = PushQueues::::builder() .high_capacity(4) .low_capacity(4) .unlimited() - .build() - .expect("failed to build PushQueues"); + .build()?; let shared_handle: Arc>> = Arc::new(OnceLock::new()); + let duplicate_handle = Arc::new(AtomicBool::new(false)); let handle_slot = Arc::clone(&shared_handle); + let duplicate_flag = Arc::clone(&duplicate_handle); let hooks = ProtocolHooks { on_connection_setup: Some(Box::new(move |handle, _ctx| { - handle_slot - .set(handle) - .unwrap_or_else(|_| panic!("push handle already captured")); + if handle_slot.set(handle).is_err() { + duplicate_flag.store(true, Ordering::Relaxed); + } })), stream_end: Some(Box::new(|_ctx: &mut ConnectionContext| { Some(terminator_frame()) @@ -59,47 +67,63 @@ impl ActorHarness { }; let shutdown = CancellationToken::new(); - let actor = ConnectionActor::with_hooks(queues, handle, None, shutdown, hooks); + let actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + None, + shutdown, + hooks, + ); let shared_handle = - Arc::try_unwrap(shared_handle).unwrap_or_else(|_| panic!("push handle still shared")); + Arc::try_unwrap(shared_handle).map_err(|_| "push handle still shared at teardown")?; let handle = shared_handle .into_inner() - .expect("connection setup hook did not run"); + .ok_or("connection setup hook did not run")?; - Self { + if duplicate_handle.load(Ordering::Relaxed) { + return Err("push handle already captured".into()); + } + + Ok(Self { actor, handle: Some(handle), - } + }) } - fn handle(&self) -> &PushHandle { - self.handle.as_ref().expect("push handle already released") + fn handle(&self) -> TestResult<&PushHandle> { + self.handle + .as_ref() + .ok_or_else(|| "push handle already released".into()) } fn release_handle(&mut self) { self.handle.take(); } - async fn run(&mut self) -> Vec { + async fn run(&mut self) -> TestResult> { let mut out = Vec::new(); - self.actor - .run(&mut out) - .await - .expect("connection actor run failed"); - out + self.actor.run(&mut out).await.map_err(|e| { + Box::new(std::io::Error::other(format!( + "connection actor run failed: {e:?}" + ))) as Box + })?; + Ok(out) } } fn parts(frame: &Envelope) -> PacketParts { frame.clone().into_parts() } #[tokio::test] -async fn client_receives_multi_packet_stream_with_terminator() { - let mut harness = ActorHarness::new(); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn client_receives_multi_packet_stream_with_terminator() -> TestResult<()> { + let mut harness = ActorHarness::new()?; let (tx, rx) = mpsc::channel(4); let correlation = Some(88_u64); for chunk in [&[1_u8][..], &[2, 3][..]] { tx.send(envelope_with_payload(STREAM_ID, None, chunk)) .await - .expect("send frame"); + .map_err(|e| format!("send frame: {e}"))?; } drop(tx); @@ -109,15 +133,19 @@ async fn client_receives_multi_packet_stream_with_terminator() { harness.release_handle(); - let out = harness.run().await; + let out = harness.run().await?; assert_eq!(out.len(), 3, "expected two frames plus terminator"); let payloads: Vec> = out.iter().map(|frame| parts(frame).payload()).collect(); - assert_eq!(payloads[0], vec![1]); - assert_eq!(payloads[1], vec![2, 3]); + assert_eq!(payloads.first(), Some(&vec![1]), "first payload mismatch"); + assert_eq!( + payloads.get(1), + Some(&vec![2, 3]), + "second payload mismatch" + ); assert_eq!( - payloads[2], - Vec::::new(), + payloads.get(2), + Some(&Vec::::new()), "terminator payload should be empty" ); @@ -125,9 +153,10 @@ async fn client_receives_multi_packet_stream_with_terminator() { assert_eq!( parts(frame).correlation_id(), correlation, - "correlation id mismatch", + "correlation id mismatch" ); } + Ok(()) } fn is_disconnect_log(record: &flexi_logger::Record) -> bool { @@ -138,9 +167,11 @@ fn is_disconnect_log(record: &flexi_logger::Record) -> bool { #[rstest] #[tokio::test] -async fn multi_packet_logs_disconnected_when_sender_dropped(mut logger: LoggerHandle) { +async fn multi_packet_logs_disconnected_when_sender_dropped( + mut logger: LoggerHandle, +) -> TestResult<()> { logger.clear(); - let mut harness = ActorHarness::new(); + let mut harness = ActorHarness::new()?; let (tx, rx) = mpsc::channel(1); let correlation = Some(41_u64); drop(tx); @@ -152,21 +183,20 @@ async fn multi_packet_logs_disconnected_when_sender_dropped(mut logger: LoggerHa harness.actor.set_fairness(interleaving_fairness()); harness - .handle() + .handle()? .push_high_priority(envelope_with_payload(11, Some(5), b"hi")) - .await - .expect("push high priority frame"); + .await?; harness.release_handle(); - let out = harness.run().await; + let out = harness.run().await?; assert_eq!(out.len(), 2, "expected push frame followed by terminator"); let last = out.last().expect("terminator missing"); assert_eq!( parts(last).correlation_id(), correlation, - "terminator correlation mismatch", + "terminator correlation mismatch" ); let mut saw_disconnect = false; @@ -177,6 +207,7 @@ async fn multi_packet_logs_disconnected_when_sender_dropped(mut logger: LoggerHa } } assert!(saw_disconnect, "missing disconnect log"); + Ok(()) } struct FrameSpec { @@ -227,32 +258,33 @@ async fn push_sequence( handle: &PushHandle, priority: PushPriority, frames: &[FrameSpec], -) { +) -> TestResult<()> { for spec in frames { let envelope = envelope_with_payload(spec.id, Some(spec.correlation), spec.payload); - let result = match priority { - PushPriority::High => handle.push_high_priority(envelope).await, - PushPriority::Low => handle.push_low_priority(envelope).await, - }; - result.expect("push frame"); + match priority { + PushPriority::High => handle.push_high_priority(envelope).await?, + PushPriority::Low => handle.push_low_priority(envelope).await?, + } } + Ok(()) } -async fn setup_stream_channel(payloads: &[&[u8]]) -> mpsc::Receiver { +async fn setup_stream_channel(payloads: &[&[u8]]) -> TestResult> { let capacity = payloads.len().max(1); let (tx, rx) = mpsc::channel(capacity); for payload in payloads { tx.send(envelope_with_payload(STREAM_ID, None, payload)) .await - .expect("send frame to multi-packet stream"); + .map_err(|e| format!("send frame to multi-packet stream: {e}"))?; } drop(tx); - rx + Ok(rx) } -async fn push_interleaved_frames(handle: &PushHandle) { - push_sequence(handle, PushPriority::High, &HIGH_PRIORITY_FRAMES).await; - push_sequence(handle, PushPriority::Low, &LOW_PRIORITY_FRAMES).await; +async fn push_interleaved_frames(handle: &PushHandle) -> TestResult<()> { + push_sequence(handle, PushPriority::High, &HIGH_PRIORITY_FRAMES).await?; + push_sequence(handle, PushPriority::Low, &LOW_PRIORITY_FRAMES).await?; + Ok(()) } fn assert_correlation_ordering(frames: &[Envelope], expected: &[Option]) { @@ -272,20 +304,20 @@ fn assert_frame_identities(frames: &[Envelope], expected: &[u32]) { } #[tokio::test] -async fn interleaved_multi_packet_and_push_frames_preserve_correlations() { - let mut harness = ActorHarness::new(); +async fn interleaved_multi_packet_and_push_frames_preserve_correlations() -> TestResult<()> { + let mut harness = ActorHarness::new()?; let stream_correlation = Some(73_u64); - let rx = setup_stream_channel(&[&[10_u8][..], &[20][..], &[30][..]]).await; + let rx = setup_stream_channel(&[&[10_u8][..], &[20][..], &[30][..]]).await?; harness .actor .set_multi_packet_with_correlation(Some(rx), stream_correlation); harness.actor.set_fairness(interleaving_fairness()); - push_interleaved_frames(harness.handle()).await; + push_interleaved_frames(harness.handle()?).await?; harness.release_handle(); - let frames = harness.run().await; + let frames = harness.run().await?; assert_correlation_ordering( &frames, @@ -305,4 +337,5 @@ async fn interleaved_multi_packet_and_push_frames_preserve_correlations() { &frames, &[2, 3, 4, 5, STREAM_ID, STREAM_ID, STREAM_ID, TERMINATOR_ID], ); + Ok(()) } diff --git a/tests/preamble.rs b/tests/preamble.rs index 1a2a9890..fc5a9edf 100644 --- a/tests/preamble.rs +++ b/tests/preamble.rs @@ -2,6 +2,7 @@ //! Tests for connection preamble reading. use std::{ + error::Error, io, sync::{Arc, Mutex}, }; @@ -9,7 +10,7 @@ use std::{ use bincode::error::DecodeError; use futures::future::BoxFuture; mod common; -use common::{factory, unused_listener}; +use common::{TestResult, factory, unused_listener}; use rstest::rstest; use tokio::{ io::{AsyncReadExt, AsyncWriteExt, duplex}, @@ -65,16 +66,18 @@ where } /// Run the provided server while executing `block`. -async fn with_running_server(server: WireframeServer, block: B) +async fn with_running_server(server: WireframeServer, block: B) -> TestResult where F: Fn() -> WireframeApp + Send + Sync + Clone + 'static, T: wireframe::preamble::Preamble, - Fut: std::future::Future, + Fut: std::future::Future, B: FnOnce(std::net::SocketAddr) -> Fut, { let listener = unused_listener(); - let server = server.bind_existing_listener(listener).expect("bind"); - let addr = server.local_addr().expect("addr"); + let server = server.bind_existing_listener(listener)?; + let addr = server + .local_addr() + .ok_or_else(|| Box::::from("server missing local addr"))?; let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); let handle = tokio::spawn(async move { server @@ -82,40 +85,49 @@ where let _ = shutdown_rx.await; }) .await - .expect("server run failed"); }); - block(addr).await; + block(addr).await?; let _ = shutdown_tx.send(()); - handle.await.expect("server join failed"); + let run_result = handle.await?; + run_result?; + Ok(()) } #[tokio::test] -async fn parse_valid_preamble() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn parse_valid_preamble() -> TestResult { let (mut client, mut server) = duplex(64); let bytes = b"TRTPHOTL\x00\x01\x00\x02"; - client.write_all(bytes).await.expect("write failed"); - client.shutdown().await.expect("shutdown failed"); - let (p, _) = read_preamble::<_, HotlinePreamble>(&mut server) - .await - .expect("valid preamble"); - eprintln!("decoded: {p:?}"); - p.validate().expect("preamble validation failed"); - assert_eq!(p.magic, HotlinePreamble::MAGIC); - assert_eq!(p.min_version, 1); - assert_eq!(p.client_version, 2); + client.write_all(bytes).await?; + client.shutdown().await?; + let (p, _) = read_preamble::<_, HotlinePreamble>(&mut server).await?; + p.validate()?; + assert_eq!(p.magic, HotlinePreamble::MAGIC, "preamble magic mismatch"); + assert_eq!(p.min_version, 1, "preamble minimum version mismatch"); + assert_eq!(p.client_version, 2, "preamble client version mismatch"); + Ok(()) } #[tokio::test] -async fn invalid_magic_is_error() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn invalid_magic_is_error() -> TestResult { let (mut client, mut server) = duplex(64); let bytes = b"WRONGMAG\x00\x01\x00\x02"; - client.write_all(bytes).await.expect("write failed"); - client.shutdown().await.expect("shutdown failed"); - let (preamble, _) = read_preamble::<_, HotlinePreamble>(&mut server) - .await - .expect("decoded"); - assert!(preamble.validate().is_err()); + client.write_all(bytes).await?; + client.shutdown().await?; + let (preamble, _) = read_preamble::<_, HotlinePreamble>(&mut server).await?; + assert!( + preamble.validate().is_err(), + "invalid magic should fail validation" + ); + Ok(()) } #[derive(Clone, Copy)] @@ -132,7 +144,7 @@ async fn server_triggers_expected_callback( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, #[case] bytes: &'static [u8], #[case] expected: ExpectedCallback, -) { +) -> TestResult { let (success_tx, success_rx) = tokio::sync::oneshot::channel::(); let (failure_tx, failure_rx) = tokio::sync::oneshot::channel::<()>(); let success_tx = std::sync::Arc::new(std::sync::Mutex::new(Some(success_tx))); @@ -145,10 +157,10 @@ async fn server_triggers_expected_callback( let success_tx = success_tx.clone(); let clone = p.clone(); Box::pin(async move { - if let Some(tx) = success_tx.lock().expect("lock poisoned").take() { + if let Some(tx) = take_sender_io(&success_tx)? { let _ = tx.send(clone); } - Ok(()) + Ok::<(), io::Error>(()) }) } }, @@ -157,28 +169,26 @@ async fn server_triggers_expected_callback( move |_, _| { let failure_tx = failure_tx.clone(); Box::pin(async move { - if let Some(tx) = failure_tx.lock().expect("lock poisoned").take() { + if let Some(tx) = take_sender_io(&failure_tx)? { let _ = tx.send(()); } - Ok(()) + Ok::<(), io::Error>(()) }) } }, ); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); - stream.write_all(bytes).await.expect("write failed"); - stream.shutdown().await.expect("shutdown failed"); + let mut stream = TcpStream::connect(addr).await?; + stream.write_all(bytes).await?; + stream.shutdown().await?; + Ok(()) }) - .await; + .await?; match expected { ExpectedCallback::Success => { - let preamble = timeout(Duration::from_secs(1), success_rx) - .await - .expect("timeout waiting for success") - .expect("success send"); + let preamble = timeout(Duration::from_secs(1), success_rx).await??; assert_eq!(preamble.magic, HotlinePreamble::MAGIC); assert!( timeout(Duration::from_millis(500), failure_rx) @@ -187,10 +197,7 @@ async fn server_triggers_expected_callback( ); } ExpectedCallback::Failure => { - timeout(Duration::from_secs(1), failure_rx) - .await - .expect("timeout waiting for failure") - .expect("failure send"); + timeout(Duration::from_secs(1), failure_rx).await??; assert!( timeout(Duration::from_millis(500), success_rx) .await @@ -198,78 +205,80 @@ async fn server_triggers_expected_callback( ); } } + Ok(()) } #[rstest] #[tokio::test] async fn success_callback_can_write_response( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, -) { +) -> TestResult { let server = server_with_handlers( factory, |_, stream| { Box::pin(async move { - stream.write_all(b"ACK").await.expect("write failed"); - stream.flush().await.expect("flush failed"); - Ok(()) + stream.write_all(b"ACK").await?; + stream.flush().await?; + Ok::<(), io::Error>(()) }) }, |_, _| Box::pin(async { Ok::<(), io::Error>(()) }), ); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); + let mut stream = TcpStream::connect(addr).await?; let bytes = b"TRTPHOTL\x00\x01\x00\x02"; - stream.write_all(bytes).await.expect("write failed"); + stream.write_all(bytes).await?; let mut buf = [0u8; 3]; - stream.read_exact(&mut buf).await.expect("read failed"); + stream.read_exact(&mut buf).await?; assert_eq!(&buf, b"ACK"); + Ok(()) }) - .await; + .await?; + Ok(()) } #[rstest] #[tokio::test] async fn failure_callback_can_write_response( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, -) { +) -> TestResult { let (failure_holder, failure_rx) = channel_holder(); let server = WireframeServer::new(factory) .with_preamble::() .on_preamble_decode_failure(move |_, stream| { let failure_holder = failure_holder.clone(); Box::pin(async move { - stream.write_all(b"ERR").await.expect("write failed"); - stream.flush().await.expect("flush failed"); - if let Some(tx) = failure_holder.lock().expect("lock").take() { + stream.write_all(b"ERR").await?; + stream.flush().await?; + if let Some(tx) = take_sender_io(&failure_holder)? { let _ = tx.send(()); } - Ok(()) + Ok::<(), io::Error>(()) }) }); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); - stream.write_all(b"BAD").await.expect("write failed"); - stream.shutdown().await.expect("shutdown failed"); + let mut stream = TcpStream::connect(addr).await?; + stream.write_all(b"BAD").await?; + stream.shutdown().await?; let mut buf = [0u8; 3]; let read = timeout(Duration::from_secs(1), stream.read_exact(&mut buf)).await; - let result = read.expect("timeout waiting for failure handler"); - result.expect("read error"); + let result = read?; + result?; assert_eq!(&buf, b"ERR"); - timeout(Duration::from_millis(200), failure_rx) - .await - .expect("timeout waiting for failure callback") - .expect("failure callback send"); + recv_within(Duration::from_millis(200), failure_rx).await?; + Ok(()) }) - .await; + .await?; + Ok(()) } #[rstest] #[tokio::test] async fn preamble_timeout_invokes_failure_handler_and_closes_connection( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, -) { +) -> TestResult { let (failure_holder, failure_rx) = channel_holder(); let server = WireframeServer::new(factory) .with_preamble::() @@ -285,45 +294,41 @@ async fn preamble_timeout_invokes_failure_handler_and_closes_connection( ), "expected timed out error, got {err:?}" ); - stream.write_all(b"ERR").await.expect("write failed"); - stream.flush().await.expect("flush failed"); - stream.shutdown().await.expect("shutdown failed"); - if let Some(tx) = failure_holder.lock().expect("lock").take() { + stream.write_all(b"ERR").await?; + stream.flush().await?; + stream.shutdown().await?; + if let Some(tx) = take_sender_io(&failure_holder)? { let _ = tx.send(()); } - Ok(()) + Ok::<(), io::Error>(()) }) }); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); - timeout(Duration::from_secs(1), failure_rx) - .await - .expect("timeout waiting for failure callback") - .expect("failure callback send"); + let mut stream = TcpStream::connect(addr).await?; + recv_within(Duration::from_secs(1), failure_rx).await?; let mut buf = [0u8; 3]; - timeout(Duration::from_millis(500), stream.read_exact(&mut buf)) - .await - .expect("did not receive timeout response in time") - .expect("read timeout response failed"); + timeout(Duration::from_millis(500), stream.read_exact(&mut buf)).await??; assert_eq!(&buf, b"ERR"); let mut eof = [0u8; 1]; let read = timeout(Duration::from_millis(200), stream.read(&mut eof)).await; - match read.expect("timeout waiting for close") { + match read? { Ok(0) => {} Ok(n) => panic!("expected connection to close, read {n} bytes"), Err(e) if e.kind() == io::ErrorKind::ConnectionReset => {} Err(e) => panic!("unexpected read error: {e:?}"), } + Ok(()) }) - .await; + .await?; + Ok(()) } #[rstest] #[tokio::test] async fn success_handler_runs_without_failure_handler( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, -) { +) -> TestResult { let (success_tx, success_rx) = tokio::sync::oneshot::channel::(); let success_tx = Arc::new(Mutex::new(Some(success_tx))); let server = WireframeServer::new(factory) @@ -334,33 +339,32 @@ async fn success_handler_runs_without_failure_handler( let success_tx = success_tx.clone(); let preamble = p.clone(); Box::pin(async move { - if let Some(tx) = success_tx.lock().expect("lock").take() { + if let Some(tx) = take_sender_io(&success_tx)? { let _ = tx.send(preamble); } - Ok(()) + Ok::<(), io::Error>(()) }) } }); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); + let mut stream = TcpStream::connect(addr).await?; let bytes = b"TRTPHOTL\x00\x01\x00\x02"; - stream.write_all(bytes).await.expect("write failed"); - stream.shutdown().await.expect("shutdown failed"); - let preamble = timeout(Duration::from_secs(1), success_rx) - .await - .expect("timeout waiting for success") - .expect("success send"); + stream.write_all(bytes).await?; + stream.shutdown().await?; + let preamble = recv_within(Duration::from_secs(1), success_rx).await?; assert_eq!(preamble.magic, HotlinePreamble::MAGIC); + Ok(()) }) - .await; + .await?; + Ok(()) } #[rstest] #[tokio::test] async fn preamble_timeout_allows_timely_preamble( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, -) { +) -> TestResult { let (success_holder, success_rx) = channel_holder(); let (failure_holder, failure_rx) = channel_holder(); let server = WireframeServer::new(factory) @@ -372,14 +376,14 @@ async fn preamble_timeout_allows_timely_preamble( let success_holder = success_holder.clone(); let clone = p.clone(); Box::pin(async move { - if let Some(tx) = success_holder.lock().expect("lock").take() { + if let Some(tx) = take_sender_io(&success_holder)? { let _ = tx.send(()); } - stream.write_all(b"OK").await.expect("write failed"); - stream.flush().await.expect("flush failed"); + stream.write_all(b"OK").await?; + stream.flush().await?; // keep connection open by not shutting down here assert_eq!(clone.magic, HotlinePreamble::MAGIC); - Ok(()) + Ok::<(), io::Error>(()) }) } }) @@ -388,23 +392,20 @@ async fn preamble_timeout_allows_timely_preamble( move |_, _| { let failure_holder = failure_holder.clone(); Box::pin(async move { - if let Some(tx) = failure_holder.lock().expect("lock").take() { + if let Some(tx) = take_sender_io(&failure_holder)? { let _ = tx.send(()); } - Ok(()) + Ok::<(), io::Error>(()) }) } }); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); + let mut stream = TcpStream::connect(addr).await?; let bytes = b"TRTPHOTL\x00\x01\x00\x02"; - stream.write_all(bytes).await.expect("write failed"); + stream.write_all(bytes).await?; - timeout(Duration::from_millis(200), success_rx) - .await - .expect("timeout waiting for success") - .expect("success send"); + recv_within(Duration::from_millis(200), success_rx).await?; assert!( timeout(Duration::from_millis(150), failure_rx) .await @@ -413,27 +414,26 @@ async fn preamble_timeout_allows_timely_preamble( ); let mut buf = [0u8; 2]; - stream - .read_exact(&mut buf) - .await - .expect("expected response from success handler"); + stream.read_exact(&mut buf).await?; assert_eq!(&buf, b"OK"); + Ok(()) }) - .await; + .await?; + Ok(()) } #[rstest] #[tokio::test] async fn failure_handler_error_is_logged_and_connection_closes( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, -) { +) -> TestResult { let (failure_holder, failure_rx) = channel_holder(); let server = WireframeServer::new(factory) .with_preamble::() .on_preamble_decode_failure(move |_, _| { let failure_holder = failure_holder.clone(); Box::pin(async move { - if let Some(tx) = failure_holder.lock().expect("lock").take() { + if let Some(tx) = take_sender_io(&failure_holder)? { let _ = tx.send(()); } Err::<(), io::Error>(io::Error::other("boom")) @@ -441,25 +441,24 @@ async fn failure_handler_error_is_logged_and_connection_closes( }); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); - stream.write_all(b"BAD").await.expect("write failed"); - stream.shutdown().await.expect("shutdown failed"); + let mut stream = TcpStream::connect(addr).await?; + stream.write_all(b"BAD").await?; + stream.shutdown().await?; - timeout(Duration::from_secs(1), failure_rx) - .await - .expect("failure handler not invoked") - .expect("failure handler send failed"); + recv_within(Duration::from_secs(1), failure_rx).await?; let mut buf = [0u8; 1]; let read = timeout(Duration::from_millis(200), stream.read(&mut buf)).await; - match read.expect("timeout waiting for close") { + match read? { Ok(0) => {} Ok(n) => panic!("expected connection close, read {n} bytes"), Err(e) if e.kind() == io::ErrorKind::ConnectionReset => {} Err(e) => panic!("unexpected read error: {e:?}"), } + Ok(()) }) - .await; + .await?; + Ok(()) } #[derive(Debug, Clone, Copy, PartialEq, Eq, bincode::Encode, bincode::Decode)] @@ -472,6 +471,17 @@ fn channel_holder() -> (Holder, oneshot::Receiver<()>) { (Arc::new(Mutex::new(Some(tx))), rx) } +fn take_sender_io(holder: &Mutex>) -> io::Result> { + holder + .lock() + .map_err(|e| io::Error::other(format!("lock poisoned: {e}"))) + .map(|mut guard| guard.take()) +} + +async fn recv_within(duration: Duration, rx: oneshot::Receiver) -> TestResult { + Ok(timeout(duration, rx).await??) +} + fn success_cb

( holder: Arc>>>, ) -> impl for<'a> Fn(&'a P, &'a mut TcpStream) -> BoxFuture<'a, io::Result<()>> + Send + Sync + 'static @@ -479,7 +489,7 @@ fn success_cb

( move |_, _| { let holder = holder.clone(); Box::pin(async move { - if let Some(tx) = holder.lock().expect("lock").take() { + if let Some(tx) = take_sender_io(&holder)? { let _ = tx.send(()); } Ok(()) @@ -496,7 +506,7 @@ fn failure_cb( move |_, _| { let holder = holder.clone(); Box::pin(async move { - if let Some(tx) = holder.lock().expect("lock").take() { + if let Some(tx) = take_sender_io(&holder)? { let _ = tx.send(()); } Ok(()) @@ -508,7 +518,7 @@ fn failure_cb( #[tokio::test] async fn callbacks_dropped_when_overriding_preamble( factory: impl Fn() -> WireframeApp + Send + Sync + Clone + 'static, -) { +) -> TestResult { let (hotline_success, hotline_success_rx) = channel_holder(); let (hotline_failure, hotline_failure_rx) = channel_holder(); let (other_success, other_success_rx) = channel_holder(); @@ -523,21 +533,19 @@ async fn callbacks_dropped_when_overriding_preamble( .on_preamble_decode_failure(failure_cb(other_failure.clone())); with_running_server(server, |addr| async move { - let mut stream = TcpStream::connect(addr).await.expect("connect failed"); + let mut stream = TcpStream::connect(addr).await?; let config = bincode::config::standard() .with_big_endian() .with_fixed_int_encoding(); - let mut bytes = bincode::encode_to_vec(OtherPreamble(1), config).expect("encode preamble"); + let mut bytes = bincode::encode_to_vec(OtherPreamble(1), config)?; bytes.resize(8, 0); - stream.write_all(&bytes).await.expect("write failed"); - stream.shutdown().await.expect("shutdown failed"); + stream.write_all(&bytes).await?; + stream.shutdown().await?; // Wait for the success callback before shutting down the server. - timeout(Duration::from_secs(1), other_success_rx) - .await - .expect("timeout waiting for other success") - .expect("other success send"); + recv_within(Duration::from_secs(1), other_success_rx).await?; + Ok(()) }) - .await; + .await?; assert!( timeout(Duration::from_millis(500), other_failure_rx) .await @@ -556,4 +564,5 @@ async fn callbacks_dropped_when_overriding_preamble( .is_err(), "hotline failure callback invoked", ); + Ok(()) } diff --git a/tests/push.rs b/tests/push.rs index 1c70d15a..0da2f2fd 100644 --- a/tests/push.rs +++ b/tests/push.rs @@ -19,21 +19,21 @@ use wireframe::push::{ }; use wireframe_testing::{push_expect, recv_expect}; +mod common; +use common::TestResult; + #[fixture] -fn queues() -> (PushQueues, PushHandle) { +fn queues() -> Result<(PushQueues, PushHandle), PushConfigError> { support::builder::() .high_capacity(2) .low_capacity(2) .rate(Some(1)) .build() - .expect("failed to build PushQueues") } #[fixture] -fn small_queues() -> (PushQueues, PushHandle) { - support::builder::() - .build() - .expect("failed to build PushQueues") +fn small_queues() -> Result<(PushQueues, PushHandle), PushConfigError> { + support::builder::().build() } /// Builder rejects rates outside the supported range. @@ -61,23 +61,28 @@ fn builder_accepts_max_rate() { /// Disabling throttling allows rapid bursts to succeed. #[tokio::test] -async fn disables_throttling_allows_burst_pushes() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn disables_throttling_allows_burst_pushes() -> TestResult<()> { time::pause(); let (_queues, handle) = support::builder::() .high_capacity(20) .low_capacity(20) .unlimited() - .build() - .expect("failed to build PushQueues"); + .build()?; for i in 0u8..10 { push_expect!(handle.push_high_priority(i)); push_expect!(handle.push_low_priority(i)); } let res = time::timeout(Duration::from_millis(10), handle.push_high_priority(99)).await; + let push_res = res.expect("push should not block when throttling disabled"); assert!( - res.is_ok(), - "push should not block when throttling disabled" + push_res.is_ok(), + "push should not error when throttling disabled" ); + Ok(()) } #[test] @@ -106,8 +111,12 @@ fn builder_rejects_zero_capacity() { /// Frames are delivered to queues matching their push priority. #[tokio::test] -async fn frames_routed_to_correct_priority_queues() { - let (mut queues, handle) = small_queues(); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn frames_routed_to_correct_priority_queues() -> TestResult<()> { + let (mut queues, handle) = small_queues()?; push_expect!(handle.push_low_priority(1u8)); push_expect!(handle.push_high_priority(2u8)); @@ -115,10 +124,19 @@ async fn frames_routed_to_correct_priority_queues() { let (prio1, frame1) = recv_expect!(queues.recv()); let (prio2, frame2) = recv_expect!(queues.recv()); - assert_eq!(prio1, PushPriority::High); - assert_eq!(frame1, 2); - assert_eq!(prio2, PushPriority::Low); - assert_eq!(frame2, 1); + assert_eq!( + prio1, + PushPriority::High, + "first frame should be high priority" + ); + assert_eq!(frame1, 2, "unexpected first frame value"); + assert_eq!( + prio2, + PushPriority::Low, + "second frame should be low priority" + ); + assert_eq!(frame2, 1, "unexpected second frame value"); + Ok(()) } /// `try_push` honours the selected queue policy when full. @@ -126,30 +144,49 @@ async fn frames_routed_to_correct_priority_queues() { /// Using [`PushPolicy::ReturnErrorIfFull`] causes `try_push` to /// return [`PushError::QueueFull`] once the queue is at capacity. #[tokio::test] -async fn try_push_respects_policy() { - let (mut queues, handle) = small_queues(); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn try_push_respects_policy() -> TestResult<()> { + let (mut queues, handle) = small_queues()?; push_expect!(handle.push_high_priority(1u8)); let result = handle.try_push(2u8, PushPriority::High, PushPolicy::ReturnErrorIfFull); - assert!(matches!(result, Err(PushError::QueueFull))); + assert!( + matches!(result, Err(PushError::QueueFull)), + "expected queue full error" + ); // drain queue to allow new push let _ = queues.recv().await; push_expect!(handle.push_high_priority(3u8)); let (_, last) = recv_expect!(queues.recv()); - assert_eq!(last, 3); + assert_eq!(last, 3, "unexpected drained frame"); + Ok(()) } /// Push attempts return `Closed` when all queues have been shut down. #[tokio::test] -async fn push_queues_error_on_closed() { - let (mut queues, handle) = small_queues(); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn push_queues_error_on_closed() -> TestResult<()> { + let (mut queues, handle) = small_queues()?; queues.close(); let res = handle.push_high_priority(42u8).await; - assert!(matches!(res, Err(PushError::Closed))); + assert!( + matches!(res, Err(PushError::Closed)), + "expected closed error on high priority push" + ); let res = handle.push_low_priority(24u8).await; - assert!(matches!(res, Err(PushError::Closed))); + assert!( + matches!(res, Err(PushError::Closed)), + "expected closed error on low priority push" + ); + Ok(()) } /// A push beyond the configured rate is blocked. @@ -159,9 +196,9 @@ async fn push_queues_error_on_closed() { #[case::high(PushPriority::High)] #[case::low(PushPriority::Low)] #[tokio::test] -async fn rate_limiter_blocks_when_exceeded(#[case] priority: PushPriority) { +async fn rate_limiter_blocks_when_exceeded(#[case] priority: PushPriority) -> TestResult<()> { time::pause(); - let (mut queues, handle) = queues(); + let (mut queues, handle) = queues()?; match priority { PushPriority::High => push_expect!(handle.push_high_priority(1u8)), @@ -186,30 +223,44 @@ async fn rate_limiter_blocks_when_exceeded(#[case] priority: PushPriority) { let (_, first) = recv_expect!(queues.recv()); let (_, second) = recv_expect!(queues.recv()); - assert_eq!((first, second), (1, 3)); + assert_eq!( + (first, second), + (1, 3), + "unexpected drained frames under rate limit" + ); + Ok(()) } /// Exceeding the rate limit succeeds after the window has passed. #[tokio::test] -async fn rate_limiter_allows_after_wait() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn rate_limiter_allows_after_wait() -> TestResult<()> { time::pause(); - let (mut queues, handle) = queues(); + let (mut queues, handle) = queues()?; push_expect!(handle.push_high_priority(1u8)); time::advance(Duration::from_secs(1)).await; push_expect!(handle.push_high_priority(2u8)); let (_, a) = recv_expect!(queues.recv()); let (_, b) = recv_expect!(queues.recv()); - assert_eq!((a, b), (1, 2)); + assert_eq!((a, b), (1, 2), "unexpected frame ordering after wait"); + Ok(()) } /// The limiter counts pushes from all priority queues. /// The token bucket is shared, so pushes from one priority reduce /// the allowance for the other. #[tokio::test] -async fn rate_limiter_shared_across_priorities() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn rate_limiter_shared_across_priorities() -> TestResult<()> { time::pause(); - let (mut queues, handle) = queues(); + let (mut queues, handle) = queues()?; push_expect!(handle.push_high_priority(1u8)); let mut fut = handle.push_low_priority(2u8).boxed(); @@ -224,40 +275,46 @@ async fn rate_limiter_shared_across_priorities() { let (prio1, frame1) = recv_expect!(queues.recv()); let (prio2, frame2) = recv_expect!(queues.recv()); - assert_eq!(prio1, PushPriority::High); - assert_eq!(frame1, 1); - assert_eq!(prio2, PushPriority::Low); - assert_eq!(frame2, 2); + assert_eq!(prio1, PushPriority::High, "first priority should be high"); + assert_eq!(frame1, 1, "unexpected first frame value"); + assert_eq!(prio2, PushPriority::Low, "second priority should be low"); + assert_eq!(frame2, 2, "unexpected second frame value"); + Ok(()) } /// Unlimited queues never block pushes. #[tokio::test] -async fn unlimited_queues_do_not_block() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn unlimited_queues_do_not_block() -> TestResult<()> { time::pause(); - let (mut queues, handle) = support::builder::() - .unlimited() - .build() - .expect("failed to build PushQueues"); + let (mut queues, handle) = support::builder::().unlimited().build()?; push_expect!(handle.push_high_priority(1u8)); let res = time::timeout(Duration::from_millis(10), handle.push_low_priority(2u8)).await; assert!(res.is_ok(), "pushes should not block when unlimited"); let (_, a) = recv_expect!(queues.recv()); let (_, b) = recv_expect!(queues.recv()); - assert_eq!((a, b), (1, 2)); + assert_eq!((a, b), (1, 2), "unexpected ordering for unlimited queues"); + Ok(()) } /// A burst up to capacity succeeds and further pushes are blocked. /// The maximum burst size equals the configured `capacity` parameter. #[tokio::test] -async fn rate_limiter_allows_burst_within_capacity_and_blocks_excess() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn rate_limiter_allows_burst_within_capacity_and_blocks_excess() -> TestResult<()> { time::pause(); let (mut queues, handle) = support::builder::() .high_capacity(4) .low_capacity(4) .rate(Some(3)) - .build() - .expect("failed to build PushQueues"); + .build()?; for i in 0u8..3 { push_expect!(handle.push_high_priority(i)); @@ -275,6 +332,10 @@ async fn rate_limiter_allows_burst_within_capacity_and_blocks_excess() { for expected in [0u8, 1u8, 2u8, 100u8] { let (_, frame) = recv_expect!(queues.recv()); - assert_eq!(frame, expected); + assert_eq!( + frame, expected, + "frames drained in unexpected order: expected {expected}, got {frame}" + ); } + Ok(()) } diff --git a/tests/push_policies.rs b/tests/push_policies.rs index cb181446..cdc6d548 100644 --- a/tests/push_policies.rs +++ b/tests/push_policies.rs @@ -3,63 +3,93 @@ mod support; +use std::io; + use futures::{FutureExt, future::BoxFuture}; use rstest::{fixture, rstest}; use serial_test::serial; -use tokio::{runtime::Runtime, sync::mpsc}; +use tokio::sync::mpsc; use wireframe::push::{PushPolicy, PushPriority, PushQueuesBuilder}; use wireframe_testing::{LoggerHandle, logger}; -/// Builds a single-thread [`Runtime`] for async tests. -#[fixture] -fn rt() -> Runtime { - tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("failed to build test runtime") -} +mod common; +use common::TestResult; +#[expect( + clippy::allow_attributes, + reason = "rstest single-line fixtures need allow to avoid unfulfilled lint expectations" +)] +#[allow( + unfulfilled_lint_expectations, + reason = "rstest occasionally misses the expected lint for single-line fixtures on stable" +)] #[expect( unused_braces, reason = "rustc false positive for single-line rstest fixtures" )] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] #[fixture] fn builder() -> PushQueuesBuilder { support::builder::() } +#[derive(Clone, Copy)] +struct PolicyCase { + policy: PushPolicy, + expect_warning: bool, + expected_msg: &'static str, +} + +type DlqSetup = fn(&mpsc::Sender, &mut Option>) -> TestResult<()>; +type DlqAssertion = for<'a> fn(&'a mut Option>) -> BoxFuture<'a, TestResult<()>>; + +#[derive(Clone, Copy)] +struct DlqCase { + setup: DlqSetup, + policy: PushPolicy, + assertion: DlqAssertion, + expected: &'static str, +} + /// Verifies how queue policies log and drop when the queue is full. #[rstest] -#[case::drop_if_full(PushPolicy::DropIfFull, false, "push queue full")] -#[case::warn_and_drop(PushPolicy::WarnAndDropIfFull, true, "push queue full")] +#[case::drop_if_full(PolicyCase { policy: PushPolicy::DropIfFull, expect_warning: false, expected_msg: "push queue full" })] +#[case::warn_and_drop(PolicyCase { policy: PushPolicy::WarnAndDropIfFull, expect_warning: true, expected_msg: "push queue full" })] #[serial(push_policies)] fn push_policy_behaviour( - rt: Runtime, mut logger: LoggerHandle, builder: PushQueuesBuilder, - #[case] policy: PushPolicy, - #[case] expect_warning: bool, - #[case] expected_msg: &str, -) { - rt.block_on(async { + #[case] case: PolicyCase, +) -> TestResult { + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build()?; + let PolicyCase { + policy, + expect_warning, + expected_msg, + } = case; + rt.block_on(async move { while logger.pop().is_some() {} - let (mut queues, handle) = builder.build().expect("failed to build PushQueues"); + let (mut queues, handle) = builder + .build() + .map_err(|e| io::Error::other(format!("build queues failed: {e}")))?; handle .push_high_priority(1u8) .await - .expect("push high priority failed"); + .map_err(|e| io::Error::other(format!("push high priority failed: {e}")))?; handle .try_push(2u8, PushPriority::High, policy) - .expect("try_push failed"); + .map_err(|e| io::Error::other(format!("try_push failed: {e}")))?; - let (_, val) = queues.recv().await.expect("recv failed"); - assert_eq!(val, 1); - assert!( - queues.recv().now_or_never().is_none(), - "queue should be empty" - ); + let (_, val) = queues + .recv() + .await + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "recv failed"))?; + if val != 1 { + return Err(io::Error::other("unexpected value dequeued").into()); + } + if queues.recv().now_or_never().is_some() { + return Err(io::Error::other("queue should be empty").into()); + } let mut found_warning = false; while let Some(record) = logger.pop() { @@ -69,109 +99,164 @@ fn push_policy_behaviour( } if expect_warning { - assert!(found_warning, "warning log not found"); - } else { - assert!(!found_warning, "unexpected warning log found"); + if !found_warning { + return Err(io::Error::other("warning log not found").into()); + } + } else if found_warning { + return Err(io::Error::other("unexpected warning log found").into()); } - }); + Ok::<(), Box>(()) + })?; + Ok(()) } /// Dropped frames are forwarded to the dead letter queue. #[rstest] -fn dropped_frame_goes_to_dlq(rt: Runtime, builder: PushQueuesBuilder) { - rt.block_on(async { +fn dropped_frame_goes_to_dlq(builder: PushQueuesBuilder) -> TestResult { + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build()?; + rt.block_on(async move { let (dlq_tx, mut dlq_rx) = mpsc::channel(1); let (mut queues, handle) = builder .unlimited() .dlq(Some(dlq_tx)) .build() - .expect("failed to build PushQueues"); + .map_err(|e| io::Error::other(format!("build queues failed: {e}")))?; handle .push_high_priority(1u8) .await - .expect("push high priority failed"); + .map_err(|e| io::Error::other(format!("push high priority failed: {e}")))?; handle .try_push(2u8, PushPriority::High, PushPolicy::DropIfFull) - .expect("try_push failed"); + .map_err(|e| io::Error::other(format!("try_push failed: {e}")))?; - let (_, val) = queues.recv().await.expect("recv failed"); - assert_eq!(val, 1); - assert_eq!(dlq_rx.recv().await.expect("dlq recv failed"), 2); - }); + let (_, val) = queues + .recv() + .await + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "recv failed"))?; + if val != 1 { + return Err(io::Error::other("unexpected dequeued value").into()); + } + let dlq_val = dlq_rx + .recv() + .await + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "dlq recv failed"))?; + if dlq_val != 2 { + return Err(io::Error::other("unexpected DLQ value").into()); + } + Ok::<(), Box>(()) + })?; + Ok(()) } /// Preloads the DLQ to simulate a full queue. -fn fill_dlq(tx: &mpsc::Sender, _rx: &mut Option>) { - tx.try_send(99).expect("send failed"); +fn fill_dlq(tx: &mpsc::Sender, _rx: &mut Option>) -> TestResult<()> { + tx.try_send(99) + .map_err(|e| io::Error::other(format!("send failed: {e}")))?; + Ok(()) } /// Drops the receiver to simulate a closed DLQ channel. -fn close_dlq(_: &mpsc::Sender, rx: &mut Option>) { drop(rx.take()); } +fn close_dlq(_: &mpsc::Sender, rx: &mut Option>) -> TestResult<()> { + if rx.is_none() { + return Err("DLQ receiver missing".into()); + } + drop(rx.take()); + Ok(()) +} /// Asserts that one message is queued and the DLQ then reports empty. -fn assert_dlq_full(rx: &mut Option>) -> BoxFuture<'_, ()> { +fn assert_dlq_full(rx: &mut Option>) -> BoxFuture<'_, TestResult<()>> { Box::pin(async move { - let receiver = rx.as_mut().expect("receiver missing"); - assert_eq!(receiver.recv().await.expect("dlq recv failed"), 99); - assert!(receiver.try_recv().is_err()); + let receiver = rx + .as_mut() + .ok_or_else(|| io::Error::other("receiver missing"))?; + let value = receiver + .recv() + .await + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "dlq recv failed"))?; + if value != 99 { + return Err(io::Error::other("unexpected DLQ value").into()); + } + if receiver.try_recv().is_ok() { + return Err(io::Error::other("expected DLQ to be empty").into()); + } + Ok(()) }) } /// Confirms no receiver is present when the DLQ is closed. -fn assert_dlq_closed(_: &mut Option>) -> BoxFuture<'_, ()> { Box::pin(async {}) } +fn assert_dlq_closed(_: &mut Option>) -> BoxFuture<'_, TestResult<()>> { + Box::pin(async { Ok(()) }) +} /// Parameterised checks for error logs when DLQ interactions fail. #[rstest] #[case::dlq_full( - fill_dlq, - PushPolicy::WarnAndDropIfFull, - assert_dlq_full, - "DLQ dropped frames" + DlqCase { + setup: fill_dlq, + policy: PushPolicy::WarnAndDropIfFull, + assertion: assert_dlq_full, + expected: "DLQ dropped frames" + } )] #[case::dlq_closed( - close_dlq, - PushPolicy::DropIfFull, - assert_dlq_closed, - "DLQ dropped frames" + DlqCase { + setup: close_dlq, + policy: PushPolicy::DropIfFull, + assertion: assert_dlq_closed, + expected: "DLQ dropped frames" + } )] #[serial(push_policies)] -fn dlq_error_scenarios( - rt: Runtime, +fn dlq_error_scenarios( mut logger: LoggerHandle, - #[case] setup: Setup, - #[case] policy: PushPolicy, - #[case] assertion: AssertFn, - #[case] expected: &str, + #[case] case: DlqCase, builder: PushQueuesBuilder, -) where - Setup: FnOnce(&mpsc::Sender, &mut Option>), - AssertFn: FnOnce(&mut Option>) -> BoxFuture<'_, ()>, -{ - rt.block_on(async { +) -> TestResult { + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build()?; + rt.block_on(async move { while logger.pop().is_some() {} + let DlqCase { + setup, + policy, + assertion, + expected, + } = case; let (dlq_tx, dlq_rx) = mpsc::channel(1); let mut dlq_rx = Some(dlq_rx); - setup(&dlq_tx, &mut dlq_rx); + setup(&dlq_tx, &mut dlq_rx) + .map_err(|e| io::Error::other(format!("DLQ setup failed: {e}")))?; let (mut queues, handle) = builder .unlimited() .dlq(Some(dlq_tx)) .build() - .expect("failed to build PushQueues"); + .map_err(|e| io::Error::other(format!("build queues failed: {e}")))?; handle .push_high_priority(1u8) .await - .expect("push high priority failed"); + .map_err(|e| io::Error::other(format!("push high priority failed: {e}")))?; handle .try_push(2u8, PushPriority::High, policy) - .expect("try_push failed"); + .map_err(|e| io::Error::other(format!("try_push failed: {e}")))?; - let (_, val) = queues.recv().await.expect("recv failed"); - assert_eq!(val, 1); + let (_, val) = queues + .recv() + .await + .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "recv failed"))?; + if val != 1 { + return Err(io::Error::other("unexpected dequeued value").into()); + } - assertion(&mut dlq_rx).await; + assertion(&mut dlq_rx) + .await + .map_err(|e| io::Error::other(format!("DLQ assertion failed: {e}")))?; let mut found = false; while let Some(record) = logger.pop() { @@ -179,6 +264,10 @@ fn dlq_error_scenarios( found = true; } } - assert!(found, "expected DLQ warning log missing"); - }); + if !found { + return Err(io::Error::other("expected DLQ warning log missing").into()); + } + Ok::<(), Box>(()) + })?; + Ok(()) } diff --git a/tests/response.rs b/tests/response.rs index dc07ee0e..5b3546f0 100644 --- a/tests/response.rs +++ b/tests/response.rs @@ -19,7 +19,7 @@ use wireframe::{ use wireframe_testing::{decode_frames, decode_frames_with_max, encode_frame, run_app}; mod common; -use common::TestApp; +use common::{TestApp, TestResult}; // Larger cap used for oversized frame tests. const LARGE_FRAME: usize = 16 * 1024 * 1024; @@ -53,18 +53,25 @@ struct Large(Vec); /// Tests that sending a response serializes and frames the data correctly, /// and that the response can be decoded and deserialized back to its original value asynchronously. #[tokio::test] -async fn send_response_encodes_and_frames() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn send_response_encodes_and_frames() -> TestResult { let app = TestApp::new().expect("failed to create app"); let mut out = Vec::new(); app.send_response(&mut out, &TestResp(7)) .await - .expect("send_response failed"); + .map_err(|e| format!("send_response failed: {e}"))?; let frames = decode_frames(out); assert_eq!(frames.len(), 1, "expected a single response frame"); - let (decoded, _) = TestResp::from_bytes(&frames[0]).expect("deserialize failed"); - assert_eq!(decoded, TestResp(7)); + let frame = frames.first().ok_or("expected frame missing")?; + let (decoded, _) = + TestResp::from_bytes(frame).map_err(|e| format!("deserialize failed: {e}"))?; + assert_eq!(decoded, TestResp(7), "decoded payload mismatch"); + Ok(()) } /// Tests that decoding with an incomplete length prefix header returns `None` and does not consume @@ -139,7 +146,10 @@ fn custom_length_roundtrip( codec .encode(frame.clone().into(), &mut buf) .expect("encode failed"); - assert_eq!(&buf[..prefix.len()], &prefix[..]); + let head = buf + .get(..prefix.len()) + .expect("encoded buffer shorter than prefix"); + assert_eq!(head, &prefix[..]); let decoded = codec .decode(&mut buf) .expect("decode failed") @@ -206,10 +216,12 @@ async fn send_response_returns_encode_error() { /// Ensures `send_response` permits frames up to the configured buffer capacity, /// exceeding the codec's default 8 MiB limit. #[tokio::test] -async fn send_response_honours_buffer_capacity() { - let app = TestApp::new() - .expect("failed to create app") - .buffer_capacity(LARGE_FRAME); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn send_response_honours_buffer_capacity() -> TestResult { + let app = TestApp::new()?.buffer_capacity(LARGE_FRAME); let payload = vec![0_u8; 9 * 1024 * 1024]; let large = Large(payload.clone()); @@ -217,39 +229,45 @@ async fn send_response_honours_buffer_capacity() { app.send_response(&mut out, &large) .await - .expect("send_response failed"); + .map_err(|e| format!("send_response failed: {e}"))?; let frames = decode_frames_with_max(out, LARGE_FRAME); assert_eq!(frames.len(), 1, "expected a single response frame"); - let (decoded, _) = Large::from_bytes(&frames[0]).expect("deserialize failed"); + let frame = frames.first().ok_or("response frame missing")?; + let (decoded, _) = Large::from_bytes(frame).map_err(|e| format!("deserialize failed: {e}"))?; assert_eq!(decoded.0.len(), payload.len()); + Ok(()) } /// Verifies inbound and outbound codecs respect the application's buffer /// capacity by round-tripping a 9 MiB payload. #[tokio::test] -async fn process_stream_honours_buffer_capacity() { - let app = TestApp::new() - .expect("failed to create app") +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn process_stream_honours_buffer_capacity() -> TestResult { + let app = TestApp::new()? .buffer_capacity(LARGE_FRAME) - .route(1, Arc::new(|_: &Envelope| Box::pin(async {}))) - .expect("route registration failed"); + .route(1, Arc::new(|_: &Envelope| Box::pin(async {})))?; let payload = vec![0_u8; 9 * 1024 * 1024]; let env = Envelope::new(1, None, payload.clone()); - let bytes = BincodeSerializer.serialize(&env).expect("serialize failed"); + let bytes = BincodeSerializer + .serialize(&env) + .map_err(|e| format!("serialize failed: {e}"))?; let mut codec = app.length_codec(); let frame = encode_frame(&mut codec, bytes); - let out = run_app(app, vec![frame], Some(10 * 1024 * 1024)) - .await - .expect("run_app failed"); + let out = run_app(app, vec![frame], Some(10 * 1024 * 1024)).await?; let frames = decode_frames_with_max(out, LARGE_FRAME); assert_eq!(frames.len(), 1, "expected a single response frame"); + let frame = frames.first().ok_or("response frame missing")?; let (resp_env, _) = BincodeSerializer - .deserialize::(&frames[0]) - .expect("deserialize failed"); + .deserialize::(frame) + .map_err(|e| format!("deserialize failed: {e}"))?; let resp_len = resp_env.into_parts().payload().len(); assert_eq!(resp_len, payload.len()); + Ok(()) } diff --git a/tests/routes.rs b/tests/routes.rs index 0d60d4bf..1e3371cb 100644 --- a/tests/routes.rs +++ b/tests/routes.rs @@ -3,12 +3,15 @@ //! //! They validate handler invocation, echo responses, and sequential processing. +mod common; + use std::sync::{ Arc, atomic::{AtomicUsize, Ordering}, }; use bytes::BytesMut; +use common::TestResult; use rstest::rstest; use tokio_util::codec::Encoder; use wireframe::{ @@ -60,122 +63,132 @@ impl Packet for TestEnvelope { #[derive(bincode::Encode, bincode::BorrowDecode, PartialEq, Debug)] struct Echo(u8); -#[rstest] #[tokio::test] -async fn handler_receives_message_and_echoes_response() { +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn handler_receives_message_and_echoes_response() -> TestResult<()> { let called = Arc::new(AtomicUsize::new(0)); let called_clone = called.clone(); - let app = TestApp::new() - .expect("failed to create app") - .route( - 1, - std::sync::Arc::new(move |_: &TestEnvelope| { - let called_inner = called_clone.clone(); - Box::pin(async move { - called_inner.fetch_add(1, Ordering::SeqCst); - // `WireframeApp` sends the envelope back automatically - }) - }), - ) - .expect("route registration failed"); - let msg_bytes = Echo(42).to_bytes().expect("encode failed"); + let app = TestApp::new()?.route( + 1, + std::sync::Arc::new(move |_: &TestEnvelope| { + let called_inner = called_clone.clone(); + Box::pin(async move { + called_inner.fetch_add(1, Ordering::SeqCst); + // `WireframeApp` sends the envelope back automatically + }) + }), + )?; + let msg_bytes = Echo(42).to_bytes()?; let env = TestEnvelope { id: 1, correlation_id: Some(99), payload: msg_bytes, }; - let out = drive_with_bincode(app, env) - .await - .expect("drive_with_bincode failed"); + let out = drive_with_bincode(app, env).await?; let frames = decode_frames(out); - assert_eq!(frames.len(), 1, "expected a single response frame"); - let (resp_env, _) = BincodeSerializer - .deserialize::(&frames[0]) - .expect("deserialize failed"); - assert_eq!(resp_env.correlation_id, Some(99)); - let (echo, _) = Echo::from_bytes(&resp_env.payload).expect("decode echo failed"); - assert_eq!(echo, Echo(42)); - assert_eq!(called.load(Ordering::SeqCst), 1); + let [first] = frames.as_slice() else { + return Err("expected a single response frame".into()); + }; + let (resp_env, _) = BincodeSerializer.deserialize::(first)?; + assert_eq!(resp_env.correlation_id, Some(99), "correlation id mismatch"); + let (echo, _) = Echo::from_bytes(&resp_env.payload)?; + assert_eq!(echo, Echo(42), "echo payload mismatch"); + assert_eq!( + called.load(Ordering::SeqCst), + 1, + "route not invoked exactly once" + ); + Ok(()) } #[tokio::test] -async fn handler_echoes_with_none_correlation_id() { - let app = TestApp::new() - .expect("failed to create app") - .route( - 1, - std::sync::Arc::new(|_: &TestEnvelope| Box::pin(async {})), - ) - .expect("route registration failed"); - - let msg_bytes = Echo(7).to_bytes().expect("encode failed"); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn handler_echoes_with_none_correlation_id() -> TestResult<()> { + let app = TestApp::new()?.route( + 1, + std::sync::Arc::new(|_: &TestEnvelope| Box::pin(async {})), + )?; + + let msg_bytes = Echo(7).to_bytes()?; let env = TestEnvelope { id: 1, correlation_id: None, payload: msg_bytes, }; - let out = drive_with_bincode(app, env).await.expect("drive failed"); + let out = drive_with_bincode(app, env).await?; let frames = decode_frames(out); - assert_eq!(frames.len(), 1, "expected a single response frame"); - let (resp_env, _) = BincodeSerializer - .deserialize::(&frames[0]) - .expect("deserialize failed"); - - assert_eq!(resp_env.correlation_id, None); - let (echo, _) = Echo::from_bytes(&resp_env.payload).expect("decode echo failed"); - assert_eq!(echo, Echo(7)); + let [first] = frames.as_slice() else { + return Err("expected a single response frame".into()); + }; + let (resp_env, _) = BincodeSerializer.deserialize::(first)?; + + assert!( + resp_env.correlation_id.is_none(), + "unexpected correlation id" + ); + let (echo, _) = Echo::from_bytes(&resp_env.payload)?; + assert_eq!(echo, Echo(7), "echo payload mismatch"); + Ok(()) } #[tokio::test] -async fn multiple_frames_processed_in_sequence() { - let app = TestApp::new() - .expect("failed to create app") - .route( - 1, - std::sync::Arc::new(|_: &TestEnvelope| Box::pin(async {})), - ) - .expect("route registration failed"); +#[expect( + clippy::panic_in_result_fn, + reason = "asserts provide clearer diagnostics in tests" +)] +async fn multiple_frames_processed_in_sequence() -> TestResult<()> { + let app = TestApp::new()?.route( + 1, + std::sync::Arc::new(|_: &TestEnvelope| Box::pin(async {})), + )?; let mut codec = new_test_codec(TEST_MAX_FRAME); let mut encoded_frames = Vec::new(); for id in 1u8..=2 { - let msg_bytes = Echo(id).to_bytes().expect("encode failed"); + let msg_bytes = Echo(id).to_bytes()?; let env = TestEnvelope { id: 1, correlation_id: Some(u64::from(id)), payload: msg_bytes, }; - let env_bytes = BincodeSerializer - .serialize(&env) - .expect("serialization failed"); + let env_bytes = BincodeSerializer.serialize(&env)?; let mut framed = BytesMut::with_capacity(env_bytes.len() + 4); - codec - .encode(env_bytes.into(), &mut framed) - .expect("encode failed"); + codec.encode(env_bytes.into(), &mut framed)?; encoded_frames.push(framed.to_vec()); } - let out = drive_with_frames(app, encoded_frames) - .await - .expect("drive_with_frames failed"); + let out = drive_with_frames(app, encoded_frames).await?; let frames = decode_frames(out); - assert_eq!(frames.len(), 2, "expected two response frames"); - let (env1, _) = BincodeSerializer - .deserialize::(&frames[0]) - .expect("deserialize failed"); - let (echo1, _) = Echo::from_bytes(&env1.payload).expect("decode echo failed"); - let (env2, _) = BincodeSerializer - .deserialize::(&frames[1]) - .expect("deserialize failed"); - let (echo2, _) = Echo::from_bytes(&env2.payload).expect("decode echo failed"); - assert_eq!(env1.correlation_id, Some(1)); - assert_eq!(env2.correlation_id, Some(2)); - assert_eq!(echo1, Echo(1)); - assert_eq!(echo2, Echo(2)); + let [first, second] = frames.as_slice() else { + return Err("expected two response frames".into()); + }; + let (env1, _) = BincodeSerializer.deserialize::(first)?; + let (echo1, _) = Echo::from_bytes(&env1.payload)?; + let (env2, _) = BincodeSerializer.deserialize::(second)?; + let (echo2, _) = Echo::from_bytes(&env2.payload)?; + assert_eq!( + env1.correlation_id, + Some(1), + "first correlation id mismatch" + ); + assert_eq!( + env2.correlation_id, + Some(2), + "second correlation id mismatch" + ); + assert_eq!(echo1, Echo(1), "first echo payload mismatch"); + assert_eq!(echo2, Echo(2), "second echo payload mismatch"); + Ok(()) } #[rstest] @@ -183,39 +196,33 @@ async fn multiple_frames_processed_in_sequence() { #[case(Some(1))] #[case(Some(2))] #[tokio::test] -async fn single_frame_propagates_correlation_id(#[case] cid: Option) { - let app = TestApp::new() - .expect("failed to create app") - .route( - 1, - std::sync::Arc::new(|_: &TestEnvelope| Box::pin(async {})), - ) - .expect("route registration failed"); - - let msg_bytes = Echo(5).to_bytes().expect("encode failed"); +async fn single_frame_propagates_correlation_id(#[case] cid: Option) -> TestResult<()> { + let app = TestApp::new()?.route( + 1, + std::sync::Arc::new(|_: &TestEnvelope| Box::pin(async {})), + )?; + + let msg_bytes = Echo(5).to_bytes()?; let env = TestEnvelope { id: 1, correlation_id: cid, payload: msg_bytes, }; - let env_bytes = BincodeSerializer.serialize(&env).expect("serialize failed"); + let env_bytes = BincodeSerializer.serialize(&env)?; - let mut framed = BytesMut::with_capacity(env_bytes.len() + 4); + let mut frame_buf = BytesMut::with_capacity(env_bytes.len() + 4); let mut codec = new_test_codec(TEST_MAX_FRAME); - codec - .encode(env_bytes.into(), &mut framed) - .expect("encode failed"); + codec.encode(env_bytes.into(), &mut frame_buf)?; - let out = drive_with_frames(app, vec![framed.to_vec()]) - .await - .expect("drive failed"); + let out = drive_with_frames(app, vec![frame_buf.to_vec()]).await?; let frames = decode_frames(out); - assert_eq!(frames.len(), 1, "expected a single response frame"); - let (resp, _) = BincodeSerializer - .deserialize::(&frames[0]) - .expect("deserialize failed"); + let [first] = frames.as_slice() else { + return Err("expected a single response frame".into()); + }; + let (resp, _) = BincodeSerializer.deserialize::(first)?; - assert_eq!(resp.correlation_id, cid); + assert_eq!(resp.correlation_id, cid, "correlation id mismatch"); + Ok(()) } #[test] diff --git a/tests/server.rs b/tests/server.rs index 2dc593bc..4c14bfd2 100644 --- a/tests/server.rs +++ b/tests/server.rs @@ -2,38 +2,63 @@ //! Tests for [`WireframeServer`] configuration. mod common; -use common::{factory, unused_listener}; +use common::{TestResult, factory, unused_listener}; use wireframe::server::WireframeServer; #[test] -fn default_worker_count_matches_cpu_count() { +fn default_worker_count_matches_cpu_count() -> TestResult { let server = WireframeServer::new(factory()); let expected = std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get); - assert_eq!(server.worker_count(), expected); + if server.worker_count() != expected { + return Err(format!( + "worker count mismatch: actual={}, expected={}", + server.worker_count(), + expected + ) + .into()); + } + Ok(()) } #[test] -fn default_workers_at_least_one() { +fn default_workers_at_least_one() -> TestResult { let server = WireframeServer::new(factory()); - assert!(server.worker_count() >= 1); + if server.worker_count() < 1 { + return Err(format!("worker count below 1: {}", server.worker_count()).into()); + } + Ok(()) } #[test] -fn workers_method_enforces_minimum() { +fn workers_method_enforces_minimum() -> TestResult { let server = WireframeServer::new(factory()).workers(0); - assert_eq!(server.worker_count(), 1); + if server.worker_count() != 1 { + return Err(format!( + "worker count should clamp to 1, got {}", + server.worker_count() + ) + .into()); + } + Ok(()) } #[test] -fn workers_accepts_large_values() { +fn workers_accepts_large_values() -> TestResult { let server = WireframeServer::new(factory()).workers(128); - assert_eq!(server.worker_count(), 128); + if server.worker_count() != 128 { + return Err(format!( + "worker count should be 128 after config, got {}", + server.worker_count() + ) + .into()); + } + Ok(()) } /// Ensure dropping the readiness receiver logs a warning and does not /// prevent the server from accepting connections. #[tokio::test] -async fn readiness_receiver_dropped() { +async fn readiness_receiver_dropped() -> TestResult { use tokio::{ net::TcpStream, sync::oneshot, @@ -64,4 +89,5 @@ async fn readiness_receiver_dropped() { // Server should still accept connections let _stream = TcpStream::connect(addr).await.expect("connect failed"); + Ok(()) } diff --git a/tests/session_registry.rs b/tests/session_registry.rs index 269e82c2..d0e6661d 100644 --- a/tests/session_registry.rs +++ b/tests/session_registry.rs @@ -1,95 +1,74 @@ -#![cfg(not(loom))] //! Tests for the `SessionRegistry`. +#![cfg(not(loom))] + use rstest::{fixture, rstest}; use wireframe::{ - push::{PushHandle, PushQueues}, + push::{PushConfigError, PushHandle, PushQueues}, session::{ConnectionId, SessionRegistry}, }; -#[expect( - unused_braces, - reason = "rustc false positive for single-line rstest fixtures" -)] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] -#[fixture] -fn registry() -> SessionRegistry { SessionRegistry::default() } +mod common; +mod support; +use common::TestResult; -#[expect( - unused_braces, - reason = "rustc false positive for single-line rstest fixtures" -)] -// allow(unfulfilled_lint_expectations): rustc occasionally fails to emit the expected -// lint for single-line rstest fixtures on stable. -#[allow(unfulfilled_lint_expectations)] #[fixture] -fn push_setup() -> (PushQueues, PushHandle) { - PushQueues::::builder() - .high_capacity(1) - .low_capacity(1) - .build() - .expect("failed to build PushQueues") +fn registry() -> SessionRegistry { + // Fixtures use the default registry to minimize setup noise. + SessionRegistry::default() +} + +fn push_setup() -> Result<(PushQueues, PushHandle), PushConfigError> { + support::builder().build() } /// Test that handles can be retrieved whilst the connection remains alive. #[rstest] #[tokio::test] -async fn handle_retrieved_while_alive( - registry: SessionRegistry, - #[from(push_setup)] setup: (PushQueues, PushHandle), -) { - let (mut queues, handle) = setup; +async fn handle_retrieved_while_alive(registry: SessionRegistry) -> TestResult<()> { + let (mut queues, handle) = push_setup()?; let id = ConnectionId::new(42); registry.insert(id, &handle); let retrieved = registry.get(&id).expect("handle should be present"); - retrieved.push_high_priority(7).await.expect("push failed"); + retrieved.push_high_priority(7).await?; let (_, val) = queues.recv().await.expect("recv failed"); assert_eq!(val, 7); + Ok(()) } /// Test that [`SessionRegistry::get`] returns `None` after the handle is dropped. #[rstest] #[tokio::test] -async fn get_returns_none_after_drop( - registry: SessionRegistry, - #[from(push_setup)] setup: (PushQueues, PushHandle), -) { - let (_queues, handle) = setup; +async fn get_returns_none_after_drop(registry: SessionRegistry) -> TestResult<()> { + let (_queues, handle) = push_setup()?; let id = ConnectionId::new(1); registry.insert(id, &handle); drop(handle); assert!(registry.get(&id).is_none()); + Ok(()) } /// Calling `get` should remove expired entries. #[rstest] #[tokio::test] -async fn get_prunes_dead_handle( - registry: SessionRegistry, - #[from(push_setup)] setup: (PushQueues, PushHandle), -) { - let (_queues, handle) = setup; +async fn get_prunes_dead_handle(registry: SessionRegistry) -> TestResult<()> { + let (_queues, handle) = push_setup()?; let id = ConnectionId::new(11); registry.insert(id, &handle); drop(handle); assert!(registry.get(&id).is_none()); assert!(!registry.active_ids().contains(&id)); + Ok(()) } /// `active_handles` returns only live sessions. #[rstest] #[tokio::test] -async fn active_handles_lists_live_connections( - registry: SessionRegistry, - #[from(push_setup)] setup1: (PushQueues, PushHandle), - #[from(push_setup)] setup2: (PushQueues, PushHandle), -) { - let (_queues1, handle1) = setup1; - let (_queues2, handle2) = setup2; +async fn active_handles_lists_live_connections(registry: SessionRegistry) -> TestResult<()> { + let (_queues1, handle1) = push_setup()?; + let (_queues2, handle2) = push_setup()?; let id1 = ConnectionId::new(21); let id2 = ConnectionId::new(22); registry.insert(id1, &handle1); @@ -98,21 +77,21 @@ async fn active_handles_lists_live_connections( let handles = registry.active_handles(); assert_eq!(handles.len(), 1); - assert_eq!(handles[0].0, id2); + let first = handles.first().expect("no active handles"); + assert_eq!(first.0, id2); + Ok(()) } /// Test that `prune` removes entries whose handles have been dropped. #[rstest] #[tokio::test] -async fn prune_removes_dead_entries( - registry: SessionRegistry, - #[from(push_setup)] setup: (PushQueues, PushHandle), -) { - let (_queues, handle) = setup; +async fn prune_removes_dead_entries(registry: SessionRegistry) -> TestResult<()> { + let (_queues, handle) = push_setup()?; let id = ConnectionId::new(5); registry.insert(id, &handle); drop(handle); registry.prune(); assert!(registry.get(&id).is_none()); + Ok(()) } diff --git a/tests/steps/correlation_steps.rs b/tests/steps/correlation_steps.rs index bd58da0a..cc99d1a7 100644 --- a/tests/steps/correlation_steps.rs +++ b/tests/steps/correlation_steps.rs @@ -1,7 +1,7 @@ //! Steps for `correlation_id` behavioural tests. use cucumber::{given, then, when}; -use crate::world::CorrelationWorld; +use crate::world::{CorrelationWorld, TestResult}; #[given(expr = "a correlation id {int}")] fn given_cid(world: &mut CorrelationWorld, id: u64) { world.set_expected(Some(id)); } @@ -10,19 +10,25 @@ fn given_cid(world: &mut CorrelationWorld, id: u64) { world.set_expected(Some(id fn given_no_correlation(world: &mut CorrelationWorld) { world.set_expected(None); } #[when("a stream of frames is processed")] -async fn when_process(world: &mut CorrelationWorld) { world.process().await; } +async fn when_process(world: &mut CorrelationWorld) -> TestResult { world.process().await } #[when("a multi-packet channel emits frames")] -async fn when_process_multi(world: &mut CorrelationWorld) { world.process_multi().await; } +async fn when_process_multi(world: &mut CorrelationWorld) -> TestResult { + world.process_multi().await +} #[then(expr = "each emitted frame uses correlation id {int}")] -fn then_verify(world: &mut CorrelationWorld, id: u64) { - assert_eq!(world.expected(), Some(id)); - world.verify(); +fn then_verify(world: &mut CorrelationWorld, id: u64) -> TestResult { + if world.expected() != Some(id) { + return Err("mismatched expected correlation id".into()); + } + world.verify() } #[then("each emitted frame has no correlation id")] -fn then_verify_absent(world: &mut CorrelationWorld) { - assert_eq!(world.expected(), None); - world.verify(); +fn then_verify_absent(world: &mut CorrelationWorld) -> TestResult { + if world.expected().is_some() { + return Err("expected correlation id should be cleared".into()); + } + world.verify() } diff --git a/tests/steps/fragment_steps.rs b/tests/steps/fragment_steps.rs index 608e397f..926007ef 100644 --- a/tests/steps/fragment_steps.rs +++ b/tests/steps/fragment_steps.rs @@ -4,82 +4,111 @@ use std::time::Duration; use cucumber::{given, then, when}; use wireframe::{FragmentHeader, FragmentIndex, MessageId}; -use crate::world::FragmentWorld; +use crate::world::{FragmentWorld, TestResult}; #[given(expr = "a fragment series for message {int}")] fn given_series(world: &mut FragmentWorld, message: u64) { world.start_series(message); } #[given(expr = "the series expects fragment index {int}")] -fn given_series_expectation(world: &mut FragmentWorld, index: u32) { - world.force_next_index(index); +fn given_series_expectation(world: &mut FragmentWorld, index: u32) -> TestResult { + world.force_next_index(index)?; + Ok(()) } #[when(expr = "fragment {int} arrives marked non-final")] -fn when_fragment_non_final(world: &mut FragmentWorld, index: u32) { - world.accept_fragment(index, false); +fn when_fragment_non_final(world: &mut FragmentWorld, index: u32) -> TestResult { + world.accept_fragment(index, false)?; + Ok(()) } #[when(expr = "fragment {int} arrives marked final")] -fn when_fragment_final(world: &mut FragmentWorld, index: u32) { - world.accept_fragment(index, true); +fn when_fragment_final(world: &mut FragmentWorld, index: u32) -> TestResult { + world.accept_fragment(index, true)?; + Ok(()) } #[when(expr = "fragment {int} from message {int} arrives marked non-final")] -fn when_fragment_other_message(world: &mut FragmentWorld, index: u32, message: u64) { - world.accept_fragment_from(message, index, false); +fn when_fragment_other_message(world: &mut FragmentWorld, index: u32, message: u64) -> TestResult { + world.accept_fragment_from(message, index, false)?; + Ok(()) } #[then("the fragment completes the message")] -fn then_fragment_completes(world: &mut FragmentWorld) { world.assert_completion(); } +fn then_fragment_completes(world: &mut FragmentWorld) -> TestResult { + world.assert_completion()?; + Ok(()) +} #[then("the fragment is rejected as out-of-order")] -fn then_fragment_out_of_order(world: &mut FragmentWorld) { world.assert_index_mismatch(); } +fn then_fragment_out_of_order(world: &mut FragmentWorld) -> TestResult { + world.assert_index_mismatch()?; + Ok(()) +} #[then("the fragment is rejected for the wrong message")] -fn then_fragment_wrong_message(world: &mut FragmentWorld) { world.assert_message_mismatch(); } +fn then_fragment_wrong_message(world: &mut FragmentWorld) -> TestResult { + world.assert_message_mismatch()?; + Ok(()) +} #[then("the fragment is rejected for index overflow")] -fn then_fragment_overflow(world: &mut FragmentWorld) { world.assert_index_overflow(); } +fn then_fragment_overflow(world: &mut FragmentWorld) -> TestResult { + world.assert_index_overflow()?; + Ok(()) +} #[then("the fragment is rejected because the series is complete")] -fn then_fragment_complete(world: &mut FragmentWorld) { world.assert_series_complete_error(); } +fn then_fragment_complete(world: &mut FragmentWorld) -> TestResult { + world.assert_series_complete_error()?; + Ok(()) +} #[given(expr = "a fragmenter capped at {int} bytes per fragment")] -fn given_fragmenter(world: &mut FragmentWorld, max_payload: usize) { - world.configure_fragmenter(max_payload); +fn given_fragmenter(world: &mut FragmentWorld, max_payload: usize) -> TestResult { + world.configure_fragmenter(max_payload)?; + Ok(()) } #[when(expr = "the fragmenter splits a payload of {int} bytes")] -fn when_fragmenter_splits(world: &mut FragmentWorld, len: usize) { world.fragment_payload(len); } +fn when_fragmenter_splits(world: &mut FragmentWorld, len: usize) -> TestResult { + world.fragment_payload(len)?; + Ok(()) +} #[then(expr = "the fragmenter produces {int} fragments")] -fn then_fragment_count(world: &mut FragmentWorld, expected: usize) { - world.assert_fragment_count(expected); +fn then_fragment_count(world: &mut FragmentWorld, expected: usize) -> TestResult { + world.assert_fragment_count(expected)?; + Ok(()) } #[then(expr = "fragment {int} carries {int} bytes")] -fn then_fragment_payload_len(world: &mut FragmentWorld, index: usize, len: usize) { - world.assert_fragment_payload_len(index, len); +fn then_fragment_payload_len(world: &mut FragmentWorld, index: usize, len: usize) -> TestResult { + world.assert_fragment_payload_len(index, len)?; + Ok(()) } #[then(expr = "fragment {int} is marked final")] -fn then_fragment_final(world: &mut FragmentWorld, index: usize) { - world.assert_fragment_final_flag(index, true); +fn then_fragment_final(world: &mut FragmentWorld, index: usize) -> TestResult { + world.assert_fragment_final_flag(index, true)?; + Ok(()) } #[then(expr = "fragment {int} is marked non-final")] -fn then_fragment_non_final(world: &mut FragmentWorld, index: usize) { - world.assert_fragment_final_flag(index, false); +fn then_fragment_non_final(world: &mut FragmentWorld, index: usize) -> TestResult { + world.assert_fragment_final_flag(index, false)?; + Ok(()) } #[then(expr = "the fragments use message id {int}")] -fn then_fragment_message_id(world: &mut FragmentWorld, message_id: u64) { - world.assert_message_id(message_id); +fn then_fragment_message_id(world: &mut FragmentWorld, message_id: u64) -> TestResult { + world.assert_message_id(message_id)?; + Ok(()) } #[given(expr = "a reassembler allowing {int} bytes with a {int}-second reassembly timeout")] -fn given_reassembler(world: &mut FragmentWorld, max_bytes: usize, timeout_secs: u64) { - world.configure_reassembler(max_bytes, timeout_secs); +fn given_reassembler(world: &mut FragmentWorld, max_bytes: usize, timeout_secs: u64) -> TestResult { + world.configure_reassembler(max_bytes, timeout_secs)?; + Ok(()) } #[when(expr = "fragment {int} for message {int} with {int} bytes arrives marked non-final")] @@ -88,9 +117,10 @@ fn when_reassembler_fragment_non_final( index: u32, message: u64, len: usize, -) { +) -> TestResult { let header = FragmentHeader::new(MessageId::new(message), FragmentIndex::new(index), false); - world.push_fragment(header, len); + world.push_fragment(header, len)?; + Ok(()) } #[when(expr = "fragment {int} for message {int} with {int} bytes arrives marked final")] @@ -99,41 +129,56 @@ fn when_reassembler_fragment_final( index: u32, message: u64, len: usize, -) { +) -> TestResult { let header = FragmentHeader::new(MessageId::new(message), FragmentIndex::new(index), true); - world.push_fragment(header, len); + world.push_fragment(header, len)?; + Ok(()) } #[when(expr = "time advances by {int} seconds")] -fn when_time_advances(world: &mut FragmentWorld, seconds: u64) { - world.advance_time(Duration::from_secs(seconds)); +fn when_time_advances(world: &mut FragmentWorld, seconds: u64) -> TestResult { + world.advance_time(Duration::from_secs(seconds))?; + Ok(()) } #[when("expired reassembly buffers are purged")] -fn when_reassembly_purged(world: &mut FragmentWorld) { world.purge_reassembly(); } +fn when_reassembly_purged(world: &mut FragmentWorld) -> TestResult { + world.purge_reassembly()?; + Ok(()) +} #[then(expr = "the reassembler outputs a payload of {int} bytes")] -fn then_reassembled_len(world: &mut FragmentWorld, expected: usize) { - world.assert_reassembled_len(expected); +fn then_reassembled_len(world: &mut FragmentWorld, expected: usize) -> TestResult { + world.assert_reassembled_len(expected)?; + Ok(()) } #[then("no message has been reassembled yet")] -fn then_no_reassembled_message(world: &mut FragmentWorld) { world.assert_no_reassembly(); } +fn then_no_reassembled_message(world: &mut FragmentWorld) -> TestResult { + world.assert_no_reassembly()?; + Ok(()) +} #[then("the reassembler reports a message-too-large error")] -fn then_reassembly_over_limit(world: &mut FragmentWorld) { world.assert_reassembly_over_limit(); } +fn then_reassembly_over_limit(world: &mut FragmentWorld) -> TestResult { + world.assert_reassembly_over_limit()?; + Ok(()) +} #[then("the reassembler reports an out-of-order fragment error")] -fn then_reassembly_out_of_order(world: &mut FragmentWorld) { - world.assert_reassembly_out_of_order(); +fn then_reassembly_out_of_order(world: &mut FragmentWorld) -> TestResult { + world.assert_reassembly_out_of_order()?; + Ok(()) } #[then(expr = "the reassembler is buffering {int} messages")] -fn then_buffered_messages(world: &mut FragmentWorld, expected: usize) { - world.assert_buffered_messages(expected); +fn then_buffered_messages(world: &mut FragmentWorld, expected: usize) -> TestResult { + world.assert_buffered_messages(expected)?; + Ok(()) } #[then(expr = "message {int} is evicted")] -fn then_message_evicted(world: &mut FragmentWorld, message: u64) { - world.assert_evicted_message(message); +fn then_message_evicted(world: &mut FragmentWorld, message: u64) -> TestResult { + world.assert_evicted_message(message)?; + Ok(()) } diff --git a/tests/steps/multi_packet_steps.rs b/tests/steps/multi_packet_steps.rs index fb7dfcc4..12791fb2 100644 --- a/tests/steps/multi_packet_steps.rs +++ b/tests/steps/multi_packet_steps.rs @@ -1,22 +1,26 @@ //! Steps for multi-packet response behavioural tests. use cucumber::{then, when}; -use crate::world::MultiPacketWorld; +use crate::world::{MultiPacketWorld, TestResult}; #[when("a handler uses the with_channel helper to emit messages")] -async fn when_multi(world: &mut MultiPacketWorld) { world.process().await; } +async fn when_multi(world: &mut MultiPacketWorld) -> TestResult { world.process().await } #[then("all messages are received in order")] fn then_multi(world: &mut MultiPacketWorld) { world.verify(); } #[when("a handler uses the with_channel helper to emit no messages")] -async fn when_multi_empty(world: &mut MultiPacketWorld) { world.process_empty().await; } +async fn when_multi_empty(world: &mut MultiPacketWorld) -> TestResult { + world.process_empty().await +} #[then("no messages are received")] fn then_multi_empty(world: &mut MultiPacketWorld) { world.verify_empty(); } #[when("a handler emits more messages than the channel capacity")] -async fn when_multi_overflow(world: &mut MultiPacketWorld) { world.process_overflow().await; } +async fn when_multi_overflow(world: &mut MultiPacketWorld) -> TestResult { + world.process_overflow().await +} #[then("overflow messages are handled according to channel policy")] fn then_multi_overflow(world: &mut MultiPacketWorld) { world.verify_overflow(); } diff --git a/tests/steps/panic_steps.rs b/tests/steps/panic_steps.rs index 17c3330a..439c193a 100644 --- a/tests/steps/panic_steps.rs +++ b/tests/steps/panic_steps.rs @@ -5,14 +5,23 @@ use cucumber::{given, then, when}; -use crate::world::PanicWorld; +use crate::world::{PanicWorld, TestResult}; #[given("a running wireframe server with a panic in connection setup")] -async fn start_server(world: &mut PanicWorld) { world.start_panic_server().await; } +async fn start_server(world: &mut PanicWorld) -> TestResult { + world.start_panic_server().await?; + Ok(()) +} #[when("I connect to the server")] #[when("I connect to the server again")] -async fn connect(world: &mut PanicWorld) { world.connect_once().await; } +async fn connect(world: &mut PanicWorld) -> TestResult { + world.connect_once().await?; + Ok(()) +} #[then("both connections succeed")] -async fn verify(world: &mut PanicWorld) { world.verify_and_shutdown().await; } +async fn verify(world: &mut PanicWorld) -> TestResult { + world.verify_and_shutdown().await?; + Ok(()) +} diff --git a/tests/steps/stream_end_steps.rs b/tests/steps/stream_end_steps.rs index eb566606..ff0485d1 100644 --- a/tests/steps/stream_end_steps.rs +++ b/tests/steps/stream_end_steps.rs @@ -1,31 +1,35 @@ //! Steps for stream terminator behavioural tests. use cucumber::{then, when}; -use crate::world::StreamEndWorld; +use crate::world::{StreamEndWorld, TestResult}; #[when("a streaming response completes")] -async fn when_stream(world: &mut StreamEndWorld) { world.process().await; } +async fn when_stream(world: &mut StreamEndWorld) -> TestResult { world.process().await } #[then("an end-of-stream frame is sent")] fn then_end(world: &mut StreamEndWorld) { world.verify(); } #[when("a multi-packet channel drains")] -async fn when_multi_channel(world: &mut StreamEndWorld) { world.process_multi().await; } +async fn when_multi_channel(world: &mut StreamEndWorld) -> TestResult { + world.process_multi().await +} #[then("a multi-packet end-of-stream frame is sent")] fn then_multi_end(world: &mut StreamEndWorld) { world.verify_multi(); } #[when("a multi-packet channel disconnects abruptly")] -fn when_multi_disconnect(world: &mut StreamEndWorld) { world.process_multi_disconnect(); } +fn when_multi_disconnect(world: &mut StreamEndWorld) -> TestResult { + world.process_multi_disconnect() +} #[when("shutdown closes a multi-packet channel")] -fn when_multi_shutdown(world: &mut StreamEndWorld) { world.process_multi_shutdown(); } +fn when_multi_shutdown(world: &mut StreamEndWorld) -> TestResult { world.process_multi_shutdown() } #[then("no multi-packet terminator is sent")] fn then_no_multi(world: &mut StreamEndWorld) { world.verify_no_multi(); } #[then(expr = "the multi-packet termination reason is {word}")] -fn then_reason(world: &mut StreamEndWorld, reason: String) { +fn then_reason(world: &mut StreamEndWorld, reason: String) -> TestResult { let reason = reason.into_boxed_str(); - world.verify_reason(reason.as_ref()); + world.verify_reason(reason.as_ref()) } diff --git a/tests/stream_end.rs b/tests/stream_end.rs index 6998c074..d4f68af4 100644 --- a/tests/stream_end.rs +++ b/tests/stream_end.rs @@ -1,6 +1,7 @@ //! Tests for explicit end-of-stream signalling. #![cfg(not(loom))] +mod common; mod support; use std::sync::Arc; @@ -10,7 +11,7 @@ use rstest::{fixture, rstest}; use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; use wireframe::{ - connection::ConnectionActor, + connection::{ConnectionActor, ConnectionChannels}, hooks::{ConnectionContext, ProtocolHooks, WireframeProtocol}, push::{PushHandle, PushQueues}, response::FrameStream, @@ -18,19 +19,20 @@ use wireframe::{ #[path = "common/terminator.rs"] mod terminator; +use common::TestResult; use terminator::Terminator; #[fixture] -fn queues() -> (PushQueues, PushHandle) { - support::builder::() - .build() - .expect("failed to build PushQueues") +fn queues() -> Result<(PushQueues, PushHandle), wireframe::push::PushConfigError> { + support::builder::().build() } #[rstest] #[tokio::test] -async fn emits_end_frame(queues: (PushQueues, PushHandle)) { - let (queues, handle) = queues; +async fn emits_end_frame( + queues: Result<(PushQueues, PushHandle), wireframe::push::PushConfigError>, +) -> TestResult<()> { + let (queues, handle) = queues?; // fixture injected above let stream: FrameStream = Box::pin(try_stream! { yield 1; @@ -38,37 +40,63 @@ async fn emits_end_frame(queues: (PushQueues, PushHandle)) { }); let shutdown = CancellationToken::new(); let hooks = ProtocolHooks::from_protocol(&Arc::new(Terminator)); - let mut actor = ConnectionActor::with_hooks(queues, handle, Some(stream), shutdown, hooks); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + Some(stream), + shutdown, + hooks, + ); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("connection actor failed: {e:?}")))?; - assert_eq!(out, vec![1, 2, 0]); + assert_eq!(out, vec![1, 2, 0], "unexpected output frames"); + Ok(()) } #[rstest] #[tokio::test] -async fn multi_packet_emits_end_frame(queues: (PushQueues, PushHandle)) { - let (queues, handle) = queues; +async fn multi_packet_emits_end_frame( + queues: Result<(PushQueues, PushHandle), wireframe::push::PushConfigError>, +) -> TestResult<()> { + let (queues, handle) = queues?; let (tx, rx) = mpsc::channel(4); - tx.send(1).await.expect("send frame"); - tx.send(2).await.expect("send frame"); + tx.send(1) + .await + .map_err(|e| std::io::Error::other(format!("send frame: {e}")))?; + tx.send(2) + .await + .map_err(|e| std::io::Error::other(format!("send frame: {e}")))?; drop(tx); let shutdown = CancellationToken::new(); let hooks = ProtocolHooks::from_protocol(&Arc::new(Terminator)); - let mut actor = ConnectionActor::with_hooks(queues, handle, None, shutdown, hooks); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + None, + shutdown, + hooks, + ); actor.set_multi_packet(Some(rx)); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("connection actor failed: {e:?}")))?; - assert_eq!(out, vec![1, 2, 0]); + assert_eq!(out, vec![1, 2, 0], "unexpected output frames"); + Ok(()) } #[rstest] #[tokio::test] -async fn multi_packet_respects_no_terminator(queues: (PushQueues, PushHandle)) { +async fn multi_packet_respects_no_terminator( + queues: Result<(PushQueues, PushHandle), wireframe::push::PushConfigError>, +) -> TestResult<()> { struct NoTerminator; impl WireframeProtocol for NoTerminator { @@ -78,45 +106,67 @@ async fn multi_packet_respects_no_terminator(queues: (PushQueues, PushHandle fn stream_end_frame(&self, _ctx: &mut ConnectionContext) -> Option { None } } - let (queues, handle) = queues; + let (queues, handle) = queues?; let (tx, rx) = mpsc::channel(2); - tx.send(9).await.expect("send frame"); + tx.send(9) + .await + .map_err(|e| std::io::Error::other(format!("send frame: {e}")))?; drop(tx); let shutdown = CancellationToken::new(); let hooks = ProtocolHooks::from_protocol(&Arc::new(NoTerminator)); - let mut actor = ConnectionActor::with_hooks(queues, handle, None, shutdown, hooks); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + None, + shutdown, + hooks, + ); actor.set_multi_packet(Some(rx)); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("connection actor failed: {e:?}")))?; - assert_eq!(out, vec![9]); + assert_eq!(out, vec![9], "unexpected output frames"); + Ok(()) } #[rstest] #[tokio::test] -async fn multi_packet_empty_channel_emits_end(queues: (PushQueues, PushHandle)) { - let (queues, handle) = queues; +async fn multi_packet_empty_channel_emits_end( + queues: Result<(PushQueues, PushHandle), wireframe::push::PushConfigError>, +) -> TestResult<()> { + let (queues, handle) = queues?; let (tx, rx) = mpsc::channel(1); drop(tx); let shutdown = CancellationToken::new(); let hooks = ProtocolHooks::from_protocol(&Arc::new(Terminator)); - let mut actor = ConnectionActor::with_hooks(queues, handle, None, shutdown, hooks); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + None, + shutdown, + hooks, + ); actor.set_multi_packet(Some(rx)); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("connection actor failed: {e:?}")))?; - assert_eq!(out, vec![0]); + assert_eq!(out, vec![0], "unexpected output frames"); + Ok(()) } #[rstest] #[tokio::test] async fn multi_packet_empty_channel_no_terminator_emits_nothing( - queues: (PushQueues, PushHandle), -) { + queues: Result<(PushQueues, PushHandle), wireframe::push::PushConfigError>, +) -> TestResult<()> { struct NoTerminator; impl WireframeProtocol for NoTerminator { @@ -126,24 +176,35 @@ async fn multi_packet_empty_channel_no_terminator_emits_nothing( fn stream_end_frame(&self, _ctx: &mut ConnectionContext) -> Option { None } } - let (queues, handle) = queues; + let (queues, handle) = queues?; let (tx, rx) = mpsc::channel(1); drop(tx); let shutdown = CancellationToken::new(); let hooks = ProtocolHooks::from_protocol(&Arc::new(NoTerminator)); - let mut actor = ConnectionActor::with_hooks(queues, handle, None, shutdown, hooks); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + None, + shutdown, + hooks, + ); actor.set_multi_packet(Some(rx)); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("connection actor failed: {e:?}")))?; - assert!(out.is_empty()); + assert_eq!(out, Vec::::new(), "expected no frames"); + Ok(()) } #[rstest] #[tokio::test] -async fn emits_no_end_frame_when_none(queues: (PushQueues, PushHandle)) { +async fn emits_no_end_frame_when_none( + queues: Result<(PushQueues, PushHandle), wireframe::push::PushConfigError>, +) -> TestResult<()> { struct NoTerminator; impl WireframeProtocol for NoTerminator { @@ -153,7 +214,7 @@ async fn emits_no_end_frame_when_none(queues: (PushQueues, PushHandle)) fn stream_end_frame(&self, _ctx: &mut ConnectionContext) -> Option { None } } - let (queues, handle) = queues; + let (queues, handle) = queues?; // fixture injected above let stream: FrameStream = Box::pin(try_stream! { yield 7; @@ -162,10 +223,19 @@ async fn emits_no_end_frame_when_none(queues: (PushQueues, PushHandle)) let shutdown = CancellationToken::new(); let hooks = ProtocolHooks::from_protocol(&Arc::new(NoTerminator)); - let mut actor = ConnectionActor::with_hooks(queues, handle, Some(stream), shutdown, hooks); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + Some(stream), + shutdown, + hooks, + ); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("connection actor failed: {e:?}")))?; - assert_eq!(out, vec![7, 8]); + assert_eq!(out, vec![7, 8], "unexpected frames"); + Ok(()) } diff --git a/tests/wireframe_protocol.rs b/tests/wireframe_protocol.rs index b4876f85..037ff111 100644 --- a/tests/wireframe_protocol.rs +++ b/tests/wireframe_protocol.rs @@ -11,6 +11,8 @@ use std::sync::{ atomic::{AtomicUsize, Ordering}, }; +mod common; +use common::TestResult; use futures::stream; use rstest::{fixture, rstest}; use tokio_util::sync::CancellationToken; @@ -18,21 +20,22 @@ use wireframe::{ ConnectionContext, WireframeProtocol, app::Envelope, - connection::ConnectionActor, - push::PushQueues, + connection::{ConnectionActor, ConnectionChannels}, + push::{PushConfigError, PushQueues}, serializer::BincodeSerializer, }; type TestApp = wireframe::app::WireframeApp; +type QueueResult = + Result<(PushQueues>, wireframe::push::PushHandle>), PushConfigError>; #[fixture] -fn queues() -> (PushQueues>, wireframe::push::PushHandle>) { +fn queues() -> QueueResult { PushQueues::>::builder() .high_capacity(8) .low_capacity(8) .unlimited() .build() - .expect("failed to build PushQueues") } struct TestProtocol { @@ -60,58 +63,66 @@ impl WireframeProtocol for TestProtocol { #[rstest] #[tokio::test] -async fn builder_produces_protocol_hooks( - queues: (PushQueues>, wireframe::push::PushHandle>), -) { +async fn builder_produces_protocol_hooks(queues: QueueResult) -> TestResult<()> { let counter = Arc::new(AtomicUsize::new(0)); let protocol = TestProtocol { counter: counter.clone(), }; - let app = TestApp::new() - .expect("failed to create app") - .with_protocol(protocol); + let app = TestApp::new()?.with_protocol(protocol); let mut hooks = app.protocol_hooks(); - let (_queues, handle) = queues; + let (_queues, handle) = queues?; hooks.on_connection_setup(handle, &mut ConnectionContext); let mut frame = vec![1u8]; hooks.before_send(&mut frame, &mut ConnectionContext); hooks.on_command_end(&mut ConnectionContext); - assert_eq!(frame, vec![1, 1]); - assert_eq!(counter.load(Ordering::SeqCst), 2); + assert_eq!(frame, vec![1, 1], "before_send did not mutate frame"); + assert_eq!( + counter.load(Ordering::SeqCst), + 2, + "expected two protocol callbacks" + ); + Ok(()) } #[rstest] #[tokio::test] -async fn connection_actor_uses_protocol_from_builder( - queues: (PushQueues>, wireframe::push::PushHandle>), -) { +async fn connection_actor_uses_protocol_from_builder(queues: QueueResult) -> TestResult<()> { let counter = Arc::new(AtomicUsize::new(0)); let protocol = TestProtocol { counter: counter.clone(), }; - let app = TestApp::new() - .expect("failed to create app") - .with_protocol(protocol); + let app = TestApp::new()?.with_protocol(protocol); let hooks = app.protocol_hooks(); - let (queues, handle) = queues; + let (queues, handle) = queues?; handle .push_high_priority(vec![1]) .await - .expect("push failed"); + .map_err(|e| std::io::Error::other(format!("push failed: {e}")))?; let stream = stream::iter(vec![Ok(vec![2u8])]); let mut actor: ConnectionActor<_, ()> = ConnectionActor::with_hooks( - queues, - handle, + ConnectionChannels::new(queues, handle), Some(Box::pin(stream)), CancellationToken::new(), hooks, ); let mut out = Vec::new(); - actor.run(&mut out).await.expect("actor run failed"); + actor + .run(&mut out) + .await + .map_err(|e| std::io::Error::other(format!("connection actor failed: {e:?}")))?; - assert_eq!(out, vec![vec![1, 1], vec![2, 1]]); - assert_eq!(counter.load(Ordering::SeqCst), 2); + assert_eq!( + out, + vec![vec![1, 1], vec![2, 1]], + "frames not mutated as expected" + ); + assert_eq!( + counter.load(Ordering::SeqCst), + 2, + "expected two protocol callbacks" + ); + Ok(()) } diff --git a/tests/world.rs b/tests/world.rs index d56211e4..5ce08561 100644 --- a/tests/world.rs +++ b/tests/world.rs @@ -5,6 +5,7 @@ mod worlds; pub use worlds::{ + common::TestResult, correlation::CorrelationWorld, fragment::FragmentWorld, multi_packet::MultiPacketWorld, diff --git a/tests/worlds/correlation.rs b/tests/worlds/correlation.rs index 910497e5..41aceced 100644 --- a/tests/worlds/correlation.rs +++ b/tests/worlds/correlation.rs @@ -14,7 +14,7 @@ use wireframe::{ response::FrameStream, }; -use super::build_small_queues; +use super::{TestResult, build_small_queues}; #[derive(Debug, Default, World)] pub struct CorrelationWorld { @@ -30,59 +30,68 @@ impl CorrelationWorld { /// Run the connection actor and collect frames for later verification. /// - /// # Panics - /// Panics if `self.expected` is `None` when the streaming scenario requires - /// a correlation id (via `expect("streaming scenario requires a correlation - /// id")`), or if running the actor fails. - pub async fn process(&mut self) { + /// # Errors + /// Returns an error if the expected correlation id is absent or if running + /// the actor fails. + pub async fn process(&mut self) -> TestResult { let cid = self .expected - .expect("streaming scenario requires a correlation id"); + .ok_or("streaming scenario requires a correlation id")?; let stream: FrameStream = Box::pin(try_stream! { yield Envelope::new(1, Some(cid), vec![1]); yield Envelope::new(1, Some(cid), vec![2]); }); - let (queues, handle) = build_small_queues::(); + let (queues, handle) = build_small_queues::()?; let shutdown = CancellationToken::new(); let mut actor = ConnectionActor::new(queues, handle, Some(stream), shutdown); - actor.run(&mut self.frames).await.expect("actor run failed"); + actor + .run(&mut self.frames) + .await + .map_err(|e| format!("actor run failed: {e:?}"))?; + Ok(()) } /// Run the connection actor for a multi-packet channel and collect frames. /// - /// # Panics - /// Panics if sending to the channel or running the actor fails. - pub async fn process_multi(&mut self) { + /// # Errors + /// Returns an error if sending frames or running the actor fails. + pub async fn process_multi(&mut self) -> TestResult { let expected = self.expected; let (tx, rx) = mpsc::channel(4); - tx.send(Envelope::new(1, None, vec![1])) - .await - .expect("send frame"); - tx.send(Envelope::new(1, Some(99), vec![2])) - .await - .expect("send frame"); + tx.send(Envelope::new(1, None, vec![1])).await?; + tx.send(Envelope::new(1, Some(99), vec![2])).await?; drop(tx); - let (queues, handle) = build_small_queues::(); + let (queues, handle) = build_small_queues::()?; let shutdown = CancellationToken::new(); let mut actor: ConnectionActor = ConnectionActor::new(queues, handle, None, shutdown); actor.set_multi_packet_with_correlation(Some(rx), expected); - actor.run(&mut self.frames).await.expect("actor run failed"); + actor + .run(&mut self.frames) + .await + .map_err(|e| format!("actor run failed: {e:?}"))?; + Ok(()) } /// Verify that all received frames respect the configured correlation expectation. /// - /// # Panics - /// Panics if any frame violates the stored correlation expectation. - pub fn verify(&self) { + /// # Errors + /// Returns an error if any frame violates the stored correlation + /// expectation. + pub fn verify(&self) -> TestResult { + let ok = match self.expected { + Some(cid) => self.frames.iter().all(|f| f.correlation_id() == Some(cid)), + None => self.frames.iter().all(|f| f.correlation_id().is_none()), + }; + + if ok { + return Ok(()); + } + match self.expected { - Some(cid) => { - assert!(self.frames.iter().all(|f| f.correlation_id() == Some(cid))); - } - None => { - assert!(self.frames.iter().all(|f| f.correlation_id().is_none())); - } + Some(cid) => Err(format!("frames missing expected correlation id {cid}").into()), + None => Err("frames unexpectedly carried correlation id".into()), } } } diff --git a/tests/worlds/fragment/mod.rs b/tests/worlds/fragment/mod.rs index 12542cc6..203c295e 100644 --- a/tests/worlds/fragment/mod.rs +++ b/tests/worlds/fragment/mod.rs @@ -26,6 +26,8 @@ use wireframe::fragment::{ ReassemblyError, }; +use super::TestResult; + #[derive(Debug, World)] pub struct FragmentWorld { series: Option, @@ -65,211 +67,238 @@ impl FragmentWorld { /// Configure a fragmenter with the provided payload cap so outbound /// fragmentation scenarios can chunk messages during behavioural tests. /// - /// # Panics - /// Panics if `max_payload` is zero. - pub fn configure_fragmenter(&mut self, max_payload: usize) { - let cap = NonZeroUsize::new(max_payload).expect("fragment cap must be non-zero"); + /// # Errors + /// Returns an error if the payload cap is zero. + pub fn configure_fragmenter(&mut self, max_payload: usize) -> TestResult { + let cap = NonZeroUsize::new(max_payload).ok_or("fragment cap must be non-zero")?; self.fragmenter = Some(Fragmenter::new(cap)); self.last_batch = None; + Ok(()) } /// Request fragmentation for a payload of `len` bytes, simulating outbound /// fragment production for the behavioural scenarios. /// - /// # Panics - /// Panics if [`configure_fragmenter`] has not been called yet. - pub fn fragment_payload(&mut self, len: usize) { - let fragmenter = self.fragmenter.as_ref().expect("fragmenter not configured"); + /// # Errors + /// Returns an error if the fragmenter is missing or fragmentation fails. + pub fn fragment_payload(&mut self, len: usize) -> TestResult { + let fragmenter = self + .fragmenter + .as_ref() + .ok_or("fragmenter not configured")?; let payload = vec![0_u8; len]; - let batch = fragmenter - .fragment_bytes(payload) - .expect("fragmentation must succeed in tests"); + let batch = fragmenter.fragment_bytes(payload)?; self.last_batch = Some(batch); + Ok(()) } /// Force the next expected fragment index for overflow scenarios. /// - /// # Panics - /// Panics if [`start_series`] has not been called. - pub fn force_next_index(&mut self, index: u32) { - self.series_mut() + /// # Errors + /// Returns an error if a fragment series has not been initialised. + pub fn force_next_index(&mut self, index: u32) -> TestResult { + self.series_mut()? .force_next_index_for_tests(FragmentIndex::new(index)); + Ok(()) } /// Feed a fragment that references the currently tracked message. /// - /// # Panics - /// Panics if [`start_series`] has not been called. - pub fn accept_fragment(&mut self, index: u32, is_last: bool) { - let message = self.series().message_id().get(); - self.accept_fragment_from(message, index, is_last); + /// # Errors + /// Returns an error if no fragment series has been initialised. + pub fn accept_fragment(&mut self, index: u32, is_last: bool) -> TestResult { + let message = self.series()?.message_id().get(); + self.accept_fragment_from(message, index, is_last) } /// Feed a fragment for an explicit message identifier. /// - /// # Panics - /// Panics if [`start_series`] has not been called. - pub fn accept_fragment_from(&mut self, message: u64, index: u32, is_last: bool) { + /// # Errors + /// Returns an error if no fragment series has been initialised. + pub fn accept_fragment_from(&mut self, message: u64, index: u32, is_last: bool) -> TestResult { let header = FragmentHeader::new(MessageId::new(message), FragmentIndex::new(index), is_last); - self.last_result = Some(self.series_mut().accept(header)); + self.last_result = Some(self.series_mut()?.accept(header)); + Ok(()) } /// Return the most recent fragment outcome. - /// - /// # Panics - /// Panics if no fragment has been processed yet. - fn last_result(&self) -> &Result { + fn last_result(&self) -> TestResult<&Result> { self.last_result .as_ref() - .expect("no fragment processed yet") + .ok_or_else(|| "no fragment processed yet".into()) } - fn batch(&self) -> &FragmentBatch { - self.last_batch.as_ref().expect("no payload fragmented yet") + fn batch(&self) -> TestResult<&FragmentBatch> { + self.last_batch + .as_ref() + .ok_or_else(|| "no payload fragmented yet".into()) } - /// Retrieve the fragment at `index`, panicking if it is missing. - fn get_fragment_at(&self, index: usize) -> &FragmentFrame { - self.batch() + /// Retrieve the fragment at `index`. + fn get_fragment_at(&self, index: usize) -> TestResult<&FragmentFrame> { + let fragment = self + .batch()? .fragments() .get(index) - .unwrap_or_else(|| panic!("fragment {index} missing")) + .ok_or_else(|| format!("fragment {index} missing"))?; + Ok(fragment) } - fn assert_error(&self, predicate: F, expected_desc: &str) + fn assert_error(&self, predicate: F, expected_desc: &str) -> TestResult where F: FnOnce(&FragmentError) -> bool, { - let err = match self.last_result() { + let err = match self.last_result()? { Err(err) => err, - Ok(status) => panic!("expected error but received {status:?}"), + Ok(status) => return Err(format!("expected error but received {status:?}").into()), }; - assert!(predicate(err), "expected {expected_desc}, got {err}"); + if !predicate(err) { + return Err(format!("expected {expected_desc}, got {err}").into()); + } + Ok(()) } /// Assert that the latest fragment completed the logical message. /// - /// # Panics - /// Panics if no fragment was processed or if the fragment failed to - /// complete the message. - pub fn assert_completion(&self) { - match self.last_result() { + /// # Errors + /// Returns an error if the fragment did not complete the message or no + /// fragment was processed. + pub fn assert_completion(&self) -> TestResult { + match self.last_result()? { Ok(FragmentStatus::Complete) => {} - Ok(status) => panic!("unexpected status: {status:?}"), - Err(err) => panic!("expected completion but got error: {err}"), + Ok(status) => return Err(format!("unexpected status: {status:?}").into()), + Err(err) => return Err(format!("expected completion but got error: {err}").into()), } - assert!( - self.series().is_complete(), - "series should be marked complete" - ); + if !self.series()?.is_complete() { + return Err("series should be marked complete".into()); + } + Ok(()) } - fn series(&self) -> &FragmentSeries { + fn series(&self) -> TestResult<&FragmentSeries> { self.series .as_ref() - .expect("fragment series not initialised") + .ok_or_else(|| "fragment series not initialised".into()) } - fn series_mut(&mut self) -> &mut FragmentSeries { + fn series_mut(&mut self) -> TestResult<&mut FragmentSeries> { self.series .as_mut() - .expect("fragment series not initialised") + .ok_or_else(|| "fragment series not initialised".into()) } /// Assert that the latest fragment failed due to an index mismatch. /// - /// # Panics - /// Panics if no fragment was processed or if the fragment failed for some - /// other reason. - pub fn assert_index_mismatch(&self) { + /// # Errors + /// Returns an error if the last fragment result does not indicate an index + /// mismatch or no fragment was processed. + pub fn assert_index_mismatch(&self) -> TestResult { self.assert_error( |err| matches!(err, FragmentError::IndexMismatch { .. }), "index mismatch", - ); + ) } /// Assert that the latest fragment failed because the message identifier /// did not match the tracked series. /// - /// # Panics - /// Panics if no fragment was processed or if the fragment failed for a - /// different reason. - pub fn assert_message_mismatch(&self) { + /// # Errors + /// Returns an error if the last fragment result is not a message mismatch + /// or no fragment was processed. + pub fn assert_message_mismatch(&self) -> TestResult { self.assert_error( |err| matches!(err, FragmentError::MessageMismatch { .. }), "message mismatch", - ); + ) } /// Assert that the latest fragment failed because the index overflowed. /// - /// # Panics - /// Panics if the series did not report an overflow. - pub fn assert_index_overflow(&self) { + /// # Errors + /// Returns an error if the last fragment result is not an overflow or no + /// fragment was processed. + pub fn assert_index_overflow(&self) -> TestResult { self.assert_error( |err| matches!(err, FragmentError::IndexOverflow { .. }), "overflow error", - ); + ) } - /// Assert that the latest fragment failed because the series was already complete. + /// Assert that the latest fragment failed because the series was already + /// complete. /// - /// # Panics - /// Panics if the series did not report a completion error. - pub fn assert_series_complete_error(&self) { + /// # Errors + /// Returns an error if the last fragment result is not a completion error + /// or no fragment was processed. + pub fn assert_series_complete_error(&self) -> TestResult { self.assert_error( |err| matches!(err, FragmentError::SeriesComplete), "series completion error", - ); + ) } /// Assert that the most recent fragmentation produced `expected` fragments /// for outbound fragmentation scenarios. /// - /// # Panics - /// Panics if no payload has been fragmented yet. - pub fn assert_fragment_count(&self, expected: usize) { - assert_eq!(self.batch().len(), expected, "unexpected fragment count"); + /// # Errors + /// Returns an error if no batch exists or the fragment count mismatches. + pub fn assert_fragment_count(&self, expected: usize) -> TestResult { + let actual = self.batch()?.len(); + if actual != expected { + return Err(format!("expected {expected} fragments, got {actual}").into()); + } + Ok(()) } /// Assert that the payload length of fragment `index` matches `expected` /// bytes for outbound fragments. /// - /// # Panics - /// Panics if no payload has been fragmented or if `index` exceeds the batch. - pub fn assert_fragment_payload_len(&self, index: usize, expected: usize) { - let fragment = self.get_fragment_at(index); - assert_eq!( - fragment.payload().len(), - expected, - "payload length mismatch" - ); + /// # Errors + /// Returns an error if the batch is missing or the payload length differs. + pub fn assert_fragment_payload_len(&self, index: usize, expected: usize) -> TestResult { + let fragment = self.get_fragment_at(index)?; + let actual = fragment.payload().len(); + if actual != expected { + return Err(format!( + "fragment {index} payload length mismatch: expected {expected}, got {actual}" + ) + .into()); + } + Ok(()) } /// Assert that outbound fragment `index` carries the expected final flag. /// - /// # Panics - /// Panics if no payload has been fragmented or if `index` exceeds the batch. - pub fn assert_fragment_final_flag(&self, index: usize, expected_final: bool) { - let fragment = self.get_fragment_at(index); - assert_eq!( - fragment.header().is_last_fragment(), - expected_final, - "fragment {index} final flag mismatch", - ); + /// # Errors + /// Returns an error if the batch is missing or the final flag mismatches. + pub fn assert_fragment_final_flag(&self, index: usize, expected_final: bool) -> TestResult { + let fragment = self.get_fragment_at(index)?; + let actual = fragment.header().is_last_fragment(); + if actual != expected_final { + return Err(format!( + "fragment {index} final flag mismatch: expected {expected_final}, got {actual}" + ) + .into()); + } + Ok(()) } /// Assert that the outbound fragment batch carries the expected message /// identifier. /// - /// # Panics - /// Panics if no payload has been fragmented yet. - pub fn assert_message_id(&self, expected: u64) { - assert_eq!( - self.batch().message_id(), - MessageId::new(expected), - "unexpected message identifier", - ); + /// # Errors + /// Returns an error if the batch is missing or the message id differs from + /// the expectation. + pub fn assert_message_id(&self, expected: u64) -> TestResult { + let actual = self.batch()?.message_id(); + let expected_id = MessageId::new(expected); + if actual != expected_id { + return Err(format!( + "unexpected message identifier: expected {expected_id:?}, got {actual:?}" + ) + .into()); + } + Ok(()) } } diff --git a/tests/worlds/fragment/reassembly.rs b/tests/worlds/fragment/reassembly.rs index 20d9a4e5..55e39a8a 100644 --- a/tests/worlds/fragment/reassembly.rs +++ b/tests/worlds/fragment/reassembly.rs @@ -9,30 +9,37 @@ use super::{ MessageId, Reassembler, ReassemblyError, + TestResult, }; impl FragmentWorld { /// Configure a reassembler with size and timeout guards. /// - /// # Panics - /// Panics if `max_message_size` is zero. - pub fn configure_reassembler(&mut self, max_message_size: usize, timeout_secs: u64) { - let size = NonZeroUsize::new(max_message_size).expect("reassembly cap must be non-zero"); + /// # Errors + /// Returns an error when the message size is zero or the configuration + /// cannot be constructed. + pub fn configure_reassembler( + &mut self, + max_message_size: usize, + timeout_secs: u64, + ) -> TestResult { + let size = NonZeroUsize::new(max_message_size).ok_or("reassembly cap must be non-zero")?; self.reassembler = Some(Reassembler::new(size, Duration::from_secs(timeout_secs))); self.last_reassembled = None; self.last_reassembly_error = None; self.last_evicted.clear(); + Ok(()) } /// Submit a fragment to the configured reassembler. /// - /// # Panics - /// Panics if the reassembler has not been configured. - pub fn push_fragment(&mut self, header: FragmentHeader, payload_len: usize) { + /// # Errors + /// Returns an error if the reassembler is missing. + pub fn push_fragment(&mut self, header: FragmentHeader, payload_len: usize) -> TestResult { let reassembler = self .reassembler .as_mut() - .expect("reassembler not configured"); + .ok_or("reassembler not configured")?; let payload = vec![0_u8; payload_len]; self.last_reassembly_error = None; self.last_reassembled = None; @@ -40,94 +47,104 @@ impl FragmentWorld { Ok(output) => self.last_reassembled = output, Err(err) => self.last_reassembly_error = Some(err), } + Ok(()) } /// Advance the simulated clock. /// - /// # Panics - /// - /// Panics if advancing the clock would overflow [`Instant`]. - pub fn advance_time(&mut self, delta: Duration) { + /// # Errors + /// Returns an error if the simulated clock would overflow. + pub fn advance_time(&mut self, delta: Duration) -> TestResult { self.now = self .now .checked_add(delta) - .expect("time advance overflowed"); + .ok_or("time advance overflowed")?; + Ok(()) } /// Purge expired partial messages based on the current clock reading. /// - /// # Panics - /// - /// Panics if the reassembler has not been configured. - pub fn purge_reassembly(&mut self) { + /// # Errors + /// Returns an error if the reassembler has not been configured. + pub fn purge_reassembly(&mut self) -> TestResult { let reassembler = self .reassembler .as_mut() - .expect("reassembler not configured"); + .ok_or("reassembler not configured")?; self.last_evicted = reassembler.purge_expired_at(self.now); + Ok(()) } - /// Assert that a message has been reassembled with the expected payload length. + /// Assert that a message has been reassembled with the expected payload + /// length. /// - /// # Panics - /// Panics if no message has been reassembled yet. - pub fn assert_reassembled_len(&self, expected_len: usize) { + /// # Errors + /// Returns an error if no message has been reassembled or the length does + /// not match the expectation. + pub fn assert_reassembled_len(&self, expected_len: usize) -> TestResult { let message = self .last_reassembled .as_ref() - .expect("no message reassembled"); - assert_eq!( - message.payload().len(), - expected_len, - "payload length mismatch" - ); + .ok_or("no message reassembled")?; + if message.payload().len() != expected_len { + return Err("payload length mismatch".into()); + } + Ok(()) } /// Assert that no message has been fully reassembled. /// - /// # Panics - /// - /// Panics if a message has already been reassembled. - pub fn assert_no_reassembly(&self) { - assert!( - self.last_reassembled.is_none(), - "unexpected reassembled message present" - ); + /// # Errors + /// Returns an error if a message has already been reassembled. + pub fn assert_no_reassembly(&self) -> TestResult { + if self.last_reassembled.is_some() { + return Err("unexpected reassembled message present".into()); + } + Ok(()) } /// Helper for asserting on the latest captured reassembly error. /// - /// # Panics - /// Panics if no reassembly error was captured or if the predicate returns false. - fn assert_reassembly_error_matches(&self, predicate: F, expected_description: &str) + /// # Errors + /// Returns an error when no reassembly error was captured or the predicate + /// does not match the error variant. + fn assert_reassembly_error_matches( + &self, + predicate: F, + expected_description: &str, + ) -> TestResult where F: FnOnce(&ReassemblyError) -> bool, { let err = self .last_reassembly_error .as_ref() - .expect("no reassembly error captured"); - assert!(predicate(err), "expected {expected_description}, got {err}"); + .ok_or("no reassembly error captured")?; + if !predicate(err) { + return Err(format!("expected {expected_description}, got {err}").into()); + } + Ok(()) } /// Assert the latest reassembly error signalled an over-limit message. /// - /// # Panics - /// - /// Panics if no reassembly error was captured. - pub fn assert_reassembly_over_limit(&self) { + /// # Errors + /// Returns an error if no reassembly error was captured or it was not a + /// message-too-large error. + pub fn assert_reassembly_over_limit(&self) -> TestResult { self.assert_reassembly_error_matches( |err| matches!(err, ReassemblyError::MessageTooLarge { .. }), "message-too-large error", - ); + ) } - /// Assert that the latest reassembly error was triggered by an out-of-order fragment. - /// - /// # Panics + /// Assert that the latest reassembly error was triggered by an out-of-order + /// fragment. /// - /// Panics if no reassembly error was captured or if the error was not an index mismatch. - pub fn assert_reassembly_out_of_order(&self) { + /// # Errors + /// Returns an error if no reassembly error was captured or it was not an + /// index-mismatch error. + pub fn assert_reassembly_out_of_order(&self) -> TestResult { self.assert_reassembly_error_matches( |err| { matches!( @@ -136,34 +153,34 @@ impl FragmentWorld { ) }, "out-of-order error", - ); + ) } /// Assert the number of buffered partial messages. /// - /// # Panics - /// Panics if the reassembler has not been configured. - pub fn assert_buffered_messages(&self, expected: usize) { + /// # Errors + /// Returns an error if the reassembler is missing or the buffered count + /// differs from the expectation. + pub fn assert_buffered_messages(&self, expected: usize) -> TestResult { let reassembler = self .reassembler .as_ref() - .expect("reassembler not configured"); - assert_eq!( - reassembler.buffered_len(), - expected, - "unexpected buffered message count" - ); + .ok_or("reassembler not configured")?; + let actual = reassembler.buffered_len(); + if actual != expected { + return Err(format!("expected {expected} buffered messages, got {actual}").into()); + } + Ok(()) } /// Assert that the most recent purge evicted a specific message identifier. /// - /// # Panics - /// - /// Panics if the purge record does not contain `message_id`. - pub fn assert_evicted_message(&self, message_id: u64) { - assert!( - self.last_evicted.contains(&MessageId::new(message_id)), - "message {message_id} was not evicted" - ); + /// # Errors + /// Returns an error if the expected message identifier was not evicted. + pub fn assert_evicted_message(&self, message_id: u64) -> TestResult { + if !self.last_evicted.contains(&MessageId::new(message_id)) { + return Err(format!("message {message_id} was not evicted").into()); + } + Ok(()) } } diff --git a/tests/worlds/mod.rs b/tests/worlds/mod.rs index ac54f948..99a2ccaa 100644 --- a/tests/worlds/mod.rs +++ b/tests/worlds/mod.rs @@ -7,8 +7,8 @@ #![cfg(not(loom))] #[path = "../common/mod.rs"] -mod common; -pub(crate) use common::unused_listener; +pub mod common; +pub use common::{TestResult, unused_listener}; #[path = "../common/terminator.rs"] mod terminator; @@ -22,11 +22,8 @@ use wireframe::{app::Envelope, push::PushQueues, serializer::BincodeSerializer}; pub(crate) type TestApp = wireframe::app::WireframeApp; pub(crate) fn build_small_queues() --> (PushQueues, wireframe::push::PushHandle) { - support::builder::() - .unlimited() - .build() - .expect("failed to build PushQueues") +-> Result<(PushQueues, wireframe::push::PushHandle), wireframe::push::PushConfigError> { + support::builder::().unlimited().build() } pub mod correlation; diff --git a/tests/worlds/multi_packet.rs b/tests/worlds/multi_packet.rs index d8925132..13a96648 100644 --- a/tests/worlds/multi_packet.rs +++ b/tests/worlds/multi_packet.rs @@ -4,12 +4,23 @@ //! Provides [`MultiPacketWorld`] to verify message ordering, back-pressure //! handling, and channel lifecycle in cucumber-based behaviour tests. +use std::{error::Error, fmt}; + use cucumber::World; use tokio::sync::mpsc::{self, error::TrySendError}; use tokio_util::sync::CancellationToken; use wireframe::{Response, connection::ConnectionActor}; -use super::build_small_queues; +use super::{TestResult, build_small_queues}; + +#[derive(Debug)] +struct WireframeRunError(wireframe::WireframeError); + +impl fmt::Display for WireframeRunError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{:?}", self.0) } +} + +impl Error for WireframeRunError {} #[derive(Debug, Default, World)] pub struct MultiPacketWorld { @@ -18,86 +29,99 @@ pub struct MultiPacketWorld { } impl MultiPacketWorld { - async fn collect_frames_from(rx: mpsc::Receiver) -> Vec { - let (queues, handle) = build_small_queues::(); + async fn collect_frames_from(rx: mpsc::Receiver) -> TestResult> { + let (queues, handle) = build_small_queues::()?; let shutdown = CancellationToken::new(); let mut actor: ConnectionActor<_, ()> = ConnectionActor::new(queues, handle, None, shutdown); actor.set_multi_packet(Some(rx)); let mut frames = Vec::new(); - actor.run(&mut frames).await.expect("actor run failed"); - frames + actor + .run(&mut frames) + .await + .map_err(WireframeRunError) + .map_err(Box::::from)?; + Ok(frames) + } + + /// Send a single byte with back-pressure then close the channel. + async fn send_with_backpressure(sender: mpsc::Sender, value: u8) -> TestResult<()> { + sender.send(value).await?; + drop(sender); + Ok(()) } /// Helper method to process messages through a multi-packet response built /// via [`Response::with_channel`]. /// - /// # Panics - /// Panics if [`Response::with_channel`] fails to produce a `MultiPacket` - /// response or if spawning or joining the producer task fails. - async fn process_messages(&mut self, messages: &[u8]) { + /// # Errors + /// Returns an error if the response cannot be converted to a multi-packet + /// stream or if producer tasks fail. + async fn process_messages(&mut self, messages: &[u8]) -> TestResult { let (sender, response): (mpsc::Sender, Response) = Response::with_channel(4); let Response::MultiPacket(rx) = response else { - panic!("helper did not return a MultiPacket response"); + return Err("helper did not return a MultiPacket response".into()); }; let payload = messages.to_vec(); - let producer = tokio::spawn(async move { - for msg in payload { - if sender.send(msg).await.is_err() { - return; - } - } - drop(sender); - }); + let producer = tokio::spawn(Self::send_payload(sender, payload)); - let frames = Self::collect_frames_from(rx).await; - producer.await.expect("producer task panicked"); + let frames = Self::collect_frames_from(rx).await?; + producer.await?; self.messages = frames; self.is_overflow_error = false; + Ok(()) + } + + /// Send each byte to the channel, stopping silently if the receiver closes + /// to simulate a producer completing without error when the consumer is + /// gone. + async fn send_payload(sender: mpsc::Sender, payload: Vec) { + for msg in payload { + if sender.send(msg).await.is_err() { + return; + } + } } /// Send messages through a multi-packet response and record them. /// - /// # Panics - /// Panics if [`Response::with_channel`] fails to produce a `MultiPacket` - /// response or if spawning or joining the producer task fails. - pub async fn process(&mut self) { self.process_messages(&[1, 2, 3]).await; } + /// # Errors + /// Returns an error if the response cannot be converted to a multi-packet + /// stream or if producer tasks fail. + pub async fn process(&mut self) -> TestResult { self.process_messages(&[1, 2, 3]).await } /// Record zero messages from a closed channel. /// - /// # Panics - /// Panics if [`Response::with_channel`] fails to produce a `MultiPacket` - /// response or if spawning or joining the producer task fails. - pub async fn process_empty(&mut self) { self.process_messages(&[]).await; } + /// # Errors + /// Returns an error if the response cannot be converted to a multi-packet + /// stream or if producer tasks fail. + pub async fn process_empty(&mut self) -> TestResult { self.process_messages(&[]).await } /// Attempt to send more messages than the channel can buffer at once. /// - /// # Panics - /// Panics if sending to the channel fails unexpectedly or the producer task panics. - pub async fn process_overflow(&mut self) { + /// # Errors + /// Returns an error if sending to the channel fails unexpectedly or the + /// producer task returns an error. + pub async fn process_overflow(&mut self) -> TestResult { let (sender, response): (mpsc::Sender, Response) = Response::with_channel(1); let Response::MultiPacket(rx) = response else { - panic!("helper did not return a MultiPacket response"); + return Err("helper did not return a MultiPacket response".into()); }; - sender.try_send(1).expect("send initial frame"); + sender.try_send(1)?; let overflow_error = matches!(sender.try_send(2), Err(TrySendError::Full(2))); - let producer = tokio::spawn(async move { - sender - .send(2) - .await - .expect("send follow-up frame after draining"); - drop(sender); - }); + let producer = tokio::spawn(Self::send_with_backpressure(sender, 2)); - let frames = Self::collect_frames_from(rx).await; - producer.await.expect("producer task panicked"); + let frames = Self::collect_frames_from(rx).await?; + // Unwrap JoinError from await, then the task's Result + producer.await??; self.messages = frames; self.is_overflow_error = overflow_error; + Ok(()) } /// Verify that no messages were received. diff --git a/tests/worlds/panic.rs b/tests/worlds/panic.rs index 86115262..66e6eeb6 100644 --- a/tests/worlds/panic.rs +++ b/tests/worlds/panic.rs @@ -10,7 +10,7 @@ use cucumber::World; use tokio::{net::TcpStream, sync::oneshot}; use wireframe::server::WireframeServer; -use super::{TestApp, unused_listener}; +use super::{TestApp, TestResult, unused_listener}; #[derive(Debug)] struct PanicServer { @@ -20,38 +20,42 @@ struct PanicServer { } impl PanicServer { - async fn spawn() -> Self { + #[expect( + clippy::expect_used, + reason = "panic world should fail loudly if the panic app cannot be built" + )] + async fn spawn() -> TestResult { let factory = || { TestApp::new() - .expect("Failed to create WireframeApp") - .on_connection_setup(|| async { panic!("boom") }) - .expect("Failed to set connection setup callback") + .and_then(|app| app.on_connection_setup(|| async { panic!("boom") })) + .expect("failed to build panic app") }; let listener = unused_listener(); let server = WireframeServer::new(factory) .workers(1) - .bind_existing_listener(listener) - .expect("bind"); - let addr = server.local_addr().expect("Failed to get server address"); + .bind_existing_listener(listener)?; + let addr = server.local_addr().ok_or("Failed to get server address")?; let (tx_shutdown, rx_shutdown) = oneshot::channel(); let (tx_ready, rx_ready) = oneshot::channel(); let handle = tokio::spawn(async move { - server + if let Err(err) = server .ready_signal(tx_ready) .run_with_shutdown(async { let _ = rx_shutdown.await; }) .await - .expect("Server task failed"); + { + tracing::error!("server task failed: {err}"); + } }); - rx_ready.await.expect("Server did not signal ready"); + rx_ready.await.map_err(|_| "Server did not signal ready")?; - Self { + Ok(Self { addr, shutdown: Some(tx_shutdown), handle, - } + }) } } @@ -80,29 +84,39 @@ pub struct PanicWorld { impl PanicWorld { /// Start a server that panics during connection setup. /// - /// # Panics - /// Panics if `TestApp::new()` fails, `.on_connection_setup(...)` fails, - /// binding the server fails, or the server task fails. - pub async fn start_panic_server(&mut self) { self.server.replace(PanicServer::spawn().await); } + /// # Errors + /// Returns an error if building the app factory or binding the server + /// fails. + pub async fn start_panic_server(&mut self) -> TestResult { + let server = PanicServer::spawn().await?; + self.server.replace(server); + Ok(()) + } /// Connect to the running server once. /// - /// # Panics - /// Panics if the server address is unknown or the connection fails. - pub async fn connect_once(&mut self) { - let addr = self.server.as_ref().expect("Server not started").addr; - TcpStream::connect(addr).await.expect("Failed to connect"); + /// # Errors + /// Returns an error if the server address is unknown or the connection + /// attempt fails. + pub async fn connect_once(&mut self) -> TestResult { + let addr = self.server.as_ref().ok_or("Server not started")?.addr; + TcpStream::connect(addr).await?; self.attempts += 1; + Ok(()) } /// Verify both connections succeeded and shut down the server. /// - /// # Panics - /// Panics if the connection attempts do not match the expected count. - pub async fn verify_and_shutdown(&mut self) { - assert_eq!(self.attempts, 2); + /// # Errors + /// Returns an error if the connection attempts do not match the expected + /// count. + pub async fn verify_and_shutdown(&mut self) -> TestResult { + if self.attempts != 2 { + return Err("expected two successful connection attempts".into()); + } // dropping PanicServer will shut it down self.server.take(); tokio::task::yield_now().await; + Ok(()) } } diff --git a/tests/worlds/stream_end.rs b/tests/worlds/stream_end.rs index 873faf68..d4d1e957 100644 --- a/tests/worlds/stream_end.rs +++ b/tests/worlds/stream_end.rs @@ -12,13 +12,13 @@ use log::Level; use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; use wireframe::{ - connection::{ConnectionActor, test_support::ActorHarness}, + connection::{ConnectionActor, ConnectionChannels, test_support::ActorHarness}, hooks::ProtocolHooks, response::FrameStream, }; use wireframe_testing::{LoggerHandle, logger}; -use super::{Terminator, build_small_queues}; +use super::{Terminator, TestResult, build_small_queues}; #[derive(Debug, Default, World)] pub struct StreamEndWorld { @@ -47,12 +47,12 @@ impl StreamEndWorld { fn finalize_test(&mut self, logger: &mut LoggerHandle) { self.capture_logs(logger); } - async fn run_actor_test(&mut self, mode: ActorMode) { + async fn run_actor_test(&mut self, mode: ActorMode) -> TestResult { let mut temp = StreamEndWorld::default(); mem::swap(self, &mut temp); let mut logger = temp.prepare_test(); - let (queues, handle) = build_small_queues::(); + let (queues, handle) = build_small_queues::()?; let shutdown = CancellationToken::new(); let hooks = ProtocolHooks::from_protocol(&Arc::new(Terminator)); @@ -62,37 +62,55 @@ impl StreamEndWorld { yield 1u8; yield 2u8; }); - let mut actor = - ConnectionActor::with_hooks(queues, handle, Some(stream), shutdown, hooks); - actor.run(&mut temp.frames).await.expect("actor run failed"); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + Some(stream), + shutdown, + hooks, + ); + actor + .run(&mut temp.frames) + .await + .map_err(|e| format!("actor run failed: {e:?}"))?; } ActorMode::MultiPacket => { let (tx, rx) = mpsc::channel(4); - tx.send(1u8).await.expect("send frame"); - tx.send(2u8).await.expect("send frame"); + tx.send(1u8).await?; + tx.send(2u8).await?; drop(tx); - let mut actor = ConnectionActor::with_hooks(queues, handle, None, shutdown, hooks); + let mut actor = ConnectionActor::with_hooks( + ConnectionChannels::new(queues, handle), + None, + shutdown, + hooks, + ); actor.set_multi_packet(Some(rx)); - actor.run(&mut temp.frames).await.expect("actor run failed"); + actor + .run(&mut temp.frames) + .await + .map_err(|e| format!("actor run failed: {e:?}"))?; } } temp.finalize_test(&mut logger); mem::swap(self, &mut temp); + Ok(()) } /// Run the connection actor and record emitted frames. /// - /// # Panics - /// Panics if the actor fails to run successfully. - pub async fn process(&mut self) { self.run_actor_test(ActorMode::Stream).await; } + /// # Errors + /// Returns an error if the actor fails to run successfully. + pub async fn process(&mut self) -> TestResult { self.run_actor_test(ActorMode::Stream).await } /// Run the connection actor with a multi-packet channel and record emitted frames. /// - /// # Panics - /// Panics if sending to the channel or running the actor fails. - pub async fn process_multi(&mut self) { self.run_actor_test(ActorMode::MultiPacket).await; } + /// # Errors + /// Returns an error if sending to the channel or running the actor fails. + pub async fn process_multi(&mut self) -> TestResult { + self.run_actor_test(ActorMode::MultiPacket).await + } fn capture_logs(&mut self, logger: &mut LoggerHandle) { while let Some(record) = logger.pop() { @@ -107,14 +125,17 @@ impl StreamEndWorld { .find(|(_, message)| message.contains("multi-packet stream closed")) } - fn run_multi_packet_harness(&mut self, mode: &MultiPacketMode, correlation_id: u64) { + fn run_multi_packet_harness( + &mut self, + mode: &MultiPacketMode, + correlation_id: u64, + ) -> TestResult { let mut temp = StreamEndWorld::default(); mem::swap(self, &mut temp); let mut logger = temp.prepare_test(); let hooks = ProtocolHooks::from_protocol(&Arc::new(Terminator)); - let mut harness = ActorHarness::new_with_state(hooks, false, true) - .expect("failed to create ActorHarness"); + let mut harness = ActorHarness::new_with_state(hooks, false, true)?; let (tx, rx) = mpsc::channel(4); harness .actor_mut() @@ -122,8 +143,8 @@ impl StreamEndWorld { match mode { MultiPacketMode::Disconnect { send_frames } => { if *send_frames { - tx.try_send(1u8).expect("send frame"); - tx.try_send(2u8).expect("send frame"); + tx.try_send(1u8)?; + tx.try_send(2u8)?; } drop(tx); logger.clear(); @@ -139,22 +160,23 @@ impl StreamEndWorld { temp.finalize_test(&mut logger); mem::swap(self, &mut temp); + Ok(()) } /// Simulate a disconnected multi-packet channel by dropping the sender before draining. /// - /// # Panics - /// Panics if creating the harness or sending frames fails. - pub fn process_multi_disconnect(&mut self) { - self.run_multi_packet_harness(&MultiPacketMode::Disconnect { send_frames: true }, 42); + /// # Errors + /// Returns an error if creating the harness or sending frames fails. + pub fn process_multi_disconnect(&mut self) -> TestResult { + self.run_multi_packet_harness(&MultiPacketMode::Disconnect { send_frames: true }, 42) } /// Trigger shutdown handling on a multi-packet channel without emitting a terminator. /// - /// # Panics - /// Panics if creating the harness fails. - pub fn process_multi_shutdown(&mut self) { - self.run_multi_packet_harness(&MultiPacketMode::Shutdown, 77); + /// # Errors + /// Returns an error if creating the harness fails. + pub fn process_multi_shutdown(&mut self) -> TestResult { + self.run_multi_packet_harness(&MultiPacketMode::Shutdown, 77) } /// Verify that a terminator frame was appended to the stream. @@ -186,23 +208,23 @@ impl StreamEndWorld { /// Verify the logged multi-packet termination reason. /// - /// # Panics - /// Panics if the closure log is missing or contains unexpected details. - pub fn verify_reason(&self, expected: &str) { + /// # Errors + /// Returns an error if the closure log is missing or contains unexpected + /// details. + pub fn verify_reason(&self, expected: &str) -> TestResult { let (level, message) = self .closure_log() - .expect("multi-packet closure log missing"); + .ok_or("multi-packet closure log missing")?; let expected_level = match expected { "disconnected" => Level::Warn, _ => Level::Info, }; - assert_eq!( - *level, expected_level, - "unexpected log level: message={message}", - ); - assert!( - message.contains(&format!("reason={expected}")), - "closure log missing reason: message={message}", - ); + if *level != expected_level { + return Err("unexpected log level for closure".into()); + } + if !message.contains(&format!("reason={expected}")) { + return Err("closure log missing reason detail".into()); + } + Ok(()) } } diff --git a/wireframe_testing/README.md b/wireframe_testing/README.md index 62edf791..fbc536ea 100644 --- a/wireframe_testing/README.md +++ b/wireframe_testing/README.md @@ -1,10 +1,11 @@ # wireframe_testing -Helper utilities for exercising [`wireframe`](https://crates.io/crates/wireframe) -applications in tests without opening real sockets. The crate runs a -`WireframeApp` against in-memory duplex streams, captures every frame the app -emits, and provides small helpers for encoding or decoding frames so assertions -stay focused on behaviour rather than plumbing. +Helper utilities for exercising +[`wireframe`](https://crates.io/crates/wireframe) applications in tests without +opening real sockets. The crate runs a `WireframeApp` against in-memory duplex +streams, captures every frame the app emits, and provides small helpers for +encoding or decoding frames so assertions stay focused on behaviour rather than +plumbing. - Drive an app with length-delimited frames or bincode-serialised payloads. - Collect multi-frame responses into a single buffer for snapshot-style