From e1585bedab15c3d3afe6a15366a35341d90ee076 Mon Sep 17 00:00:00 2001 From: Leo dev Date: Sun, 14 Sep 2025 11:08:02 +0200 Subject: [PATCH] Custom websocket --- Cargo.toml | 2 + cli/Cargo.lock | 8 +++ src/client.rs | 180 +++++++++++++++++++++++++++++++++++++++---------- src/lib.rs | 2 +- src/types.rs | 3 +- 5 files changed, 157 insertions(+), 38 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index bd412a9..fc91b01 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,12 +5,14 @@ edition = "2024" [dependencies] anyhow = "1.0.99" +base64 = "0.22.1" chrono = "0.4.42" libloading = { version = "0.8.8", optional = true } once_cell = "1.21.3" rusqlite = "0.37.0" serde = { version = "1.0.219", features = ["serde_derive"] } serde_json = "1.0.143" +sha1 = "0.10.6" tungstenite = "0.27.0" [features] diff --git a/cli/Cargo.lock b/cli/Cargo.lock index d83c247..76abc67 100644 --- a/cli/Cargo.lock +++ b/cli/Cargo.lock @@ -23,6 +23,12 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" +[[package]] +name = "base64" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" + [[package]] name = "bitflags" version = "2.9.4" @@ -539,12 +545,14 @@ name = "voxa-server" version = "0.1.0" dependencies = [ "anyhow", + "base64", "chrono", "libloading", "once_cell", "rusqlite", "serde", "serde_json", + "sha1", "tungstenite", ] diff --git a/src/client.rs b/src/client.rs index 7155d42..d20988f 100644 --- a/src/client.rs +++ b/src/client.rs @@ -1,30 +1,77 @@ use std::{ hash::{Hash, Hasher}, + io::{Read, Write}, net::TcpStream, - sync::{Arc, Mutex}, }; -use anyhow::Error; -use tungstenite::{Message, Utf8Bytes, WebSocket, accept}; +use anyhow::anyhow; +use serde::Serialize; -use crate::types::{ClientMessage, ServerMessage, WsMessage, data::ResponseError}; +use crate::types::{ClientMessage, WsMessage}; -#[derive(Clone)] -pub struct Client(Arc>>); +pub mod handshake { + use base64::Engine; + use base64::engine::general_purpose::STANDARD as Base64; + use sha1::{Digest, Sha1}; + use std::io::{Read, Write}; + use std::net::TcpStream; + + const WS_GUID: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; + + pub fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> { + let mut buffer = [0; 1024]; + let size = stream.read(&mut buffer)?; + let request = String::from_utf8_lossy(&buffer[..size]); + + let key_line = request + .lines() + .find(|line| line.to_lowercase().starts_with("sec-websocket-key")) + .ok_or_else(|| { + std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key") + })?; + + let key = key_line.splitn(2, ':').nth(1).unwrap().trim(); + + let mut hasher = Sha1::new(); + hasher.update(key.as_bytes()); + hasher.update(WS_GUID.as_bytes()); + let hash = hasher.finalize(); + + let accept_key = Base64.encode(hash); + + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\n\ + Upgrade: websocket\r\n\ + Connection: Upgrade\r\n\ + Sec-WebSocket-Accept: {}\r\n\r\n", + accept_key + ); + + stream.write_all(response.as_bytes())?; + stream.flush()?; + + Ok(()) + } +} + +pub struct Client(TcpStream); impl Client { - pub fn new_ws(ws: WebSocket) -> Self { - Self(Arc::new(Mutex::new(ws))) + pub fn new(mut stream: TcpStream) -> crate::Result { + handshake::handle_websocket_handshake(&mut stream)?; + Ok(Client(stream)) } +} - pub fn new_tcp(ws: TcpStream) -> crate::Result { - Ok(Self::new_ws(accept(ws)?)) +impl Clone for Client { + fn clone(&self) -> Self { + Client(self.0.try_clone().expect("failed to clone TcpStream")) } } impl PartialEq for Client { fn eq(&self, other: &Self) -> bool { - Arc::ptr_eq(&self.0, &other.0) + self.0.peer_addr().unwrap() == other.0.peer_addr().unwrap() } } @@ -32,42 +79,105 @@ impl Eq for Client {} impl Hash for Client { fn hash(&self, state: &mut H) { - std::ptr::hash(Arc::as_ptr(&self.0), state) + self.0.peer_addr().unwrap().hash(state); } } impl Client { + /// Read a full WebSocket message, handling fragmentation (FIN) pub fn read(&self) -> crate::Result>> { - match self.0.lock().unwrap().read()? { - Message::Text(t) => { - let v = t.to_string(); - match serde_json::from_str(&v) { - Ok(f) => Ok(Some(WsMessage::Message(f))), - Err(_) => Ok(Some(WsMessage::String(v))), + let mut stream = &self.0; + let mut message_payload = Vec::new(); + let mut final_frame = false; + + while !final_frame { + let mut header = [0u8; 2]; + if stream.read_exact(&mut header).is_err() { + return Ok(None); // connection closed + } + + 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; + + // Extended payload lengths + if payload_len == 126 { + let mut ext_len = [0u8; 2]; + stream.read_exact(&mut ext_len)?; + payload_len = u16::from_be_bytes(ext_len) as u64; + } else if payload_len == 127 { + let mut ext_len = [0u8; 8]; + stream.read_exact(&mut ext_len)?; + payload_len = u64::from_be_bytes(ext_len); + } + + // Mask key (client → server) + let mut mask = [0u8; 4]; + if masked { + stream.read_exact(&mut mask)?; + } + + // Read payload + let mut payload = vec![0u8; payload_len as usize]; + stream.read_exact(&mut payload)?; + + if masked { + for i in 0..payload.len() { + payload[i] ^= mask[i % 4]; } } - Message::Binary(b) => Ok(Some(WsMessage::Binary(b))), + match opcode { + 0x0 | 0x1 | 0x2 => { + // Continuation / Text / Binary + message_payload.extend(payload); + } + 0x8 => return Ok(None), // Close + 0x9 => continue, // Ping → ignore + 0xA => continue, // Pong → ignore + _ => return Err(anyhow!("Unsupported WebSocket opcode: {}", opcode).into()), + } - Message::Close(_) => Ok(None), - - m => Err(Error::msg(format!("Invalid websocket format: {m}"))), + final_frame = fin; } + + // Try parsing JSON into ClientMessage + let message = match String::from_utf8(message_payload.clone()) { + Ok(text) => match serde_json::from_str(&text) { + Ok(msg) => WsMessage::Message(msg), + Err(_) => WsMessage::String(text), + }, + Err(_) => WsMessage::Binary(message_payload), + }; + + Ok(Some(message)) } - pub fn send(&self, m: ServerMessage) -> crate::Result<()> { - self.0 - .lock() - .unwrap() - .send(Message::Text(Utf8Bytes::from(serde_json::to_string(&m)?))) - .map_err(|e| e.into()) - } + /// Send a JSON-serializable object as a WebSocket text frame + pub fn send(&self, m: T) -> crate::Result<()> { + let payload = serde_json::to_string(&m)?; + let payload_bytes = payload.as_bytes(); - pub fn send_err(&self, m: ResponseError) -> crate::Result<()> { - self.0 - .lock() - .unwrap() - .send(Message::Text(Utf8Bytes::from(serde_json::to_string(&m)?))) - .map_err(|e| e.into()) + let mut stream = self.0.try_clone()?; + let mut header = Vec::new(); + header.push(0x81); // FIN=1, opcode=0x1 (text) + + let len = payload_bytes.len(); + if len < 126 { + header.push(len as u8); + } else if len <= 65535 { + header.push(126); + header.extend_from_slice(&(len as u16).to_be_bytes()); + } else { + header.push(127); + header.extend_from_slice(&(len as u64).to_be_bytes()); + } + + stream.write_all(&header)?; + stream.write_all(payload_bytes)?; + stream.flush()?; + + Ok(()) } } diff --git a/src/lib.rs b/src/lib.rs index 0863ed5..120b5a7 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -116,7 +116,7 @@ impl Server { fn handle_client(self: &Arc, stream: TcpStream) -> anyhow::Result<()> { Self::LOGGER.info(format!("New connection: {}", stream.peer_addr()?)); // Initialize client - let client = Client::new_tcp(stream)?; + let client = Client::new(stream)?; // Insert to the set of all connected clients self.clients.lock().unwrap().insert(client.clone()); diff --git a/src/types.rs b/src/types.rs index 9694079..b4f0855 100644 --- a/src/types.rs +++ b/src/types.rs @@ -1,5 +1,4 @@ use serde::{Deserialize, Serialize}; -use tungstenite::Bytes; /// Messages sent *from the client* (user’s app) to the server #[derive(Debug, Clone, Serialize, Deserialize)] @@ -58,7 +57,7 @@ pub enum ServerMessage { #[derive(Debug, Clone)] pub enum WsMessage Deserialize<'de>> { Message(T), - Binary(Bytes), + Binary(Vec), String(String), }