Refactored types

This commit is contained in:
2026-02-19 00:23:47 +01:00
parent c5dd9f019d
commit 8d67252800
6 changed files with 76 additions and 59 deletions
+10 -9
View File
@@ -1,6 +1,6 @@
use std::sync::Arc; use std::sync::Arc;
use session_rs::{SessionFrame, ws::WebSocket}; use session_rs::ws::{Frame, WebSocket};
#[tokio::main(flavor = "current_thread")] #[tokio::main(flavor = "current_thread")]
async fn main() -> session_rs::Result<()> { async fn main() -> session_rs::Result<()> {
@@ -10,13 +10,14 @@ async fn main() -> session_rs::Result<()> {
let read_session = Arc::clone(&session); let read_session = Arc::clone(&session);
tokio::spawn(async move { tokio::spawn(async move {
loop { loop {
match read_session.read().await { println!("{:?}", read_session.read().await);
Ok(SessionFrame::Text(text)) => { // match read_session.read().await {
println!("Server says: {}", text); // Ok(Frame::Text(text)) => {
} // println!("Server says: {}", text);
Ok(_) => {} // }
Err(_) => break, // Ok(_) => {}
} // Err(_) => break,
// }
} }
}); });
@@ -24,7 +25,7 @@ async fn main() -> session_rs::Result<()> {
for i in 0..5 { for i in 0..5 {
println!("sending"); println!("sending");
let msg = serde_json::json!({ "hello": i }); 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; tokio::time::sleep(std::time::Duration::from_secs(1)).await;
} }
+3 -3
View File
@@ -1,7 +1,7 @@
use std::sync::Arc; use std::sync::Arc;
use tokio::net::TcpListener; use tokio::net::TcpListener;
use session_rs::{SessionFrame, ws::WebSocket}; use session_rs::ws::{Frame, WebSocket};
#[tokio::main(flavor = "current_thread")] #[tokio::main(flavor = "current_thread")]
async fn main() -> session_rs::Result<()> { async fn main() -> session_rs::Result<()> {
@@ -27,11 +27,11 @@ async fn main() -> session_rs::Result<()> {
// Read loop // Read loop
loop { loop {
match session.read().await { match session.read().await {
Ok(SessionFrame::Text(text)) => { Ok(Frame::Text(text)) => {
println!("Received text: {}", text); println!("Received text: {}", text);
// Echo back // 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); eprintln!("Send error: {:?}", e);
break; break;
} }
+8 -21
View File
@@ -1,27 +1,20 @@
use std::string::FromUtf8Error;
pub mod server; pub mod server;
pub mod session; pub mod session;
pub mod ws; pub mod ws;
pub enum SessionFrame {
Text(String),
Binary(Vec<u8>),
Ping,
Pong,
Close,
}
pub type Result<T> = std::result::Result<T, Error>; pub type Result<T> = std::result::Result<T, Error>;
#[derive(Debug)] #[derive(Debug)]
pub enum Error { pub enum Error {
Io(std::io::Error), WebSocket(ws::Error),
Json(serde_json::Error), Json(serde_json::Error),
InvalidFrame(String), Io(std::io::Error),
HandshakeFailed(String), }
ConnectionClosed,
Utf8(FromUtf8Error), impl From<ws::Error> for Error {
fn from(value: ws::Error) -> Self {
Self::WebSocket(value)
}
} }
impl From<std::io::Error> for Error { impl From<std::io::Error> for Error {
@@ -35,9 +28,3 @@ impl From<serde_json::Error> for Error {
Self::Json(value) Self::Json(value)
} }
} }
impl From<FromUtf8Error> for Error {
fn from(value: FromUtf8Error) -> Self {
Self::Utf8(value)
}
}
+24
View File
@@ -0,0 +1,24 @@
use std::string::FromUtf8Error;
pub type Result<T> = std::result::Result<T, Error>;
#[derive(Debug)]
pub enum Error {
Io(std::io::Error),
InvalidFrame(String),
HandshakeFailed(String),
Utf8(FromUtf8Error),
ConnectionClosed,
}
impl From<std::io::Error> for Error {
fn from(value: std::io::Error) -> Self {
Self::Io(value)
}
}
impl From<FromUtf8Error> for Error {
fn from(value: FromUtf8Error) -> Self {
Self::Utf8(value)
}
}
+4 -4
View File
@@ -82,7 +82,7 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu
} }
impl WebSocket { impl WebSocket {
pub async fn handshake(mut stream: TcpStream) -> crate::Result<Self> { pub async fn handshake(mut stream: TcpStream) -> super::Result<Self> {
handle_websocket_handshake(&mut stream).await?; handle_websocket_handshake(&mut stream).await?;
let (read, write) = stream.into_split(); let (read, write) = stream.into_split();
@@ -96,7 +96,7 @@ impl WebSocket {
} }
/// Connect to a WebSocket server and perform the handshake /// Connect to a WebSocket server and perform the handshake
pub async fn connect(addr: &str, path: &str) -> crate::Result<Self> { pub async fn connect(addr: &str, path: &str) -> super::Result<Self> {
// 1. TCP connect // 1. TCP connect
let mut stream = TcpStream::connect(addr).await?; let mut stream = TcpStream::connect(addr).await?;
@@ -123,7 +123,7 @@ impl WebSocket {
let mut status_line = String::new(); let mut status_line = String::new();
reader.read_line(&mut status_line).await?; reader.read_line(&mut status_line).await?;
if !status_line.starts_with("HTTP/1.1 101") { 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: {}", "Expected 101 Switching Protocols, got: {}",
status_line.trim_end() status_line.trim_end()
))); )));
@@ -153,7 +153,7 @@ impl WebSocket {
base64::encode(sha1.finalize()) base64::encode(sha1.finalize())
}; };
if sec_accept.as_deref() != Some(expected.as_str()) { if sec_accept.as_deref() != Some(expected.as_str()) {
return Err(crate::Error::HandshakeFailed( return Err(super::Error::HandshakeFailed(
"Sec-WebSocket-Accept mismatch".into(), "Sec-WebSocket-Accept mismatch".into(),
)); ));
} }
+27 -22
View File
@@ -1,4 +1,6 @@
pub mod error;
pub mod handshake; pub mod handshake;
pub use error::{Error, Result};
use std::{ use std::{
hash::{Hash, Hasher}, hash::{Hash, Hasher},
@@ -9,7 +11,14 @@ use tokio::{
sync::Mutex, sync::Mutex,
}; };
use crate::SessionFrame; #[derive(Debug, Clone)]
pub enum Frame {
Text(String),
Binary(Vec<u8>),
Ping,
Pong,
Close,
}
pub struct WebSocket { pub struct WebSocket {
pub(crate) reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>, pub(crate) reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>,
@@ -44,7 +53,7 @@ impl Hash for WebSocket {
} }
impl 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 writer = self.writer.lock().await;
let mut header = Vec::with_capacity(10); let mut header = Vec::with_capacity(10);
@@ -86,23 +95,23 @@ impl WebSocket {
} }
impl WebSocket { impl WebSocket {
pub async fn send<T: serde::Serialize>(&self, msg: &T) -> crate::Result<()> { pub async fn send(&self, msg: &str) -> Result<()> {
self.send_frame(0x1, &serde_json::to_vec(msg)?).await 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 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 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 self.send_frame(0xA, &[]).await
} }
pub async fn close(&self) -> crate::Result<()> { pub async fn close(&self) -> Result<()> {
self.send_frame(0x8, &[]).await self.send_frame(0x8, &[]).await
} }
@@ -123,7 +132,7 @@ impl WebSocket {
impl WebSocket { impl WebSocket {
/// Read a full WebSocket frame (handling masking and control frames) /// Read a full WebSocket frame (handling masking and control frames)
/// Returns (opcode, payload) /// Returns (opcode, payload)
pub async fn read_frame(&self) -> crate::Result<(bool, u8, Vec<u8>)> { pub async fn read_frame(&self) -> Result<(bool, u8, Vec<u8>)> {
let mut reader = self.reader.lock().await; let mut reader = self.reader.lock().await;
// --- 1. Read first 2-byte header --- // --- 1. Read first 2-byte header ---
@@ -150,7 +159,7 @@ impl WebSocket {
if !masked && !self.mask_payload { if !masked && !self.mask_payload {
// Per spec, client-to-server frames MUST be masked // Per spec, client-to-server frames MUST be masked
self.close().await.ok(); self.close().await.ok();
return Err(crate::Error::InvalidFrame( return Err(Error::InvalidFrame(
"Received unmasked frame from client".into(), "Received unmasked frame from client".into(),
)); ));
} }
@@ -171,7 +180,7 @@ impl WebSocket {
Ok((fin, opcode, payload)) Ok((fin, opcode, payload))
} }
pub async fn read(&self) -> crate::Result<SessionFrame> { pub async fn read(&self) -> Result<Frame> {
let (fin, opcode, mut payload) = self.read_frame().await?; let (fin, opcode, mut payload) = self.read_frame().await?;
if !fin { if !fin {
@@ -194,9 +203,7 @@ impl WebSocket {
0xA => {} 0xA => {}
_ => { _ => {
self.close().await.ok(); self.close().await.ok();
return Err(crate::Error::InvalidFrame(format!( return Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}")));
"Unknown opcode: {opcode}"
)));
} }
} }
} }
@@ -206,29 +213,27 @@ impl WebSocket {
// Close // Close
0x8 => { 0x8 => {
self.close().await.ok(); self.close().await.ok();
Ok(SessionFrame::Close) Ok(Frame::Close)
} }
// Ping // Ping
0x9 => { 0x9 => {
self.send_pong().await.ok(); self.send_pong().await.ok();
Ok(SessionFrame::Ping) Ok(Frame::Ping)
} }
// Pong // Pong
0xA => Ok(SessionFrame::Pong), 0xA => Ok(Frame::Pong),
// Text // Text
0x1 => Ok(SessionFrame::Text(String::from_utf8(payload)?)), 0x1 => Ok(Frame::Text(String::from_utf8(payload)?)),
// Binary // Binary
0x2 => Ok(SessionFrame::Binary(payload)), 0x2 => Ok(Frame::Binary(payload)),
_ => { _ => {
self.close().await.ok(); self.close().await.ok();
Err(crate::Error::InvalidFrame(format!( Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}")))
"Unknown opcode: {opcode}"
)))
} }
} }
} }