From bad32bfc3611f5fdb1e0add1d4393ec1f5a68dd8 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 04:15:31 +0100 Subject: [PATCH] Async handler --- examples/client.rs | 4 +++- src/lib.rs | 6 +++++- src/session.rs | 23 ++++++++++++++++------- 3 files changed, 24 insertions(+), 9 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 3d7564d..9cfbd34 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -18,7 +18,9 @@ async fn main() -> session_rs::Result<()> { session.request::(()).await?; - session.on::(|i, d| println!("Ok {i} {d:?}")).await; + session + .on::(async |i, d| println!("Ok {i} {d:?}")) + .await; tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; diff --git a/src/lib.rs b/src/lib.rs index 9913bc4..0101b66 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,5 @@ +use std::pin::Pin; + use serde::{Deserialize, Serialize}; pub mod server; @@ -5,10 +7,12 @@ pub mod session; pub mod ws; pub type Result = std::result::Result; +pub type BoxFuture<'a> = Pin + Send + 'a>>; +pub type MethodHandler = Box BoxFuture<'static> + Send + Sync>; pub trait Method { const NAME: &'static str; - type Request: Serialize + for<'de> Deserialize<'de>; + type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync; type Response: Serialize + for<'de> Deserialize<'de>; } diff --git a/src/session.rs b/src/session.rs index 8c739b0..2b5ea9b 100644 --- a/src/session.rs +++ b/src/session.rs @@ -3,7 +3,7 @@ use std::{collections::HashMap, sync::Arc}; use serde::{Deserialize, Serialize}; use tokio::sync::Mutex; -use crate::{GenericMethod, Method, ws::WebSocket}; +use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; #[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "lowercase", tag = "type")] @@ -27,7 +27,7 @@ pub enum Message { pub struct Session { pub ws: WebSocket, id: Arc>, - methods: Arc>>>, + methods: Arc>>, } impl Session { @@ -68,7 +68,7 @@ impl Session { match msg { Message::Request { id, method, data } => { if let Some(m) = s.methods.lock().await.get(&method) { - (m)(id, data) + (m)(id, data).await } } _ => {} @@ -81,13 +81,22 @@ impl Session { }); } - pub async fn on(&self, handler: impl Fn(u32, M::Request) + Send + Sync + 'static) { + pub async fn on + Send + 'static>( + &self, + handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static, + ) { + let handler = Arc::new(handler); + self.methods.lock().await.insert( M::NAME.to_string(), Box::new(move |id, value| { - if let Ok(req) = serde_json::from_value(value) { - (handler)(id, req) - } + let handler = Arc::clone(&handler); + + Box::pin(async move { + if let Ok(req) = serde_json::from_value(value) { + handler(id, req).await; + } + }) }), ); }