This commit is contained in:
2026-02-23 05:28:19 +01:00
parent 5bd9dac052
commit a26a6b75a0
5 changed files with 72 additions and 8 deletions
Generated
+1 -1
View File
@@ -314,7 +314,7 @@ dependencies = [
[[package]] [[package]]
name = "session-rs" name = "session-rs"
version = "0.1.0" version = "0.1.1"
dependencies = [ dependencies = [
"base64", "base64",
"rand", "rand",
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "session-rs" name = "session-rs"
version = "0.1.0" version = "0.1.1"
edition = "2024" edition = "2024"
description = "A lightweight async WebSocket protocol" description = "A lightweight async WebSocket protocol"
license = "Apache-2.0" license = "Apache-2.0"
+1 -1
View File
@@ -18,7 +18,7 @@ async fn main() -> session_rs::Result<()> {
server server
.session_loop(async |session, _| { .session_loop(async |session, _| {
session session
.on::<Data, _>(async |_, req| { .on_request::<Data, _>(async |_, req| {
println!("Msg from client: {req}"); println!("Msg from client: {req}");
if req == "invalid_data" { if req == "invalid_data" {
+2 -1
View File
@@ -7,7 +7,8 @@ pub mod session;
pub mod ws; pub mod ws;
pub type Result<T> = std::result::Result<T, Error>; pub type Result<T> = std::result::Result<T, Error>;
pub type BoxFuture<'a> = Pin<Box<dyn Future<Output = Option<(bool, serde_json::Value)>> + Send + 'a>>; pub type BoxFuture<'a, T = Option<(bool, serde_json::Value)>> =
Pin<Box<dyn Future<Output = T> + Send + 'a>>;
pub type MethodHandler = Box<dyn Fn(u32, serde_json::Value) -> BoxFuture<'static> + Send + Sync>; pub type MethodHandler = Box<dyn Fn(u32, serde_json::Value) -> BoxFuture<'static> + Send + Sync>;
pub trait Method { pub trait Method {
+67 -4
View File
@@ -3,7 +3,9 @@ use std::{collections::HashMap, sync::Arc};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tokio::sync::Mutex; use tokio::sync::Mutex;
use tokio::sync::broadcast; use tokio::sync::broadcast;
use tokio::time::timeout;
use crate::BoxFuture;
use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket};
#[derive(Debug, Serialize, Deserialize)] #[derive(Debug, Serialize, Deserialize)]
@@ -32,7 +34,10 @@ pub struct Session {
pub ws: WebSocket, pub ws: WebSocket,
id: Arc<Mutex<u32>>, id: Arc<Mutex<u32>>,
methods: Arc<Mutex<HashMap<String, MethodHandler>>>, methods: Arc<Mutex<HashMap<String, MethodHandler>>>,
on_close_fn:
Arc<Mutex<Option<Box<dyn Fn() -> BoxFuture<'static, Result<(), String>> + Send + Sync>>>>,
tx: broadcast::Sender<(u32, bool, serde_json::Value)>, tx: broadcast::Sender<(u32, bool, serde_json::Value)>,
pong_tx: broadcast::Sender<()>,
} }
impl Session { impl Session {
@@ -41,18 +46,25 @@ impl Session {
ws: self.ws.clone(), ws: self.ws.clone(),
id: self.id.clone(), id: self.id.clone(),
methods: self.methods.clone(), methods: self.methods.clone(),
on_close_fn: self.on_close_fn.clone(),
tx: self.tx.clone(), tx: self.tx.clone(),
pong_tx: self.pong_tx.clone(),
} }
} }
} }
impl Session { impl Session {
pub fn from_ws(ws: WebSocket) -> Self { pub fn from_ws(ws: WebSocket) -> Self {
let (tx, _) = broadcast::channel(8192);
let (pong_tx, _) = broadcast::channel(16);
Self { Self {
ws, ws,
id: Arc::new(Mutex::new(0)), id: Arc::new(Mutex::new(0)),
methods: Arc::new(Mutex::new(HashMap::new())), methods: Arc::new(Mutex::new(HashMap::new())),
tx: broadcast::channel(8192).0, on_close_fn: Arc::new(Mutex::new(None)),
tx,
pong_tx,
} }
} }
@@ -95,14 +107,45 @@ impl Session {
_ => {} _ => {}
} }
} }
Ok(crate::ws::Frame::Pong) => {
let _ = s.pong_tx.send(());
}
Ok(_) => {} Ok(_) => {}
Err(_) => break, Err(_) => {
s.trigger_close().await;
break;
}
}
}
});
}
pub fn start_ping(&self, interval: tokio::time::Duration, timeout_dur: tokio::time::Duration) {
let s = self.clone();
tokio::spawn(async move {
let mut pong_rx = s.pong_tx.subscribe();
loop {
tokio::time::sleep(interval).await;
if s.ws.send_ping().await.is_err() {
s.trigger_close().await;
break;
}
let result = timeout(timeout_dur, pong_rx.recv()).await;
if result.is_err() {
// timeout expired
let _ = s.close().await;
s.trigger_close().await;
break;
} }
} }
}); });
} }
pub async fn on< pub async fn on_request<
M: Method, M: Method,
Fut: Future<Output = Result<M::Response, M::Error>> + Send + 'static, Fut: Future<Output = Result<M::Response, M::Error>> + Send + 'static,
>( >(
@@ -127,6 +170,18 @@ impl Session {
}), }),
); );
} }
pub async fn on_close<Fut>(&self, handler: impl Fn() -> Fut + Send + Sync + 'static)
where
Fut: Future<Output = Result<(), String>> + Send + 'static,
{
let handler = Arc::new(handler);
*self.on_close_fn.lock().await = Some(Box::new(move || {
let handler = handler.clone();
Box::pin(async move { handler().await })
}));
}
} }
impl Session { impl Session {
@@ -192,7 +247,15 @@ impl Session {
.await .await
} }
async fn trigger_close(&self) {
if let Some(handler) = self.on_close_fn.lock().await.as_ref() {
let _ = handler().await;
}
}
pub async fn close(&self) -> crate::Result<()> { pub async fn close(&self) -> crate::Result<()> {
Ok(self.ws.close().await?) let res = self.ws.close().await;
self.trigger_close().await;
Ok(res?)
} }
} }