Async session
This commit is contained in:
+1
-1
@@ -9,4 +9,4 @@ rand = "0.10.0"
|
|||||||
serde = "1.0.228"
|
serde = "1.0.228"
|
||||||
serde_json = "1.0.149"
|
serde_json = "1.0.149"
|
||||||
sha1 = "0.10.6"
|
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
@@ -1,90 +1,102 @@
|
|||||||
use std::{
|
|
||||||
hash::{Hash, Hasher},
|
|
||||||
io::{self, Read, Write},
|
|
||||||
net::TcpStream,
|
|
||||||
};
|
|
||||||
|
|
||||||
use serde::{Deserialize, Serialize};
|
|
||||||
|
|
||||||
use crate::SessionMessage;
|
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 {
|
impl Session {
|
||||||
/// Create a client
|
pub async fn new(mut stream: TcpStream) -> crate::Result<Self> {
|
||||||
pub fn new(mut stream: TcpStream) -> crate::Result<Self> {
|
crate::handshake::handle_websocket_handshake(&mut stream).await?;
|
||||||
crate::handshake::handle_websocket_handshake(&mut stream)?;
|
|
||||||
stream.set_read_timeout(Some(std::time::Duration::from_secs(10)))?;
|
let (read, write) = stream.into_split();
|
||||||
stream.set_write_timeout(Some(std::time::Duration::from_secs(10)))?;
|
|
||||||
Ok(Session(stream, rand::random()))
|
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 async fn send_bin(&self, payload: &[u8]) -> crate::Result<()> {
|
||||||
pub fn send_close(&self) -> crate::Result<()> {
|
self.send_frame(0x2, payload).await
|
||||||
let mut stream = self.0.try_clone()?;
|
}
|
||||||
stream.write_all(&[0x88])?;
|
|
||||||
stream.flush()?;
|
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(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Send a ping (no payload)
|
pub async fn send_pong(&self) -> crate::Result<()> {
|
||||||
fn send_ping(&self) -> crate::Result<()> {
|
let mut writer = self.writer.lock().await;
|
||||||
let mut stream = self.0.try_clone()?;
|
// FIN + opcode = 0x8A (pong), payload length = 0
|
||||||
// FIN + opcode (ping = 0x89), payload length = 0x00
|
writer.write_all(&[0x8A, 0x00]).await?;
|
||||||
stream.write_all(&[0x89, 0x00])?;
|
writer.flush().await?;
|
||||||
stream.flush()?;
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Send a pong (no payload)
|
pub fn start_ping(self: Arc<Self>) {
|
||||||
fn send_pong(&self) -> crate::Result<()> {
|
tokio::task::spawn(async move {
|
||||||
let mut stream = self.0.try_clone()?;
|
let mut interval = tokio::time::interval(std::time::Duration::from_secs(15));
|
||||||
// FIN + opcode (pong = 0x8A), payload length = 0x00
|
loop {
|
||||||
stream.write_all(&[0x8A, 0x00])?;
|
interval.tick().await;
|
||||||
stream.flush()?;
|
if self.send_ping().await.is_err() {
|
||||||
Ok(())
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Send a text/binary frame (server->client must NOT mask)
|
async fn send_frame(&self, opcode: u8, payload: &[u8]) -> crate::Result<()> {
|
||||||
pub fn send<T: Serialize>(&self, m: T) -> crate::Result<()> {
|
let mut writer = self.writer.lock().await;
|
||||||
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()?;
|
|
||||||
|
|
||||||
let mut header = Vec::with_capacity(10);
|
let mut header = Vec::with_capacity(10);
|
||||||
|
header.push(0x80 | opcode);
|
||||||
// FIN=1, opcode=2 (binary)
|
|
||||||
header.push(0x82);
|
|
||||||
|
|
||||||
let len = payload.len();
|
let len = payload.len();
|
||||||
|
|
||||||
if len < 126 {
|
if len < 126 {
|
||||||
header.push(len as u8); // mask bit = 0
|
header.push(len as u8);
|
||||||
} else if len <= 0xFFFF {
|
} else if len <= 0xFFFF {
|
||||||
header.push(126);
|
header.push(126);
|
||||||
header.extend_from_slice(&(len as u16).to_be_bytes());
|
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());
|
header.extend_from_slice(&(len as u64).to_be_bytes());
|
||||||
}
|
}
|
||||||
|
|
||||||
stream.write_all(&header)?;
|
writer.write_all(&header).await?;
|
||||||
stream.write_all(payload)?;
|
writer.write_all(payload).await?;
|
||||||
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)?;
|
|
||||||
Ok(())
|
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);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
Reference in New Issue
Block a user