From a26a6b75a0a32faab218e2f7de8e9077b11abe9c Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Mon, 23 Feb 2026 05:28:19 +0100 Subject: [PATCH] Pinging --- Cargo.lock | 2 +- Cargo.toml | 2 +- examples/server.rs | 2 +- src/lib.rs | 3 +- src/session.rs | 71 +++++++++++++++++++++++++++++++++++++++++++--- 5 files changed, 72 insertions(+), 8 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 743c953..5482269 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -314,7 +314,7 @@ dependencies = [ [[package]] name = "session-rs" -version = "0.1.0" +version = "0.1.1" dependencies = [ "base64", "rand", diff --git a/Cargo.toml b/Cargo.toml index 308b35e..39dbf61 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "session-rs" -version = "0.1.0" +version = "0.1.1" edition = "2024" description = "A lightweight async WebSocket protocol" license = "Apache-2.0" diff --git a/examples/server.rs b/examples/server.rs index 77583b1..f5e1362 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -18,7 +18,7 @@ async fn main() -> session_rs::Result<()> { server .session_loop(async |session, _| { session - .on::(async |_, req| { + .on_request::(async |_, req| { println!("Msg from client: {req}"); if req == "invalid_data" { diff --git a/src/lib.rs b/src/lib.rs index 310676f..b4fd668 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -7,7 +7,8 @@ pub mod session; pub mod ws; pub type Result = std::result::Result; -pub type BoxFuture<'a> = Pin> + Send + 'a>>; +pub type BoxFuture<'a, T = Option<(bool, serde_json::Value)>> = + Pin + Send + 'a>>; pub type MethodHandler = Box BoxFuture<'static> + Send + Sync>; pub trait Method { diff --git a/src/session.rs b/src/session.rs index d20c7ca..932c60e 100644 --- a/src/session.rs +++ b/src/session.rs @@ -3,7 +3,9 @@ use std::{collections::HashMap, sync::Arc}; use serde::{Deserialize, Serialize}; use tokio::sync::Mutex; use tokio::sync::broadcast; +use tokio::time::timeout; +use crate::BoxFuture; use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; #[derive(Debug, Serialize, Deserialize)] @@ -32,7 +34,10 @@ pub struct Session { pub ws: WebSocket, id: Arc>, methods: Arc>>, + on_close_fn: + Arc BoxFuture<'static, Result<(), String>> + Send + Sync>>>>, tx: broadcast::Sender<(u32, bool, serde_json::Value)>, + pong_tx: broadcast::Sender<()>, } impl Session { @@ -41,18 +46,25 @@ impl Session { ws: self.ws.clone(), id: self.id.clone(), methods: self.methods.clone(), + on_close_fn: self.on_close_fn.clone(), tx: self.tx.clone(), + pong_tx: self.pong_tx.clone(), } } } impl Session { pub fn from_ws(ws: WebSocket) -> Self { + let (tx, _) = broadcast::channel(8192); + let (pong_tx, _) = broadcast::channel(16); + Self { ws, id: Arc::new(Mutex::new(0)), 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(_) => {} - 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, Fut: Future> + Send + 'static, >( @@ -127,6 +170,18 @@ impl Session { }), ); } + + pub async fn on_close(&self, handler: impl Fn() -> Fut + Send + Sync + 'static) + where + Fut: Future> + 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 { @@ -192,7 +247,15 @@ impl Session { .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<()> { - Ok(self.ws.close().await?) + let res = self.ws.close().await; + self.trigger_close().await; + Ok(res?) } }