Fixing issues

This commit is contained in:
2026-02-18 22:15:37 +01:00
parent 9ea2978c8d
commit 5edbe8c8a1
4 changed files with 46 additions and 38 deletions
+2 -2
View File
@@ -4,7 +4,7 @@ use session_rs::session::Session;
#[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::new_server("127.0.0.1:8080", "/").await?); let session = Arc::new(Session::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);
@@ -27,7 +27,7 @@ async fn main() -> session_rs::Result<()> {
for i in 0..5 { for i in 0..5 {
let msg = serde_json::json!({ "hello": i }); let msg = serde_json::json!({ "hello": i });
session.send(&msg).await?; session.send(&msg).await?;
tokio::time::sleep(std::time::Duration::from_secs(1)).await; tokio::time::sleep(std::time::Duration::from_secs(1000)).await;
} }
session.close().await?; session.close().await?;
+5 -2
View File
@@ -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::new_client(stream).await { let session = match Session::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);
@@ -51,7 +51,10 @@ async fn main() -> session_rs::Result<()> {
} }
} }
Ok(None) => {} Ok(None) => {}
Err(_) => break, // connection closed Err(e) => {
eprintln!("{e:?}");
break;
}
} }
} }
+6 -4
View File
@@ -82,20 +82,21 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu
} }
impl Session { impl Session {
pub async fn new_client(mut stream: TcpStream) -> crate::Result<Self> { pub async fn handshake(mut stream: TcpStream) -> crate::Result<Self> {
crate::handshake::handle_websocket_handshake(&mut stream).await?; crate::handshake::handle_websocket_handshake(&mut stream).await?;
let (read, write) = stream.into_split(); let (read, write) = stream.into_split();
Ok(Self { Ok(Self {
id: rand::random(),
reader: Arc::new(Mutex::new(read)), reader: Arc::new(Mutex::new(read)),
writer: Arc::new(Mutex::new(write)), writer: Arc::new(Mutex::new(write)),
id: rand::random(), mask_payload: false,
}) })
} }
/// Connect to a WebSocket server and perform the handshake /// Connect to a WebSocket server and perform the handshake
pub async fn new_server(addr: &str, path: &str) -> crate::Result<Self> { pub async fn connect(addr: &str, path: &str) -> crate::Result<Self> {
// 1. TCP connect // 1. TCP connect
let mut stream = TcpStream::connect(addr).await?; let mut stream = TcpStream::connect(addr).await?;
@@ -161,9 +162,10 @@ impl Session {
let (read, write) = stream.into_split(); let (read, write) = stream.into_split();
Ok(Self { Ok(Self {
id: rand::random(),
reader: Arc::new(Mutex::new(read)), reader: Arc::new(Mutex::new(read)),
writer: Arc::new(Mutex::new(write)), writer: Arc::new(Mutex::new(write)),
id: rand::random(), mask_payload: true,
}) })
} }
} }
+31 -28
View File
@@ -4,7 +4,6 @@ use std::{
}; };
use tokio::{ use tokio::{
io::{AsyncReadExt, AsyncWriteExt}, io::{AsyncReadExt, AsyncWriteExt},
net::TcpStream,
sync::Mutex, sync::Mutex,
}; };
@@ -12,6 +11,7 @@ pub struct Session {
pub(crate) reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>, pub(crate) reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>,
pub(crate) writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>, pub(crate) writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>,
pub(crate) id: u64, pub(crate) id: u64,
pub(crate) mask_payload: bool,
} }
impl Clone for Session { impl Clone for Session {
@@ -19,6 +19,7 @@ impl Clone for Session {
Session { Session {
reader: self.reader.clone(), reader: self.reader.clone(),
writer: self.writer.clone(), writer: self.writer.clone(),
mask_payload: self.mask_payload.clone(),
id: self.id, id: self.id,
} }
} }
@@ -43,29 +44,46 @@ impl Session {
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);
header.push(0x80 | opcode); let mask_bit = if self.mask_payload { 0x80 } else { 0x00 };
header.push(0x80 | opcode); // FIN + opcode
let len = payload.len(); let len = payload.len();
if len < 126 { if len < 126 {
header.push(len as u8); header.push((len as u8) | mask_bit);
} else if len <= 0xFFFF { } else if len <= 0xFFFF {
header.push(126); header.push(126 | mask_bit);
header.extend_from_slice(&(len as u16).to_be_bytes()); header.extend_from_slice(&(len as u16).to_be_bytes());
} else { } else {
header.push(127); header.push(127 | mask_bit);
header.extend_from_slice(&(len as u64).to_be_bytes()); 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(&header).await?;
writer.write_all(payload).await?; writer.write_all(payload).await?;
}
writer.flush().await?;
Ok(()) Ok(())
} }
} }
impl Session { impl Session {
pub async fn send<T: serde::Serialize>(&self, msg: &T) -> crate::Result<()> { pub async fn send<T: serde::Serialize>(&self, msg: &T) -> crate::Result<()> {
let payload = serde_json::to_vec(msg)?; self.send_frame(0x1, &serde_json::to_vec(msg)?).await
self.send_frame(0x1, &payload).await
} }
pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> { pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> {
@@ -73,19 +91,15 @@ impl Session {
} }
pub async fn send_ping(&self) -> crate::Result<()> { pub async fn send_ping(&self) -> crate::Result<()> {
let mut writer = self.writer.lock().await; self.send_frame(0x9, &[]).await
// FIN + opcode = 0x89 (ping), payload length = 0
writer.write_all(&[0x89, 0x00]).await?;
writer.flush().await?;
Ok(())
} }
pub async fn send_pong(&self) -> crate::Result<()> { pub async fn send_pong(&self) -> crate::Result<()> {
let mut writer = self.writer.lock().await; self.send_frame(0xA, &[]).await
// FIN + opcode = 0x8A (pong), payload length = 0 }
writer.write_all(&[0x8A, 0x00]).await?;
writer.flush().await?; pub async fn close(&self) -> crate::Result<()> {
Ok(()) self.send_frame(0x8, &[]).await
} }
pub fn start_ping(self: Arc<Self>) { pub fn start_ping(self: Arc<Self>) {
@@ -99,17 +113,6 @@ impl Session {
} }
}); });
} }
/// Send a close frame and flush
pub async fn close(&self) -> crate::Result<()> {
let mut writer = self.writer.lock().await;
// FIN=1, opcode=0x8 (close), payload length=0
let frame = [0x88, 0x00];
writer.write_all(&frame).await?;
writer.flush().await?;
Ok(())
}
} }
impl Session { impl Session {