From 3c254b75ff86714c3b906c0984b8ce4610fb219d Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 21:24:00 +0100 Subject: [PATCH] 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()))