Refactored types
This commit is contained in:
+10
-9
@@ -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
@@ -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
@@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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}"
|
|
||||||
)))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user