Async session

This commit is contained in:
2026-02-18 21:33:48 +01:00
parent 3c254b75ff
commit dd1d261353
2 changed files with 86 additions and 260 deletions
+1 -1
View File
@@ -9,4 +9,4 @@ rand = "0.10.0"
serde = "1.0.228"
serde_json = "1.0.149"
sha1 = "0.10.6"
tokio = { version = "1.49.0", features = ["io-util", "net"] }
tokio = { version = "1.49.0", features = ["io-util", "net", "rt", "sync", "time"] }
+82 -256
View File
@@ -1,90 +1,102 @@
use std::{
hash::{Hash, Hasher},
io::{self, Read, Write},
net::TcpStream,
};
use serde::{Deserialize, Serialize};
use crate::SessionMessage;
pub struct Session(TcpStream, u64);
use std::{
hash::{Hash, Hasher},
sync::Arc,
};
use tokio::{io::AsyncWriteExt, net::TcpStream, sync::Mutex};
pub struct Session {
reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>,
writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>,
id: u64,
}
impl Session {
/// Create a client
pub fn new(mut stream: TcpStream) -> crate::Result<Self> {
crate::handshake::handle_websocket_handshake(&mut stream)?;
stream.set_read_timeout(Some(std::time::Duration::from_secs(10)))?;
stream.set_write_timeout(Some(std::time::Duration::from_secs(10)))?;
Ok(Session(stream, rand::random()))
pub async fn new(mut stream: TcpStream) -> crate::Result<Self> {
crate::handshake::handle_websocket_handshake(&mut stream).await?;
let (read, write) = stream.into_split();
Ok(Self {
reader: Arc::new(Mutex::new(read)),
writer: Arc::new(Mutex::new(write)),
id: rand::random(),
})
}
}
impl Clone for Session {
fn clone(&self) -> Self {
Session {
reader: self.reader.clone(),
writer: self.writer.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 {
pub async fn send<T: serde::Serialize>(&self, msg: &T) -> crate::Result<()> {
let payload = serde_json::to_vec(msg)?;
self.send_frame(0x1, &payload).await
}
/// Send a close frame and flush.
pub fn send_close(&self) -> crate::Result<()> {
let mut stream = self.0.try_clone()?;
stream.write_all(&[0x88])?;
stream.flush()?;
pub async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> {
self.send_frame(0x2, payload).await
}
pub async fn send_ping(&self) -> crate::Result<()> {
let mut writer = self.writer.lock().await;
// FIN + opcode = 0x89 (ping), payload length = 0
writer.write_all(&[0x89, 0x00]).await?;
writer.flush().await?;
Ok(())
}
/// Send a ping (no payload)
fn send_ping(&self) -> crate::Result<()> {
let mut stream = self.0.try_clone()?;
// FIN + opcode (ping = 0x89), payload length = 0x00
stream.write_all(&[0x89, 0x00])?;
stream.flush()?;
pub async fn send_pong(&self) -> crate::Result<()> {
let mut writer = self.writer.lock().await;
// FIN + opcode = 0x8A (pong), payload length = 0
writer.write_all(&[0x8A, 0x00]).await?;
writer.flush().await?;
Ok(())
}
/// Send a pong (no payload)
fn send_pong(&self) -> crate::Result<()> {
let mut stream = self.0.try_clone()?;
// FIN + opcode (pong = 0x8A), payload length = 0x00
stream.write_all(&[0x8A, 0x00])?;
stream.flush()?;
Ok(())
pub fn start_ping(self: Arc<Self>) {
tokio::task::spawn(async move {
let mut interval = tokio::time::interval(std::time::Duration::from_secs(15));
loop {
interval.tick().await;
if self.send_ping().await.is_err() {
break;
}
}
});
}
/// Send a text/binary frame (server->client must NOT mask)
pub fn send<T: Serialize>(&self, m: T) -> crate::Result<()> {
let mut stream = self.0.try_clone()?;
let payload = serde_json::to_string(&m)?;
let payload_bytes = payload.as_bytes();
let len = payload_bytes.len();
let mut header = Vec::new();
header.push(0x81); // FIN=1, opcode=0x1 (text)
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(())
}
/// Send a binary WebSocket frame (server -> client)
pub fn send_bin(&self, payload: &[u8]) -> crate::Result<()> {
let mut stream = self.0.try_clone()?;
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);
// FIN=1, opcode=2 (binary)
header.push(0x82);
header.push(0x80 | opcode);
let len = payload.len();
if len < 126 {
header.push(len as u8); // mask bit = 0
header.push(len as u8);
} else if len <= 0xFFFF {
header.push(126);
header.extend_from_slice(&(len as u16).to_be_bytes());
@@ -93,194 +105,8 @@ impl Session {
header.extend_from_slice(&(len as u64).to_be_bytes());
}
stream.write_all(&header)?;
stream.write_all(payload)?;
stream.flush()?;
Ok(())
}
/// Read a full WebSocket message, handling fragmentation and control frames.
///
/// Returns:
/// - Ok(Some(WsMessage)) on an application message (text/binary)
/// - Ok(None) if the connection should be closed (close received / read EOF)
/// - Err on protocol or IO errors.
pub fn read_t<T: Serialize + for<'de> Deserialize<'de>>(
&self,
) -> crate::Result<Option<SessionMessage<T>>> {
let mut stream = self.0.try_clone()?;
let mut message_payload = Vec::new();
let mut expecting_continuation = false;
let mut message_type: Option<u8> = None; // 0x1 for text, 0x2 for binary
loop {
// Read 2-byte header
let mut header = [0u8; 2];
match stream.read_exact(&mut header) {
Ok(_) => {}
Err(e) => match e.kind() {
io::ErrorKind::WouldBlock | io::ErrorKind::TimedOut => {
self.send_ping()?;
continue;
}
io::ErrorKind::UnexpectedEof | io::ErrorKind::BrokenPipe => return Ok(None),
_ => return Err(e.into()),
},
}
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 length
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
let mut mask = [0u8; 4];
if masked {
stream.read_exact(&mut mask)?;
} else {
let _ = self.send_close();
return Ok(None);
}
// Control frame checks
if matches!(opcode, 0x8 | 0x9 | 0xA) {
if payload_len > 125 {
let _ = self.send_close();
return Ok(None);
}
if !fin {
let _ = self.send_close();
return Ok(None);
}
}
// Read payload
let mut payload = vec![0u8; payload_len as usize];
if payload_len > 0 {
stream.read_exact(&mut payload)?;
for i in 0..payload.len() {
payload[i] ^= mask[i % 4];
}
}
match opcode {
0x0 => {
// Continuation
if !expecting_continuation {
let _ = self.send_close();
return Ok(None);
}
message_payload.extend(payload);
if fin {
break;
}
}
0x1 => {
// Text
if expecting_continuation {
let _ = self.send_close();
return Ok(None);
}
message_payload.extend(payload);
message_type = Some(0x1);
if fin {
break;
} else {
expecting_continuation = true;
}
}
0x2 => {
// Binary
if expecting_continuation {
let _ = self.send_close();
return Ok(None);
}
message_payload.extend(payload);
message_type = Some(0x2);
if fin {
break;
} else {
expecting_continuation = true;
}
}
0x8 => {
// Close
let _ = self.send_close();
return Ok(None);
}
0x9 => {
// Ping
self.send_pong()?;
continue;
}
0xA => {
// Pong
continue;
}
_ => {
let _ = self.send_close();
return Ok(None);
}
}
}
// Convert payload into proper message type
let message = match message_type {
Some(0x1) => {
// Text frame → try JSON, otherwise keep text
match String::from_utf8(message_payload.clone()) {
Ok(text) => match serde_json::from_str(&text) {
Ok(msg) => SessionMessage::SessionMessage(msg),
Err(e) => return Err(crate::Error::Json(e)),
},
Err(_) => SessionMessage::Binary(message_payload),
}
}
Some(0x2) => SessionMessage::Binary(message_payload),
_ => return Ok(None), // Should not happen
};
Ok(Some(message))
}
pub fn close(&self) -> crate::Result<()> {
self.0.shutdown(std::net::Shutdown::Both)?;
writer.write_all(&header).await?;
writer.write_all(payload).await?;
Ok(())
}
}
impl Clone for Session {
fn clone(&self) -> Self {
Session(
self.0.try_clone().expect("failed to clone TcpStream"),
self.1.clone(),
)
}
}
impl PartialEq for Session {
fn eq(&self, other: &Self) -> bool {
self.1 == other.1
}
}
impl Eq for Session {}
impl Hash for Session {
fn hash<H: Hasher>(&self, state: &mut H) {
self.1.hash(state);
}
}