From 09c1fbc2323de18a600f736b0f3211575ef0b785 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 21:49:11 +0100 Subject: [PATCH] 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))) + } +}