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}" + ))) + } + } + } +}