diff --git a/examples/client.rs b/examples/client.rs index 515158f..0cbb402 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -1,6 +1,6 @@ use std::sync::Arc; -use session_rs::ws::WebSocket; +use session_rs::{SessionFrame, ws::WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -11,13 +11,10 @@ async fn main() -> session_rs::Result<()> { tokio::spawn(async move { loop { match read_session.read().await { - Ok(Some((opcode, payload))) => { - if opcode == 0x1 { - let text = String::from_utf8(payload).unwrap_or_default(); - println!("Server says: {}", text); - } + Ok(SessionFrame::Text(text)) => { + println!("Server says: {}", text); } - Ok(None) => {} + Ok(_) => {} Err(_) => break, } } diff --git a/examples/server.rs b/examples/server.rs index 3b78c5f..edf75b2 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::WebSocket; +use session_rs::{SessionFrame, ws::WebSocket}; #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -26,21 +26,17 @@ async fn main() -> session_rs::Result<()> { // Read loop loop { - match session.read_frame().await { - Ok(Some((opcode, payload))) => { - if opcode == 0x1 { - // Text frame → parse JSON if possible - let text = String::from_utf8(payload).unwrap_or_default(); - println!("Received text: {}", text); + match session.read().await { + Ok(SessionFrame::Text(text)) => { + println!("Received text: {}", text); - // Echo back - if let Err(e) = session.send(&serde_json::json!({"echo": text})).await { - eprintln!("Send error: {:?}", e); - break; - } + // Echo back + if let Err(e) = session.send(&serde_json::json!({"echo": text})).await { + eprintln!("Send error: {:?}", e); + break; } } - Ok(None) => {} + Ok(_) => {} Err(e) => { eprintln!("{e:?}"); break; diff --git a/src/lib.rs b/src/lib.rs index 208c505..7308492 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,13 +1,15 @@ +use std::string::FromUtf8Error; + pub mod server; pub mod session; pub mod ws; -pub enum SessionFrame { - Typed(T), +pub enum SessionFrame { + Text(String), Binary(Vec), Ping, Pong, - Close + Close, } pub type Result = std::result::Result; @@ -19,6 +21,7 @@ pub enum Error { InvalidFrame(String), HandshakeFailed(String), ConnectionClosed, + Utf8(FromUtf8Error), } impl From for Error { @@ -32,3 +35,9 @@ impl From for Error { Self::Json(value) } } + +impl From for Error { + fn from(value: FromUtf8Error) -> Self { + Self::Utf8(value) + } +} diff --git a/src/ws/mod.rs b/src/ws/mod.rs index c42df73..3f87e19 100644 --- a/src/ws/mod.rs +++ b/src/ws/mod.rs @@ -171,8 +171,36 @@ impl WebSocket { Ok((fin, opcode, payload)) } - pub async fn read(&self) -> crate::Result> { - let (fin, opcode, payload) = self.read_frame().await?; + pub async fn read(&self) -> crate::Result { + let (fin, opcode, mut payload) = self.read_frame().await?; + + if !fin { + // Continuation loop + while let (fin, o, mut p) = self.read_frame().await? + && !fin + { + match o { + // Continuation + 0x0 => payload.append(&mut p), + // Close + 0x8 => { + self.close().await.ok(); + } + // Ping + 0x9 => { + self.send_pong().await.ok(); + } + // Pong + 0xA => {} + _ => { + self.close().await.ok(); + return Err(crate::Error::InvalidFrame(format!( + "Unknown opcode: {opcode}" + ))); + } + } + } + } match opcode { // Close @@ -190,16 +218,11 @@ impl WebSocket { // Pong 0xA => Ok(SessionFrame::Pong), - // Continuation - 0x0 => Ok(SessionFrame::Pong), - // Text - // 0x1 => { - - // }, + 0x1 => Ok(SessionFrame::Text(String::from_utf8(payload)?)), // Binary - 0x2 => Ok(SessionFrame::Pong), + 0x2 => Ok(SessionFrame::Binary(payload)), _ => { self.close().await.ok();