From 5edbe8c8a19799fa56951aaaaddabe08f4926995 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Wed, 18 Feb 2026 22:15:37 +0100 Subject: [PATCH] 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 {