From d11d4ea08b3102e83df7dc1b16e0cba1d79e05a5 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Fri, 25 Sep 2026 05:35:57 +0200 Subject: [PATCH] Make the protocol transport-agnostic and use tokio-tungstenite Replace the hand-rolled WebSocket implementation with a transport boundary (Frame over any Sink/Stream) plus adapters for tokio-tungstenite (server/client features) and axum (axum feature). The JSON wire format is unchanged, so existing peers keep working. Fixes: - unbounded frame lengths were allocated up front; the server now enforces message/frame size limits (1 MiB default) - responses arriving before `request` subscribed were lost - an accept() error ended `session_loop` - requests sent right after connecting could arrive before handlers were registered; the receiver now starts after `on_conn` returns - panics in the receive loop skipped `on_close` and leaked sessions; `on_close` now runs exactly once and handler panics fail only their request - unknown methods and invalid data got no reply; they now get an error response - slow handlers blocked pongs and responses Adds `on_notification`, `request_timeout`, `closed`, `is_closed`, `id`, `ServerConfig`, `from_transport`, `from_tungstenite` and `from_axum`, integration tests, and an axum example. Bumps to 0.2.0 since `Session::connect` now takes a URL and the `ws` module is removed. Co-Authored-By: Claude Opus 5.5 --- Cargo.lock | 988 ++++++++++++++++++++++++++++++++++++++++++-- Cargo.toml | 45 +- README.md | 119 ++++-- examples/axum.rs | 41 ++ examples/client.rs | 6 +- src/axum.rs | 39 ++ src/client.rs | 13 + src/lib.rs | 58 ++- src/server.rs | 128 ++++-- src/session.rs | 555 +++++++++++++++++-------- src/transport.rs | 36 ++ src/tungstenite.rs | 47 +++ src/ws/error.rs | 31 -- src/ws/handshake.rs | 225 ---------- src/ws/mod.rs | 251 ----------- tests/protocol.rs | 452 ++++++++++++++++++++ 16 files changed, 2230 insertions(+), 804 deletions(-) create mode 100644 examples/axum.rs create mode 100644 src/axum.rs create mode 100644 src/client.rs create mode 100644 src/transport.rs create mode 100644 src/tungstenite.rs delete mode 100644 src/ws/error.rs delete mode 100644 src/ws/handshake.rs delete mode 100644 src/ws/mod.rs create mode 100644 tests/protocol.rs diff --git a/Cargo.lock b/Cargo.lock index 43db25f..48a6f12 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -8,6 +8,67 @@ version = "1.0.101" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5f0e0fee31ef5ed1ba1316088939cea399010ed7731dba877ed44aeb407a75ea" +[[package]] +name = "atomic-waker" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" + +[[package]] +name = "axum" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90" +dependencies = [ + "axum-core", + "base64", + "bytes", + "form_urlencoded", + "futures-util", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-util", + "itoa", + "matchit", + "memchr", + "mime", + "percent-encoding", + "pin-project-lite", + "serde_core", + "serde_json", + "serde_path_to_error", + "serde_urlencoded", + "sha1 0.10.6", + "sync_wrapper", + "tokio", + "tokio-tungstenite 0.29.0", + "tower", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "axum-core" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "sync_wrapper", + "tower-layer", + "tower-service", + "tracing", +] + [[package]] name = "base64" version = "0.22.1" @@ -29,12 +90,31 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-buffer" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2f6c7dbe95a6ed67ad9f18e57daf93a2f034c524b99fd2b76d18fdfeb6660aa" +dependencies = [ + "hybrid-array", +] + [[package]] name = "bytes" version = "1.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b35204fbdc0b3f4446b89fc1ac2cf84a8a68971995d0bf2e925ec7cd960f9cb3" +[[package]] +name = "cc" +version = "1.4.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "54413ede23c2daf518f35156dfde027feb2374004d63bd497f983c8db9c0e313" +dependencies = [ + "find-msvc-tools", + "shlex", +] + [[package]] name = "cfg-if" version = "1.0.4" @@ -43,15 +123,37 @@ checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" [[package]] name = "chacha20" -version = "0.10.0" +version = "0.10.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601" +checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" dependencies = [ "cfg-if", "cpufeatures 0.3.0", - "rand_core", + "rand_core 0.10.0", ] +[[package]] +name = "const-oid" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" + +[[package]] +name = "core-foundation" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2a6cd9ae233e7f62ba4e9353e81a88df7fc8a5987b8d445b4d90c879bd156f6" +dependencies = [ + "core-foundation-sys", + "libc", +] + +[[package]] +name = "core-foundation-sys" +version = "0.8.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" + [[package]] name = "cpufeatures" version = "0.2.17" @@ -80,14 +182,40 @@ dependencies = [ "typenum", ] +[[package]] +name = "crypto-common" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce6e4c961d6cd6c9a86db418387425e8bdeaf05b3c8bc1411e6dca4c252f1453" +dependencies = [ + "hybrid-array", +] + +[[package]] +name = "data-encoding" +version = "2.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4583a4551df46e2792f82ceeac45e850d2e2d5debba0b91f102385cda5b11f06" + [[package]] name = "digest" version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer", - "crypto-common", + "block-buffer 0.10.4", + "crypto-common 0.1.7", +] + +[[package]] +name = "digest" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" +dependencies = [ + "block-buffer 0.12.1", + "const-oid", + "crypto-common 0.2.2", ] [[package]] @@ -96,12 +224,98 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "fastrand" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" + +[[package]] +name = "find-msvc-tools" +version = "0.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ef25905e51abafe4dcea6c15fec58c57b601cdbd0ee53d22ea1d3016c587d39b" + [[package]] name = "foldhash" version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" +[[package]] +name = "foreign-types" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f6f339eb8adc052cd2ca78910fda869aefa38d22d5cb648e6485e4d3fc06f3b1" +dependencies = [ + "foreign-types-shared", +] + +[[package]] +name = "foreign-types-shared" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "00b0228411908ca8685dba7fc2cdd70ec9990a6e753e89b6ac91a84c40fbaf4b" + +[[package]] +name = "form_urlencoded" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" +dependencies = [ + "percent-encoding", +] + +[[package]] +name = "futures-channel" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b1f9e3d69d39e4862ffed03ed071a76f9a13ba1d9109d355b0f0aa6b15e393c4" +dependencies = [ + "futures-core", +] + +[[package]] +name = "futures-core" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e" + +[[package]] +name = "futures-sink" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1944426bf7d03f1d14f708785e4b33efd750b36d48a157b836b3efc15ede8e1d" + +[[package]] +name = "futures-task" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cd417de3d1d015fc3bfd2b1ea46dfc7bab72ef86f1cc7cc9c78e728b34a6d1fd" + +[[package]] +name = "futures-util" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc" +dependencies = [ + "futures-core", + "futures-sink", + "futures-task", + "pin-project-lite", + "slab", +] + [[package]] name = "generic-array" version = "0.14.7" @@ -112,6 +326,29 @@ dependencies = [ "version_check", ] +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "libc", + "wasi", +] + +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi", + "wasip2", +] + [[package]] name = "getrandom" version = "0.4.1" @@ -121,7 +358,7 @@ dependencies = [ "cfg-if", "libc", "r-efi", - "rand_core", + "rand_core 0.10.0", "wasip2", "wasip3", ] @@ -147,6 +384,95 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "http" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "918d3568bebf352712bc2ef3d46a8bcf1a75b373be6539de198e9105cbbf9ce0" +dependencies = [ + "bytes", + "itoa", +] + +[[package]] +name = "http-body" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ca2a8f2913ee65f60facd6a5905613afaa448497a0230cc41ce022d93290bc2c" +dependencies = [ + "bytes", + "http", +] + +[[package]] +name = "http-body-util" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23169fe34a5fbcdd3f3862e78fb9b6fccd5f02a6dc6f732547005d45631ce71c" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "pin-project-lite", +] + +[[package]] +name = "httparse" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" + +[[package]] +name = "httpdate" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" + +[[package]] +name = "hybrid-array" +version = "0.4.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3944cf8cf766b40e2a1a333ee5e9b563f854d5fa49d6a8ca2764e97c6eddb214" +dependencies = [ + "typenum", +] + +[[package]] +name = "hyper" +version = "1.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "27b501faa50e7a26c3d3560ca625132f4078a17771f4810baf70475ae48cbe43" +dependencies = [ + "atomic-waker", + "bytes", + "futures-channel", + "futures-core", + "http", + "http-body", + "httparse", + "httpdate", + "itoa", + "pin-project-lite", + "smallvec", + "tokio", +] + +[[package]] +name = "hyper-util" +version = "0.1.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddc03d96684f9226b8a787cdb71488417b53ab5ea8fdb1dac946cb9431cc8bff" +dependencies = [ + "bytes", + "http", + "http-body", + "hyper", + "pin-project-lite", + "tokio", + "tower-service", +] + [[package]] name = "id-arena" version = "2.3.0" @@ -183,18 +509,36 @@ version = "0.2.182" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6800badb6cb2082ffd7b6a67e6125bb39f18782f793520caee8cb8846be06112" +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + [[package]] name = "log" version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "matchit" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" + [[package]] name = "memchr" version = "2.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" +[[package]] +name = "mime" +version = "0.3.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" + [[package]] name = "mio" version = "1.1.1" @@ -206,12 +550,99 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "native-tls" +version = "0.2.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "465500e14ea162429d264d44189adc38b199b62b1c21eea9f69e4b73cb03bbf2" +dependencies = [ + "libc", + "log", + "openssl", + "openssl-probe", + "openssl-sys", + "schannel", + "security-framework", + "security-framework-sys", + "tempfile", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "openssl" +version = "0.10.81" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77823a27f0babb03091cb9ed9ef80af3b39dbc82f97e8fa530374b7dafd87a45" +dependencies = [ + "bitflags", + "cfg-if", + "foreign-types", + "libc", + "openssl-macros", + "openssl-sys", +] + +[[package]] +name = "openssl-macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a948666b637a0f465e8564c73e89d4dde00d72d4d473cc972f390fc3dcee7d9c" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.116", +] + +[[package]] +name = "openssl-probe" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" + +[[package]] +name = "openssl-sys" +version = "0.9.117" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b47e7e6bb2c38cd930d25a23b40fa52e068c10e85f3e03a7f5ba5aaca5713695" +dependencies = [ + "cc", + "libc", + "pkg-config", + "vcpkg", +] + +[[package]] +name = "percent-encoding" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" + [[package]] name = "pin-project-lite" version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3b3cff922bd51709b605d9ead9aa71031d81447142d828eb4a6eba76fe619f9b" +[[package]] +name = "pkg-config" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548" + +[[package]] +name = "ppv-lite86" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +dependencies = [ + "zerocopy", +] + [[package]] name = "prettyplease" version = "0.2.37" @@ -219,7 +650,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" dependencies = [ "proc-macro2", - "syn", + "syn 2.0.116", ] [[package]] @@ -248,13 +679,42 @@ checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" [[package]] name = "rand" -version = "0.10.0" +version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bc266eb313df6c5c09c1c7b1fbe2510961e5bcd3add930c1e31f7ed9da0feff8" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" +dependencies = [ + "rand_chacha", + "rand_core 0.9.5", +] + +[[package]] +name = "rand" +version = "0.10.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "65c9fb96cbc91e3478eaae79a69fcd3f1ae4ad052e471fe6732fff548984b4af" dependencies = [ "chacha20", - "getrandom", - "rand_core", + "getrandom 0.4.1", + "rand_core 0.10.0", +] + +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", ] [[package]] @@ -263,6 +723,104 @@ version = "0.10.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0c8d0fd677905edcbeedbf2edb6494d676f0e98d54d5cf9bda0b061cb8fb8aba" +[[package]] +name = "ring" +version = "0.17.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" +dependencies = [ + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", +] + +[[package]] +name = "rustix" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "891efababe418670775f199f0d233d84843c227a0949a883ce15b37c78d6629d" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + +[[package]] +name = "rustls" +version = "0.23.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d41d731c7d2f962d1ccc364cec258de3c0e93b38c2fb3ba97ac74513048d634" +dependencies = [ + "once_cell", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-pki-types" +version = "1.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2f4925028c7eb5d1fcdaf196971378ed9d2c1c4efc7dc5d011256f76c99c0a96" +dependencies = [ + "zeroize", +] + +[[package]] +name = "rustls-webpki" +version = "0.103.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f3c3cf1d8b1e7d4927e2d154c3fcb02979afb9939629c62cd9048d4f07b60ac2" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", +] + +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "schannel" +version = "0.1.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91c1b7e4904c873ef0710c1f407dde2e6287de2bebc1bbbf7d430bb7cbffd939" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "security-framework" +version = "3.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" +dependencies = [ + "bitflags", + "core-foundation", + "core-foundation-sys", + "libc", + "security-framework-sys", +] + +[[package]] +name = "security-framework-sys" +version = "2.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2691df843ecc5d231c0b14ece2acc3efb62c0a398c7e1d875f3983ce020e3" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "semver" version = "1.0.27" @@ -296,7 +854,7 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.116", ] [[package]] @@ -313,15 +871,38 @@ dependencies = [ ] [[package]] -name = "session-rs" -version = "0.1.3" +name = "serde_path_to_error" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457" dependencies = [ - "base64", - "rand", + "itoa", + "serde", + "serde_core", +] + +[[package]] +name = "serde_urlencoded" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3491c14715ca2294c4d6a88f15e84739788c1d030eed8c110436aafdaa2f3fd" +dependencies = [ + "form_urlencoded", + "itoa", + "ryu", + "serde", +] + +[[package]] +name = "session-rs" +version = "0.2.0" +dependencies = [ + "axum", + "futures-util", "serde", "serde_json", - "sha1", "tokio", + "tokio-tungstenite 0.30.0", ] [[package]] @@ -332,9 +913,38 @@ checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" dependencies = [ "cfg-if", "cpufeatures 0.2.17", - "digest", + "digest 0.10.7", ] +[[package]] +name = "sha1" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "aacc4cc499359472b4abe1bf11d0b12e688af9a805fa5e3016f9a386dc2d0214" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.3", +] + +[[package]] +name = "shlex" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" + +[[package]] +name = "slab" +version = "0.4.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c790de23124f9ab44544d7ac05d60440adc586479ce501c1d6d7da3cd8c9cf5" + +[[package]] +name = "smallvec" +version = "1.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba467056f1b547ed52077911161fc86985becbc60e8e1857c8a144dab0def891" + [[package]] name = "socket2" version = "0.6.1" @@ -345,6 +955,12 @@ dependencies = [ "windows-sys 0.60.2", ] +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + [[package]] name = "syn" version = "2.0.116" @@ -356,6 +972,56 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "syn" +version = "3.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8593e8e72159ed2257d083c7a454a85cbf854f37a0966d8d483aff8c8a3ebcee" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "sync_wrapper" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" + +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.1", + "once_cell", + "rustix", + "windows-sys 0.61.2", +] + +[[package]] +name = "thiserror" +version = "2.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09e52cb86a36cede5cb101bf8908837b3e4c6e5e59fe7fd85c23fb56200d189e" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe5197923287db20a58125f0bc85c062f7f2c892de97b18c356f9efb14b28524" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.6", +] + [[package]] name = "tokio" version = "1.49.0" @@ -379,7 +1045,140 @@ checksum = "af407857209536a95c8e56f8231ef2c2e2aff839b22e07a1ffcbc617e9db9fa5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.116", +] + +[[package]] +name = "tokio-native-tls" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbae76ab933c85776efabc971569dd6119c580d8f5d448769dec1764bf796ef2" +dependencies = [ + "native-tls", + "tokio", +] + +[[package]] +name = "tokio-rustls" +version = "0.26.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b0c85f2c3ef0b1cd58b36682f4b17aaa995f0e5db534d85692b4903abce21f67" +dependencies = [ + "rustls", + "tokio", +] + +[[package]] +name = "tokio-tungstenite" +version = "0.29.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f72a05e828585856dacd553fba484c242c46e391fb0e58917c942ee9202915c" +dependencies = [ + "futures-util", + "log", + "tokio", + "tungstenite 0.29.0", +] + +[[package]] +name = "tokio-tungstenite" +version = "0.30.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "17a073bfed563fa236697a068031408a93cd9522e08abf9933ead3e73411bd71" +dependencies = [ + "futures-util", + "log", + "native-tls", + "rustls", + "rustls-pki-types", + "tokio", + "tokio-native-tls", + "tokio-rustls", + "tungstenite 0.30.0", + "webpki-roots 0.26.11", +] + +[[package]] +name = "tower" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" +dependencies = [ + "futures-core", + "futures-util", + "pin-project-lite", + "sync_wrapper", + "tokio", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tower-layer" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e" + +[[package]] +name = "tower-service" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3" + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "log", + "pin-project-lite", + "tracing-core", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", +] + +[[package]] +name = "tungstenite" +version = "0.29.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c01152af293afb9c7c2a57e4b559c5620b421f6d133261c60dd2d0cdb38e6b8" +dependencies = [ + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand 0.9.5", + "sha1 0.10.6", + "thiserror", +] + +[[package]] +name = "tungstenite" +version = "0.30.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e48ac77174b19c110a50ab2128b24215ac9cb40e0e12e093fb602d175c569d22" +dependencies = [ + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "native-tls", + "rand 0.10.3", + "rustls", + "rustls-pki-types", + "sha1 0.11.0", + "thiserror", ] [[package]] @@ -400,6 +1199,18 @@ version = "0.2.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + +[[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + [[package]] name = "version_check" version = "0.9.5" @@ -464,19 +1275,46 @@ dependencies = [ "semver", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.9", +] + +[[package]] +name = "webpki-roots" +version = "1.0.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7dcd9d09a39985f5344844e66b0c530a33843579125f23e21e9f0f220850f22a" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "windows-link" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets 0.52.6", +] + [[package]] name = "windows-sys" version = "0.60.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb" dependencies = [ - "windows-targets", + "windows-targets 0.53.5", ] [[package]] @@ -488,6 +1326,22 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm 0.52.6", + "windows_aarch64_msvc 0.52.6", + "windows_i686_gnu 0.52.6", + "windows_i686_gnullvm 0.52.6", + "windows_i686_msvc 0.52.6", + "windows_x86_64_gnu 0.52.6", + "windows_x86_64_gnullvm 0.52.6", + "windows_x86_64_msvc 0.52.6", +] + [[package]] name = "windows-targets" version = "0.53.5" @@ -495,58 +1349,106 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3" dependencies = [ "windows-link", - "windows_aarch64_gnullvm", - "windows_aarch64_msvc", - "windows_i686_gnu", - "windows_i686_gnullvm", - "windows_i686_msvc", - "windows_x86_64_gnu", - "windows_x86_64_gnullvm", - "windows_x86_64_msvc", + "windows_aarch64_gnullvm 0.53.1", + "windows_aarch64_msvc 0.53.1", + "windows_i686_gnu 0.53.1", + "windows_i686_gnullvm 0.53.1", + "windows_i686_msvc 0.53.1", + "windows_x86_64_gnu 0.53.1", + "windows_x86_64_gnullvm 0.53.1", + "windows_x86_64_msvc 0.53.1", ] +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + [[package]] name = "windows_aarch64_gnullvm" version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53" +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + [[package]] name = "windows_aarch64_msvc" version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006" +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + [[package]] name = "windows_i686_gnu" version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3" +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + [[package]] name = "windows_i686_gnullvm" version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c" +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + [[package]] name = "windows_i686_msvc" version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2" +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + [[package]] name = "windows_x86_64_gnu" version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499" +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + [[package]] name = "windows_x86_64_gnullvm" version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1" +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + [[package]] name = "windows_x86_64_msvc" version = "0.53.1" @@ -583,7 +1485,7 @@ dependencies = [ "heck", "indexmap", "prettyplease", - "syn", + "syn 2.0.116", "wasm-metadata", "wit-bindgen-core", "wit-component", @@ -599,7 +1501,7 @@ dependencies = [ "prettyplease", "proc-macro2", "quote", - "syn", + "syn 2.0.116", "wit-bindgen-core", "wit-bindgen-rust", ] @@ -641,6 +1543,32 @@ dependencies = [ "wasmparser", ] +[[package]] +name = "zerocopy" +version = "0.8.58" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c17e8fafad82b542ff3717217ecdc736231b59e387768c9630123b4ce4d2db44" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.58" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "595f56e044df4f46a0c9a626f65c3d99eb8488f7e8a8baa12dd76326d9710bf2" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.116", +] + +[[package]] +name = "zeroize" +version = "1.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" + [[package]] name = "zmij" version = "1.0.21" diff --git a/Cargo.toml b/Cargo.toml index c62b502..4c6a873 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,21 +1,38 @@ [package] name = "session-rs" -version = "0.1.3" +version = "0.2.0" edition = "2024" -description = "A lightweight async WebSocket protocol" +description = "A lightweight async request/response and notification protocol over WebSockets" license = "Apache-2.0" +repository = "https://git.selimaj.dev/selimaj-dev/session-rs" + +[features] +default = ["server", "client"] +# `SessionServer`: a standalone tokio-tungstenite WebSocket server. +server = ["dep:tokio-tungstenite", "tokio/net"] +# `Session::connect`: a tokio-tungstenite WebSocket client (ws:// only without a TLS feature). +client = ["dep:tokio-tungstenite", "tokio-tungstenite/connect"] +# `wss://` support for the client. +rustls = ["client", "tokio-tungstenite/rustls-tls-webpki-roots"] +native-tls = ["client", "tokio-tungstenite/native-tls"] +# `Session::from_axum`: run a session on an axum WebSocket upgrade. +axum = ["dep:axum"] [dependencies] -base64 = "0.22.1" -rand = "0.10.0" -serde = { version = "1.0.228", features = ["serde_derive"] } +futures-util = { version = "0.3.34", default-features = false, features = ["sink", "std"] } +serde = { version = "1.0.228", features = ["derive"] } serde_json = "1.0.149" -sha1 = "0.10.6" -tokio = { version = "1.49.0", features = [ - "io-util", - "macros", - "net", - "rt", - "sync", - "time", -] } +tokio = { version = "1.49.0", features = ["macros", "rt", "sync", "time"] } +tokio-tungstenite = { version = "0.30.0", default-features = false, features = ["handshake"], optional = true } +axum = { version = "0.8.9", default-features = false, features = ["ws"], optional = true } + +[dev-dependencies] +tokio = { version = "1.49.0", features = ["macros", "rt-multi-thread", "net", "time", "io-util"] } +axum = { version = "0.8.9", features = ["ws"] } + +[[example]] +name = "axum" +required-features = ["axum"] + +[package.metadata.docs.rs] +all-features = true diff --git a/README.md b/README.md index 7a7535a..88596c5 100644 --- a/README.md +++ b/README.md @@ -6,20 +6,12 @@ ## Introduction -This library provides **type-safe WebSocket communication** with a request-response and notification system built on top of a flexible protocol. -It ensures compile-time guarantees for message structure, reduces runtime errors, and simplifies building Rust client/server applications. +`session-rs` is a small request/response + notification protocol that runs over WebSockets, with typed methods on both ends. -- **Dynamic Methods**: Each message includes a method enum for type safety. -- **Typed Requests & Responses**: Automatic serialization and deserialization. -- **Optional Notifications**: Send asynchronous notifications across sessions. - -## Features - -- Fully typed WebSocket sessions -- Type-safe request/response mechanism -- Optional typed notifications (Todo) -- Lightweight, minimal runtime overhead -- Async-first with Tokio support +- **Typed methods**: requests, responses and errors are (de)serialized for you. +- **Both directions**: either peer can send requests and notifications. +- **Transport-agnostic**: the protocol runs over any `Sink`/`Stream` of frames. Adapters ship for [tokio-tungstenite](https://docs.rs/tokio-tungstenite) and [axum](https://docs.rs/axum). +- **Bounded**: the built-in server limits message and frame sizes (1 MiB by default). ## Installation @@ -27,9 +19,16 @@ It ensures compile-time guarantees for message structure, reduces runtime errors cargo add session-rs ``` ---- +| Feature | Default | Enables | +| --- | --- | --- | +| `server` | yes | `SessionServer`, a standalone WebSocket server | +| `client` | yes | `Session::connect` for `ws://` URLs | +| `rustls` / `native-tls` | no | `wss://` URLs in `Session::connect` | +| `axum` | no | `Session::from_axum` for axum WebSocket upgrades | -### **Basic Example (client)** +## Usage + +Define a method once and share it between both peers: ```rust #[derive(Debug, Serialize, Deserialize)] @@ -41,44 +40,74 @@ impl Method for Data { type Response = String; type Error = String; } - -let session = Session::connect("127.0.0.1:8080", "/").await?; - -session.start_receiver(); - -session - .request::("Hello from client".to_string()) - .await?; ``` -### **Basic Example (server)** +### Server ```rust -#[derive(Debug, Serialize, Deserialize)] -struct Data; - -impl Method for Data { - const NAME: &'static str = "data"; - type Request = String; - type Response = String; - type Error = String; -} - let server = SessionServer::bind("127.0.0.1:8080").await?; server .session_loop(async |session, addr| { - // This will run on every new client + // Runs for every new client. Register handlers here; the session + // starts reading once this returns, so no message arrives early. + session + .on_request::(async |_id, req| Ok(format!("echo: {req}"))) + .await; Ok(()) - }).await; + }) + .await?; ``` +Use `SessionServer::with_config(ServerConfig { .. })` to change the size limits or the handshake timeout. + +### Client + +```rust +let session = Session::connect("ws://127.0.0.1:8080").await?; +session.start_receiver(); + +let reply = session.request::("Hello".to_string()).await?; // Ok("echo: Hello") +``` + +### axum + +With the `axum` feature, a session can share a router with ordinary HTTP routes, such as a health check: + +```rust +async fn ws(upgrade: WebSocketUpgrade) -> Response { + upgrade.max_message_size(1 << 20).on_upgrade(async |socket| { + let session = Session::from_axum(socket); + session.on_request::(async |_, req| Ok(req)).await; + session.start_receiver(); + }) +} + +let app = Router::new() + .route("/", get(ws)) + .route("/health", get(async || "ok")); +``` + +### Other transports + +`Session::from_transport(sink, stream)` accepts any `Sink` and `Stream>`, and `Session::from_tungstenite` wraps an existing tokio-tungstenite stream (e.g. one accepted over TLS). The transport must answer pings itself. + +### Semantics + +- Incoming requests and notifications are handled one at a time, in arrival order. Responses to your own requests are delivered independently, so a handler can `request` from its peer. +- A request for an unknown method, or with data that doesn't deserialize, gets an error response instead of no reply. +- A handler that panics fails only its own request. +- `request` fails with `Error::ConnectionClosed` if the session closes first; `request_timeout` adds a deadline. +- `on_close` runs exactly once, whichever side closes. `start_ping(interval, timeout)` closes peers that stop answering pings. + ## Protocol +Every message is a JSON text frame with a `type` tag. + #### Request -The request `id` is separated from the peer, and will increment only on it's requests. +The `id` is chosen by the sender and increments per peer. ```json { "type": "request", "id": 1, "method": "data", "data": "Hello from client" } @@ -86,16 +115,22 @@ The request `id` is separated from the peer, and will increment only on it's req #### Response -The response `id` **must** remain the same as the request. +A response **must** carry the id of the request it answers. ```json { "type": "response", "id": 1, "result": "Hello from server" } ``` -#### Notifications - -A notification is a method that doesn't need validation or output, it simply notifies a peer for a specific information +#### Error response ```json -{ "type": "notification", "result": "Hello from server" } +{ "type": "errorresponse", "id": 1, "error": "Invalid data" } +``` + +#### Notification + +A notification is fire-and-forget and gets no response. + +```json +{ "type": "notification", "method": "data", "data": "Hello from server" } ``` diff --git a/examples/axum.rs b/examples/axum.rs new file mode 100644 index 0000000..750eb19 --- /dev/null +++ b/examples/axum.rs @@ -0,0 +1,41 @@ +//! The same server as `examples/server.rs`, mounted in an axum router next to +//! ordinary HTTP routes. + +use axum::{Router, extract::WebSocketUpgrade, response::Response, routing::get}; +use serde::{Deserialize, Serialize}; +use session_rs::{Method, Session}; + +#[derive(Debug, Serialize, Deserialize)] +struct Data; + +impl Method for Data { + const NAME: &'static str = "data"; + type Request = String; + type Response = String; + type Error = String; +} + +async fn ws(upgrade: WebSocketUpgrade) -> Response { + upgrade.max_message_size(1 << 20).on_upgrade(async |socket| { + let session = Session::from_axum(socket); + + session + .on_request::(async |_, req| { + println!("Msg from client: {req}"); + Ok(format!("Hello from axum, you said {req:?}")) + }) + .await; + + session.start_receiver(); + }) +} + +#[tokio::main] +async fn main() -> std::io::Result<()> { + let app = Router::new() + .route("/", get(ws)) + .route("/health", get(async || "ok")); + + let listener = tokio::net::TcpListener::bind("127.0.0.1:8080").await?; + axum::serve(listener, app).await +} diff --git a/examples/client.rs b/examples/client.rs index 14a1484..b0dd9b4 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,5 +1,5 @@ use serde::{Deserialize, Serialize}; -use session_rs::{Method, session::Session}; +use session_rs::{Method, Session}; #[derive(Debug, Serialize, Deserialize)] struct Data; @@ -13,7 +13,7 @@ impl Method for Data { #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { - let session = Session::connect("127.0.0.1:8080", "/").await?; + let session = Session::connect("ws://127.0.0.1:8080").await?; session.start_receiver(); @@ -31,8 +31,6 @@ async fn main() -> session_rs::Result<()> { .await? ); - tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; - session.close().await?; Ok(()) } diff --git a/src/axum.rs b/src/axum.rs new file mode 100644 index 0000000..efa7085 --- /dev/null +++ b/src/axum.rs @@ -0,0 +1,39 @@ +use axum::extract::ws::{Message, WebSocket}; +use futures_util::{SinkExt, StreamExt}; + +use crate::{Frame, Session}; + +fn to_ws(frame: Frame) -> Message { + match frame { + Frame::Text(text) => Message::Text(text.into()), + Frame::Binary(data) => Message::Binary(data.into()), + Frame::Ping(data) => Message::Ping(data.into()), + Frame::Pong(data) => Message::Pong(data.into()), + Frame::Close => Message::Close(None), + } +} + +fn from_ws(msg: Message) -> Frame { + match msg { + Message::Text(text) => Frame::Text(text.as_str().to_owned()), + Message::Binary(data) => Frame::Binary(data.to_vec()), + Message::Ping(data) => Frame::Ping(data.to_vec()), + Message::Pong(data) => Frame::Pong(data.to_vec()), + Message::Close(_) => Frame::Close, + } +} + +impl Session { + /// Run a session over an axum WebSocket upgrade. + /// + /// Message size limits are set on the upgrade, e.g. + /// `ws.max_message_size(1 << 20).on_upgrade(...)`. + pub fn from_axum(socket: WebSocket) -> Self { + let (sink, stream) = socket.split(); + + Session::from_transport( + sink.with(|frame| async move { Ok::<_, axum::Error>(to_ws(frame)) }), + stream.map(|msg| msg.map(from_ws)), + ) + } +} diff --git a/src/client.rs b/src/client.rs new file mode 100644 index 0000000..69a51a3 --- /dev/null +++ b/src/client.rs @@ -0,0 +1,13 @@ +use crate::{Error, Session}; + +impl Session { + /// Connect to a `ws://` (or, with the `rustls`/`native-tls` feature, + /// `wss://`) URL. Call [`Session::start_receiver`] after registering handlers. + pub async fn connect(url: &str) -> crate::Result { + let (ws, _) = tokio_tungstenite::connect_async(url) + .await + .map_err(|e| Error::Transport(Box::new(e)))?; + + Ok(Self::from_tungstenite(ws)) + } +} diff --git a/src/lib.rs b/src/lib.rs index 16b1f2b..3049929 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,16 +1,35 @@ -use std::{pin::Pin, sync::Arc}; +//! A small request/response + notification protocol over WebSockets. +//! +//! The protocol layer ([`Session`]) is transport-agnostic: it runs over any +//! [`Sink`](futures_util::Sink)/[`Stream`](futures_util::Stream) pair of +//! [`Frame`]s. Adapters are provided for tokio-tungstenite (`server`/`client` +//! features) and axum (`axum` feature). + +use std::pin::Pin; use serde::{Deserialize, Serialize}; +#[cfg(feature = "axum")] +mod axum; +#[cfg(feature = "client")] +mod client; +#[cfg(feature = "server")] pub mod server; pub mod session; -pub mod ws; +pub mod transport; +#[cfg(any(feature = "server", feature = "client"))] +mod tungstenite; + +pub use session::{Message, Session}; +pub use transport::Frame; pub type Result = std::result::Result; -pub type BoxFuture<'a, T = Option<(bool, serde_json::Value)>> = - Pin + Send + 'a>>; -pub type MethodHandler = Arc BoxFuture<'static> + Send + Sync>; +pub type BoxFuture<'a, T> = Pin + Send + 'a>>; +pub type BoxError = Box; +/// A named, typed RPC. Implement this on a marker type and use it with +/// [`Session::request`], [`Session::on_request`], [`Session::notify`] and +/// [`Session::on_notification`]. pub trait Method { const NAME: &'static str; type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync; @@ -18,6 +37,7 @@ pub trait Method { type Error: Serialize + for<'de> Deserialize<'de>; } +/// Untyped method used internally for raw JSON values. pub struct GenericMethod; impl Method for GenericMethod { @@ -29,18 +49,30 @@ impl Method for GenericMethod { #[derive(Debug)] pub enum Error { - WebSocket(ws::Error), + /// The underlying transport (WebSocket) failed. + Transport(BoxError), Json(serde_json::Error), Io(std::io::Error), - RecvError(tokio::sync::broadcast::error::RecvError), + /// The session is closed; no more messages can be sent or received. + ConnectionClosed, + /// A request or handshake did not complete in time. + Timeout, } -impl From for Error { - fn from(value: ws::Error) -> Self { - Self::WebSocket(value) +impl std::fmt::Display for Error { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Error::Transport(e) => write!(f, "transport error: {e}"), + Error::Json(e) => write!(f, "json error: {e}"), + Error::Io(e) => write!(f, "io error: {e}"), + Error::ConnectionClosed => write!(f, "connection closed"), + Error::Timeout => write!(f, "timed out"), + } } } +impl std::error::Error for Error {} + impl From for Error { fn from(value: std::io::Error) -> Self { Self::Io(value) @@ -53,8 +85,8 @@ impl From for Error { } } -impl From for Error { - fn from(value: tokio::sync::broadcast::error::RecvError) -> Self { - Self::RecvError(value) +impl From for Error { + fn from(_: tokio::time::error::Elapsed) -> Self { + Self::Timeout } } diff --git a/src/server.rs b/src/server.rs index 5ac89b8..c563cee 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,62 +1,124 @@ -use std::{net::SocketAddr, sync::Arc}; +use std::{net::SocketAddr, sync::Arc, time::Duration}; -use tokio::{net::TcpListener, time::timeout}; +use tokio::net::{TcpListener, TcpStream, ToSocketAddrs}; +use tokio_tungstenite::tungstenite::protocol::WebSocketConfig; -use crate::{session::Session, ws::WebSocket}; +use crate::{Error, session::Session}; +/// Limits applied to every accepted connection. +#[derive(Debug, Clone)] +pub struct ServerConfig { + /// Largest message (after reassembling fragments) a peer may send. + pub max_message_size: usize, + /// Largest single frame a peer may send. + pub max_frame_size: usize, + /// How long a client gets to complete the WebSocket handshake. + pub handshake_timeout: Duration, +} + +impl Default for ServerConfig { + fn default() -> Self { + Self { + max_message_size: 1 << 20, + max_frame_size: 1 << 20, + handshake_timeout: Duration::from_secs(5), + } + } +} + +/// A standalone WebSocket server that hands each connection to a callback as a +/// [`Session`]. pub struct SessionServer { listener: TcpListener, + config: ServerConfig, } impl SessionServer { - pub async fn bind(addr: &str) -> crate::Result { - Ok(Self { - listener: TcpListener::bind(addr).await?, - }) + pub async fn bind(addr: impl ToSocketAddrs) -> crate::Result { + Ok(Self::from_listener(TcpListener::bind(addr).await?)) } + pub fn from_listener(listener: TcpListener) -> Self { + Self { + listener, + config: ServerConfig::default(), + } + } + + pub fn with_config(mut self, config: ServerConfig) -> Self { + self.config = config; + self + } + + pub fn local_addr(&self) -> std::io::Result { + self.listener.local_addr() + } + + /// Accept one connection. The receiver is not started: register handlers, + /// then call [`Session::start_receiver`]. pub async fn accept(&self) -> crate::Result<(Session, SocketAddr)> { let (stream, addr) = self.listener.accept().await?; - - let ws = WebSocket::handshake(stream).await?; - - Ok((Session::from_ws(ws), addr)) + Ok((handshake(stream, &self.config).await?, addr)) } + /// Accept connections forever, running `on_conn` for each one. + /// + /// `on_conn` should register handlers and return; the receiver starts once + /// it does, so no message is processed before its handler exists. If it + /// returns an error the connection is closed. pub async fn session_loop(&self, on_conn: F) -> crate::Result<()> where F: Fn(Session, SocketAddr) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { - let conn_handler = Arc::new(on_conn); + let on_conn = Arc::new(on_conn); loop { - let (stream, addr) = self.listener.accept().await?; - let conn_handler = conn_handler.clone(); + let (stream, addr) = match self.listener.accept().await { + Ok(conn) => conn, + Err(e) => { + // e.g. out of file descriptors: back off instead of exiting. + eprintln!("Accept failed: {e}"); + tokio::time::sleep(Duration::from_millis(100)).await; + continue; + } + }; + + let on_conn = on_conn.clone(); + let config = self.config.clone(); tokio::spawn(async move { - match timeout( - tokio::time::Duration::from_secs(5), - WebSocket::handshake(stream), - ) - .await - { - Ok(Ok(ws)) => { - let session = Session::from_ws(ws); - session.start_receiver(); + let session = match handshake(stream, &config).await { + Ok(session) => session, + Err(e) => { + eprintln!("Handshake failed from {addr}: {e}"); + return; + } + }; - if let Err(e) = conn_handler(session, addr).await { - eprintln!("Connection error: {:?}", e); - } - } - Ok(Err(e)) => { - eprintln!("Handshake failed from {}: {:?}", addr, e); - } - Err(_) => { - eprintln!("Handshake failed from {}: Handshake Timeout", addr); - } + if let Err(e) = on_conn(session.clone(), addr).await { + eprintln!("Connection error from {addr}: {e}"); + let _ = session.close().await; + return; } + + session.start_receiver(); }); } } } + +async fn handshake(stream: TcpStream, config: &ServerConfig) -> crate::Result { + let ws_config = WebSocketConfig::default() + .max_message_size(Some(config.max_message_size)) + .max_frame_size(Some(config.max_frame_size)); + + let ws = tokio::time::timeout( + config.handshake_timeout, + tokio_tungstenite::accept_async_with_config(stream, Some(ws_config)), + ) + .await? + .map_err(|e| Error::Transport(Box::new(e)))?; + + Ok(Session::from_tungstenite(ws)) +} diff --git a/src/session.rs b/src/session.rs index 2443a85..e351bb1 100644 --- a/src/session.rs +++ b/src/session.rs @@ -1,14 +1,25 @@ +use std::collections::HashMap; use std::hash::Hash; -use std::{collections::HashMap, sync::Arc}; +use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; +use futures_util::{Sink, SinkExt, Stream, StreamExt}; use serde::{Deserialize, Serialize}; -use tokio::sync::Mutex; -use tokio::sync::broadcast; -use tokio::time::timeout; +use serde_json::Value; +use tokio::sync::{mpsc, oneshot, watch}; -use crate::BoxFuture; -use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; +use crate::transport::{self, BoxSink, BoxStream, Frame}; +use crate::{BoxError, BoxFuture, Error, GenericMethod, Method}; +/// A protocol message, serialized as JSON in a text frame. +/// +/// ```json +/// { "type": "request", "id": 1, "method": "data", "data": "..." } +/// { "type": "response", "id": 1, "result": "..." } +/// { "type": "errorresponse", "id": 1, "error": "..." } +/// { "type": "notification", "method": "data", "data": "..." } +/// ``` #[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "lowercase", tag = "type")] pub enum Message { @@ -31,126 +42,246 @@ pub enum Message { }, } +type RequestHandler = Arc BoxFuture<'static, Result> + Send + Sync>; +type NotificationHandler = Arc BoxFuture<'static, ()> + Send + Sync>; +type CloseHandler = Arc BoxFuture<'static, Result<(), String>> + Send + Sync>; + +/// Frames queued for the writer task before `send` applies backpressure. +const OUTGOING_BUFFER: usize = 256; +/// Requests/notifications queued for handlers before the reader applies backpressure. +const INCOMING_BUFFER: usize = 64; + +static NEXT_SESSION_ID: AtomicU64 = AtomicU64::new(1); + +/// One peer connection speaking the session protocol. +/// +/// Cloning is cheap and every clone refers to the same connection; equality and +/// hashing are by connection, so sessions can be stored in sets and maps. +/// +/// Incoming requests and notifications are handled one at a time, in arrival +/// order. Responses to this side's own requests are delivered independently, +/// so a handler may itself `request` from the peer. +#[derive(Clone)] pub struct Session { - pub ws: WebSocket, - id: Arc>, - methods: Arc>>, - on_close_fn: - Arc BoxFuture<'static, Result<(), String>> + Send + Sync>>>>, - tx: broadcast::Sender<(u32, bool, serde_json::Value)>, - pong_tx: broadcast::Sender<()>, + inner: Arc, +} + +struct Inner { + id: u64, + next_request_id: AtomicU32, + outgoing: mpsc::Sender, + /// The read half, until `start_receiver` takes it. + stream: Mutex>, + requests: Mutex>, + notifications: Mutex>, + pending: Mutex>>>, + on_close: Mutex>, + closed: AtomicBool, + closed_tx: watch::Sender, + /// Bumped on every pong, for `start_ping`. + pongs: watch::Sender, +} + +enum Incoming { + Request { id: u32, method: String, data: Value }, + Notification { method: String, data: Value }, } impl Session { - pub fn clone(&self) -> Self { - Self { - ws: self.ws.clone(), - id: self.id.clone(), - methods: self.methods.clone(), - on_close_fn: self.on_close_fn.clone(), - tx: self.tx.clone(), - pong_tx: self.pong_tx.clone(), - } + /// Build a session over any frame sink/stream pair. + /// + /// The transport is expected to answer pings itself (tungstenite and axum + /// both do). Nothing is read until [`Session::start_receiver`] is called, + /// so handlers can be registered first. + pub fn from_transport(sink: Si, stream: St) -> Self + where + Si: Sink + Send + 'static, + SiE: Into + 'static, + St: Stream> + Send + 'static, + StE: Into + 'static, + { + let (sink, stream) = transport::boxed(sink, stream); + let (outgoing, outgoing_rx) = mpsc::channel(OUTGOING_BUFFER); + let (closed_tx, _) = watch::channel(false); + + let session = Self { + inner: Arc::new(Inner { + id: NEXT_SESSION_ID.fetch_add(1, Ordering::Relaxed), + next_request_id: AtomicU32::new(0), + outgoing, + stream: Mutex::new(Some(stream)), + requests: Mutex::new(HashMap::new()), + notifications: Mutex::new(HashMap::new()), + pending: Mutex::new(HashMap::new()), + on_close: Mutex::new(None), + closed: AtomicBool::new(false), + closed_tx, + pongs: watch::channel(0).0, + }), + }; + + tokio::spawn(write_loop( + outgoing_rx, + sink, + session.inner.closed_tx.subscribe(), + Arc::downgrade(&session.inner), + )); + + session } -} -impl Session { - pub fn from_ws(ws: WebSocket) -> Self { - let (tx, _) = broadcast::channel(8192); - let (pong_tx, _) = broadcast::channel(16); + /// A process-unique id for this connection. + pub fn id(&self) -> u64 { + self.inner.id + } - Self { - ws, - id: Arc::new(Mutex::new(0)), - methods: Arc::new(Mutex::new(HashMap::new())), - on_close_fn: Arc::new(Mutex::new(None)), - tx, - pong_tx, - } + pub fn is_closed(&self) -> bool { + self.inner.closed.load(Ordering::SeqCst) } - pub async fn connect(addr: &str, path: &str) -> crate::Result { - Ok(Self::from_ws(WebSocket::connect(addr, path).await?)) + /// Resolves once the session is closed, from either side. + pub async fn closed(&self) { + wait_closed(&mut self.inner.closed_tx.subscribe()).await; } } impl Session { + /// Start reading from the peer. Calling it again has no effect. pub fn start_receiver(&self) { + let Some(stream) = self.inner.stream.lock().unwrap().take() else { + return; + }; + + let (dispatch_tx, dispatch_rx) = mpsc::channel(INCOMING_BUFFER); + tokio::spawn(self.clone().dispatch_loop(dispatch_rx)); + tokio::spawn(self.clone().read_loop(stream, dispatch_tx)); + } + + /// Ping the peer every `interval` and close the session if no pong arrives + /// within `timeout`. + pub fn start_ping(&self, interval: Duration, timeout: Duration) { let s = self.clone(); + tokio::spawn(async move { + let mut closed = s.inner.closed_tx.subscribe(); + let mut pongs = s.inner.pongs.subscribe(); + loop { - match s.ws.read().await { - Ok(crate::ws::Frame::Text(text)) => { - let Ok(msg) = serde_json::from_str::>(&text) else { - continue; - }; + tokio::select! { + _ = tokio::time::sleep(interval) => {} + _ = wait_closed(&mut closed) => return, + } - match msg { - Message::Request { id, method, data } => { - let handler = { - let methods = s.methods.lock().await; - methods.get(&method).cloned() - }; + pongs.mark_unchanged(); - if let Some(m) = handler { - if let Some((err, res)) = (m)(id, data).await { - if err { - s.respond_error(id, res) - .await - .expect("Failed to respond"); - } else { - s.respond(id, res).await.expect("Failed to respond"); - } - } - } - } - Message::Response { id, result } => { - s.tx.send((id, false, result)).unwrap(); - } - Message::ErrorResponse { id, error } => { - s.tx.send((id, true, error)).unwrap(); - } - _ => {} - } - } - Ok(crate::ws::Frame::Pong) => { - let _ = s.pong_tx.send(()); - } - Ok(_) => {} - Err(_) => { - s.trigger_close().await; + if s.inner.outgoing.send(Frame::Ping(Vec::new())).await.is_err() { + break; + } + + match tokio::time::timeout(timeout, pongs.changed()).await { + Ok(Ok(())) => {} + _ => break, + } + } + + s.shutdown().await; + }); + } + + async fn read_loop(self, mut stream: BoxStream, dispatch: mpsc::Sender) { + let mut closed = self.inner.closed_tx.subscribe(); + + loop { + let frame = tokio::select! { + frame = stream.next() => frame, + _ = wait_closed(&mut closed) => break, + }; + + match frame { + Some(Ok(Frame::Text(text))) => { + if !self.handle_text(&text, &dispatch).await { break; } } - } - }); - } - pub fn start_ping(&self, interval: tokio::time::Duration, timeout_dur: tokio::time::Duration) { - let s = self.clone(); - - tokio::spawn(async move { - let mut pong_rx = s.pong_tx.subscribe(); - - loop { - tokio::time::sleep(interval).await; - - if s.ws.send_ping().await.is_err() { - s.trigger_close().await; - break; - } - - let result = timeout(timeout_dur, pong_rx.recv()).await; - - if result.is_err() { - // timeout expired - let _ = s.close().await; - s.trigger_close().await; - break; + Some(Ok(Frame::Pong(_))) => { + self.inner.pongs.send_modify(|n| *n = n.wrapping_add(1)); } + Some(Ok(Frame::Ping(_) | Frame::Binary(_))) => {} + Some(Ok(Frame::Close)) | Some(Err(_)) | None => break, } - }); + } + + self.shutdown().await; } + /// Returns false once the dispatcher is gone. + async fn handle_text(&self, text: &str, dispatch: &mpsc::Sender) -> bool { + let Ok(msg) = serde_json::from_str::>(text) else { + return true; + }; + + let incoming = match msg { + Message::Request { id, method, data } => Incoming::Request { id, method, data }, + Message::Notification { method, data } => Incoming::Notification { method, data }, + Message::Response { id, result } => { + self.complete(id, Ok(result)); + return true; + } + Message::ErrorResponse { id, error } => { + self.complete(id, Err(error)); + return true; + } + }; + + dispatch.send(incoming).await.is_ok() + } + + fn complete(&self, id: u32, result: Result) { + if let Some(tx) = self.inner.pending.lock().unwrap().remove(&id) { + let _ = tx.send(result); + } + } + + async fn dispatch_loop(self, mut rx: mpsc::Receiver) { + while let Some(incoming) = rx.recv().await { + match incoming { + Incoming::Request { id, method, data } => { + let handler = self.inner.requests.lock().unwrap().get(&method).cloned(); + + let result = match handler { + // Spawned so a panicking handler fails this request, not the session. + Some(handler) => tokio::spawn(handler(id, data)) + .await + .unwrap_or_else(|_| Err(Value::from("Handler panicked"))), + None => Err(Value::from(format!("Unknown method: {method}"))), + }; + + let reply = match result { + Ok(v) => self.respond(id, v).await, + Err(e) => self.respond_error(id, e).await, + }; + + if reply.is_err() { + break; + } + } + Incoming::Notification { method, data } => { + let handler = self.inner.notifications.lock().unwrap().get(&method).cloned(); + + if let Some(handler) = handler { + let _ = tokio::spawn(handler(data)).await; + } + } + } + } + } +} + +impl Session { + /// Handle requests for `M`. Replaces any previous handler for the method. + /// + /// Requests whose data does not deserialize into `M::Request` are answered + /// with an error response. pub async fn on_request< M: Method, Fut: Future> + Send + 'static, @@ -158,57 +289,89 @@ impl Session { &self, handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static, ) { - let handler = Arc::new(handler); + let handler: RequestHandler = Arc::new(move |id, value| { + let fut = serde_json::from_value::(value).map(|req| handler(id, req)); - self.methods.lock().await.insert( - M::NAME.to_string(), - Arc::new(move |id, value| { - let handler = Arc::clone(&handler); + Box::pin(async move { + match fut { + Err(e) => Err(Value::from(format!("Invalid request data: {e}"))), + Ok(fut) => match fut.await { + Ok(res) => serde_json::to_value(res) + .map_err(|e| Value::from(format!("Failed to serialize response: {e}"))), + Err(err) => Err(serde_json::to_value(err).unwrap_or_else(|e| { + Value::from(format!("Failed to serialize error: {e}")) + })), + }, + } + }) + }); - Box::pin(async move { - Some( - match handler(id, serde_json::from_value(value).ok()?).await { - Ok(v) => (false, serde_json::to_value(v).ok()?), - Err(v) => (true, serde_json::to_value(v).ok()?), - }, - ) - }) - }), - ); + self.inner + .requests + .lock() + .unwrap() + .insert(M::NAME.to_string(), handler); } + /// Handle notifications for `M`. Notifications with invalid data are dropped. + pub async fn on_notification + Send + 'static>( + &self, + handler: impl Fn(M::Request) -> Fut + Send + Sync + 'static, + ) { + let handler: NotificationHandler = Arc::new(move |value| { + let fut = serde_json::from_value::(value).map(&handler); + + Box::pin(async move { + if let Ok(fut) = fut { + fut.await; + } + }) + }); + + self.inner + .notifications + .lock() + .unwrap() + .insert(M::NAME.to_string(), handler); + } + + /// Run `handler` once when the session closes, from either side. pub async fn on_close(&self, handler: impl Fn() -> Fut + Send + Sync + 'static) where Fut: Future> + Send + 'static, { - let handler = Arc::new(handler); - - *self.on_close_fn.lock().await = Some(Box::new(move || { - let handler = handler.clone(); - Box::pin(async move { handler().await }) - })); + *self.inner.on_close.lock().unwrap() = Some(Arc::new(move || Box::pin(handler()))); } } impl Session { - pub async fn send(&self, data: &Message) -> crate::Result<()> { - self.ws - .send_text_payload(&serde_json::to_vec(&data)?) - .await?; - Ok(()) - } + pub async fn send(&self, msg: &Message) -> crate::Result<()> { + if self.is_closed() { + return Err(Error::ConnectionClosed); + } - pub async fn use_id(&self) -> u32 { - let mut id = self.id.lock().await; - *id += 1; - *id + let text = serde_json::to_string(msg)?; + + self.inner + .outgoing + .send(Frame::Text(text)) + .await + .map_err(|_| Error::ConnectionClosed) } + /// Send a request and wait for the peer's response. + /// + /// Fails with [`Error::ConnectionClosed`] if the session closes first. pub async fn request( &self, req: M::Request, ) -> crate::Result> { - let id = self.use_id().await; + let id = self.inner.next_request_id.fetch_add(1, Ordering::Relaxed).wrapping_add(1); + let (tx, rx) = oneshot::channel(); + + // Registered before sending so a fast response can't be missed. + self.inner.pending.lock().unwrap().insert(id, tx); + let _guard = PendingGuard { inner: &self.inner, id }; self.send::(&Message::Request { id, @@ -217,30 +380,27 @@ impl Session { }) .await?; - let mut rx = self.tx.subscribe(); - - loop { - let r = rx.recv().await?; - - if r.0 == id { - break Ok(if r.1 { - Err(serde_json::from_value(r.2)?) - } else { - Ok(serde_json::from_value(r.2)?) - }); - } - } - } - - pub async fn respond(&self, to: u32, val: serde_json::Value) -> crate::Result<()> { - self.send::(&Message::Response { - id: to, - result: val, + Ok(match rx.await.map_err(|_| Error::ConnectionClosed)? { + Ok(v) => Ok(serde_json::from_value(v)?), + Err(e) => Err(serde_json::from_value(e)?), }) - .await } - pub async fn respond_error(&self, to: u32, val: serde_json::Value) -> crate::Result<()> { + /// [`Session::request`] with a deadline; fails with [`Error::Timeout`]. + pub async fn request_timeout( + &self, + req: M::Request, + timeout: Duration, + ) -> crate::Result> { + tokio::time::timeout(timeout, self.request::(req)).await? + } + + pub async fn respond(&self, to: u32, val: Value) -> crate::Result<()> { + self.send::(&Message::Response { id: to, result: val }) + .await + } + + pub async fn respond_error(&self, to: u32, val: Value) -> crate::Result<()> { self.send::(&Message::ErrorResponse { id: to, error: val }) .await } @@ -253,28 +413,101 @@ impl Session { .await } - async fn trigger_close(&self) { - if let Some(handler) = self.on_close_fn.lock().await.as_ref() { + /// Close the connection. Queued messages are flushed first. + pub async fn close(&self) -> crate::Result<()> { + self.shutdown().await; + Ok(()) + } + + /// Marks the session closed, fails pending requests and runs `on_close`, + /// exactly once. + async fn shutdown(&self) { + if self.inner.closed.swap(true, Ordering::SeqCst) { + return; + } + + self.inner.closed_tx.send_replace(true); + + // Dropping the senders fails every waiting `request`. + drop(std::mem::take(&mut *self.inner.pending.lock().unwrap())); + + let handler = self.inner.on_close.lock().unwrap().clone(); + if let Some(handler) = handler { let _ = handler().await; } } +} - pub async fn close(&self) -> crate::Result<()> { - let res = self.ws.close().await; - self.trigger_close().await; - Ok(res?) +/// Removes a pending request if its `request` future is dropped early. +struct PendingGuard<'a> { + inner: &'a Inner, + id: u32, +} + +impl Drop for PendingGuard<'_> { + fn drop(&mut self) { + self.inner.pending.lock().unwrap().remove(&self.id); + } +} + +/// Waits for the closed flag. The `watch::Ref` is dropped inside, so callers +/// can use this in `select!` without holding a non-`Send` guard. +async fn wait_closed(closed: &mut watch::Receiver) { + let _ = closed.wait_for(|c| *c).await; +} + +async fn write_loop( + mut rx: mpsc::Receiver, + mut sink: BoxSink, + mut closed: watch::Receiver, + session: std::sync::Weak, +) { + loop { + tokio::select! { + biased; + frame = rx.recv() => { + let Some(frame) = frame else { break }; + if sink.send(frame).await.is_err() { + break; + } + } + _ = wait_closed(&mut closed) => { + while let Ok(frame) = rx.try_recv() { + if sink.send(frame).await.is_err() { + break; + } + } + let _ = sink.send(Frame::Close).await; + break; + } + } + } + + let _ = sink.close().await; + + if let Some(inner) = session.upgrade() { + Session { inner }.shutdown().await; + } +} + +impl std::fmt::Debug for Session { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Session") + .field("id", &self.inner.id) + .field("closed", &self.is_closed()) + .finish() } } impl Hash for Session { fn hash(&self, state: &mut H) { - self.ws.id.hash(state); + self.inner.id.hash(state); } } impl PartialEq for Session { fn eq(&self, other: &Self) -> bool { - self.ws.id == other.ws.id + self.inner.id == other.inner.id } } diff --git a/src/transport.rs b/src/transport.rs new file mode 100644 index 0000000..eaabedc --- /dev/null +++ b/src/transport.rs @@ -0,0 +1,36 @@ +//! The transport boundary between the protocol and a WebSocket implementation. + +use futures_util::{Sink, SinkExt, Stream, StreamExt}; + +use crate::BoxError; + +/// A WebSocket-level frame as seen by the protocol layer. +/// +/// Adapters translate their library's message type to and from this. Only +/// `Text` carries protocol messages; `Ping`/`Pong` drive liveness checks and +/// `Close` ends the session. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum Frame { + Text(String), + Binary(Vec), + Ping(Vec), + Pong(Vec), + Close, +} + +pub(crate) type BoxSink = std::pin::Pin + Send>>; +pub(crate) type BoxStream = + std::pin::Pin> + Send>>; + +pub(crate) fn boxed(sink: Si, stream: St) -> (BoxSink, BoxStream) +where + Si: Sink + Send + 'static, + SiE: Into + 'static, + St: Stream> + Send + 'static, + StE: Into + 'static, +{ + ( + Box::pin(sink.sink_map_err(Into::into)), + Box::pin(stream.map(|r| r.map_err(Into::into))), + ) +} diff --git a/src/tungstenite.rs b/src/tungstenite.rs new file mode 100644 index 0000000..9b8bcb8 --- /dev/null +++ b/src/tungstenite.rs @@ -0,0 +1,47 @@ +use futures_util::{SinkExt, StreamExt}; +use tokio::io::{AsyncRead, AsyncWrite}; +use tokio_tungstenite::{WebSocketStream, tungstenite}; + +use crate::{Frame, Session}; + +fn to_ws(frame: Frame) -> tungstenite::Message { + match frame { + Frame::Text(text) => tungstenite::Message::Text(text.into()), + Frame::Binary(data) => tungstenite::Message::Binary(data.into()), + Frame::Ping(data) => tungstenite::Message::Ping(data.into()), + Frame::Pong(data) => tungstenite::Message::Pong(data.into()), + Frame::Close => tungstenite::Message::Close(None), + } +} + +fn from_ws(msg: tungstenite::Message) -> Option { + Some(match msg { + tungstenite::Message::Text(text) => Frame::Text(text.as_str().to_owned()), + tungstenite::Message::Binary(data) => Frame::Binary(data.to_vec()), + tungstenite::Message::Ping(data) => Frame::Ping(data.to_vec()), + tungstenite::Message::Pong(data) => Frame::Pong(data.to_vec()), + tungstenite::Message::Close(_) => Frame::Close, + tungstenite::Message::Frame(_) => return None, + }) +} + +impl Session { + /// Run a session over an established tokio-tungstenite WebSocket, e.g. one + /// accepted with a custom TLS acceptor or handshake callback. + pub fn from_tungstenite(ws: WebSocketStream) -> Self + where + S: AsyncRead + AsyncWrite + Unpin + Send + 'static, + { + let (sink, stream) = ws.split(); + + Session::from_transport( + sink.with(|frame| async move { Ok::<_, tungstenite::Error>(to_ws(frame)) }), + stream.filter_map(|msg| async move { + match msg { + Ok(msg) => from_ws(msg).map(Ok), + Err(e) => Some(Err(e)), + } + }), + ) + } +} diff --git a/src/ws/error.rs b/src/ws/error.rs deleted file mode 100644 index 4c426f3..0000000 --- a/src/ws/error.rs +++ /dev/null @@ -1,31 +0,0 @@ -use std::string::FromUtf8Error; - -pub type Result = std::result::Result; - -#[derive(Debug)] -pub enum Error { - Io(std::io::Error), - InvalidFrame(String), - HandshakeFailed(String), - Utf8(FromUtf8Error), - ConnectionClosed, - Elapsed, -} - -impl From for Error { - fn from(value: std::io::Error) -> Self { - Self::Io(value) - } -} - -impl From for Error { - fn from(value: FromUtf8Error) -> Self { - Self::Utf8(value) - } -} - -impl From for Error { - fn from(_: tokio::time::error::Elapsed) -> Self { - Self::Elapsed - } -} diff --git a/src/ws/handshake.rs b/src/ws/handshake.rs deleted file mode 100644 index 0a3f644..0000000 --- a/src/ws/handshake.rs +++ /dev/null @@ -1,225 +0,0 @@ -use base64::Engine; -use base64::engine::general_purpose::STANDARD as Base64; -use sha1::{Digest, Sha1}; -use std::{collections::HashMap, sync::Arc}; -use tokio::{ - io::{AsyncBufReadExt, AsyncWriteExt, BufReader}, - net::TcpStream, - sync::Mutex, - time::{Duration, timeout}, -}; - -use super::WebSocket; - -pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> { - let (read_half, mut write_half) = stream.split(); - let mut reader = BufReader::new(read_half); - - // ---- 1. Read request line with timeout ---- - let mut request_line = String::new(); - timeout(Duration::from_secs(5), reader.read_line(&mut request_line)).await??; - - let request_line = request_line.trim_end(); - - if !request_line.starts_with("GET") { - write_half - .write_all( - b"HTTP/1.1 405 Method Not Allowed\r\n\ - Content-Length: 0\r\n\ - Connection: close\r\n\r\n", - ) - .await?; - write_half.shutdown().await?; - return Ok(()); - } - - // ---- 2. Read headers with timeout ---- - let mut headers = HashMap::new(); - - loop { - let mut line = String::new(); - timeout(Duration::from_secs(5), reader.read_line(&mut line)).await??; - - if line == "\r\n" { - break; - } - - if let Some((k, v)) = line.split_once(':') { - headers.insert(k.trim().to_lowercase(), v.trim().to_string()); - } - } - - // ---- 3. Check if this is a WebSocket upgrade ---- - let is_upgrade = headers - .get("upgrade") - .map(|v| v.eq_ignore_ascii_case("websocket")) - .unwrap_or(false); - - let has_connection_upgrade = headers - .get("connection") - .map(|v| v.to_lowercase().contains("upgrade")) - .unwrap_or(false); - - if !is_upgrade || !has_connection_upgrade { - // Normal HTTP response (important for browsers) - let body = b"OK"; - - write_half - .write_all( - format!( - "HTTP/1.1 200 OK\r\n\ - Content-Type: text/plain\r\n\ - Content-Length: {}\r\n\ - Connection: close\r\n\ - \r\n", - body.len() - ) - .as_bytes(), - ) - .await?; - - write_half.write_all(body).await?; - write_half.flush().await?; - write_half.shutdown().await?; - - return Ok(()); - } - - // ---- 4. Validate required headers ---- - let key = headers.get("sec-websocket-key").ok_or_else(|| { - std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key") - })?; - - let version_ok = headers - .get("sec-websocket-version") - .map(|v| v == "13") - .unwrap_or(false); - - if !version_ok { - write_half - .write_all( - b"HTTP/1.1 426 Upgrade Required\r\n\ - Sec-WebSocket-Version: 13\r\n\ - Content-Length: 0\r\n\ - Connection: close\r\n\r\n", - ) - .await?; - write_half.shutdown().await?; - return Ok(()); - } - - // ---- 5. Generate Sec-WebSocket-Accept ---- - let mut hasher = Sha1::new(); - hasher.update(key.as_bytes()); - hasher.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - - let accept = Base64.encode(hasher.finalize()); - - // ---- 6. Send upgrade response ---- - let response = format!( - "HTTP/1.1 101 Switching Protocols\r\n\ - Upgrade: websocket\r\n\ - Connection: Upgrade\r\n\ - Sec-WebSocket-Accept: {}\r\n\ - \r\n", - accept - ); - - write_half.write_all(response.as_bytes()).await?; - write_half.flush().await?; - - Ok(()) -} - -impl WebSocket { - pub async fn handshake(mut stream: TcpStream) -> super::Result { - handle_websocket_handshake(&mut stream).await?; - - let (read, write) = stream.into_split(); - - Ok(Self { - id: rand::random(), - reader: Arc::new(Mutex::new(read)), - writer: Arc::new(Mutex::new(write)), - is_server: false, - }) - } - - /// Connect to a WebSocket server and perform the handshake - pub async fn connect(addr: &str, path: &str) -> super::Result { - // 1. TCP connect - let mut stream = TcpStream::connect(addr).await?; - - // 2. Generate Sec-WebSocket-Key - let key_bytes: [u8; 16] = rand::random(); - let key = base64::prelude::BASE64_STANDARD.encode(&key_bytes); - - // 3. Send HTTP Upgrade request - let request = format!( - "GET {} HTTP/1.1\r\n\ - Host: {}\r\n\ - Upgrade: websocket\r\n\ - Connection: Upgrade\r\n\ - Sec-WebSocket-Key: {}\r\n\ - Sec-WebSocket-Version: 13\r\n\ - \r\n", - path, addr, key - ); - stream.write_all(request.as_bytes()).await?; - stream.flush().await?; - - // 4. Read HTTP response - let mut reader = BufReader::new(&mut stream); - let mut status_line = String::new(); - timeout( - tokio::time::Duration::from_secs(5), - reader.read_line(&mut status_line), - ) - .await??; - if !status_line.starts_with("HTTP/1.1 101") { - return Err(super::Error::HandshakeFailed(format!( - "Expected 101 Switching Protocols, got: {}", - status_line.trim_end() - ))); - } - - // Read headers - let mut sec_accept = None; - loop { - let mut line = String::new(); - reader.read_line(&mut line).await?; - let line = line.trim_end(); - if line.is_empty() { - break; // end of headers - } - if let Some((k, v)) = line.split_once(':') { - if k.eq_ignore_ascii_case("sec-websocket-accept") { - sec_accept = Some(v.trim().to_string()); - } - } - } - - // 5. Verify Sec-WebSocket-Accept - let expected = { - let mut sha1 = Sha1::new(); - sha1.update(key.as_bytes()); - sha1.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - base64::prelude::BASE64_STANDARD.encode(sha1.finalize()) - }; - if sec_accept.as_deref() != Some(expected.as_str()) { - return Err(super::Error::HandshakeFailed( - "Sec-WebSocket-Accept mismatch".into(), - )); - } - - // 6. Upgrade succeeded, split stream - let (read, write) = stream.into_split(); - - Ok(Self { - id: rand::random(), - reader: Arc::new(Mutex::new(read)), - writer: Arc::new(Mutex::new(write)), - is_server: true, - }) - } -} diff --git a/src/ws/mod.rs b/src/ws/mod.rs deleted file mode 100644 index 6a6de66..0000000 --- a/src/ws/mod.rs +++ /dev/null @@ -1,251 +0,0 @@ -pub mod error; -pub mod handshake; -pub use error::{Error, Result}; - -use std::{ - hash::{Hash, Hasher}, - sync::Arc, -}; -use tokio::{ - io::{AsyncReadExt, AsyncWriteExt}, - sync::Mutex, -}; - -#[derive(Debug, Clone)] -pub enum Frame { - Text(String), - Binary(Vec), - Ping, - Pong, - Close, -} - -pub struct WebSocket { - pub(crate) reader: Arc>, - pub(crate) writer: Arc>, - pub(crate) id: u64, - pub(crate) is_server: bool, -} - -impl Clone for WebSocket { - fn clone(&self) -> Self { - WebSocket { - reader: self.reader.clone(), - writer: self.writer.clone(), - is_server: self.is_server.clone(), - id: self.id, - } - } -} - -impl PartialEq for WebSocket { - fn eq(&self, other: &Self) -> bool { - self.id == other.id - } -} - -impl Eq for WebSocket {} - -impl Hash for WebSocket { - fn hash(&self, state: &mut H) { - self.id.hash(state); - } -} - -impl WebSocket { - async fn send_frame(&self, opcode: u8, payload: &[u8]) -> Result<()> { - let mut writer = self.writer.lock().await; - - let mut header = Vec::with_capacity(10); - let mask_bit = if self.is_server { 0x80 } else { 0x00 }; - header.push(0x80 | opcode); // FIN + opcode - - let len = payload.len(); - if len < 126 { - header.push((len as u8) | mask_bit); - } else if len <= 0xFFFF { - header.push(126 | mask_bit); - header.extend_from_slice(&(len as u16).to_be_bytes()); - } else { - header.push(127 | mask_bit); - header.extend_from_slice(&(len as u64).to_be_bytes()); - } - - if self.is_server { - // Generate 4-byte mask key - let mask_key: [u8; 4] = rand::random(); - header.extend_from_slice(&mask_key); - - // Mask the payload - let mut masked_payload = payload.to_vec(); - for i in 0..masked_payload.len() { - masked_payload[i] ^= mask_key[i % 4]; - } - - writer.write_all(&header).await?; - writer.write_all(&masked_payload).await?; - } else { - writer.write_all(&header).await?; - writer.write_all(payload).await?; - } - - writer.flush().await?; - Ok(()) - } -} - -impl WebSocket { - pub async fn send(&self, msg: &str) -> Result<()> { - self.send_frame(0x1, msg.as_bytes()).await - } - - pub async fn send_text_payload(&self, payload: &[u8]) -> Result<()> { - self.send_frame(0x1, payload).await - } - - pub async fn send_bin(&self, payload: &[u8]) -> Result<()> { - self.send_frame(0x2, payload).await - } - - pub async fn send_ping(&self) -> Result<()> { - self.send_frame(0x9, &[]).await - } - - pub async fn send_pong(&self) -> Result<()> { - self.send_frame(0xA, &[]).await - } - - pub async fn close(&self) -> Result<()> { - self.send_frame(0x8, &[]).await - } - - pub fn start_ping_loop(&self) { - let s = self.clone(); - tokio::task::spawn(async move { - let mut interval = tokio::time::interval(std::time::Duration::from_secs(15)); - loop { - interval.tick().await; - if s.send_ping().await.is_err() { - break; - } - } - }); - } -} - -impl WebSocket { - /// Read a full WebSocket frame (handling masking and control frames) - /// Returns (opcode, payload) - pub async fn read_frame(&self) -> Result<(bool, u8, Vec)> { - let mut reader = self.reader.lock().await; - - // --- 1. Read first 2-byte header --- - let mut header = [0u8; 2]; - reader.read_exact(&mut header).await?; - - let fin = header[0] & 0x80 != 0; - let opcode = header[0] & 0x0F; - let masked = header[1] & 0x80 != 0; - let mut payload_len = (header[1] & 0x7F) as u64; - - // --- 2. Read extended payload length if necessary --- - if payload_len == 126 { - let mut buf = [0u8; 2]; - reader.read_exact(&mut buf).await?; - payload_len = u16::from_be_bytes(buf) as u64; - } else if payload_len == 127 { - let mut buf = [0u8; 8]; - reader.read_exact(&mut buf).await?; - payload_len = u64::from_be_bytes(buf); - } - - let payload = if masked { - // --- 3. Read mask key --- - let mut mask = [0u8; 4]; - reader.read_exact(&mut mask).await?; - let mut payload = vec![0u8; payload_len as usize]; - if payload_len > 0 { - reader.read_exact(&mut payload).await?; - for i in 0..payload.len() { - payload[i] ^= mask[i % 4]; - } - } - payload - } else { - // Per spec, client-to-server frames MUST be masked - if !self.is_server { - self.close().await.ok(); - return Err(Error::InvalidFrame( - "Received unmasked frame from client".into(), - )); - } - - let mut payload = vec![0u8; payload_len as usize]; - if payload_len > 0 { - reader.read_exact(&mut payload).await?; - } - payload - }; - - // --- 6. Return opcode + payload --- - Ok((fin, opcode, payload)) - } - - pub async fn read(&self) -> Result { - let (fin, opcode, mut payload) = self.read_frame().await?; - - if !fin { - // Continuation loop - while let (fin, o, mut p) = self.read_frame().await? - && !fin - { - match o { - // Continuation - 0x0 => payload.append(&mut p), - // Close - 0x8 => { - self.close().await.ok(); - } - // Ping - 0x9 => { - self.send_pong().await.ok(); - } - // Pong - 0xA => {} - _ => { - self.close().await.ok(); - return Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}"))); - } - } - } - } - - match opcode { - // Close - 0x8 => { - self.close().await.ok(); - Ok(Frame::Close) - } - - // Ping - 0x9 => { - self.send_pong().await.ok(); - Ok(Frame::Ping) - } - - // Pong - 0xA => Ok(Frame::Pong), - - // Text - 0x1 => Ok(Frame::Text(String::from_utf8(payload)?)), - - // Binary - 0x2 => Ok(Frame::Binary(payload)), - - _ => { - self.close().await.ok(); - Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}"))) - } - } - } -} diff --git a/tests/protocol.rs b/tests/protocol.rs new file mode 100644 index 0000000..9374b8f --- /dev/null +++ b/tests/protocol.rs @@ -0,0 +1,452 @@ +#![cfg(all(feature = "server", feature = "client"))] + +use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::time::Duration; + +use futures_util::{SinkExt, StreamExt}; +use serde::{Deserialize, Serialize}; +use session_rs::server::SessionServer; +use session_rs::{Error, Method, Session}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpStream; +use tokio_tungstenite::tungstenite::Message as WsMessage; + +macro_rules! method { + ($ty:ident, $name:literal, $req:ty, $res:ty) => { + #[derive(Debug, Serialize, Deserialize)] + struct $ty; + + impl Method for $ty { + const NAME: &'static str = $name; + type Request = $req; + type Response = $res; + type Error = String; + } + }; +} + +method!(Echo, "echo", String, String); +method!(Fail, "fail", String, String); +method!(Panic, "panic", (), ()); +method!(Numbers, "numbers", Vec, u32); +method!(AskBack, "ask_back", String, String); +method!(Notice, "notice", String, ()); +method!(Silent, "silent", (), ()); + +const WAIT: Duration = Duration::from_secs(5); + +async fn register_handlers(session: &Session) { + session.on_request::(async |_, s| Ok(s)).await; + session.on_request::(async |_, s| Err(format!("failed: {s}"))).await; + session + .on_request::(async |_, ()| -> Result<(), String> { panic!("boom") }) + .await; + session + .on_request::(async |_, v| Ok(v.iter().sum())) + .await; + session + .on_request::({ + let session = session.clone(); + move |_, s| { + let session = session.clone(); + async move { + // A handler requesting from its own peer must not deadlock. + let reply = session.request::(format!("back:{s}")).await; + reply.map_err(|e| e.to_string())? + } + } + }) + .await; + session + .on_request::(async |_, ()| { + tokio::time::sleep(Duration::from_secs(60)).await; + Ok(()) + }) + .await; +} + +async fn start_server() -> SocketAddr { + start_server_with(|_| {}).await +} + +/// Starts a server; `on_session` sees every accepted session after its +/// handlers are registered. +async fn start_server_with(on_session: impl Fn(Session) + Send + Sync + 'static) -> SocketAddr { + let server = SessionServer::bind("127.0.0.1:0").await.unwrap(); + let addr = server.local_addr().unwrap(); + let on_session = Arc::new(on_session); + + tokio::spawn(async move { + server + .session_loop(move |session, _| { + let on_session = on_session.clone(); + async move { + // Registering late must not lose a request sent right after connecting. + tokio::time::sleep(Duration::from_millis(50)).await; + register_handlers(&session).await; + on_session(session); + Ok(()) + } + }) + .await + }); + + addr +} + +async fn connect(addr: SocketAddr) -> Session { + let session = Session::connect(&format!("ws://{addr}")).await.unwrap(); + register_handlers(&session).await; + session.start_receiver(); + session +} + +async fn raw_client(addr: SocketAddr) -> tokio_tungstenite::WebSocketStream> { + tokio_tungstenite::connect_async(format!("ws://{addr}")).await.unwrap().0 +} + +async fn next_text( + ws: &mut tokio_tungstenite::WebSocketStream>, +) -> serde_json::Value { + loop { + match tokio::time::timeout(WAIT, ws.next()).await.unwrap().unwrap().unwrap() { + WsMessage::Text(t) => return serde_json::from_str(&t).unwrap(), + _ => continue, + } + } +} + +#[tokio::test] +async fn request_response_and_error() { + let client = connect(start_server().await).await; + + assert_eq!(client.request::("hi".into()).await.unwrap(), Ok("hi".into())); + assert_eq!(client.request::(vec![1, 2, 3]).await.unwrap(), Ok(6)); + assert_eq!( + client.request::("x".into()).await.unwrap(), + Err("failed: x".into()) + ); +} + +#[tokio::test] +async fn concurrent_requests_are_matched_by_id() { + let client = connect(start_server().await).await; + + let replies = futures_util::future::join_all( + (0..200).map(|i| { + let client = client.clone(); + async move { client.request::(i.to_string()).await.unwrap() } + }), + ) + .await; + + for (i, reply) in replies.into_iter().enumerate() { + assert_eq!(reply, Ok(i.to_string())); + } +} + +#[tokio::test] +async fn wire_format_matches_protocol() { + let mut ws = raw_client(start_server().await).await; + + ws.send(WsMessage::Text( + r#"{"type":"request","id":7,"method":"echo","data":"hi"}"#.into(), + )) + .await + .unwrap(); + assert_eq!( + next_text(&mut ws).await, + serde_json::json!({"type": "response", "id": 7, "result": "hi"}) + ); + + ws.send(WsMessage::Text( + r#"{"type":"request","id":8,"method":"fail","data":"x"}"#.into(), + )) + .await + .unwrap(); + assert_eq!( + next_text(&mut ws).await, + serde_json::json!({"type": "errorresponse", "id": 8, "error": "failed: x"}) + ); +} + +#[tokio::test] +async fn unknown_method_and_bad_data_get_error_responses() { + let mut ws = raw_client(start_server().await).await; + + ws.send(WsMessage::Text( + r#"{"type":"request","id":1,"method":"nope","data":null}"#.into(), + )) + .await + .unwrap(); + assert_eq!( + next_text(&mut ws).await, + serde_json::json!({"type": "errorresponse", "id": 1, "error": "Unknown method: nope"}) + ); + + ws.send(WsMessage::Text( + r#"{"type":"request","id":2,"method":"numbers","data":"not a list"}"#.into(), + )) + .await + .unwrap(); + let reply = next_text(&mut ws).await; + assert_eq!(reply["type"], "errorresponse"); + assert_eq!(reply["id"], 2); + assert!(reply["error"].as_str().unwrap().starts_with("Invalid request data")); + + // Garbage text is ignored, not fatal. + ws.send(WsMessage::Text("not json".into())).await.unwrap(); + ws.send(WsMessage::Text( + r#"{"type":"request","id":3,"method":"echo","data":"still here"}"#.into(), + )) + .await + .unwrap(); + assert_eq!(next_text(&mut ws).await["result"], "still here"); +} + +#[tokio::test] +async fn panicking_handler_fails_only_that_request() { + let client = connect(start_server().await).await; + + assert_eq!( + client.request::(()).await.unwrap(), + Err("Handler panicked".into()) + ); + assert_eq!(client.request::("ok".into()).await.unwrap(), Ok("ok".into())); +} + +#[tokio::test] +async fn handler_can_request_from_its_peer() { + let client = connect(start_server().await).await; + + assert_eq!( + client.request::("x".into()).await.unwrap(), + Ok("back:x".into()) + ); +} + +#[tokio::test] +async fn notifications_reach_the_peer() { + let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel(); + let addr = start_server_with(move |session| { + let session = session.clone(); + tokio::spawn(async move { + session.notify::("hello".into()).await.unwrap(); + }); + }) + .await; + + let client = Session::connect(&format!("ws://{addr}")).await.unwrap(); + client + .on_notification::(move |msg| { + let tx = tx.clone(); + async move { + tx.send(msg).unwrap(); + } + }) + .await; + client.start_receiver(); + + assert_eq!( + tokio::time::timeout(WAIT, rx.recv()).await.unwrap(), + Some("hello".into()) + ); +} + +#[tokio::test] +async fn close_fails_pending_requests_and_runs_on_close_once() { + let closes = Arc::new(AtomicUsize::new(0)); + let (session_tx, mut session_rx) = tokio::sync::mpsc::unbounded_channel(); + let addr = start_server_with(move |session| { + session_tx.send(session).unwrap(); + }) + .await; + + let client = connect(addr).await; + client + .on_close({ + let closes = closes.clone(); + move || { + closes.fetch_add(1, Ordering::SeqCst); + async { Ok(()) } + } + }) + .await; + + let pending = tokio::spawn({ + let client = client.clone(); + async move { client.request::(()).await } + }); + + let server_side = tokio::time::timeout(WAIT, session_rx.recv()).await.unwrap().unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + server_side.close().await.unwrap(); + + let result = tokio::time::timeout(WAIT, pending).await.unwrap().unwrap(); + assert!(matches!(result, Err(Error::ConnectionClosed)), "{result:?}"); + + tokio::time::timeout(WAIT, client.closed()).await.unwrap(); + client.close().await.unwrap(); + assert_eq!(closes.load(Ordering::SeqCst), 1); + assert!(matches!( + client.request::("late".into()).await, + Err(Error::ConnectionClosed) + )); +} + +#[tokio::test] +async fn request_timeout() { + let client = connect(start_server().await).await; + + let result = client + .request_timeout::((), Duration::from_millis(100)) + .await; + assert!(matches!(result, Err(Error::Timeout)), "{result:?}"); +} + +#[tokio::test] +async fn oversized_message_closes_only_that_connection() { + let addr = start_server().await; + let mut ws = raw_client(addr).await; + + let _ = ws.send(WsMessage::Text("x".repeat(2 << 20).into())).await; + let closed = tokio::time::timeout(WAIT, async { + loop { + match ws.next().await { + None | Some(Err(_)) | Some(Ok(WsMessage::Close(_))) => break, + _ => {} + } + } + }) + .await; + assert!(closed.is_ok(), "server kept the connection open"); + + let client = connect(addr).await; + assert_eq!(client.request::("alive".into()).await.unwrap(), Ok("alive".into())); +} + +/// Opens a TCP connection and completes a WebSocket handshake by hand. +async fn raw_handshake(addr: SocketAddr) -> TcpStream { + let mut tcp = TcpStream::connect(addr).await.unwrap(); + tcp.write_all( + format!( + "GET / HTTP/1.1\r\nHost: {addr}\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\ + Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nSec-WebSocket-Version: 13\r\n\r\n" + ) + .as_bytes(), + ) + .await + .unwrap(); + + let mut response = Vec::new(); + while !response.ends_with(b"\r\n\r\n") { + response.push(tcp.read_u8().await.unwrap()); + } + assert!(response.starts_with(b"HTTP/1.1 101")); + tcp +} + +#[tokio::test] +async fn huge_frame_length_header_is_rejected_without_allocating() { + let addr = start_server().await; + let mut tcp = raw_handshake(addr).await; + + // Masked text frame claiming a 2^62-byte payload. session-rs 0.1 tried to + // allocate this up front. + let mut frame = vec![0x81, 0x80 | 127]; + frame.extend_from_slice(&(1u64 << 62).to_be_bytes()); + frame.extend_from_slice(&[1, 2, 3, 4]); + tcp.write_all(&frame).await.unwrap(); + + let mut buf = [0u8; 1024]; + let closed = tokio::time::timeout(WAIT, async { + loop { + match tcp.read(&mut buf).await { + Ok(0) | Err(_) => break, + Ok(_) => {} + } + } + }) + .await; + assert!(closed.is_ok(), "server kept the connection open"); + + let client = connect(addr).await; + assert_eq!(client.request::("alive".into()).await.unwrap(), Ok("alive".into())); +} + +#[tokio::test] +async fn ping_timeout_closes_unresponsive_peer() { + let closes = Arc::new(AtomicUsize::new(0)); + let addr = start_server_with({ + let closes = closes.clone(); + move |session| { + let closes = closes.clone(); + tokio::spawn(async move { + session + .on_close(move || { + closes.fetch_add(1, Ordering::SeqCst); + async { Ok(()) } + }) + .await; + session.start_ping(Duration::from_millis(50), Duration::from_millis(100)); + }); + } + }) + .await; + + // A tungstenite client answers pings, so it must stay connected. + let client = connect(addr).await; + tokio::time::sleep(Duration::from_millis(400)).await; + assert_eq!(client.request::("alive".into()).await.unwrap(), Ok("alive".into())); + assert_eq!(closes.load(Ordering::SeqCst), 0); + + // A raw TCP peer never reads, so it never pongs. + let _tcp = raw_handshake(addr).await; + tokio::time::timeout(WAIT, async { + while closes.load(Ordering::SeqCst) == 0 { + tokio::time::sleep(Duration::from_millis(20)).await; + } + }) + .await + .expect("unresponsive peer was not closed"); +} + +#[cfg(feature = "axum")] +#[tokio::test] +async fn axum_adapter_serves_sessions_next_to_http_routes() { + use axum::{Router, extract::WebSocketUpgrade, routing::get}; + + let app = Router::new() + .route( + "/", + get(async |upgrade: WebSocketUpgrade| { + upgrade.on_upgrade(async |socket| { + let session = Session::from_axum(socket); + register_handlers(&session).await; + session.start_receiver(); + }) + }), + ) + .route("/health", get(async || "ok")); + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await }); + + let client = connect(addr).await; + assert_eq!(client.request::("hi".into()).await.unwrap(), Ok("hi".into())); + assert_eq!( + client.request::("y".into()).await.unwrap(), + Ok("back:y".into()) + ); + + let mut tcp = TcpStream::connect(addr).await.unwrap(); + tcp.write_all(b"GET /health HTTP/1.1\r\nHost: x\r\nConnection: close\r\n\r\n") + .await + .unwrap(); + let mut body = String::new(); + tcp.read_to_string(&mut body).await.unwrap(); + assert!(body.starts_with("HTTP/1.1 200") && body.ends_with("ok"), "{body}"); +} -- 2.54.0