diff --git a/examples/client.rs b/examples/client.rs index 0cbb402..f743c59 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,6 +1,6 @@ use std::sync::Arc; -use session_rs::{SessionFrame, ws::WebSocket}; +use session_rs::ws::{Frame, WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -10,13 +10,14 @@ async fn main() -> session_rs::Result<()> { let read_session = Arc::clone(&session); tokio::spawn(async move { loop { - match read_session.read().await { - Ok(SessionFrame::Text(text)) => { - println!("Server says: {}", text); - } - Ok(_) => {} - Err(_) => break, - } + println!("{:?}", read_session.read().await); + // match read_session.read().await { + // Ok(Frame::Text(text)) => { + // println!("Server says: {}", text); + // } + // Ok(_) => {} + // Err(_) => break, + // } } }); @@ -24,7 +25,7 @@ async fn main() -> session_rs::Result<()> { for i in 0..5 { println!("sending"); let msg = serde_json::json!({ "hello": i }); - session.send(&msg).await?; + session.send(&msg.to_string()).await?; tokio::time::sleep(std::time::Duration::from_secs(1)).await; } diff --git a/examples/server.rs b/examples/server.rs index edf75b2..ce36cc8 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1,7 +1,7 @@ use std::sync::Arc; use tokio::net::TcpListener; -use session_rs::{SessionFrame, ws::WebSocket}; +use session_rs::ws::{Frame, WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -27,11 +27,11 @@ async fn main() -> session_rs::Result<()> { // Read loop loop { match session.read().await { - Ok(SessionFrame::Text(text)) => { + Ok(Frame::Text(text)) => { println!("Received text: {}", text); // Echo back - if let Err(e) = session.send(&serde_json::json!({"echo": text})).await { + if let Err(e) = session.send(&serde_json::json!({"echo": text}).to_string()).await { eprintln!("Send error: {:?}", e); break; } diff --git a/src/lib.rs b/src/lib.rs index 7308492..627be42 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,27 +1,20 @@ -use std::string::FromUtf8Error; - pub mod server; pub mod session; pub mod ws; -pub enum SessionFrame { - Text(String), - Binary(Vec), - Ping, - Pong, - Close, -} - pub type Result = std::result::Result; #[derive(Debug)] pub enum Error { - Io(std::io::Error), + WebSocket(ws::Error), Json(serde_json::Error), - InvalidFrame(String), - HandshakeFailed(String), - ConnectionClosed, - Utf8(FromUtf8Error), + Io(std::io::Error), +} + +impl From for Error { + fn from(value: ws::Error) -> Self { + Self::WebSocket(value) + } } impl From for Error { @@ -35,9 +28,3 @@ impl From for Error { Self::Json(value) } } - -impl From for Error { - fn from(value: FromUtf8Error) -> Self { - Self::Utf8(value) - } -} diff --git a/src/ws/error.rs b/src/ws/error.rs new file mode 100644 index 0000000..9aa2299 --- /dev/null +++ b/src/ws/error.rs @@ -0,0 +1,24 @@ +use std::string::FromUtf8Error; + +pub type Result = std::result::Result; + +#[derive(Debug)] +pub enum Error { + Io(std::io::Error), + InvalidFrame(String), + HandshakeFailed(String), + Utf8(FromUtf8Error), + ConnectionClosed, +} + +impl From for Error { + fn from(value: std::io::Error) -> Self { + Self::Io(value) + } +} + +impl From for Error { + fn from(value: FromUtf8Error) -> Self { + Self::Utf8(value) + } +} diff --git a/src/ws/handshake.rs b/src/ws/handshake.rs index ec36903..4f92ec6 100644 --- a/src/ws/handshake.rs +++ b/src/ws/handshake.rs @@ -82,7 +82,7 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu } impl WebSocket { - pub async fn handshake(mut stream: TcpStream) -> crate::Result { + pub async fn handshake(mut stream: TcpStream) -> super::Result { handle_websocket_handshake(&mut stream).await?; let (read, write) = stream.into_split(); @@ -96,7 +96,7 @@ impl WebSocket { } /// Connect to a WebSocket server and perform the handshake - pub async fn connect(addr: &str, path: &str) -> crate::Result { + pub async fn connect(addr: &str, path: &str) -> super::Result { // 1. TCP connect let mut stream = TcpStream::connect(addr).await?; @@ -123,7 +123,7 @@ impl WebSocket { let mut status_line = String::new(); reader.read_line(&mut status_line).await?; if !status_line.starts_with("HTTP/1.1 101") { - return Err(crate::Error::HandshakeFailed(format!( + return Err(super::Error::HandshakeFailed(format!( "Expected 101 Switching Protocols, got: {}", status_line.trim_end() ))); @@ -153,7 +153,7 @@ impl WebSocket { base64::encode(sha1.finalize()) }; if sec_accept.as_deref() != Some(expected.as_str()) { - return Err(crate::Error::HandshakeFailed( + return Err(super::Error::HandshakeFailed( "Sec-WebSocket-Accept mismatch".into(), )); } diff --git a/src/ws/mod.rs b/src/ws/mod.rs index 3f87e19..a72bf41 100644 --- a/src/ws/mod.rs +++ b/src/ws/mod.rs @@ -1,4 +1,6 @@ +pub mod error; pub mod handshake; +pub use error::{Error, Result}; use std::{ hash::{Hash, Hasher}, @@ -9,7 +11,14 @@ use tokio::{ sync::Mutex, }; -use crate::SessionFrame; +#[derive(Debug, Clone)] +pub enum Frame { + Text(String), + Binary(Vec), + Ping, + Pong, + Close, +} pub struct WebSocket { pub(crate) reader: Arc>, @@ -44,7 +53,7 @@ impl Hash for WebSocket { } impl WebSocket { - async fn send_frame(&self, opcode: u8, payload: &[u8]) -> crate::Result<()> { + async fn send_frame(&self, opcode: u8, payload: &[u8]) -> Result<()> { let mut writer = self.writer.lock().await; let mut header = Vec::with_capacity(10); @@ -86,23 +95,23 @@ impl WebSocket { } 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(&self, msg: &str) -> Result<()> { + self.send_frame(0x1, msg.as_bytes()).await } - pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { + pub async fn send_bin(&self, payload: &[u8]) -> Result<()> { self.send_frame(0x2, payload).await } - pub async fn send_ping(&self) -> crate::Result<()> { + pub async fn send_ping(&self) -> Result<()> { self.send_frame(0x9, &[]).await } - pub async fn send_pong(&self) -> crate::Result<()> { + pub async fn send_pong(&self) -> Result<()> { self.send_frame(0xA, &[]).await } - pub async fn close(&self) -> crate::Result<()> { + pub async fn close(&self) -> Result<()> { self.send_frame(0x8, &[]).await } @@ -123,7 +132,7 @@ impl WebSocket { 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)> { + pub async fn read_frame(&self) -> Result<(bool, u8, Vec)> { let mut reader = self.reader.lock().await; // --- 1. Read first 2-byte header --- @@ -150,7 +159,7 @@ impl WebSocket { if !masked && !self.mask_payload { // Per spec, client-to-server frames MUST be masked self.close().await.ok(); - return Err(crate::Error::InvalidFrame( + return Err(Error::InvalidFrame( "Received unmasked frame from client".into(), )); } @@ -171,7 +180,7 @@ impl WebSocket { Ok((fin, opcode, payload)) } - pub async fn read(&self) -> crate::Result { + pub async fn read(&self) -> Result { let (fin, opcode, mut payload) = self.read_frame().await?; if !fin { @@ -194,9 +203,7 @@ impl WebSocket { 0xA => {} _ => { self.close().await.ok(); - return Err(crate::Error::InvalidFrame(format!( - "Unknown opcode: {opcode}" - ))); + return Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}"))); } } } @@ -206,29 +213,27 @@ impl WebSocket { // Close 0x8 => { self.close().await.ok(); - Ok(SessionFrame::Close) + Ok(Frame::Close) } // Ping 0x9 => { self.send_pong().await.ok(); - Ok(SessionFrame::Ping) + Ok(Frame::Ping) } // Pong - 0xA => Ok(SessionFrame::Pong), + 0xA => Ok(Frame::Pong), // Text - 0x1 => Ok(SessionFrame::Text(String::from_utf8(payload)?)), + 0x1 => Ok(Frame::Text(String::from_utf8(payload)?)), // Binary - 0x2 => Ok(SessionFrame::Binary(payload)), + 0x2 => Ok(Frame::Binary(payload)), _ => { self.close().await.ok(); - Err(crate::Error::InvalidFrame(format!( - "Unknown opcode: {opcode}" - ))) + Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}"))) } } }