Refactoring to ws
This commit is contained in:
+3
-3
@@ -1,16 +1,16 @@
|
|||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
use session_rs::session::Session;
|
use session_rs::ws::WebSocket;
|
||||||
|
|
||||||
#[tokio::main(flavor = "current_thread")]
|
#[tokio::main(flavor = "current_thread")]
|
||||||
async fn main() -> session_rs::Result<()> {
|
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
|
// Spawn read loop
|
||||||
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_frame().await {
|
match read_session.read().await {
|
||||||
Ok(Some((opcode, payload))) => {
|
Ok(Some((opcode, payload))) => {
|
||||||
if opcode == 0x1 {
|
if opcode == 0x1 {
|
||||||
let text = String::from_utf8(payload).unwrap_or_default();
|
let text = String::from_utf8(payload).unwrap_or_default();
|
||||||
|
|||||||
+2
-2
@@ -1,7 +1,7 @@
|
|||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use tokio::net::TcpListener;
|
use tokio::net::TcpListener;
|
||||||
|
|
||||||
use session_rs::session::Session;
|
use session_rs::session::WebSocket;
|
||||||
|
|
||||||
#[tokio::main(flavor = "current_thread")]
|
#[tokio::main(flavor = "current_thread")]
|
||||||
async fn main() -> session_rs::Result<()> {
|
async fn main() -> session_rs::Result<()> {
|
||||||
@@ -14,7 +14,7 @@ async fn main() -> session_rs::Result<()> {
|
|||||||
|
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
// Wrap session in Arc so tasks can share it
|
// 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),
|
Ok(s) => Arc::new(s),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
eprintln!("Handshake failed: {:?}", e);
|
eprintln!("Handshake failed: {:?}", e);
|
||||||
|
|||||||
+1
-1
@@ -1,6 +1,6 @@
|
|||||||
pub mod handshake;
|
|
||||||
pub mod server;
|
pub mod server;
|
||||||
pub mod session;
|
pub mod session;
|
||||||
|
pub mod ws;
|
||||||
|
|
||||||
pub enum SessionFrame<T> {
|
pub enum SessionFrame<T> {
|
||||||
Typed(T),
|
Typed(T),
|
||||||
|
|||||||
-206
@@ -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<Mutex<tokio::net::tcp::OwnedReadHalf>>,
|
|
||||||
pub(crate) writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>,
|
|
||||||
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<H: Hasher>(&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<T: serde::Serialize>(&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<u8>)> {
|
|
||||||
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<T>(&self) -> crate::Result<SessionFrame<T>> {
|
|
||||||
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}"
|
|
||||||
)))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ use tokio::{
|
|||||||
sync::Mutex,
|
sync::Mutex,
|
||||||
};
|
};
|
||||||
|
|
||||||
use crate::session::Session;
|
use super::WebSocket;
|
||||||
|
|
||||||
pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> {
|
pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> {
|
||||||
let (read_half, mut write_half) = stream.split();
|
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(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Session {
|
impl WebSocket {
|
||||||
pub async fn handshake(mut stream: TcpStream) -> crate::Result<Self> {
|
pub async fn handshake(mut stream: TcpStream) -> crate::Result<Self> {
|
||||||
crate::handshake::handle_websocket_handshake(&mut stream).await?;
|
handle_websocket_handshake(&mut stream).await?;
|
||||||
|
|
||||||
let (read, write) = stream.into_split();
|
let (read, write) = stream.into_split();
|
||||||
|
|
||||||
+212
@@ -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<Mutex<tokio::net::tcp::OwnedReadHalf>>,
|
||||||
|
pub(crate) writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>,
|
||||||
|
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<H: Hasher>(&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<T: serde::Serialize>(&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<u8>)> {
|
||||||
|
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<T>(&self) -> crate::Result<SessionFrame<T>> {
|
||||||
|
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}"
|
||||||
|
)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user