From e325ca02cedbea56cd3c5a027f97e08aa8274416 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 20:51:35 +0100 Subject: [PATCH 01/27] Simple structure --- .gitignore | 5 +++++ Cargo.lock | 7 +++++++ Cargo.toml | 6 ++++++ examples/client.rs | 1 + examples/server.rs | 1 + src/lib.rs | 14 ++++++++++++++ 6 files changed, 34 insertions(+) create mode 100644 Cargo.lock create mode 100644 Cargo.toml create mode 100644 examples/client.rs create mode 100644 examples/server.rs create mode 100644 src/lib.rs diff --git a/.gitignore b/.gitignore index ad67955..0728338 100644 --- a/.gitignore +++ b/.gitignore @@ -19,3 +19,8 @@ target # and can be added to the global gitignore or merged into this file. For a more nuclear # option (not recommended) you can uncomment the following to ignore the entire idea folder. #.idea/ + + +# Added by cargo + +/target diff --git a/Cargo.lock b/Cargo.lock new file mode 100644 index 0000000..45f11d5 --- /dev/null +++ b/Cargo.lock @@ -0,0 +1,7 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "session-rs" +version = "0.1.0" diff --git a/Cargo.toml b/Cargo.toml new file mode 100644 index 0000000..2f38e0f --- /dev/null +++ b/Cargo.toml @@ -0,0 +1,6 @@ +[package] +name = "session-rs" +version = "0.1.0" +edition = "2024" + +[dependencies] diff --git a/examples/client.rs b/examples/client.rs new file mode 100644 index 0000000..e71fdf5 --- /dev/null +++ b/examples/client.rs @@ -0,0 +1 @@ +fn main() {} \ No newline at end of file diff --git a/examples/server.rs b/examples/server.rs new file mode 100644 index 0000000..f328e4d --- /dev/null +++ b/examples/server.rs @@ -0,0 +1 @@ +fn main() {} diff --git a/src/lib.rs b/src/lib.rs new file mode 100644 index 0000000..b93cf3f --- /dev/null +++ b/src/lib.rs @@ -0,0 +1,14 @@ +pub fn add(left: u64, right: u64) -> u64 { + left + right +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn it_works() { + let result = add(2, 2); + assert_eq!(result, 4); + } +} From 58d38c563b52932d3b8f6dcdeca3aa51ef583e59 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 20:54:14 +0100 Subject: [PATCH 02/27] Simple structure --- src/lib.rs | 16 ++-------------- src/server.rs | 1 + src/session.rs | 1 + 3 files changed, 4 insertions(+), 14 deletions(-) create mode 100644 src/server.rs create mode 100644 src/session.rs diff --git a/src/lib.rs b/src/lib.rs index b93cf3f..1543da2 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,14 +1,2 @@ -pub fn add(left: u64, right: u64) -> u64 { - left + right -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn it_works() { - let result = add(2, 2); - assert_eq!(result, 4); - } -} +pub mod session; +pub mod server; diff --git a/src/server.rs b/src/server.rs new file mode 100644 index 0000000..7754b94 --- /dev/null +++ b/src/server.rs @@ -0,0 +1 @@ +pub struct SessionServer(); diff --git a/src/session.rs b/src/session.rs new file mode 100644 index 0000000..704d549 --- /dev/null +++ b/src/session.rs @@ -0,0 +1 @@ +pub struct Session(); From d430d30ea201cf4da8cd4e84bf2ba681de816d30 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 21:16:13 +0100 Subject: [PATCH 03/27] Session --- Cargo.lock | 485 +++++++++++++++++++++++++++++++++++++++++++++++++ Cargo.toml | 5 + src/lib.rs | 26 ++- src/session.rs | 399 +++++++++++++++++++++++++++++++++++++++- 4 files changed, 913 insertions(+), 2 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 45f11d5..a7b9a42 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,491 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "anyhow" +version = "1.0.101" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5f0e0fee31ef5ed1ba1316088939cea399010ed7731dba877ed44aeb407a75ea" + +[[package]] +name = "base64" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" + +[[package]] +name = "bitflags" +version = "2.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" + +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "chacha20" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "rand_core", +] + +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", +] + +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", +] + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + +[[package]] +name = "getrandom" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "139ef39800118c7683f2fd3c98c1b23c09ae076556b435f8e9064ae108aaeeec" +dependencies = [ + "cfg-if", + "libc", + "r-efi", + "rand_core", + "wasip2", + "wasip3", +] + +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "foldhash", +] + +[[package]] +name = "hashbrown" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + +[[package]] +name = "id-arena" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d3067d79b975e8844ca9eb072e16b31c3c1c36928edf9c6789548c524d0d954" + +[[package]] +name = "indexmap" +version = "2.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7714e70437a7dc3ac8eb7e6f8df75fd8eb422675fc7678aff7364301092b1017" +dependencies = [ + "equivalent", + "hashbrown 0.16.1", + "serde", + "serde_core", +] + +[[package]] +name = "itoa" +version = "1.0.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2" + +[[package]] +name = "leb128fmt" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2" + +[[package]] +name = "libc" +version = "0.2.182" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6800badb6cb2082ffd7b6a67e6125bb39f18782f793520caee8cb8846be06112" + +[[package]] +name = "log" +version = "0.4.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" + +[[package]] +name = "memchr" +version = "2.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" + +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "21b2ebcf727b7760c461f091f9f0f539b77b8e87f2fd88131e7f1b433b3cece4" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + +[[package]] +name = "rand" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc266eb313df6c5c09c1c7b1fbe2510961e5bcd3add930c1e31f7ed9da0feff8" +dependencies = [ + "chacha20", + "getrandom", + "rand_core", +] + +[[package]] +name = "rand_core" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c8d0fd677905edcbeedbf2edb6494d676f0e98d54d5cf9bda0b061cb8fb8aba" + +[[package]] +name = "semver" +version = "1.0.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d767eb0aabc880b29956c35734170f26ed551a859dbd361d140cdbeca61ab1e2" + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.149" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + [[package]] name = "session-rs" version = "0.1.0" +dependencies = [ + "base64", + "rand", + "serde", + "serde_json", + "sha1", +] + +[[package]] +name = "sha1" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest", +] + +[[package]] +name = "syn" +version = "2.0.116" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3df424c70518695237746f84cede799c9c58fcb37450d7b23716568cc8bc69cb" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "typenum" +version = "1.19.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb" + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "unicode-xid" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" + +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + +[[package]] +name = "wasip2" +version = "1.0.2+wasi-0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9517f9239f02c069db75e65f174b3da828fe5f5b945c4dd26bd25d89c03ebcf5" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasip3" +version = "0.4.0+wasi-0.3.0-rc-2026-01-06" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5428f8bf88ea5ddc08faddef2ac4a67e390b88186c703ce6dbd955e1c145aca5" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasm-encoder" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "990065f2fe63003fe337b932cfb5e3b80e0b4d0f5ff650e6985b1048f62c8319" +dependencies = [ + "leb128fmt", + "wasmparser", +] + +[[package]] +name = "wasm-metadata" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909" +dependencies = [ + "anyhow", + "indexmap", + "wasm-encoder", + "wasmparser", +] + +[[package]] +name = "wasmparser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe" +dependencies = [ + "bitflags", + "hashbrown 0.15.5", + "indexmap", + "semver", +] + +[[package]] +name = "wit-bindgen" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7249219f66ced02969388cf2bb044a09756a083d0fab1e566056b04d9fbcaa5" +dependencies = [ + "wit-bindgen-rust-macro", +] + +[[package]] +name = "wit-bindgen-core" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ea61de684c3ea68cb082b7a88508a8b27fcc8b797d738bfc99a82facf1d752dc" +dependencies = [ + "anyhow", + "heck", + "wit-parser", +] + +[[package]] +name = "wit-bindgen-rust" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7c566e0f4b284dd6561c786d9cb0142da491f46a9fbed79ea69cdad5db17f21" +dependencies = [ + "anyhow", + "heck", + "indexmap", + "prettyplease", + "syn", + "wasm-metadata", + "wit-bindgen-core", + "wit-component", +] + +[[package]] +name = "wit-bindgen-rust-macro" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c0f9bfd77e6a48eccf51359e3ae77140a7f50b1e2ebfe62422d8afdaffab17a" +dependencies = [ + "anyhow", + "prettyplease", + "proc-macro2", + "quote", + "syn", + "wit-bindgen-core", + "wit-bindgen-rust", +] + +[[package]] +name = "wit-component" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2" +dependencies = [ + "anyhow", + "bitflags", + "indexmap", + "log", + "serde", + "serde_derive", + "serde_json", + "wasm-encoder", + "wasm-metadata", + "wasmparser", + "wit-parser", +] + +[[package]] +name = "wit-parser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ecc8ac4bc1dc3381b7f59c34f00b67e18f910c2c0f50015669dde7def656a736" +dependencies = [ + "anyhow", + "id-arena", + "indexmap", + "log", + "semver", + "serde", + "serde_derive", + "serde_json", + "unicode-xid", + "wasmparser", +] + +[[package]] +name = "zmij" +version = "1.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" diff --git a/Cargo.toml b/Cargo.toml index 2f38e0f..fd21142 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,3 +4,8 @@ version = "0.1.0" edition = "2024" [dependencies] +base64 = "0.22.1" +rand = "0.10.0" +serde = "1.0.228" +serde_json = "1.0.149" +sha1 = "0.10.6" diff --git a/src/lib.rs b/src/lib.rs index 1543da2..a149d19 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,2 +1,26 @@ -pub mod session; pub mod server; +pub mod session; + +pub enum SessionMessage { + SessionMessage(T), + Binary(Vec), +} + +pub type Result = std::result::Result; + +pub enum Error { + Io(std::io::Error), + Json(serde_json::Error), +} + +impl From for Error { + fn from(value: std::io::Error) -> Self { + Self::Io(value) + } +} + +impl From for Error { + fn from(value: serde_json::Error) -> Self { + Self::Json(value) + } +} diff --git a/src/session.rs b/src/session.rs index 704d549..e9a9242 100644 --- a/src/session.rs +++ b/src/session.rs @@ -1 +1,398 @@ -pub struct Session(); +use std::{ + hash::{Hash, Hasher}, + io::{self, Read, Write}, + net::TcpStream, +}; + +use serde::{Deserialize, Serialize}; + +use crate::SessionMessage; + +pub mod handshake { + use base64::Engine; + use base64::engine::general_purpose::STANDARD as Base64; + use sha1::{Digest, Sha1}; + use std::collections::HashMap; + use std::io::{BufRead, BufReader, Write}; + use std::net::TcpStream; + + const WS_GUID: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; + + pub fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> { + let mut reader = BufReader::new(stream.try_clone()?); + let mut request_line = String::new(); + reader.read_line(&mut request_line)?; + + // Trim CRLF to make sure comparisons are clean + let request_line = request_line.trim_end(); + + // Allow HEAD (used by Render for health checks) + if request_line.starts_with("HEAD") { + let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n"; + stream.write_all(response.as_bytes())?; + stream.flush()?; + return Ok(()); + } + + // Only proceed if it’s a GET + if !request_line.starts_with("GET") { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!("Invalid HTTP method: {request_line}"), + )); + } + + // Read headers + let mut headers = HashMap::new(); + let mut line = String::new(); + loop { + line.clear(); + let bytes = reader.read_line(&mut line)?; + if bytes == 0 || line == "\r\n" { + break; + } + if let Some((k, v)) = line.split_once(':') { + headers.insert(k.trim().to_lowercase(), v.trim().to_string()); + } + } + + // Check if it's actually a WebSocket upgrade request + let is_websocket_upgrade = headers + .get("upgrade") + .map(|v| v.eq_ignore_ascii_case("websocket")) + .unwrap_or(false); + + if !is_websocket_upgrade { + // Not a WebSocket request — probably a normal HTTP GET (e.g. health check) + let response = + "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\n\r\nOK"; + stream.write_all(response.as_bytes())?; + stream.flush()?; + return Ok(()); + } + + // Validate "Connection: Upgrade" + if !headers + .get("connection") + .map(|v| v.to_lowercase().contains("upgrade")) + .unwrap_or(false) + { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "Missing or invalid Connection header", + )); + } + + // Validate WebSocket key + let key = headers.get("sec-websocket-key").ok_or_else(|| { + std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key") + })?; + + // Validate version + if let Some(ver) = headers.get("sec-websocket-version") { + if ver.trim() != "13" { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!("Unsupported Sec-WebSocket-Version: {}", ver), + )); + } + } + + // Compute accept key + let mut hasher = Sha1::new(); + hasher.update(key.as_bytes()); + hasher.update(WS_GUID.as_bytes()); + let hash = hasher.finalize(); + let accept_key = Base64.encode(hash); + + // Send 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_key + ); + + stream.write_all(response.as_bytes())?; + stream.flush()?; + Ok(()) + } +} + +pub struct Session(TcpStream, u64); + +impl Session { + /// Create a client + pub fn new(mut stream: TcpStream) -> crate::Result { + handshake::handle_websocket_handshake(&mut stream)?; + stream.set_read_timeout(Some(std::time::Duration::from_secs(10)))?; + stream.set_write_timeout(Some(std::time::Duration::from_secs(10)))?; + Ok(Session(stream, rand::random())) + } + + /// Send a close frame and flush. + pub fn send_close(&self) -> crate::Result<()> { + let mut stream = self.0.try_clone()?; + stream.write_all(&[0x88])?; + stream.flush()?; + Ok(()) + } + + /// Send a ping (no payload) + fn send_ping(&self) -> crate::Result<()> { + let mut stream = self.0.try_clone()?; + // FIN + opcode (ping = 0x89), payload length = 0x00 + stream.write_all(&[0x89, 0x00])?; + stream.flush()?; + Ok(()) + } + + /// Send a pong (no payload) + fn send_pong(&self) -> crate::Result<()> { + let mut stream = self.0.try_clone()?; + // FIN + opcode (pong = 0x8A), payload length = 0x00 + stream.write_all(&[0x8A, 0x00])?; + stream.flush()?; + Ok(()) + } + + /// Send a text/binary frame (server->client must NOT mask) + pub fn send(&self, m: T) -> crate::Result<()> { + let mut stream = self.0.try_clone()?; + + let payload = serde_json::to_string(&m)?; + let payload_bytes = payload.as_bytes(); + let len = payload_bytes.len(); + + let mut header = Vec::new(); + header.push(0x81); // FIN=1, opcode=0x1 (text) + + if len < 126 { + header.push(len as u8); + } else if len <= 65535 { + header.push(126); + header.extend_from_slice(&(len as u16).to_be_bytes()); + } else { + header.push(127); + header.extend_from_slice(&(len as u64).to_be_bytes()); + } + + stream.write_all(&header)?; + stream.write_all(payload_bytes)?; + stream.flush()?; + Ok(()) + } + + /// Send a binary WebSocket frame (server -> client) + pub fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { + let mut stream = self.0.try_clone()?; + + let mut header = Vec::with_capacity(10); + + // FIN=1, opcode=2 (binary) + header.push(0x82); + + let len = payload.len(); + + if len < 126 { + header.push(len as u8); // mask bit = 0 + } else if len <= 0xFFFF { + header.push(126); + header.extend_from_slice(&(len as u16).to_be_bytes()); + } else { + header.push(127); + header.extend_from_slice(&(len as u64).to_be_bytes()); + } + + stream.write_all(&header)?; + stream.write_all(payload)?; + stream.flush()?; + + Ok(()) + } + + /// Read a full WebSocket message, handling fragmentation and control frames. + /// + /// Returns: + /// - Ok(Some(WsMessage)) on an application message (text/binary) + /// - Ok(None) if the connection should be closed (close received / read EOF) + /// - Err on protocol or IO errors. + pub fn read_t Deserialize<'de>>( + &self, + ) -> crate::Result>> { + let mut stream = self.0.try_clone()?; + + let mut message_payload = Vec::new(); + let mut expecting_continuation = false; + let mut message_type: Option = None; // 0x1 for text, 0x2 for binary + + loop { + // Read 2-byte header + let mut header = [0u8; 2]; + match stream.read_exact(&mut header) { + Ok(_) => {} + Err(e) => match e.kind() { + io::ErrorKind::WouldBlock | io::ErrorKind::TimedOut => { + self.send_ping()?; + continue; + } + io::ErrorKind::UnexpectedEof | io::ErrorKind::BrokenPipe => return Ok(None), + _ => return Err(e.into()), + }, + } + + 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; + + // Extended payload length + if payload_len == 126 { + let mut ext_len = [0u8; 2]; + stream.read_exact(&mut ext_len)?; + payload_len = u16::from_be_bytes(ext_len) as u64; + } else if payload_len == 127 { + let mut ext_len = [0u8; 8]; + stream.read_exact(&mut ext_len)?; + payload_len = u64::from_be_bytes(ext_len); + } + + // Mask key + let mut mask = [0u8; 4]; + if masked { + stream.read_exact(&mut mask)?; + } else { + let _ = self.send_close(); + return Ok(None); + } + + // Control frame checks + if matches!(opcode, 0x8 | 0x9 | 0xA) { + if payload_len > 125 { + let _ = self.send_close(); + return Ok(None); + } + if !fin { + let _ = self.send_close(); + return Ok(None); + } + } + + // Read payload + let mut payload = vec![0u8; payload_len as usize]; + if payload_len > 0 { + stream.read_exact(&mut payload)?; + for i in 0..payload.len() { + payload[i] ^= mask[i % 4]; + } + } + + match opcode { + 0x0 => { + // Continuation + if !expecting_continuation { + let _ = self.send_close(); + return Ok(None); + } + message_payload.extend(payload); + if fin { + break; + } + } + 0x1 => { + // Text + if expecting_continuation { + let _ = self.send_close(); + return Ok(None); + } + message_payload.extend(payload); + message_type = Some(0x1); + if fin { + break; + } else { + expecting_continuation = true; + } + } + 0x2 => { + // Binary + if expecting_continuation { + let _ = self.send_close(); + return Ok(None); + } + message_payload.extend(payload); + message_type = Some(0x2); + if fin { + break; + } else { + expecting_continuation = true; + } + } + 0x8 => { + // Close + let _ = self.send_close(); + return Ok(None); + } + 0x9 => { + // Ping + self.send_pong()?; + continue; + } + 0xA => { + // Pong + continue; + } + _ => { + let _ = self.send_close(); + return Ok(None); + } + } + } + + // Convert payload into proper message type + let message = match message_type { + Some(0x1) => { + // Text frame → try JSON, otherwise keep text + match String::from_utf8(message_payload.clone()) { + Ok(text) => match serde_json::from_str(&text) { + Ok(msg) => SessionMessage::SessionMessage(msg), + Err(e) => return Err(crate::Error::Json(e)), + }, + Err(_) => SessionMessage::Binary(message_payload), + } + } + Some(0x2) => SessionMessage::Binary(message_payload), + _ => return Ok(None), // Should not happen + }; + + Ok(Some(message)) + } + + pub fn close(&self) -> crate::Result<()> { + self.0.shutdown(std::net::Shutdown::Both)?; + Ok(()) + } +} + +impl Clone for Session { + fn clone(&self) -> Self { + Session( + self.0.try_clone().expect("failed to clone TcpStream"), + self.1.clone(), + ) + } +} + +impl PartialEq for Session { + fn eq(&self, other: &Self) -> bool { + self.1 == other.1 + } +} + +impl Eq for Session {} + +impl Hash for Session { + fn hash(&self, state: &mut H) { + self.1.hash(state); + } +} From 3c254b75ff86714c3b906c0984b8ce4610fb219d Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 21:24:00 +0100 Subject: [PATCH 04/27] Async handshake --- Cargo.lock | 143 +++++++++++++++++++++++++++++++++++++++++++++++ Cargo.toml | 1 + src/handshake.rs | 77 +++++++++++++++++++++++++ src/lib.rs | 1 + src/session.rs | 114 +------------------------------------ 5 files changed, 223 insertions(+), 113 deletions(-) create mode 100644 src/handshake.rs diff --git a/Cargo.lock b/Cargo.lock index a7b9a42..a1ec559 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -29,6 +29,12 @@ dependencies = [ "generic-array", ] +[[package]] +name = "bytes" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b35204fbdc0b3f4446b89fc1ac2cf84a8a68971995d0bf2e925ec7cd960f9cb3" + [[package]] name = "cfg-if" version = "1.0.4" @@ -189,6 +195,23 @@ version = "2.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" +[[package]] +name = "mio" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a69bcab0ad47271a0234d9422b131806bf3968021e5dc9328caf2d4cd58557fc" +dependencies = [ + "libc", + "wasi", + "windows-sys 0.61.2", +] + +[[package]] +name = "pin-project-lite" +version = "0.2.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b3cff922bd51709b605d9ead9aa71031d81447142d828eb4a6eba76fe619f9b" + [[package]] name = "prettyplease" version = "0.2.37" @@ -297,6 +320,7 @@ dependencies = [ "serde", "serde_json", "sha1", + "tokio", ] [[package]] @@ -310,6 +334,16 @@ dependencies = [ "digest", ] +[[package]] +name = "socket2" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "17129e116933cf371d018bb80ae557e889637989d8638274fb25622827b03881" +dependencies = [ + "libc", + "windows-sys 0.60.2", +] + [[package]] name = "syn" version = "2.0.116" @@ -321,6 +355,20 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "tokio" +version = "1.49.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72a2903cd7736441aac9df9d7688bd0ce48edccaadf181c3b90be801e81d3d86" +dependencies = [ + "bytes", + "libc", + "mio", + "pin-project-lite", + "socket2", + "windows-sys 0.61.2", +] + [[package]] name = "typenum" version = "1.19.0" @@ -345,6 +393,12 @@ version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" +[[package]] +name = "wasi" +version = "0.11.1+wasi-snapshot-preview1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" + [[package]] name = "wasip2" version = "1.0.2+wasi-0.2.9" @@ -397,6 +451,95 @@ dependencies = [ "semver", ] +[[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.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.53.5" +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", +] + +[[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.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006" + +[[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.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c" + +[[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.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499" + +[[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.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" + [[package]] name = "wit-bindgen" version = "0.51.0" diff --git a/Cargo.toml b/Cargo.toml index fd21142..c73eaf7 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -9,3 +9,4 @@ rand = "0.10.0" serde = "1.0.228" serde_json = "1.0.149" sha1 = "0.10.6" +tokio = { version = "1.49.0", features = ["io-util", "net"] } diff --git a/src/handshake.rs b/src/handshake.rs new file mode 100644 index 0000000..a3c98e8 --- /dev/null +++ b/src/handshake.rs @@ -0,0 +1,77 @@ +use tokio::{ + io::{AsyncBufReadExt, AsyncWriteExt, BufReader}, + net::TcpStream, +}; + +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); + + let mut request_line = String::new(); + reader.read_line(&mut request_line).await?; + let request_line = request_line.trim_end(); + + if request_line.starts_with("HEAD") { + write_half + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n") + .await?; + return Ok(()); + } + + if !request_line.starts_with("GET") { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "Invalid HTTP method", + )); + } + + use std::collections::HashMap; + let mut headers = HashMap::new(); + let mut line = String::new(); + + loop { + line.clear(); + 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()); + } + } + + if headers + .get("upgrade") + .map(|v| !v.eq_ignore_ascii_case("websocket")) + .unwrap_or(true) + { + write_half + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nOK") + .await?; + return Ok(()); + } + + let key = headers + .get("sec-websocket-key") + .ok_or_else(|| std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing key"))?; + + use base64::Engine; + use base64::engine::general_purpose::STANDARD as Base64; + use sha1::{Digest, Sha1}; + + let mut hasher = Sha1::new(); + hasher.update(key.as_bytes()); + hasher.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = Base64.encode(hasher.finalize()); + + 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?; + Ok(()) +} diff --git a/src/lib.rs b/src/lib.rs index a149d19..54aaca5 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,4 @@ +pub mod handshake; pub mod server; pub mod session; diff --git a/src/session.rs b/src/session.rs index e9a9242..3aa1096 100644 --- a/src/session.rs +++ b/src/session.rs @@ -8,124 +8,12 @@ use serde::{Deserialize, Serialize}; use crate::SessionMessage; -pub mod handshake { - use base64::Engine; - use base64::engine::general_purpose::STANDARD as Base64; - use sha1::{Digest, Sha1}; - use std::collections::HashMap; - use std::io::{BufRead, BufReader, Write}; - use std::net::TcpStream; - - const WS_GUID: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; - - pub fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> { - let mut reader = BufReader::new(stream.try_clone()?); - let mut request_line = String::new(); - reader.read_line(&mut request_line)?; - - // Trim CRLF to make sure comparisons are clean - let request_line = request_line.trim_end(); - - // Allow HEAD (used by Render for health checks) - if request_line.starts_with("HEAD") { - let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n"; - stream.write_all(response.as_bytes())?; - stream.flush()?; - return Ok(()); - } - - // Only proceed if it’s a GET - if !request_line.starts_with("GET") { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - format!("Invalid HTTP method: {request_line}"), - )); - } - - // Read headers - let mut headers = HashMap::new(); - let mut line = String::new(); - loop { - line.clear(); - let bytes = reader.read_line(&mut line)?; - if bytes == 0 || line == "\r\n" { - break; - } - if let Some((k, v)) = line.split_once(':') { - headers.insert(k.trim().to_lowercase(), v.trim().to_string()); - } - } - - // Check if it's actually a WebSocket upgrade request - let is_websocket_upgrade = headers - .get("upgrade") - .map(|v| v.eq_ignore_ascii_case("websocket")) - .unwrap_or(false); - - if !is_websocket_upgrade { - // Not a WebSocket request — probably a normal HTTP GET (e.g. health check) - let response = - "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\n\r\nOK"; - stream.write_all(response.as_bytes())?; - stream.flush()?; - return Ok(()); - } - - // Validate "Connection: Upgrade" - if !headers - .get("connection") - .map(|v| v.to_lowercase().contains("upgrade")) - .unwrap_or(false) - { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - "Missing or invalid Connection header", - )); - } - - // Validate WebSocket key - let key = headers.get("sec-websocket-key").ok_or_else(|| { - std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key") - })?; - - // Validate version - if let Some(ver) = headers.get("sec-websocket-version") { - if ver.trim() != "13" { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - format!("Unsupported Sec-WebSocket-Version: {}", ver), - )); - } - } - - // Compute accept key - let mut hasher = Sha1::new(); - hasher.update(key.as_bytes()); - hasher.update(WS_GUID.as_bytes()); - let hash = hasher.finalize(); - let accept_key = Base64.encode(hash); - - // Send 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_key - ); - - stream.write_all(response.as_bytes())?; - stream.flush()?; - Ok(()) - } -} - pub struct Session(TcpStream, u64); impl Session { /// Create a client pub fn new(mut stream: TcpStream) -> crate::Result { - handshake::handle_websocket_handshake(&mut stream)?; + crate::handshake::handle_websocket_handshake(&mut stream)?; stream.set_read_timeout(Some(std::time::Duration::from_secs(10)))?; stream.set_write_timeout(Some(std::time::Duration::from_secs(10)))?; Ok(Session(stream, rand::random())) From dd1d261353e4ebfeceeafed199d54200e370d042 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 21:33:48 +0100 Subject: [PATCH 05/27] Async session --- Cargo.toml | 2 +- src/session.rs | 344 ++++++++++++------------------------------------- 2 files changed, 86 insertions(+), 260 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index c73eaf7..8ed7179 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -9,4 +9,4 @@ rand = "0.10.0" serde = "1.0.228" serde_json = "1.0.149" sha1 = "0.10.6" -tokio = { version = "1.49.0", features = ["io-util", "net"] } +tokio = { version = "1.49.0", features = ["io-util", "net", "rt", "sync", "time"] } diff --git a/src/session.rs b/src/session.rs index 3aa1096..27a42d3 100644 --- a/src/session.rs +++ b/src/session.rs @@ -1,90 +1,102 @@ -use std::{ - hash::{Hash, Hasher}, - io::{self, Read, Write}, - net::TcpStream, -}; - -use serde::{Deserialize, Serialize}; - use crate::SessionMessage; -pub struct Session(TcpStream, u64); +use std::{ + hash::{Hash, Hasher}, + sync::Arc, +}; +use tokio::{io::AsyncWriteExt, net::TcpStream, sync::Mutex}; + +pub struct Session { + reader: Arc>, + writer: Arc>, + id: u64, +} impl Session { - /// Create a client - pub fn new(mut stream: TcpStream) -> crate::Result { - crate::handshake::handle_websocket_handshake(&mut stream)?; - stream.set_read_timeout(Some(std::time::Duration::from_secs(10)))?; - stream.set_write_timeout(Some(std::time::Duration::from_secs(10)))?; - Ok(Session(stream, rand::random())) + pub async fn new(mut stream: TcpStream) -> crate::Result { + crate::handshake::handle_websocket_handshake(&mut stream).await?; + + let (read, write) = stream.into_split(); + + Ok(Self { + reader: Arc::new(Mutex::new(read)), + writer: Arc::new(Mutex::new(write)), + id: rand::random(), + }) } +} - /// Send a close frame and flush. - pub fn send_close(&self) -> crate::Result<()> { - let mut stream = self.0.try_clone()?; - stream.write_all(&[0x88])?; - stream.flush()?; - Ok(()) - } - - /// Send a ping (no payload) - fn send_ping(&self) -> crate::Result<()> { - let mut stream = self.0.try_clone()?; - // FIN + opcode (ping = 0x89), payload length = 0x00 - stream.write_all(&[0x89, 0x00])?; - stream.flush()?; - Ok(()) - } - - /// Send a pong (no payload) - fn send_pong(&self) -> crate::Result<()> { - let mut stream = self.0.try_clone()?; - // FIN + opcode (pong = 0x8A), payload length = 0x00 - stream.write_all(&[0x8A, 0x00])?; - stream.flush()?; - Ok(()) - } - - /// Send a text/binary frame (server->client must NOT mask) - pub fn send(&self, m: T) -> crate::Result<()> { - let mut stream = self.0.try_clone()?; - - let payload = serde_json::to_string(&m)?; - let payload_bytes = payload.as_bytes(); - let len = payload_bytes.len(); - - let mut header = Vec::new(); - header.push(0x81); // FIN=1, opcode=0x1 (text) - - if len < 126 { - header.push(len as u8); - } else if len <= 65535 { - header.push(126); - header.extend_from_slice(&(len as u16).to_be_bytes()); - } else { - header.push(127); - header.extend_from_slice(&(len as u64).to_be_bytes()); +impl Clone for Session { + fn clone(&self) -> Self { + Session { + reader: self.reader.clone(), + writer: self.writer.clone(), + id: self.id, } + } +} - stream.write_all(&header)?; - stream.write_all(payload_bytes)?; - stream.flush()?; +impl PartialEq for Session { + fn eq(&self, other: &Self) -> bool { + self.id == other.id + } +} + +impl Eq for Session {} + +impl Hash for Session { + fn hash(&self, state: &mut H) { + self.id.hash(state); + } +} + +impl Session { + pub async fn send(&self, msg: &T) -> crate::Result<()> { + let payload = serde_json::to_vec(msg)?; + self.send_frame(0x1, &payload).await + } + + pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { + self.send_frame(0x2, payload).await + } + + pub async fn send_ping(&self) -> crate::Result<()> { + let mut writer = self.writer.lock().await; + // FIN + opcode = 0x89 (ping), payload length = 0 + writer.write_all(&[0x89, 0x00]).await?; + writer.flush().await?; Ok(()) } - /// Send a binary WebSocket frame (server -> client) - pub fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { - let mut stream = self.0.try_clone()?; + pub async fn send_pong(&self) -> crate::Result<()> { + let mut writer = self.writer.lock().await; + // FIN + opcode = 0x8A (pong), payload length = 0 + writer.write_all(&[0x8A, 0x00]).await?; + writer.flush().await?; + Ok(()) + } + + pub fn start_ping(self: Arc) { + tokio::task::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_secs(15)); + loop { + interval.tick().await; + if self.send_ping().await.is_err() { + break; + } + } + }); + } + + async fn send_frame(&self, opcode: u8, payload: &[u8]) -> crate::Result<()> { + let mut writer = self.writer.lock().await; let mut header = Vec::with_capacity(10); - - // FIN=1, opcode=2 (binary) - header.push(0x82); + header.push(0x80 | opcode); let len = payload.len(); - if len < 126 { - header.push(len as u8); // mask bit = 0 + header.push(len as u8); } else if len <= 0xFFFF { header.push(126); header.extend_from_slice(&(len as u16).to_be_bytes()); @@ -93,194 +105,8 @@ impl Session { header.extend_from_slice(&(len as u64).to_be_bytes()); } - stream.write_all(&header)?; - stream.write_all(payload)?; - stream.flush()?; - - Ok(()) - } - - /// Read a full WebSocket message, handling fragmentation and control frames. - /// - /// Returns: - /// - Ok(Some(WsMessage)) on an application message (text/binary) - /// - Ok(None) if the connection should be closed (close received / read EOF) - /// - Err on protocol or IO errors. - pub fn read_t Deserialize<'de>>( - &self, - ) -> crate::Result>> { - let mut stream = self.0.try_clone()?; - - let mut message_payload = Vec::new(); - let mut expecting_continuation = false; - let mut message_type: Option = None; // 0x1 for text, 0x2 for binary - - loop { - // Read 2-byte header - let mut header = [0u8; 2]; - match stream.read_exact(&mut header) { - Ok(_) => {} - Err(e) => match e.kind() { - io::ErrorKind::WouldBlock | io::ErrorKind::TimedOut => { - self.send_ping()?; - continue; - } - io::ErrorKind::UnexpectedEof | io::ErrorKind::BrokenPipe => return Ok(None), - _ => return Err(e.into()), - }, - } - - 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; - - // Extended payload length - if payload_len == 126 { - let mut ext_len = [0u8; 2]; - stream.read_exact(&mut ext_len)?; - payload_len = u16::from_be_bytes(ext_len) as u64; - } else if payload_len == 127 { - let mut ext_len = [0u8; 8]; - stream.read_exact(&mut ext_len)?; - payload_len = u64::from_be_bytes(ext_len); - } - - // Mask key - let mut mask = [0u8; 4]; - if masked { - stream.read_exact(&mut mask)?; - } else { - let _ = self.send_close(); - return Ok(None); - } - - // Control frame checks - if matches!(opcode, 0x8 | 0x9 | 0xA) { - if payload_len > 125 { - let _ = self.send_close(); - return Ok(None); - } - if !fin { - let _ = self.send_close(); - return Ok(None); - } - } - - // Read payload - let mut payload = vec![0u8; payload_len as usize]; - if payload_len > 0 { - stream.read_exact(&mut payload)?; - for i in 0..payload.len() { - payload[i] ^= mask[i % 4]; - } - } - - match opcode { - 0x0 => { - // Continuation - if !expecting_continuation { - let _ = self.send_close(); - return Ok(None); - } - message_payload.extend(payload); - if fin { - break; - } - } - 0x1 => { - // Text - if expecting_continuation { - let _ = self.send_close(); - return Ok(None); - } - message_payload.extend(payload); - message_type = Some(0x1); - if fin { - break; - } else { - expecting_continuation = true; - } - } - 0x2 => { - // Binary - if expecting_continuation { - let _ = self.send_close(); - return Ok(None); - } - message_payload.extend(payload); - message_type = Some(0x2); - if fin { - break; - } else { - expecting_continuation = true; - } - } - 0x8 => { - // Close - let _ = self.send_close(); - return Ok(None); - } - 0x9 => { - // Ping - self.send_pong()?; - continue; - } - 0xA => { - // Pong - continue; - } - _ => { - let _ = self.send_close(); - return Ok(None); - } - } - } - - // Convert payload into proper message type - let message = match message_type { - Some(0x1) => { - // Text frame → try JSON, otherwise keep text - match String::from_utf8(message_payload.clone()) { - Ok(text) => match serde_json::from_str(&text) { - Ok(msg) => SessionMessage::SessionMessage(msg), - Err(e) => return Err(crate::Error::Json(e)), - }, - Err(_) => SessionMessage::Binary(message_payload), - } - } - Some(0x2) => SessionMessage::Binary(message_payload), - _ => return Ok(None), // Should not happen - }; - - Ok(Some(message)) - } - - pub fn close(&self) -> crate::Result<()> { - self.0.shutdown(std::net::Shutdown::Both)?; + writer.write_all(&header).await?; + writer.write_all(payload).await?; Ok(()) } } - -impl Clone for Session { - fn clone(&self) -> Self { - Session( - self.0.try_clone().expect("failed to clone TcpStream"), - self.1.clone(), - ) - } -} - -impl PartialEq for Session { - fn eq(&self, other: &Self) -> bool { - self.1 == other.1 - } -} - -impl Eq for Session {} - -impl Hash for Session { - fn hash(&self, state: &mut H) { - self.1.hash(state); - } -} From 09c1fbc2323de18a600f736b0f3211575ef0b785 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 21:49:11 +0100 Subject: [PATCH 06/27] async Read frame --- Cargo.lock | 12 ++++ Cargo.toml | 2 +- examples/client.rs | 3 +- examples/server.rs | 3 +- src/lib.rs | 6 +- src/session.rs | 135 ++++++++++++++++++++++++++++++++++++++------- 6 files changed, 136 insertions(+), 25 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index a1ec559..3b82341 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -366,9 +366,21 @@ dependencies = [ "mio", "pin-project-lite", "socket2", + "tokio-macros", "windows-sys 0.61.2", ] +[[package]] +name = "tokio-macros" +version = "2.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af407857209536a95c8e56f8231ef2c2e2aff839b22e07a1ffcbc617e9db9fa5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "typenum" version = "1.19.0" diff --git a/Cargo.toml b/Cargo.toml index 8ed7179..7a353f2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -9,4 +9,4 @@ rand = "0.10.0" serde = "1.0.228" serde_json = "1.0.149" sha1 = "0.10.6" -tokio = { version = "1.49.0", features = ["io-util", "net", "rt", "sync", "time"] } +tokio = { version = "1.49.0", features = ["io-util", "macros", "net", "rt", "sync", "time"] } diff --git a/examples/client.rs b/examples/client.rs index e71fdf5..ab9323e 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1 +1,2 @@ -fn main() {} \ No newline at end of file +#[tokio::main(flavor = "current_thread")] +async fn main() {} diff --git a/examples/server.rs b/examples/server.rs index f328e4d..ab9323e 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1 +1,2 @@ -fn main() {} +#[tokio::main(flavor = "current_thread")] +async fn main() {} diff --git a/src/lib.rs b/src/lib.rs index 54aaca5..935aa31 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -2,8 +2,8 @@ pub mod handshake; pub mod server; pub mod session; -pub enum SessionMessage { - SessionMessage(T), +pub enum SessionFrame { + Typed(T), Binary(Vec), } @@ -12,6 +12,8 @@ pub type Result = std::result::Result; pub enum Error { Io(std::io::Error), Json(serde_json::Error), + InvalidFrame(String), + ConnectionClosed, } impl From for Error { diff --git a/src/session.rs b/src/session.rs index 27a42d3..333af59 100644 --- a/src/session.rs +++ b/src/session.rs @@ -1,10 +1,12 @@ -use crate::SessionMessage; - use std::{ hash::{Hash, Hasher}, sync::Arc, }; -use tokio::{io::AsyncWriteExt, net::TcpStream, sync::Mutex}; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpStream, + sync::Mutex, +}; pub struct Session { reader: Arc>, @@ -50,6 +52,30 @@ impl Hash for Session { } } +impl Session { + async fn send_frame(&self, opcode: u8, payload: &[u8]) -> crate::Result<()> { + let mut writer = self.writer.lock().await; + + let mut header = Vec::with_capacity(10); + header.push(0x80 | opcode); + + let len = payload.len(); + if len < 126 { + header.push(len as u8); + } else if len <= 0xFFFF { + header.push(126); + header.extend_from_slice(&(len as u16).to_be_bytes()); + } else { + header.push(127); + header.extend_from_slice(&(len as u64).to_be_bytes()); + } + + writer.write_all(&header).await?; + writer.write_all(payload).await?; + Ok(()) + } +} + impl Session { pub async fn send(&self, msg: &T) -> crate::Result<()> { let payload = serde_json::to_vec(msg)?; @@ -88,25 +114,94 @@ impl Session { }); } - async fn send_frame(&self, opcode: u8, payload: &[u8]) -> crate::Result<()> { + /// Send a close frame and flush + pub async fn close(&self) -> crate::Result<()> { let mut writer = self.writer.lock().await; - let mut header = Vec::with_capacity(10); - header.push(0x80 | opcode); - - let len = payload.len(); - if len < 126 { - header.push(len as u8); - } else if len <= 0xFFFF { - header.push(126); - header.extend_from_slice(&(len as u16).to_be_bytes()); - } else { - header.push(127); - header.extend_from_slice(&(len as u64).to_be_bytes()); - } - - writer.write_all(&header).await?; - writer.write_all(payload).await?; + // FIN=1, opcode=0x8 (close), payload length=0 + let frame = [0x88, 0x00]; + writer.write_all(&frame).await?; + writer.flush().await?; Ok(()) } } + +impl Session { + /// Read a full WebSocket frame (handling masking and control frames) + /// Returns (opcode, payload) + pub async fn read_frame(&self) -> crate::Result)>> { + 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); + } + + // --- 3. Read mask key --- + if !masked { + // Per spec, client-to-server frames MUST be masked + self.close().await.ok(); + return Err(crate::Error::InvalidFrame( + "Received unmasked frame from client".into(), + )); + } + + let mut mask = [0u8; 4]; + reader.read_exact(&mut mask).await?; + + // --- 4. Read payload --- + 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]; + } + } + + // --- 5. Handle control frames immediately --- + match opcode { + 0x8 => { + // Close + self.close().await.ok(); + return Err(crate::Error::ConnectionClosed); + } + 0x9 => { + // Ping + self.send_pong().await.ok(); + return Ok(None); + } + 0xA => { + // Pong, ignore + return Ok(None); + } + 0x0 | 0x1 | 0x2 => { + // Continuation / Text / Binary → valid payload + } + _ => { + self.close().await.ok(); + return Err(crate::Error::InvalidFrame(format!( + "Unknown opcode: {}", + opcode + ))); + } + } + + // --- 6. Return opcode + payload --- + Ok(Some((opcode, payload))) + } +} From 495de7e87d542681e0b6b96dc7eeed7892fbf2bb Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 21:54:16 +0100 Subject: [PATCH 07/27] Server client example --- examples/client.rs | 37 ++++++++++++++++++++++++++- examples/server.rs | 62 +++++++++++++++++++++++++++++++++++++++++++++- src/lib.rs | 1 + 3 files changed, 98 insertions(+), 2 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index ab9323e..5d09db5 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,2 +1,37 @@ +use std::sync::Arc; +use tokio::net::TcpStream; + +use session_rs::session::Session; + #[tokio::main(flavor = "current_thread")] -async fn main() {} +async fn main() -> session_rs::Result<()> { + let stream = TcpStream::connect("127.0.0.1:8080").await?; + let session = Arc::new(Session::new(stream).await?); + + // Spawn read loop + let read_session = Arc::clone(&session); + tokio::spawn(async move { + loop { + match read_session.read_frame().await { + Ok(Some((opcode, payload))) => { + if opcode == 0x1 { + let text = String::from_utf8(payload).unwrap_or_default(); + println!("Server says: {}", text); + } + } + Ok(None) => {} + Err(_) => break, + } + } + }); + + // Send a few messages + for i in 0..5 { + let msg = serde_json::json!({ "hello": i }); + session.send(&msg).await?; + tokio::time::sleep(std::time::Duration::from_secs(1)).await; + } + + session.close().await?; + Ok(()) +} diff --git a/examples/server.rs b/examples/server.rs index ab9323e..35b294c 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1,2 +1,62 @@ +use std::sync::Arc; +use tokio::net::TcpListener; + +use session_rs::session::Session; + #[tokio::main(flavor = "current_thread")] -async fn main() {} +async fn main() -> session_rs::Result<()> { + let listener = TcpListener::bind("127.0.0.1:8080").await?; + println!("Server listening on ws://127.0.0.1:8080"); + + loop { + let (stream, addr) = listener.accept().await?; + println!("New connection: {}", addr); + + tokio::spawn(async move { + // Wrap session in Arc so tasks can share it + let session = match Session::new(stream).await { + Ok(s) => Arc::new(s), + Err(e) => { + eprintln!("Handshake failed: {:?}", e); + return; + } + }; + + // Simple ping loop + let ping_session = Arc::clone(&session); + tokio::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_secs(15)); + loop { + interval.tick().await; + if ping_session.send_ping().await.is_err() { + break; + } + } + }); + + // Read loop + loop { + match session.read_frame().await { + Ok(Some((opcode, payload))) => { + if opcode == 0x1 { + // Text frame → parse JSON if possible + let text = String::from_utf8(payload).unwrap_or_default(); + println!("Received text: {}", text); + + // Echo back + if let Err(e) = session.send(&serde_json::json!({"echo": text})).await { + eprintln!("Send error: {:?}", e); + break; + } + } + } + Ok(None) => {} + Err(_) => break, // connection closed + } + } + + let _ = session.close().await; + println!("Connection {} closed", addr); + }); + } +} diff --git a/src/lib.rs b/src/lib.rs index 935aa31..64c4fe1 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -9,6 +9,7 @@ pub enum SessionFrame { pub type Result = std::result::Result; +#[derive(Debug)] pub enum Error { Io(std::io::Error), Json(serde_json::Error), From 9ea2978c8ddca7edf49341076249a3bd3750fd14 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 22:01:45 +0100 Subject: [PATCH 08/27] Working client --- examples/client.rs | 4 +- examples/server.rs | 2 +- src/handshake.rs | 92 ++++++++++++++++++++++++++++++++++++++++++++++ src/lib.rs | 1 + src/session.rs | 20 ++-------- 5 files changed, 98 insertions(+), 21 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 5d09db5..853832b 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,12 +1,10 @@ use std::sync::Arc; -use tokio::net::TcpStream; use session_rs::session::Session; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { - let stream = TcpStream::connect("127.0.0.1:8080").await?; - let session = Arc::new(Session::new(stream).await?); + let session = Arc::new(Session::new_server("127.0.0.1:8080", "/").await?); // Spawn read loop let read_session = Arc::clone(&session); diff --git a/examples/server.rs b/examples/server.rs index 35b294c..5ab6336 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -14,7 +14,7 @@ async fn main() -> session_rs::Result<()> { tokio::spawn(async move { // Wrap session in Arc so tasks can share it - let session = match Session::new(stream).await { + let session = match Session::new_client(stream).await { Ok(s) => Arc::new(s), Err(e) => { eprintln!("Handshake failed: {:?}", e); diff --git a/src/handshake.rs b/src/handshake.rs index a3c98e8..56edc63 100644 --- a/src/handshake.rs +++ b/src/handshake.rs @@ -1,8 +1,13 @@ +use sha1::{Digest, Sha1}; +use std::sync::Arc; use tokio::{ io::{AsyncBufReadExt, AsyncWriteExt, BufReader}, net::TcpStream, + sync::Mutex, }; +use crate::session::Session; + 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); @@ -75,3 +80,90 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu write_half.write_all(response.as_bytes()).await?; Ok(()) } + +impl Session { + pub async fn new_client(mut stream: TcpStream) -> crate::Result { + crate::handshake::handle_websocket_handshake(&mut stream).await?; + + let (read, write) = stream.into_split(); + + Ok(Self { + reader: Arc::new(Mutex::new(read)), + writer: Arc::new(Mutex::new(write)), + id: rand::random(), + }) + } + + /// Connect to a WebSocket server and perform the handshake + pub async fn new_server(addr: &str, path: &str) -> crate::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::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(); + reader.read_line(&mut status_line).await?; + if !status_line.starts_with("HTTP/1.1 101") { + return Err(crate::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::encode(sha1.finalize()) + }; + if sec_accept.as_deref() != Some(expected.as_str()) { + return Err(crate::Error::HandshakeFailed( + "Sec-WebSocket-Accept mismatch".into(), + )); + } + + // 6. Upgrade succeeded, split stream + let (read, write) = stream.into_split(); + + Ok(Self { + reader: Arc::new(Mutex::new(read)), + writer: Arc::new(Mutex::new(write)), + id: rand::random(), + }) + } +} diff --git a/src/lib.rs b/src/lib.rs index 64c4fe1..0af8d9a 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -14,6 +14,7 @@ pub enum Error { Io(std::io::Error), Json(serde_json::Error), InvalidFrame(String), + HandshakeFailed(String), ConnectionClosed, } diff --git a/src/session.rs b/src/session.rs index 333af59..732f745 100644 --- a/src/session.rs +++ b/src/session.rs @@ -9,23 +9,9 @@ use tokio::{ }; pub struct Session { - reader: Arc>, - writer: Arc>, - id: u64, -} - -impl Session { - pub async fn new(mut stream: TcpStream) -> crate::Result { - crate::handshake::handle_websocket_handshake(&mut stream).await?; - - let (read, write) = stream.into_split(); - - Ok(Self { - reader: Arc::new(Mutex::new(read)), - writer: Arc::new(Mutex::new(write)), - id: rand::random(), - }) - } + pub(crate) reader: Arc>, + pub(crate) writer: Arc>, + pub(crate) id: u64, } impl Clone for Session { From 5edbe8c8a19799fa56951aaaaddabe08f4926995 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 22:15:37 +0100 Subject: [PATCH 09/27] Fixing issues --- examples/client.rs | 4 +-- examples/server.rs | 7 ++++-- src/handshake.rs | 10 +++++--- src/session.rs | 63 ++++++++++++++++++++++++---------------------- 4 files changed, 46 insertions(+), 38 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 853832b..14ed2ab 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -4,7 +4,7 @@ use session_rs::session::Session; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { - let session = Arc::new(Session::new_server("127.0.0.1:8080", "/").await?); + let session = Arc::new(Session::connect("127.0.0.1:8080", "/").await?); // Spawn read loop let read_session = Arc::clone(&session); @@ -27,7 +27,7 @@ async fn main() -> session_rs::Result<()> { for i in 0..5 { let msg = serde_json::json!({ "hello": i }); session.send(&msg).await?; - tokio::time::sleep(std::time::Duration::from_secs(1)).await; + tokio::time::sleep(std::time::Duration::from_secs(1000)).await; } session.close().await?; diff --git a/examples/server.rs b/examples/server.rs index 5ab6336..f57d7ab 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -14,7 +14,7 @@ async fn main() -> session_rs::Result<()> { tokio::spawn(async move { // Wrap session in Arc so tasks can share it - let session = match Session::new_client(stream).await { + let session = match Session::handshake(stream).await { Ok(s) => Arc::new(s), Err(e) => { eprintln!("Handshake failed: {:?}", e); @@ -51,7 +51,10 @@ async fn main() -> session_rs::Result<()> { } } Ok(None) => {} - Err(_) => break, // connection closed + Err(e) => { + eprintln!("{e:?}"); + break; + } } } diff --git a/src/handshake.rs b/src/handshake.rs index 56edc63..bfb024b 100644 --- a/src/handshake.rs +++ b/src/handshake.rs @@ -82,20 +82,21 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu } impl Session { - pub async fn new_client(mut stream: TcpStream) -> crate::Result { + pub async fn handshake(mut stream: TcpStream) -> crate::Result { crate::handshake::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)), - id: rand::random(), + mask_payload: false, }) } /// Connect to a WebSocket server and perform the handshake - pub async fn new_server(addr: &str, path: &str) -> crate::Result { + pub async fn connect(addr: &str, path: &str) -> crate::Result { // 1. TCP connect let mut stream = TcpStream::connect(addr).await?; @@ -161,9 +162,10 @@ impl Session { let (read, write) = stream.into_split(); Ok(Self { + id: rand::random(), reader: Arc::new(Mutex::new(read)), writer: Arc::new(Mutex::new(write)), - id: rand::random(), + mask_payload: true, }) } } diff --git a/src/session.rs b/src/session.rs index 732f745..710db5f 100644 --- a/src/session.rs +++ b/src/session.rs @@ -4,7 +4,6 @@ use std::{ }; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, - net::TcpStream, sync::Mutex, }; @@ -12,6 +11,7 @@ pub struct Session { pub(crate) reader: Arc>, pub(crate) writer: Arc>, pub(crate) id: u64, + pub(crate) mask_payload: bool, } impl Clone for Session { @@ -19,6 +19,7 @@ impl Clone for Session { Session { reader: self.reader.clone(), writer: self.writer.clone(), + mask_payload: self.mask_payload.clone(), id: self.id, } } @@ -43,29 +44,46 @@ impl Session { let mut writer = self.writer.lock().await; let mut header = Vec::with_capacity(10); - header.push(0x80 | opcode); + let mask_bit = if self.mask_payload { 0x80 } else { 0x00 }; + header.push(0x80 | opcode); // FIN + opcode let len = payload.len(); if len < 126 { - header.push(len as u8); + header.push((len as u8) | mask_bit); } else if len <= 0xFFFF { - header.push(126); + header.push(126 | mask_bit); header.extend_from_slice(&(len as u16).to_be_bytes()); } else { - header.push(127); + header.push(127 | mask_bit); header.extend_from_slice(&(len as u64).to_be_bytes()); } - writer.write_all(&header).await?; - writer.write_all(payload).await?; + if self.mask_payload { + // 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 Session { pub async fn send(&self, msg: &T) -> crate::Result<()> { - let payload = serde_json::to_vec(msg)?; - self.send_frame(0x1, &payload).await + self.send_frame(0x1, &serde_json::to_vec(msg)?).await } pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { @@ -73,19 +91,15 @@ impl Session { } pub async fn send_ping(&self) -> crate::Result<()> { - let mut writer = self.writer.lock().await; - // FIN + opcode = 0x89 (ping), payload length = 0 - writer.write_all(&[0x89, 0x00]).await?; - writer.flush().await?; - Ok(()) + self.send_frame(0x9, &[]).await } pub async fn send_pong(&self) -> crate::Result<()> { - let mut writer = self.writer.lock().await; - // FIN + opcode = 0x8A (pong), payload length = 0 - writer.write_all(&[0x8A, 0x00]).await?; - writer.flush().await?; - Ok(()) + self.send_frame(0xA, &[]).await + } + + pub async fn close(&self) -> crate::Result<()> { + self.send_frame(0x8, &[]).await } pub fn start_ping(self: Arc) { @@ -99,17 +113,6 @@ impl Session { } }); } - - /// Send a close frame and flush - pub async fn close(&self) -> crate::Result<()> { - let mut writer = self.writer.lock().await; - - // FIN=1, opcode=0x8 (close), payload length=0 - let frame = [0x88, 0x00]; - writer.write_all(&frame).await?; - writer.flush().await?; - Ok(()) - } } impl Session { From a14a989691e8e4e73489bf20b1e9c79809141ac4 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 22:55:58 +0100 Subject: [PATCH 10/27] Fixed pinging --- examples/client.rs | 3 ++- examples/server.rs | 12 +----------- src/session.rs | 5 +++-- 3 files changed, 6 insertions(+), 14 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 14ed2ab..112bf1e 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -25,9 +25,10 @@ async fn main() -> session_rs::Result<()> { // Send a few messages for i in 0..5 { + println!("sending"); let msg = serde_json::json!({ "hello": i }); session.send(&msg).await?; - tokio::time::sleep(std::time::Duration::from_secs(1000)).await; + tokio::time::sleep(std::time::Duration::from_secs(1)).await; } session.close().await?; diff --git a/examples/server.rs b/examples/server.rs index f57d7ab..234e722 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -22,17 +22,7 @@ async fn main() -> session_rs::Result<()> { } }; - // Simple ping loop - let ping_session = Arc::clone(&session); - tokio::spawn(async move { - let mut interval = tokio::time::interval(std::time::Duration::from_secs(15)); - loop { - interval.tick().await; - if ping_session.send_ping().await.is_err() { - break; - } - } - }); + session.start_ping_loop(); // Read loop loop { diff --git a/src/session.rs b/src/session.rs index 710db5f..8a49da1 100644 --- a/src/session.rs +++ b/src/session.rs @@ -102,12 +102,13 @@ impl Session { self.send_frame(0x8, &[]).await } - pub fn start_ping(self: Arc) { + 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 self.send_ping().await.is_err() { + if s.send_ping().await.is_err() { break; } } From ce1e0a05d4f534673564ffaafa55bef28d2a7f65 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 23:14:08 +0100 Subject: [PATCH 11/27] Fixed mask --- examples/server.rs | 3 --- src/session.rs | 2 +- 2 files changed, 1 insertion(+), 4 deletions(-) diff --git a/examples/server.rs b/examples/server.rs index 234e722..d603fed 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -47,9 +47,6 @@ async fn main() -> session_rs::Result<()> { } } } - - let _ = session.close().await; - println!("Connection {} closed", addr); }); } } diff --git a/src/session.rs b/src/session.rs index 8a49da1..969ad35 100644 --- a/src/session.rs +++ b/src/session.rs @@ -143,7 +143,7 @@ impl Session { } // --- 3. Read mask key --- - if !masked { + if !masked && !self.mask_payload { // Per spec, client-to-server frames MUST be masked self.close().await.ok(); return Err(crate::Error::InvalidFrame( From 7940867c6a836f6580f2874eebb5d839f01d0b04 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 23:43:37 +0100 Subject: [PATCH 12/27] Refactoring --- src/lib.rs | 3 +++ src/session.rs | 51 +++++++++++++++++++++++++++++--------------------- 2 files changed, 33 insertions(+), 21 deletions(-) diff --git a/src/lib.rs b/src/lib.rs index 0af8d9a..96c8850 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -5,6 +5,9 @@ pub mod session; pub enum SessionFrame { Typed(T), Binary(Vec), + Ping, + Pong, + Close } pub type Result = std::result::Result; diff --git a/src/session.rs b/src/session.rs index 969ad35..6d96b2e 100644 --- a/src/session.rs +++ b/src/session.rs @@ -7,6 +7,8 @@ use tokio::{ sync::Mutex, }; +use crate::SessionFrame; + pub struct Session { pub(crate) reader: Arc>, pub(crate) writer: Arc>, @@ -119,14 +121,14 @@ impl Session { impl Session { /// Read a full WebSocket frame (handling masking and control frames) /// Returns (opcode, payload) - pub async fn read_frame(&self) -> crate::Result)>> { + pub async fn read_frame(&self) -> crate::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 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; @@ -163,35 +165,42 @@ impl Session { } } - // --- 5. Handle control frames immediately --- + // --- 6. Return opcode + payload --- + Ok((fin, opcode, payload)) + } + + pub async fn read(&self) -> crate::Result> { + let (fin, opcode, payload) = self.read_frame().await?; + match opcode { + // Close 0x8 => { - // Close self.close().await.ok(); - return Err(crate::Error::ConnectionClosed); + Ok(SessionFrame::Close) } + + // Ping 0x9 => { - // Ping self.send_pong().await.ok(); - return Ok(None); - } - 0xA => { - // Pong, ignore - return Ok(None); - } - 0x0 | 0x1 | 0x2 => { - // Continuation / Text / Binary → valid payload + Ok(SessionFrame::Ping) } + + // Pong, ignore + 0xA => Ok(SessionFrame::Pong), + + // Continuation / Text / Binary → valid payload + 0x0 => Ok(None), + + 0x1 => Ok(None), + + 0x2 => Ok(None), + _ => { self.close().await.ok(); - return Err(crate::Error::InvalidFrame(format!( - "Unknown opcode: {}", - opcode - ))); + Err(crate::Error::InvalidFrame(format!( + "Unknown opcode: {opcode}" + ))) } } - - // --- 6. Return opcode + payload --- - Ok(Some((opcode, payload))) } } From 77c9b20e6ba91ab0aabedb65007cf85d278fd0fd Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 23:51:54 +0100 Subject: [PATCH 13/27] Refactoring to ws --- examples/client.rs | 6 +- examples/server.rs | 4 +- src/lib.rs | 2 +- src/session.rs | 206 ------------------------------------ src/{ => ws}/handshake.rs | 6 +- src/ws/mod.rs | 212 ++++++++++++++++++++++++++++++++++++++ 6 files changed, 221 insertions(+), 215 deletions(-) rename src/{ => ws}/handshake.rs (97%) create mode 100644 src/ws/mod.rs diff --git a/examples/client.rs b/examples/client.rs index 112bf1e..515158f 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,16 +1,16 @@ use std::sync::Arc; -use session_rs::session::Session; +use session_rs::ws::WebSocket; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { - let session = Arc::new(Session::connect("127.0.0.1:8080", "/").await?); + let session = Arc::new(WebSocket::connect("127.0.0.1:8080", "/").await?); // Spawn read loop let read_session = Arc::clone(&session); tokio::spawn(async move { loop { - match read_session.read_frame().await { + match read_session.read().await { Ok(Some((opcode, payload))) => { if opcode == 0x1 { let text = String::from_utf8(payload).unwrap_or_default(); diff --git a/examples/server.rs b/examples/server.rs index d603fed..3b78c5f 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1,7 +1,7 @@ use std::sync::Arc; use tokio::net::TcpListener; -use session_rs::session::Session; +use session_rs::session::WebSocket; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -14,7 +14,7 @@ async fn main() -> session_rs::Result<()> { tokio::spawn(async move { // Wrap session in Arc so tasks can share it - let session = match Session::handshake(stream).await { + let session = match WebSocket::handshake(stream).await { Ok(s) => Arc::new(s), Err(e) => { eprintln!("Handshake failed: {:?}", e); diff --git a/src/lib.rs b/src/lib.rs index 96c8850..208c505 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,6 +1,6 @@ -pub mod handshake; pub mod server; pub mod session; +pub mod ws; pub enum SessionFrame { Typed(T), diff --git a/src/session.rs b/src/session.rs index 6d96b2e..e69de29 100644 --- a/src/session.rs +++ b/src/session.rs @@ -1,206 +0,0 @@ -use std::{ - hash::{Hash, Hasher}, - sync::Arc, -}; -use tokio::{ - io::{AsyncReadExt, AsyncWriteExt}, - sync::Mutex, -}; - -use crate::SessionFrame; - -pub struct Session { - pub(crate) reader: Arc>, - pub(crate) writer: Arc>, - pub(crate) id: u64, - pub(crate) mask_payload: bool, -} - -impl Clone for Session { - fn clone(&self) -> Self { - Session { - reader: self.reader.clone(), - writer: self.writer.clone(), - mask_payload: self.mask_payload.clone(), - id: self.id, - } - } -} - -impl PartialEq for Session { - fn eq(&self, other: &Self) -> bool { - self.id == other.id - } -} - -impl Eq for Session {} - -impl Hash for Session { - fn hash(&self, state: &mut H) { - self.id.hash(state); - } -} - -impl Session { - async fn send_frame(&self, opcode: u8, payload: &[u8]) -> crate::Result<()> { - let mut writer = self.writer.lock().await; - - let mut header = Vec::with_capacity(10); - let mask_bit = if self.mask_payload { 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.mask_payload { - // 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 Session { - pub async fn send(&self, msg: &T) -> crate::Result<()> { - self.send_frame(0x1, &serde_json::to_vec(msg)?).await - } - - pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { - self.send_frame(0x2, payload).await - } - - pub async fn send_ping(&self) -> crate::Result<()> { - self.send_frame(0x9, &[]).await - } - - pub async fn send_pong(&self) -> crate::Result<()> { - self.send_frame(0xA, &[]).await - } - - pub async fn close(&self) -> crate::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 Session { - /// Read a full WebSocket frame (handling masking and control frames) - /// Returns (opcode, payload) - pub async fn read_frame(&self) -> crate::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); - } - - // --- 3. Read mask key --- - if !masked && !self.mask_payload { - // Per spec, client-to-server frames MUST be masked - self.close().await.ok(); - return Err(crate::Error::InvalidFrame( - "Received unmasked frame from client".into(), - )); - } - - let mut mask = [0u8; 4]; - reader.read_exact(&mut mask).await?; - - // --- 4. Read payload --- - 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]; - } - } - - // --- 6. Return opcode + payload --- - Ok((fin, opcode, payload)) - } - - pub async fn read(&self) -> crate::Result> { - let (fin, opcode, payload) = self.read_frame().await?; - - match opcode { - // Close - 0x8 => { - self.close().await.ok(); - Ok(SessionFrame::Close) - } - - // Ping - 0x9 => { - self.send_pong().await.ok(); - Ok(SessionFrame::Ping) - } - - // Pong, ignore - 0xA => Ok(SessionFrame::Pong), - - // Continuation / Text / Binary → valid payload - 0x0 => Ok(None), - - 0x1 => Ok(None), - - 0x2 => Ok(None), - - _ => { - self.close().await.ok(); - Err(crate::Error::InvalidFrame(format!( - "Unknown opcode: {opcode}" - ))) - } - } - } -} diff --git a/src/handshake.rs b/src/ws/handshake.rs similarity index 97% rename from src/handshake.rs rename to src/ws/handshake.rs index bfb024b..ec36903 100644 --- a/src/handshake.rs +++ b/src/ws/handshake.rs @@ -6,7 +6,7 @@ use tokio::{ sync::Mutex, }; -use crate::session::Session; +use super::WebSocket; pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> { let (read_half, mut write_half) = stream.split(); @@ -81,9 +81,9 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu Ok(()) } -impl Session { +impl WebSocket { pub async fn handshake(mut stream: TcpStream) -> crate::Result { - crate::handshake::handle_websocket_handshake(&mut stream).await?; + handle_websocket_handshake(&mut stream).await?; let (read, write) = stream.into_split(); diff --git a/src/ws/mod.rs b/src/ws/mod.rs new file mode 100644 index 0000000..c42df73 --- /dev/null +++ b/src/ws/mod.rs @@ -0,0 +1,212 @@ +pub mod handshake; + +use std::{ + hash::{Hash, Hasher}, + sync::Arc, +}; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + sync::Mutex, +}; + +use crate::SessionFrame; + +pub struct WebSocket { + pub(crate) reader: Arc>, + pub(crate) writer: Arc>, + pub(crate) id: u64, + pub(crate) mask_payload: bool, +} + +impl Clone for WebSocket { + fn clone(&self) -> Self { + WebSocket { + reader: self.reader.clone(), + writer: self.writer.clone(), + mask_payload: self.mask_payload.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]) -> crate::Result<()> { + let mut writer = self.writer.lock().await; + + let mut header = Vec::with_capacity(10); + let mask_bit = if self.mask_payload { 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.mask_payload { + // 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: &T) -> crate::Result<()> { + self.send_frame(0x1, &serde_json::to_vec(msg)?).await + } + + pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { + self.send_frame(0x2, payload).await + } + + pub async fn send_ping(&self) -> crate::Result<()> { + self.send_frame(0x9, &[]).await + } + + pub async fn send_pong(&self) -> crate::Result<()> { + self.send_frame(0xA, &[]).await + } + + pub async fn close(&self) -> crate::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) -> crate::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); + } + + // --- 3. Read mask key --- + if !masked && !self.mask_payload { + // Per spec, client-to-server frames MUST be masked + self.close().await.ok(); + return Err(crate::Error::InvalidFrame( + "Received unmasked frame from client".into(), + )); + } + + let mut mask = [0u8; 4]; + reader.read_exact(&mut mask).await?; + + // --- 4. Read payload --- + 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]; + } + } + + // --- 6. Return opcode + payload --- + Ok((fin, opcode, payload)) + } + + pub async fn read(&self) -> crate::Result> { + let (fin, opcode, payload) = self.read_frame().await?; + + match opcode { + // Close + 0x8 => { + self.close().await.ok(); + Ok(SessionFrame::Close) + } + + // Ping + 0x9 => { + self.send_pong().await.ok(); + Ok(SessionFrame::Ping) + } + + // Pong + 0xA => Ok(SessionFrame::Pong), + + // Continuation + 0x0 => Ok(SessionFrame::Pong), + + // Text + // 0x1 => { + + // }, + + // Binary + 0x2 => Ok(SessionFrame::Pong), + + _ => { + self.close().await.ok(); + Err(crate::Error::InvalidFrame(format!( + "Unknown opcode: {opcode}" + ))) + } + } + } +} From c5dd9f019de582a9590adbdf5639f5812f1e33b1 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 00:09:55 +0100 Subject: [PATCH 14/27] Working WebSocket --- examples/client.rs | 11 ++++------- examples/server.rs | 22 +++++++++------------- src/lib.rs | 15 ++++++++++++--- src/ws/mod.rs | 41 ++++++++++++++++++++++++++++++++--------- 4 files changed, 57 insertions(+), 32 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 515158f..0cbb402 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,6 +1,6 @@ use std::sync::Arc; -use session_rs::ws::WebSocket; +use session_rs::{SessionFrame, ws::WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -11,13 +11,10 @@ async fn main() -> session_rs::Result<()> { tokio::spawn(async move { loop { match read_session.read().await { - Ok(Some((opcode, payload))) => { - if opcode == 0x1 { - let text = String::from_utf8(payload).unwrap_or_default(); - println!("Server says: {}", text); - } + Ok(SessionFrame::Text(text)) => { + println!("Server says: {}", text); } - Ok(None) => {} + Ok(_) => {} Err(_) => break, } } diff --git a/examples/server.rs b/examples/server.rs index 3b78c5f..edf75b2 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1,7 +1,7 @@ use std::sync::Arc; use tokio::net::TcpListener; -use session_rs::session::WebSocket; +use session_rs::{SessionFrame, ws::WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -26,21 +26,17 @@ async fn main() -> session_rs::Result<()> { // Read loop loop { - match session.read_frame().await { - Ok(Some((opcode, payload))) => { - if opcode == 0x1 { - // Text frame → parse JSON if possible - let text = String::from_utf8(payload).unwrap_or_default(); - println!("Received text: {}", text); + match session.read().await { + Ok(SessionFrame::Text(text)) => { + println!("Received text: {}", text); - // Echo back - if let Err(e) = session.send(&serde_json::json!({"echo": text})).await { - eprintln!("Send error: {:?}", e); - break; - } + // Echo back + if let Err(e) = session.send(&serde_json::json!({"echo": text})).await { + eprintln!("Send error: {:?}", e); + break; } } - Ok(None) => {} + Ok(_) => {} Err(e) => { eprintln!("{e:?}"); break; diff --git a/src/lib.rs b/src/lib.rs index 208c505..7308492 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,13 +1,15 @@ +use std::string::FromUtf8Error; + pub mod server; pub mod session; pub mod ws; -pub enum SessionFrame { - Typed(T), +pub enum SessionFrame { + Text(String), Binary(Vec), Ping, Pong, - Close + Close, } pub type Result = std::result::Result; @@ -19,6 +21,7 @@ pub enum Error { InvalidFrame(String), HandshakeFailed(String), ConnectionClosed, + Utf8(FromUtf8Error), } impl From for Error { @@ -32,3 +35,9 @@ impl From for Error { Self::Json(value) } } + +impl From for Error { + fn from(value: FromUtf8Error) -> Self { + Self::Utf8(value) + } +} diff --git a/src/ws/mod.rs b/src/ws/mod.rs index c42df73..3f87e19 100644 --- a/src/ws/mod.rs +++ b/src/ws/mod.rs @@ -171,8 +171,36 @@ impl WebSocket { Ok((fin, opcode, payload)) } - pub async fn read(&self) -> crate::Result> { - let (fin, opcode, payload) = self.read_frame().await?; + pub async fn read(&self) -> crate::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(crate::Error::InvalidFrame(format!( + "Unknown opcode: {opcode}" + ))); + } + } + } + } match opcode { // Close @@ -190,16 +218,11 @@ impl WebSocket { // Pong 0xA => Ok(SessionFrame::Pong), - // Continuation - 0x0 => Ok(SessionFrame::Pong), - // Text - // 0x1 => { - - // }, + 0x1 => Ok(SessionFrame::Text(String::from_utf8(payload)?)), // Binary - 0x2 => Ok(SessionFrame::Pong), + 0x2 => Ok(SessionFrame::Binary(payload)), _ => { self.close().await.ok(); From 8d672528003067f339c5ace417328b149cd0ee29 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 00:23:47 +0100 Subject: [PATCH 15/27] Refactored types --- examples/client.rs | 19 +++++++++--------- examples/server.rs | 6 +++--- src/lib.rs | 29 ++++++++------------------- src/ws/error.rs | 24 ++++++++++++++++++++++ src/ws/handshake.rs | 8 ++++---- src/ws/mod.rs | 49 +++++++++++++++++++++++++-------------------- 6 files changed, 76 insertions(+), 59 deletions(-) create mode 100644 src/ws/error.rs diff --git a/examples/client.rs b/examples/client.rs index 0cbb402..f743c59 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,6 +1,6 @@ use std::sync::Arc; -use session_rs::{SessionFrame, ws::WebSocket}; +use session_rs::ws::{Frame, WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -10,13 +10,14 @@ async fn main() -> session_rs::Result<()> { let read_session = Arc::clone(&session); tokio::spawn(async move { loop { - match read_session.read().await { - Ok(SessionFrame::Text(text)) => { - println!("Server says: {}", text); - } - Ok(_) => {} - Err(_) => break, - } + println!("{:?}", read_session.read().await); + // match read_session.read().await { + // Ok(Frame::Text(text)) => { + // println!("Server says: {}", text); + // } + // Ok(_) => {} + // Err(_) => break, + // } } }); @@ -24,7 +25,7 @@ async fn main() -> session_rs::Result<()> { for i in 0..5 { println!("sending"); let msg = serde_json::json!({ "hello": i }); - session.send(&msg).await?; + session.send(&msg.to_string()).await?; tokio::time::sleep(std::time::Duration::from_secs(1)).await; } diff --git a/examples/server.rs b/examples/server.rs index edf75b2..ce36cc8 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1,7 +1,7 @@ use std::sync::Arc; use tokio::net::TcpListener; -use session_rs::{SessionFrame, ws::WebSocket}; +use session_rs::ws::{Frame, WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -27,11 +27,11 @@ async fn main() -> session_rs::Result<()> { // Read loop loop { match session.read().await { - Ok(SessionFrame::Text(text)) => { + Ok(Frame::Text(text)) => { println!("Received text: {}", text); // Echo back - if let Err(e) = session.send(&serde_json::json!({"echo": text})).await { + if let Err(e) = session.send(&serde_json::json!({"echo": text}).to_string()).await { eprintln!("Send error: {:?}", e); break; } diff --git a/src/lib.rs b/src/lib.rs index 7308492..627be42 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,27 +1,20 @@ -use std::string::FromUtf8Error; - pub mod server; pub mod session; pub mod ws; -pub enum SessionFrame { - Text(String), - Binary(Vec), - Ping, - Pong, - Close, -} - pub type Result = std::result::Result; #[derive(Debug)] pub enum Error { - Io(std::io::Error), + WebSocket(ws::Error), Json(serde_json::Error), - InvalidFrame(String), - HandshakeFailed(String), - ConnectionClosed, - Utf8(FromUtf8Error), + Io(std::io::Error), +} + +impl From for Error { + fn from(value: ws::Error) -> Self { + Self::WebSocket(value) + } } impl From for Error { @@ -35,9 +28,3 @@ impl From for Error { Self::Json(value) } } - -impl From for Error { - fn from(value: FromUtf8Error) -> Self { - Self::Utf8(value) - } -} diff --git a/src/ws/error.rs b/src/ws/error.rs new file mode 100644 index 0000000..9aa2299 --- /dev/null +++ b/src/ws/error.rs @@ -0,0 +1,24 @@ +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, +} + +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) + } +} diff --git a/src/ws/handshake.rs b/src/ws/handshake.rs index ec36903..4f92ec6 100644 --- a/src/ws/handshake.rs +++ b/src/ws/handshake.rs @@ -82,7 +82,7 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu } impl WebSocket { - pub async fn handshake(mut stream: TcpStream) -> crate::Result { + pub async fn handshake(mut stream: TcpStream) -> super::Result { handle_websocket_handshake(&mut stream).await?; let (read, write) = stream.into_split(); @@ -96,7 +96,7 @@ impl WebSocket { } /// Connect to a WebSocket server and perform the handshake - pub async fn connect(addr: &str, path: &str) -> crate::Result { + pub async fn connect(addr: &str, path: &str) -> super::Result { // 1. TCP connect let mut stream = TcpStream::connect(addr).await?; @@ -123,7 +123,7 @@ impl WebSocket { let mut status_line = String::new(); reader.read_line(&mut status_line).await?; if !status_line.starts_with("HTTP/1.1 101") { - return Err(crate::Error::HandshakeFailed(format!( + return Err(super::Error::HandshakeFailed(format!( "Expected 101 Switching Protocols, got: {}", status_line.trim_end() ))); @@ -153,7 +153,7 @@ impl WebSocket { base64::encode(sha1.finalize()) }; if sec_accept.as_deref() != Some(expected.as_str()) { - return Err(crate::Error::HandshakeFailed( + return Err(super::Error::HandshakeFailed( "Sec-WebSocket-Accept mismatch".into(), )); } diff --git a/src/ws/mod.rs b/src/ws/mod.rs index 3f87e19..a72bf41 100644 --- a/src/ws/mod.rs +++ b/src/ws/mod.rs @@ -1,4 +1,6 @@ +pub mod error; pub mod handshake; +pub use error::{Error, Result}; use std::{ hash::{Hash, Hasher}, @@ -9,7 +11,14 @@ use tokio::{ sync::Mutex, }; -use crate::SessionFrame; +#[derive(Debug, Clone)] +pub enum Frame { + Text(String), + Binary(Vec), + Ping, + Pong, + Close, +} pub struct WebSocket { pub(crate) reader: Arc>, @@ -44,7 +53,7 @@ impl Hash for WebSocket { } impl WebSocket { - async fn send_frame(&self, opcode: u8, payload: &[u8]) -> crate::Result<()> { + async fn send_frame(&self, opcode: u8, payload: &[u8]) -> Result<()> { let mut writer = self.writer.lock().await; let mut header = Vec::with_capacity(10); @@ -86,23 +95,23 @@ impl WebSocket { } impl WebSocket { - pub async fn send(&self, msg: &T) -> crate::Result<()> { - self.send_frame(0x1, &serde_json::to_vec(msg)?).await + pub async fn send(&self, msg: &str) -> Result<()> { + self.send_frame(0x1, msg.as_bytes()).await } - pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { + pub async fn send_bin(&self, payload: &[u8]) -> Result<()> { self.send_frame(0x2, payload).await } - pub async fn send_ping(&self) -> crate::Result<()> { + pub async fn send_ping(&self) -> Result<()> { self.send_frame(0x9, &[]).await } - pub async fn send_pong(&self) -> crate::Result<()> { + pub async fn send_pong(&self) -> Result<()> { self.send_frame(0xA, &[]).await } - pub async fn close(&self) -> crate::Result<()> { + pub async fn close(&self) -> Result<()> { self.send_frame(0x8, &[]).await } @@ -123,7 +132,7 @@ impl WebSocket { impl WebSocket { /// Read a full WebSocket frame (handling masking and control frames) /// Returns (opcode, payload) - pub async fn read_frame(&self) -> crate::Result<(bool, u8, Vec)> { + pub async fn read_frame(&self) -> Result<(bool, u8, Vec)> { let mut reader = self.reader.lock().await; // --- 1. Read first 2-byte header --- @@ -150,7 +159,7 @@ impl WebSocket { if !masked && !self.mask_payload { // Per spec, client-to-server frames MUST be masked self.close().await.ok(); - return Err(crate::Error::InvalidFrame( + return Err(Error::InvalidFrame( "Received unmasked frame from client".into(), )); } @@ -171,7 +180,7 @@ impl WebSocket { Ok((fin, opcode, payload)) } - pub async fn read(&self) -> crate::Result { + pub async fn read(&self) -> Result { let (fin, opcode, mut payload) = self.read_frame().await?; if !fin { @@ -194,9 +203,7 @@ impl WebSocket { 0xA => {} _ => { self.close().await.ok(); - return Err(crate::Error::InvalidFrame(format!( - "Unknown opcode: {opcode}" - ))); + return Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}"))); } } } @@ -206,29 +213,27 @@ impl WebSocket { // Close 0x8 => { self.close().await.ok(); - Ok(SessionFrame::Close) + Ok(Frame::Close) } // Ping 0x9 => { self.send_pong().await.ok(); - Ok(SessionFrame::Ping) + Ok(Frame::Ping) } // Pong - 0xA => Ok(SessionFrame::Pong), + 0xA => Ok(Frame::Pong), // Text - 0x1 => Ok(SessionFrame::Text(String::from_utf8(payload)?)), + 0x1 => Ok(Frame::Text(String::from_utf8(payload)?)), // Binary - 0x2 => Ok(SessionFrame::Binary(payload)), + 0x2 => Ok(Frame::Binary(payload)), _ => { self.close().await.ok(); - Err(crate::Error::InvalidFrame(format!( - "Unknown opcode: {opcode}" - ))) + Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}"))) } } } From 60e5a54261229ed7d71de3b9e86ee7d93c6f42f3 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 00:43:56 +0100 Subject: [PATCH 16/27] Refactored masking --- examples/client.rs | 15 ++++++------- src/ws/handshake.rs | 4 ++-- src/ws/mod.rs | 53 +++++++++++++++++++++++++-------------------- 3 files changed, 39 insertions(+), 33 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index f743c59..84b4abb 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -10,14 +10,13 @@ async fn main() -> session_rs::Result<()> { let read_session = Arc::clone(&session); tokio::spawn(async move { loop { - println!("{:?}", read_session.read().await); - // match read_session.read().await { - // Ok(Frame::Text(text)) => { - // println!("Server says: {}", text); - // } - // Ok(_) => {} - // Err(_) => break, - // } + match read_session.read().await { + Ok(Frame::Text(text)) => { + println!("Server says: {}", text); + } + Ok(_) => {} + Err(_) => break, + } } }); diff --git a/src/ws/handshake.rs b/src/ws/handshake.rs index 4f92ec6..b9bfcb0 100644 --- a/src/ws/handshake.rs +++ b/src/ws/handshake.rs @@ -91,7 +91,7 @@ impl WebSocket { id: rand::random(), reader: Arc::new(Mutex::new(read)), writer: Arc::new(Mutex::new(write)), - mask_payload: false, + is_server: false, }) } @@ -165,7 +165,7 @@ impl WebSocket { id: rand::random(), reader: Arc::new(Mutex::new(read)), writer: Arc::new(Mutex::new(write)), - mask_payload: true, + is_server: true, }) } } diff --git a/src/ws/mod.rs b/src/ws/mod.rs index a72bf41..6f6227a 100644 --- a/src/ws/mod.rs +++ b/src/ws/mod.rs @@ -24,7 +24,7 @@ pub struct WebSocket { pub(crate) reader: Arc>, pub(crate) writer: Arc>, pub(crate) id: u64, - pub(crate) mask_payload: bool, + pub(crate) is_server: bool, } impl Clone for WebSocket { @@ -32,7 +32,7 @@ impl Clone for WebSocket { WebSocket { reader: self.reader.clone(), writer: self.writer.clone(), - mask_payload: self.mask_payload.clone(), + is_server: self.is_server.clone(), id: self.id, } } @@ -57,7 +57,7 @@ impl WebSocket { let mut writer = self.writer.lock().await; let mut header = Vec::with_capacity(10); - let mask_bit = if self.mask_payload { 0x80 } else { 0x00 }; + let mask_bit = if self.is_server { 0x80 } else { 0x00 }; header.push(0x80 | opcode); // FIN + opcode let len = payload.len(); @@ -71,7 +71,7 @@ impl WebSocket { header.extend_from_slice(&(len as u64).to_be_bytes()); } - if self.mask_payload { + if self.is_server { // Generate 4-byte mask key let mask_key: [u8; 4] = rand::random(); header.extend_from_slice(&mask_key); @@ -155,26 +155,33 @@ impl WebSocket { payload_len = u64::from_be_bytes(buf); } - // --- 3. Read mask key --- - if !masked && !self.mask_payload { - // Per spec, client-to-server frames MUST be masked - self.close().await.ok(); - return Err(Error::InvalidFrame( - "Received unmasked frame from client".into(), - )); - } - - let mut mask = [0u8; 4]; - reader.read_exact(&mut mask).await?; - - // --- 4. Read payload --- - 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]; + 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)) From f9b311b9b6631db446905a2f0dad5d311be6919a Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 00:51:51 +0100 Subject: [PATCH 17/27] removed annoying base64 warning --- src/ws/handshake.rs | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/src/ws/handshake.rs b/src/ws/handshake.rs index b9bfcb0..6b347fd 100644 --- a/src/ws/handshake.rs +++ b/src/ws/handshake.rs @@ -1,3 +1,4 @@ +use base64::Engine; use sha1::{Digest, Sha1}; use std::sync::Arc; use tokio::{ @@ -102,7 +103,7 @@ impl WebSocket { // 2. Generate Sec-WebSocket-Key let key_bytes: [u8; 16] = rand::random(); - let key = base64::encode(&key_bytes); + let key = base64::prelude::BASE64_STANDARD.encode(&key_bytes); // 3. Send HTTP Upgrade request let request = format!( @@ -150,7 +151,7 @@ impl WebSocket { let mut sha1 = Sha1::new(); sha1.update(key.as_bytes()); sha1.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - base64::encode(sha1.finalize()) + base64::prelude::BASE64_STANDARD.encode(sha1.finalize()) }; if sec_accept.as_deref() != Some(expected.as_str()) { return Err(super::Error::HandshakeFailed( From f7f1e97671e08889af13f243f4daaa194fb02035 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 01:50:37 +0100 Subject: [PATCH 18/27] Simple session --- Cargo.lock | 1 + Cargo.toml | 2 +- src/session.rs | 90 ++++++++++++++++++++++++++++++++++++++++++++++++++ src/ws/mod.rs | 4 +++ 4 files changed, 96 insertions(+), 1 deletion(-) diff --git a/Cargo.lock b/Cargo.lock index 3b82341..743c953 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -276,6 +276,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" dependencies = [ "serde_core", + "serde_derive", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 7a353f2..70031df 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -6,7 +6,7 @@ edition = "2024" [dependencies] base64 = "0.22.1" rand = "0.10.0" -serde = "1.0.228" +serde = { version = "1.0.228", features = ["serde_derive"] } serde_json = "1.0.149" sha1 = "0.10.6" tokio = { version = "1.49.0", features = ["io-util", "macros", "net", "rt", "sync", "time"] } diff --git a/src/session.rs b/src/session.rs index e69de29..e7bcd9c 100644 --- a/src/session.rs +++ b/src/session.rs @@ -0,0 +1,90 @@ +use std::{marker::PhantomData, sync::Arc}; + +use serde::{Deserialize, Serialize}; +use tokio::sync::Mutex; + +use crate::ws::WebSocket; + +#[derive(Debug, Serialize, Deserialize)] +pub enum SessionMessageKind { + Request, + Response, + Notification, +} + +#[derive(Debug, Serialize, Deserialize)] +pub struct SessionMessage { + id: u32, + kind: SessionMessageKind, + data: T, +} + +pub struct Session< + Req: Serialize + for<'a> Deserialize<'a>, + Res: Serialize + for<'a> Deserialize<'a>, + PeerReq: Serialize + for<'a> Deserialize<'a>, + PeerRes: Serialize + for<'a> Deserialize<'a>, + Notification: Serialize + for<'a> Deserialize<'a>, +> { + _pd: ( + PhantomData, + PhantomData, + PhantomData, + PhantomData, + PhantomData, + ), + pub ws: WebSocket, + id: Arc>, +} + +impl< + Req: Serialize + for<'a> Deserialize<'a>, + Res: Serialize + for<'a> Deserialize<'a>, + PeerReq: Serialize + for<'a> Deserialize<'a>, + PeerRes: Serialize + for<'a> Deserialize<'a>, + Notification: Serialize + for<'a> Deserialize<'a>, +> Session +{ + pub async fn send_id( + &self, + id: u32, + kind: SessionMessageKind, + data: &T, + ) -> crate::Result<()> { + self.ws + .send_text_payload(&serde_json::to_vec(&SessionMessage { id, kind, data })?) + .await?; + + Ok(()) + } + + pub async fn send( + &self, + kind: SessionMessageKind, + data: &T, + ) -> crate::Result<()> { + self.send_id( + { + let mut i = self.id.lock().await; + *i += 1; + *i + }, + kind, + data, + ) + .await + } + + pub async fn request(&self, data: &Req) -> crate::Result<()> { + self.send(SessionMessageKind::Request, data).await + } + + pub async fn respond(&self, to_req: &SessionMessage, data: &Res) -> crate::Result<()> { + self.send_id(to_req.id, SessionMessageKind::Response, data) + .await + } + + pub async fn notify(&self, data: &Res) -> crate::Result<()> { + self.send(SessionMessageKind::Notification, data).await + } +} diff --git a/src/ws/mod.rs b/src/ws/mod.rs index 6f6227a..6a6de66 100644 --- a/src/ws/mod.rs +++ b/src/ws/mod.rs @@ -99,6 +99,10 @@ impl WebSocket { 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 } From 0253f15931e5170d7fd4bd9744e9938cb4db2d41 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 02:26:32 +0100 Subject: [PATCH 19/27] Session types --- examples/client.rs | 44 ++++++++++++++++++------------------- src/session.rs | 54 ++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 22 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 84b4abb..36db0c9 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,32 +1,32 @@ -use std::sync::Arc; +use serde::{Deserialize, Serialize}; +use session_rs::session::Session; -use session_rs::ws::{Frame, WebSocket}; +#[derive(Debug, Serialize, Deserialize)] +struct Data {} + +type Communication = Session; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { - let session = Arc::new(WebSocket::connect("127.0.0.1:8080", "/").await?); + let session = Communication::connect("127.0.0.1:8080", "/").await?; // Spawn read loop - let read_session = Arc::clone(&session); - tokio::spawn(async move { - loop { - match read_session.read().await { - Ok(Frame::Text(text)) => { - println!("Server says: {}", text); - } - Ok(_) => {} - Err(_) => break, - } - } - }); + // tokio::spawn({ + // let session = session.clone(); + // async move { + // loop { + // match session.read().await { + // Ok(Frame::Text(text)) => { + // println!("Server says: {}", text); + // } + // Ok(_) => {} + // Err(_) => break, + // } + // } + // } + // }); - // Send a few messages - for i in 0..5 { - println!("sending"); - let msg = serde_json::json!({ "hello": i }); - session.send(&msg.to_string()).await?; - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - } + session.request(&Data {}).await?; session.close().await?; Ok(()) diff --git a/src/session.rs b/src/session.rs index e7bcd9c..ba0ebac 100644 --- a/src/session.rs +++ b/src/session.rs @@ -37,6 +37,56 @@ pub struct Session< id: Arc>, } +impl< + Req: Serialize + for<'a> Deserialize<'a>, + Res: Serialize + for<'a> Deserialize<'a>, + PeerReq: Serialize + for<'a> Deserialize<'a>, + PeerRes: Serialize + for<'a> Deserialize<'a>, + Notification: Serialize + for<'a> Deserialize<'a>, +> Session +{ + pub fn clone(&self) -> Self { + Self { + _pd: ( + PhantomData, + PhantomData, + PhantomData, + PhantomData, + PhantomData, + ), + ws: self.ws.clone(), + id: self.id.clone(), + } + } +} + +impl< + Req: Serialize + for<'a> Deserialize<'a>, + Res: Serialize + for<'a> Deserialize<'a>, + PeerReq: Serialize + for<'a> Deserialize<'a>, + PeerRes: Serialize + for<'a> Deserialize<'a>, + Notification: Serialize + for<'a> Deserialize<'a>, +> Session +{ + pub fn from_ws(ws: WebSocket) -> Self { + Self { + _pd: ( + PhantomData, + PhantomData, + PhantomData, + PhantomData, + PhantomData, + ), + ws, + id: Arc::new(Mutex::new(0)), + } + } + + pub async fn connect(addr: &str, path: &str) -> crate::Result { + Ok(Self::from_ws(WebSocket::connect(addr, path).await?)) + } +} + impl< Req: Serialize + for<'a> Deserialize<'a>, Res: Serialize + for<'a> Deserialize<'a>, @@ -87,4 +137,8 @@ impl< pub async fn notify(&self, data: &Res) -> crate::Result<()> { self.send(SessionMessageKind::Notification, data).await } + + pub async fn close(&self) -> crate::Result<()> { + Ok(self.ws.close().await?) + } } From 52b77dea83a5f054de2d421cc959930bd11a61bb Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 03:28:37 +0100 Subject: [PATCH 20/27] Working session structure --- examples/client.rs | 43 +++++++------ examples/server.rs | 2 +- src/lib.rs | 8 +++ src/session.rs | 149 ++++++++++++++++----------------------------- 4 files changed, 86 insertions(+), 116 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 36db0c9..4239b5c 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,32 +1,37 @@ use serde::{Deserialize, Serialize}; -use session_rs::session::Session; +use session_rs::{Method, session::Session, ws::Frame}; #[derive(Debug, Serialize, Deserialize)] struct Data {} -type Communication = Session; +impl Method for Data { + const NAME: &'static str = "data"; + type Request = (); + type Response = (); +} #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { - let session = Communication::connect("127.0.0.1:8080", "/").await?; + let session = Session::connect("127.0.0.1:8080", "/").await?; - // Spawn read loop - // tokio::spawn({ - // let session = session.clone(); - // async move { - // loop { - // match session.read().await { - // Ok(Frame::Text(text)) => { - // println!("Server says: {}", text); - // } - // Ok(_) => {} - // Err(_) => break, - // } - // } - // } - // }); + tokio::spawn({ + let session = session.clone(); + async move { + loop { + match session.ws.read().await { + Ok(Frame::Text(text)) => { + println!("Server says: {}", text); + } + Ok(_) => {} + Err(_) => break, + } + } + } + }); - session.request(&Data {}).await?; + session.request::(()).await?; + + tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; session.close().await?; Ok(()) diff --git a/examples/server.rs b/examples/server.rs index ce36cc8..d515a16 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -31,7 +31,7 @@ async fn main() -> session_rs::Result<()> { println!("Received text: {}", text); // Echo back - if let Err(e) = session.send(&serde_json::json!({"echo": text}).to_string()).await { + if let Err(e) = session.send(&text).await { eprintln!("Send error: {:?}", e); break; } diff --git a/src/lib.rs b/src/lib.rs index 627be42..4761d38 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,9 +1,17 @@ +use serde::{Deserialize, Serialize}; + pub mod server; pub mod session; pub mod ws; pub type Result = std::result::Result; +pub trait Method { + const NAME: &'static str; + type Request: Serialize + for<'de> Deserialize<'de>; + type Response: Serialize + for<'de> Deserialize<'de>; +} + #[derive(Debug)] pub enum Error { WebSocket(ws::Error), diff --git a/src/session.rs b/src/session.rs index ba0ebac..0235c95 100644 --- a/src/session.rs +++ b/src/session.rs @@ -1,82 +1,45 @@ -use std::{marker::PhantomData, sync::Arc}; +use std::sync::Arc; use serde::{Deserialize, Serialize}; use tokio::sync::Mutex; -use crate::ws::WebSocket; +use crate::{Method, ws::WebSocket}; #[derive(Debug, Serialize, Deserialize)] -pub enum SessionMessageKind { - Request, - Response, - Notification, +pub enum Message { + Request { + id: u32, + method: String, + data: M::Request, + }, + Response { + id: u32, + error: bool, + result: M::Response, + }, + Notification { + method: String, + data: M::Request, + }, } -#[derive(Debug, Serialize, Deserialize)] -pub struct SessionMessage { - id: u32, - kind: SessionMessageKind, - data: T, -} - -pub struct Session< - Req: Serialize + for<'a> Deserialize<'a>, - Res: Serialize + for<'a> Deserialize<'a>, - PeerReq: Serialize + for<'a> Deserialize<'a>, - PeerRes: Serialize + for<'a> Deserialize<'a>, - Notification: Serialize + for<'a> Deserialize<'a>, -> { - _pd: ( - PhantomData, - PhantomData, - PhantomData, - PhantomData, - PhantomData, - ), +pub struct Session { pub ws: WebSocket, id: Arc>, } -impl< - Req: Serialize + for<'a> Deserialize<'a>, - Res: Serialize + for<'a> Deserialize<'a>, - PeerReq: Serialize + for<'a> Deserialize<'a>, - PeerRes: Serialize + for<'a> Deserialize<'a>, - Notification: Serialize + for<'a> Deserialize<'a>, -> Session -{ +impl Session { pub fn clone(&self) -> Self { Self { - _pd: ( - PhantomData, - PhantomData, - PhantomData, - PhantomData, - PhantomData, - ), ws: self.ws.clone(), id: self.id.clone(), } } } -impl< - Req: Serialize + for<'a> Deserialize<'a>, - Res: Serialize + for<'a> Deserialize<'a>, - PeerReq: Serialize + for<'a> Deserialize<'a>, - PeerRes: Serialize + for<'a> Deserialize<'a>, - Notification: Serialize + for<'a> Deserialize<'a>, -> Session -{ +impl Session { pub fn from_ws(ws: WebSocket) -> Self { Self { - _pd: ( - PhantomData, - PhantomData, - PhantomData, - PhantomData, - PhantomData, - ), ws, id: Arc::new(Mutex::new(0)), } @@ -87,55 +50,49 @@ impl< } } -impl< - Req: Serialize + for<'a> Deserialize<'a>, - Res: Serialize + for<'a> Deserialize<'a>, - PeerReq: Serialize + for<'a> Deserialize<'a>, - PeerRes: Serialize + for<'a> Deserialize<'a>, - Notification: Serialize + for<'a> Deserialize<'a>, -> Session -{ - pub async fn send_id( - &self, - id: u32, - kind: SessionMessageKind, - data: &T, - ) -> crate::Result<()> { +impl Session { + pub async fn send(&self, data: &Message) -> crate::Result<()> { self.ws - .send_text_payload(&serde_json::to_vec(&SessionMessage { id, kind, data })?) + .send_text_payload(&serde_json::to_vec(&data)?) .await?; - Ok(()) } - pub async fn send( - &self, - kind: SessionMessageKind, - data: &T, - ) -> crate::Result<()> { - self.send_id( - { - let mut i = self.id.lock().await; - *i += 1; - *i - }, - kind, - data, - ) + pub async fn use_id(&self) -> u32 { + let mut id = self.id.lock().await; + *id += 1; + *id + } + + pub async fn request(&self, req: M::Request) -> crate::Result<()> { + self.send::(&Message::Request { + id: self.use_id().await, + method: M::NAME.to_string(), + data: req, + }) .await } - pub async fn request(&self, data: &Req) -> crate::Result<()> { - self.send(SessionMessageKind::Request, data).await + pub async fn respond( + &self, + to: u32, + error: bool, + res: M::Response, + ) -> crate::Result<()> { + self.send::(&Message::Response { + id: to, + error, + result: res, + }) + .await } - pub async fn respond(&self, to_req: &SessionMessage, data: &Res) -> crate::Result<()> { - self.send_id(to_req.id, SessionMessageKind::Response, data) - .await - } - - pub async fn notify(&self, data: &Res) -> crate::Result<()> { - self.send(SessionMessageKind::Notification, data).await + pub async fn notify(&self, data: M::Request) -> crate::Result<()> { + self.send::(&Message::Notification { + method: M::NAME.to_string(), + data, + }) + .await } pub async fn close(&self) -> crate::Result<()> { From 095903d3275926ed35c0219be6da74e7180ea608 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 03:31:53 +0100 Subject: [PATCH 21/27] Improved json structure --- src/session.rs | 1 + 1 file changed, 1 insertion(+) diff --git a/src/session.rs b/src/session.rs index 0235c95..07980e5 100644 --- a/src/session.rs +++ b/src/session.rs @@ -6,6 +6,7 @@ use tokio::sync::Mutex; use crate::{Method, ws::WebSocket}; #[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "lowercase", tag = "type")] pub enum Message { Request { id: u32, From 0940731bd19f879ce4d890e0e68b87808718ea56 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 04:05:22 +0100 Subject: [PATCH 22/27] Request receiver --- examples/client.rs | 21 +++++---------------- src/lib.rs | 8 ++++++++ src/session.rs | 46 ++++++++++++++++++++++++++++++++++++++++++++-- 3 files changed, 57 insertions(+), 18 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 4239b5c..3d7564d 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,8 +1,8 @@ use serde::{Deserialize, Serialize}; -use session_rs::{Method, session::Session, ws::Frame}; +use session_rs::{Method, session::Session}; #[derive(Debug, Serialize, Deserialize)] -struct Data {} +struct Data; impl Method for Data { const NAME: &'static str = "data"; @@ -14,23 +14,12 @@ impl Method for Data { async fn main() -> session_rs::Result<()> { let session = Session::connect("127.0.0.1:8080", "/").await?; - tokio::spawn({ - let session = session.clone(); - async move { - loop { - match session.ws.read().await { - Ok(Frame::Text(text)) => { - println!("Server says: {}", text); - } - Ok(_) => {} - Err(_) => break, - } - } - } - }); + session.start_receiver(); session.request::(()).await?; + session.on::(|i, d| println!("Ok {i} {d:?}")).await; + tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; session.close().await?; diff --git a/src/lib.rs b/src/lib.rs index 4761d38..9913bc4 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -12,6 +12,14 @@ pub trait Method { type Response: Serialize + for<'de> Deserialize<'de>; } +pub struct GenericMethod; + +impl Method for GenericMethod { + const NAME: &'static str = "generic_do_not_use"; + type Request = serde_json::Value; + type Response = serde_json::Value; +} + #[derive(Debug)] pub enum Error { WebSocket(ws::Error), diff --git a/src/session.rs b/src/session.rs index 07980e5..8c739b0 100644 --- a/src/session.rs +++ b/src/session.rs @@ -1,9 +1,9 @@ -use std::sync::Arc; +use std::{collections::HashMap, sync::Arc}; use serde::{Deserialize, Serialize}; use tokio::sync::Mutex; -use crate::{Method, ws::WebSocket}; +use crate::{GenericMethod, Method, ws::WebSocket}; #[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "lowercase", tag = "type")] @@ -27,6 +27,7 @@ pub enum Message { pub struct Session { pub ws: WebSocket, id: Arc>, + methods: Arc>>>, } impl Session { @@ -34,6 +35,7 @@ impl Session { Self { ws: self.ws.clone(), id: self.id.clone(), + methods: self.methods.clone(), } } } @@ -43,6 +45,7 @@ impl Session { Self { ws, id: Arc::new(Mutex::new(0)), + methods: Arc::new(Mutex::new(HashMap::new())), } } @@ -51,6 +54,45 @@ impl Session { } } +impl Session { + pub fn start_receiver(&self) { + let s = self.clone(); + tokio::spawn(async move { + loop { + match s.ws.read().await { + Ok(crate::ws::Frame::Text(text)) => { + let Ok(msg) = serde_json::from_str::>(&text) else { + continue; + }; + + match msg { + Message::Request { id, method, data } => { + if let Some(m) = s.methods.lock().await.get(&method) { + (m)(id, data) + } + } + _ => {} + } + } + Ok(_) => {} + Err(_) => break, + } + } + }); + } + + pub async fn on(&self, handler: impl Fn(u32, M::Request) + Send + Sync + 'static) { + self.methods.lock().await.insert( + M::NAME.to_string(), + Box::new(move |id, value| { + if let Ok(req) = serde_json::from_value(value) { + (handler)(id, req) + } + }), + ); + } +} + impl Session { pub async fn send(&self, data: &Message) -> crate::Result<()> { self.ws From bad32bfc3611f5fdb1e0add1d4393ec1f5a68dd8 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 04:15:31 +0100 Subject: [PATCH 23/27] Async handler --- examples/client.rs | 4 +++- src/lib.rs | 6 +++++- src/session.rs | 23 ++++++++++++++++------- 3 files changed, 24 insertions(+), 9 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 3d7564d..9cfbd34 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -18,7 +18,9 @@ async fn main() -> session_rs::Result<()> { session.request::(()).await?; - session.on::(|i, d| println!("Ok {i} {d:?}")).await; + session + .on::(async |i, d| println!("Ok {i} {d:?}")) + .await; tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; diff --git a/src/lib.rs b/src/lib.rs index 9913bc4..0101b66 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,5 @@ +use std::pin::Pin; + use serde::{Deserialize, Serialize}; pub mod server; @@ -5,10 +7,12 @@ pub mod session; pub mod ws; pub type Result = std::result::Result; +pub type BoxFuture<'a> = Pin + Send + 'a>>; +pub type MethodHandler = Box BoxFuture<'static> + Send + Sync>; pub trait Method { const NAME: &'static str; - type Request: Serialize + for<'de> Deserialize<'de>; + type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync; type Response: Serialize + for<'de> Deserialize<'de>; } diff --git a/src/session.rs b/src/session.rs index 8c739b0..2b5ea9b 100644 --- a/src/session.rs +++ b/src/session.rs @@ -3,7 +3,7 @@ use std::{collections::HashMap, sync::Arc}; use serde::{Deserialize, Serialize}; use tokio::sync::Mutex; -use crate::{GenericMethod, Method, ws::WebSocket}; +use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; #[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "lowercase", tag = "type")] @@ -27,7 +27,7 @@ pub enum Message { pub struct Session { pub ws: WebSocket, id: Arc>, - methods: Arc>>>, + methods: Arc>>, } impl Session { @@ -68,7 +68,7 @@ impl Session { match msg { Message::Request { id, method, data } => { if let Some(m) = s.methods.lock().await.get(&method) { - (m)(id, data) + (m)(id, data).await } } _ => {} @@ -81,13 +81,22 @@ impl Session { }); } - pub async fn on(&self, handler: impl Fn(u32, M::Request) + Send + Sync + 'static) { + pub async fn on + Send + 'static>( + &self, + handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static, + ) { + let handler = Arc::new(handler); + self.methods.lock().await.insert( M::NAME.to_string(), Box::new(move |id, value| { - if let Ok(req) = serde_json::from_value(value) { - (handler)(id, req) - } + let handler = Arc::clone(&handler); + + Box::pin(async move { + if let Ok(req) = serde_json::from_value(value) { + handler(id, req).await; + } + }) }), ); } From e09623e8b9b333e91a97f8bac2caea8f15f349b4 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 04:46:08 +0100 Subject: [PATCH 24/27] Respond awaiter --- examples/client.rs | 3 ++- src/lib.rs | 9 +++++++ src/session.rs | 61 ++++++++++++++++++++++++++++++++++------------ 3 files changed, 57 insertions(+), 16 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 9cfbd34..0a7e647 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -8,6 +8,7 @@ impl Method for Data { const NAME: &'static str = "data"; type Request = (); type Response = (); + type Error = (); } #[tokio::main(flavor = "current_thread")] @@ -16,7 +17,7 @@ async fn main() -> session_rs::Result<()> { session.start_receiver(); - session.request::(()).await?; + println!("{:?}", session.request::(()).await?); session .on::(async |i, d| println!("Ok {i} {d:?}")) diff --git a/src/lib.rs b/src/lib.rs index 0101b66..86365df 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -14,6 +14,7 @@ pub trait Method { const NAME: &'static str; type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync; type Response: Serialize + for<'de> Deserialize<'de>; + type Error: Serialize + for<'de> Deserialize<'de>; } pub struct GenericMethod; @@ -22,6 +23,7 @@ impl Method for GenericMethod { const NAME: &'static str = "generic_do_not_use"; type Request = serde_json::Value; type Response = serde_json::Value; + type Error = serde_json::Value; } #[derive(Debug)] @@ -29,6 +31,7 @@ pub enum Error { WebSocket(ws::Error), Json(serde_json::Error), Io(std::io::Error), + RecvError(tokio::sync::broadcast::error::RecvError), } impl From for Error { @@ -48,3 +51,9 @@ impl From for Error { Self::Json(value) } } + +impl From for Error { + fn from(value: tokio::sync::broadcast::error::RecvError) -> Self { + Self::RecvError(value) + } +} diff --git a/src/session.rs b/src/session.rs index 2b5ea9b..7aec8d1 100644 --- a/src/session.rs +++ b/src/session.rs @@ -2,6 +2,7 @@ use std::{collections::HashMap, sync::Arc}; use serde::{Deserialize, Serialize}; use tokio::sync::Mutex; +use tokio::sync::broadcast; use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; @@ -15,9 +16,12 @@ pub enum Message { }, Response { id: u32, - error: bool, result: M::Response, }, + ErrorResponse { + id: u32, + error: M::Error, + }, Notification { method: String, data: M::Request, @@ -28,6 +32,7 @@ pub struct Session { pub ws: WebSocket, id: Arc>, methods: Arc>>, + tx: broadcast::Sender<(u32, bool, serde_json::Value)>, } impl Session { @@ -36,6 +41,7 @@ impl Session { ws: self.ws.clone(), id: self.id.clone(), methods: self.methods.clone(), + tx: self.tx.clone(), } } } @@ -46,6 +52,7 @@ impl Session { ws, id: Arc::new(Mutex::new(0)), methods: Arc::new(Mutex::new(HashMap::new())), + tx: broadcast::channel(8192).0, } } @@ -71,6 +78,12 @@ impl Session { (m)(id, data).await } } + Message::Response { id, result } => { + s.tx.send((id, false, result)).unwrap(); + } + Message::ErrorResponse { id, error } => { + s.tx.send((id, true, error)).unwrap(); + } _ => {} } } @@ -116,27 +129,45 @@ impl Session { *id } - pub async fn request(&self, req: M::Request) -> crate::Result<()> { + pub async fn request( + &self, + req: M::Request, + ) -> crate::Result> { + let id = self.use_id().await; + self.send::(&Message::Request { - id: self.use_id().await, + id, method: M::NAME.to_string(), data: req, }) + .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, res: M::Response) -> crate::Result<()> { + self.send::(&Message::Response { + id: to, + result: res, + }) .await } - pub async fn respond( - &self, - to: u32, - error: bool, - res: M::Response, - ) -> crate::Result<()> { - self.send::(&Message::Response { - id: to, - error, - result: res, - }) - .await + pub async fn respond_error(&self, to: u32, err: M::Error) -> crate::Result<()> { + self.send::(&Message::ErrorResponse { id: to, error: err }) + .await } pub async fn notify(&self, data: M::Request) -> crate::Result<()> { From 9af16feda6d279e562b161d2430d7675c256c542 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 05:05:44 +0100 Subject: [PATCH 25/27] Making progress towoards responses --- examples/client.rs | 6 +++++- examples/server.rs | 48 ++++++++++++++++++---------------------------- src/lib.rs | 2 +- src/session.rs | 23 +++++++++++++++++----- 4 files changed, 43 insertions(+), 36 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 0a7e647..da7847b 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -20,7 +20,11 @@ async fn main() -> session_rs::Result<()> { println!("{:?}", session.request::(()).await?); session - .on::(async |i, d| println!("Ok {i} {d:?}")) + .on::(async |i, d| { + println!("Ok {i} {d:?}"); + + Ok(()) + }) .await; tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; diff --git a/examples/server.rs b/examples/server.rs index d515a16..770f39d 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1,7 +1,17 @@ -use std::sync::Arc; +use serde::{Deserialize, Serialize}; use tokio::net::TcpListener; -use session_rs::ws::{Frame, WebSocket}; +use session_rs::{Method, session::Session, ws::WebSocket}; + +#[derive(Debug, Serialize, Deserialize)] +struct Data; + +impl Method for Data { + const NAME: &'static str = "data"; + type Request = (); + type Response = (); + type Error = (); +} #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -14,35 +24,15 @@ async fn main() -> session_rs::Result<()> { tokio::spawn(async move { // Wrap session in Arc so tasks can share it - let session = match WebSocket::handshake(stream).await { - Ok(s) => Arc::new(s), - Err(e) => { - eprintln!("Handshake failed: {:?}", e); - return; - } - }; + let session = Session::from_ws( + WebSocket::handshake(stream) + .await + .expect("Failed to initialize websocket"), + ); - session.start_ping_loop(); + session.start_receiver(); - // Read loop - loop { - match session.read().await { - Ok(Frame::Text(text)) => { - println!("Received text: {}", text); - - // Echo back - if let Err(e) = session.send(&text).await { - eprintln!("Send error: {:?}", e); - break; - } - } - Ok(_) => {} - Err(e) => { - eprintln!("{e:?}"); - break; - } - } - } + session.on::(async |_, _| Ok(())).await; }); } } diff --git a/src/lib.rs b/src/lib.rs index 86365df..310676f 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -7,7 +7,7 @@ pub mod session; pub mod ws; pub type Result = std::result::Result; -pub type BoxFuture<'a> = Pin + Send + 'a>>; +pub type BoxFuture<'a> = Pin> + Send + 'a>>; pub type MethodHandler = Box BoxFuture<'static> + Send + Sync>; pub trait Method { diff --git a/src/session.rs b/src/session.rs index 7aec8d1..ef8e1a3 100644 --- a/src/session.rs +++ b/src/session.rs @@ -75,7 +75,14 @@ impl Session { match msg { Message::Request { id, method, data } => { if let Some(m) = s.methods.lock().await.get(&method) { - (m)(id, data).await + (m)(id, data).await; + // let s = s.clone(); + // if let Ok(req) = serde_json::from_value(value) { + // match handler(id, req).await { + // Ok(res) => s.respond::(id, res).await, + // Err(res) => s.respond_error::(id, res).await, + // }; + // } } } Message::Response { id, result } => { @@ -94,7 +101,10 @@ impl Session { }); } - pub async fn on + Send + 'static>( + pub async fn on< + M: Method, + Fut: Future> + Send + 'static, + >( &self, handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static, ) { @@ -106,9 +116,12 @@ impl Session { let handler = Arc::clone(&handler); Box::pin(async move { - if let Ok(req) = serde_json::from_value(value) { - handler(id, req).await; - } + 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()?), + }, + ) }) }), ); From d1bed32ca395d3b55dc24fb2f7701c856893c8b2 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 05:11:36 +0100 Subject: [PATCH 26/27] Working response --- src/session.rs | 27 ++++++++++++++------------- 1 file changed, 14 insertions(+), 13 deletions(-) diff --git a/src/session.rs b/src/session.rs index ef8e1a3..d20c7ca 100644 --- a/src/session.rs +++ b/src/session.rs @@ -75,14 +75,15 @@ impl Session { match msg { Message::Request { id, method, data } => { if let Some(m) = s.methods.lock().await.get(&method) { - (m)(id, data).await; - // let s = s.clone(); - // if let Ok(req) = serde_json::from_value(value) { - // match handler(id, req).await { - // Ok(res) => s.respond::(id, res).await, - // Err(res) => s.respond_error::(id, res).await, - // }; - // } + 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 } => { @@ -170,16 +171,16 @@ impl Session { } } - pub async fn respond(&self, to: u32, res: M::Response) -> crate::Result<()> { - self.send::(&Message::Response { + pub async fn respond(&self, to: u32, val: serde_json::Value) -> crate::Result<()> { + self.send::(&Message::Response { id: to, - result: res, + result: val, }) .await } - pub async fn respond_error(&self, to: u32, err: M::Error) -> crate::Result<()> { - self.send::(&Message::ErrorResponse { id: to, error: err }) + pub async fn respond_error(&self, to: u32, val: serde_json::Value) -> crate::Result<()> { + self.send::(&Message::ErrorResponse { id: to, error: val }) .await } From 2da9c1abed37fcca376d70255f201cdf8ddc89db Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 05:17:34 +0100 Subject: [PATCH 27/27] Improved examples --- examples/client.rs | 26 +++++++++++++++----------- examples/server.rs | 18 ++++++++++++++---- 2 files changed, 29 insertions(+), 15 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index da7847b..14a1484 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -6,9 +6,9 @@ struct Data; impl Method for Data { const NAME: &'static str = "data"; - type Request = (); - type Response = (); - type Error = (); + type Request = String; + type Response = String; + type Error = String; } #[tokio::main(flavor = "current_thread")] @@ -17,15 +17,19 @@ async fn main() -> session_rs::Result<()> { session.start_receiver(); - println!("{:?}", session.request::(()).await?); + println!( + "Hi: {:?}", + session + .request::("Hello from client".to_string()) + .await? + ); - session - .on::(async |i, d| { - println!("Ok {i} {d:?}"); - - Ok(()) - }) - .await; + println!( + "Invalid data response: {:?}", + session + .request::("invalid_data".to_string()) + .await? + ); tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; diff --git a/examples/server.rs b/examples/server.rs index 770f39d..0d6fb8e 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -8,9 +8,9 @@ struct Data; impl Method for Data { const NAME: &'static str = "data"; - type Request = (); - type Response = (); - type Error = (); + type Request = String; + type Response = String; + type Error = String; } #[tokio::main(flavor = "current_thread")] @@ -32,7 +32,17 @@ async fn main() -> session_rs::Result<()> { session.start_receiver(); - session.on::(async |_, _| Ok(())).await; + session + .on::(async |_, req| { + println!("Msg from client: {req}"); + + if req == "invalid_data" { + return Err("Invalid data".to_string()); + } + + Ok("Hello from server".to_string()) + }) + .await; }); } }