From 1537b617097d6b3e6eb19a4bc29c3f974b419f6a Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Mon, 30 Mar 2026 13:27:44 +0200 Subject: [PATCH 1/3] Fixed handshake --- src/ws/error.rs | 7 +++ src/ws/handshake.rs | 109 ++++++++++++++++++++++++++++++++------------ 2 files changed, 88 insertions(+), 28 deletions(-) diff --git a/src/ws/error.rs b/src/ws/error.rs index 9aa2299..4c426f3 100644 --- a/src/ws/error.rs +++ b/src/ws/error.rs @@ -9,6 +9,7 @@ pub enum Error { HandshakeFailed(String), Utf8(FromUtf8Error), ConnectionClosed, + Elapsed, } impl From for Error { @@ -22,3 +23,9 @@ impl From for Error { Self::Utf8(value) } } + +impl From for Error { + fn from(_: tokio::time::error::Elapsed) -> Self { + Self::Elapsed + } +} diff --git a/src/ws/handshake.rs b/src/ws/handshake.rs index 6b347fd..0a3f644 100644 --- a/src/ws/handshake.rs +++ b/src/ws/handshake.rs @@ -1,10 +1,12 @@ use base64::Engine; +use base64::engine::general_purpose::STANDARD as Base64; use sha1::{Digest, Sha1}; -use std::sync::Arc; +use std::{collections::HashMap, sync::Arc}; use tokio::{ io::{AsyncBufReadExt, AsyncWriteExt, BufReader}, net::TcpStream, sync::Mutex, + time::{Duration, timeout}, }; use super::WebSocket; @@ -13,72 +15,119 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu let (read_half, mut write_half) = stream.split(); let mut reader = BufReader::new(read_half); + // ---- 1. Read request line with timeout ---- let mut request_line = String::new(); - reader.read_line(&mut request_line).await?; + timeout(Duration::from_secs(5), reader.read_line(&mut request_line)).await??; + let request_line = request_line.trim_end(); - if request_line.starts_with("HEAD") { + if !request_line.starts_with("GET") { write_half - .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n") + .write_all( + b"HTTP/1.1 405 Method Not Allowed\r\n\ + Content-Length: 0\r\n\ + Connection: close\r\n\r\n", + ) .await?; + write_half.shutdown().await?; return Ok(()); } - if !request_line.starts_with("GET") { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - "Invalid HTTP method", - )); - } - - use std::collections::HashMap; + // ---- 2. Read headers with timeout ---- let mut headers = HashMap::new(); - let mut line = String::new(); loop { - line.clear(); - reader.read_line(&mut line).await?; + let mut line = String::new(); + timeout(Duration::from_secs(5), reader.read_line(&mut line)).await??; + if line == "\r\n" { break; } + if let Some((k, v)) = line.split_once(':') { headers.insert(k.trim().to_lowercase(), v.trim().to_string()); } } - if headers + // ---- 3. Check if this is a WebSocket upgrade ---- + let is_upgrade = headers .get("upgrade") - .map(|v| !v.eq_ignore_ascii_case("websocket")) - .unwrap_or(true) - { + .map(|v| v.eq_ignore_ascii_case("websocket")) + .unwrap_or(false); + + let has_connection_upgrade = headers + .get("connection") + .map(|v| v.to_lowercase().contains("upgrade")) + .unwrap_or(false); + + if !is_upgrade || !has_connection_upgrade { + // Normal HTTP response (important for browsers) + let body = b"OK"; + write_half - .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nOK") + .write_all( + format!( + "HTTP/1.1 200 OK\r\n\ + Content-Type: text/plain\r\n\ + Content-Length: {}\r\n\ + Connection: close\r\n\ + \r\n", + body.len() + ) + .as_bytes(), + ) .await?; + + write_half.write_all(body).await?; + write_half.flush().await?; + write_half.shutdown().await?; + return Ok(()); } - let key = headers - .get("sec-websocket-key") - .ok_or_else(|| std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing key"))?; + // ---- 4. Validate required headers ---- + let key = headers.get("sec-websocket-key").ok_or_else(|| { + std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key") + })?; - use base64::Engine; - use base64::engine::general_purpose::STANDARD as Base64; - use sha1::{Digest, Sha1}; + let version_ok = headers + .get("sec-websocket-version") + .map(|v| v == "13") + .unwrap_or(false); + if !version_ok { + write_half + .write_all( + b"HTTP/1.1 426 Upgrade Required\r\n\ + Sec-WebSocket-Version: 13\r\n\ + Content-Length: 0\r\n\ + Connection: close\r\n\r\n", + ) + .await?; + write_half.shutdown().await?; + return Ok(()); + } + + // ---- 5. Generate Sec-WebSocket-Accept ---- let mut hasher = Sha1::new(); hasher.update(key.as_bytes()); hasher.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = Base64.encode(hasher.finalize()); + // ---- 6. Send upgrade response ---- let response = format!( "HTTP/1.1 101 Switching Protocols\r\n\ Upgrade: websocket\r\n\ Connection: Upgrade\r\n\ - Sec-WebSocket-Accept: {}\r\n\r\n", + Sec-WebSocket-Accept: {}\r\n\ + \r\n", accept ); write_half.write_all(response.as_bytes()).await?; + write_half.flush().await?; + Ok(()) } @@ -122,7 +171,11 @@ impl WebSocket { // 4. Read HTTP response let mut reader = BufReader::new(&mut stream); let mut status_line = String::new(); - reader.read_line(&mut status_line).await?; + timeout( + tokio::time::Duration::from_secs(5), + reader.read_line(&mut status_line), + ) + .await??; if !status_line.starts_with("HTTP/1.1 101") { return Err(super::Error::HandshakeFailed(format!( "Expected 101 Switching Protocols, got: {}", From ec984ffaa3826efff6bcb7dbdf1ccf67cd3b8587 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Mon, 30 Mar 2026 13:32:28 +0200 Subject: [PATCH 2/3] Fixed Session loop --- src/server.rs | 37 +++++++++++++++++++++++++++++-------- 1 file changed, 29 insertions(+), 8 deletions(-) diff --git a/src/server.rs b/src/server.rs index 4b4db24..5ac89b8 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,6 +1,6 @@ use std::{net::SocketAddr, sync::Arc}; -use tokio::net::TcpListener; +use tokio::{net::TcpListener, time::timeout}; use crate::{session::Session, ws::WebSocket}; @@ -23,19 +23,40 @@ impl SessionServer { Ok((Session::from_ws(ws), addr)) } - pub async fn session_loop> + Send + 'static>( - &self, - on_conn: impl Fn(Session, SocketAddr) -> Fut + 'static, - ) -> crate::Result<()> { + pub async fn session_loop(&self, on_conn: F) -> crate::Result<()> + where + F: Fn(Session, SocketAddr) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { let conn_handler = Arc::new(on_conn); loop { - let (session, addr) = self.accept().await?; + let (stream, addr) = self.listener.accept().await?; let conn_handler = conn_handler.clone(); - session.start_receiver(); + tokio::spawn(async move { + match timeout( + tokio::time::Duration::from_secs(5), + WebSocket::handshake(stream), + ) + .await + { + Ok(Ok(ws)) => { + let session = Session::from_ws(ws); + session.start_receiver(); - tokio::spawn(conn_handler(session, addr)); + if let Err(e) = conn_handler(session, addr).await { + eprintln!("Connection error: {:?}", e); + } + } + Ok(Err(e)) => { + eprintln!("Handshake failed from {}: {:?}", addr, e); + } + Err(_) => { + eprintln!("Handshake failed from {}: Handshake Timeout", addr); + } + } + }); } } } From a853de031de78263dccd57d27af03aa5bff0d376 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Mon, 30 Mar 2026 13:39:13 +0200 Subject: [PATCH 3/3] fixed session --- src/lib.rs | 4 ++-- src/session.rs | 9 +++++++-- 2 files changed, 9 insertions(+), 4 deletions(-) diff --git a/src/lib.rs b/src/lib.rs index b4fd668..16b1f2b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,4 +1,4 @@ -use std::pin::Pin; +use std::{pin::Pin, sync::Arc}; use serde::{Deserialize, Serialize}; @@ -9,7 +9,7 @@ pub mod ws; pub type Result = std::result::Result; pub type BoxFuture<'a, T = Option<(bool, serde_json::Value)>> = Pin + Send + 'a>>; -pub type MethodHandler = Box BoxFuture<'static> + Send + Sync>; +pub type MethodHandler = Arc BoxFuture<'static> + Send + Sync>; pub trait Method { const NAME: &'static str; diff --git a/src/session.rs b/src/session.rs index 1dd6e96..2443a85 100644 --- a/src/session.rs +++ b/src/session.rs @@ -87,7 +87,12 @@ impl Session { match msg { Message::Request { id, method, data } => { - if let Some(m) = s.methods.lock().await.get(&method) { + let handler = { + let methods = s.methods.lock().await; + methods.get(&method).cloned() + }; + + if let Some(m) = handler { if let Some((err, res)) = (m)(id, data).await { if err { s.respond_error(id, res) @@ -157,7 +162,7 @@ impl Session { self.methods.lock().await.insert( M::NAME.to_string(), - Box::new(move |id, value| { + Arc::new(move |id, value| { let handler = Arc::clone(&handler); Box::pin(async move {