From e09623e8b9b333e91a97f8bac2caea8f15f349b4 Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Thu, 19 Feb 2026 04:46:08 +0100 Subject: [PATCH] Respond awaiter --- examples/client.rs | 3 ++- src/lib.rs | 9 +++++++ src/session.rs | 61 ++++++++++++++++++++++++++++++++++------------ 3 files changed, 57 insertions(+), 16 deletions(-) diff --git a/examples/client.rs b/examples/client.rs index 9cfbd34..0a7e647 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -8,6 +8,7 @@ impl Method for Data { const NAME: &'static str = "data"; type Request = (); type Response = (); + type Error = (); } #[tokio::main(flavor = "current_thread")] @@ -16,7 +17,7 @@ async fn main() -> session_rs::Result<()> { session.start_receiver(); - session.request::(()).await?; + println!("{:?}", session.request::(()).await?); session .on::(async |i, d| println!("Ok {i} {d:?}")) diff --git a/src/lib.rs b/src/lib.rs index 0101b66..86365df 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -14,6 +14,7 @@ pub trait Method { const NAME: &'static str; type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync; type Response: Serialize + for<'de> Deserialize<'de>; + type Error: Serialize + for<'de> Deserialize<'de>; } pub struct GenericMethod; @@ -22,6 +23,7 @@ impl Method for GenericMethod { const NAME: &'static str = "generic_do_not_use"; type Request = serde_json::Value; type Response = serde_json::Value; + type Error = serde_json::Value; } #[derive(Debug)] @@ -29,6 +31,7 @@ pub enum Error { WebSocket(ws::Error), Json(serde_json::Error), Io(std::io::Error), + RecvError(tokio::sync::broadcast::error::RecvError), } impl From for Error { @@ -48,3 +51,9 @@ impl From for Error { Self::Json(value) } } + +impl From for Error { + fn from(value: tokio::sync::broadcast::error::RecvError) -> Self { + Self::RecvError(value) + } +} diff --git a/src/session.rs b/src/session.rs index 2b5ea9b..7aec8d1 100644 --- a/src/session.rs +++ b/src/session.rs @@ -2,6 +2,7 @@ use std::{collections::HashMap, sync::Arc}; use serde::{Deserialize, Serialize}; use tokio::sync::Mutex; +use tokio::sync::broadcast; use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; @@ -15,9 +16,12 @@ pub enum Message { }, Response { id: u32, - error: bool, result: M::Response, }, + ErrorResponse { + id: u32, + error: M::Error, + }, Notification { method: String, data: M::Request, @@ -28,6 +32,7 @@ pub struct Session { pub ws: WebSocket, id: Arc>, methods: Arc>>, + tx: broadcast::Sender<(u32, bool, serde_json::Value)>, } impl Session { @@ -36,6 +41,7 @@ impl Session { ws: self.ws.clone(), id: self.id.clone(), methods: self.methods.clone(), + tx: self.tx.clone(), } } } @@ -46,6 +52,7 @@ impl Session { ws, id: Arc::new(Mutex::new(0)), methods: Arc::new(Mutex::new(HashMap::new())), + tx: broadcast::channel(8192).0, } } @@ -71,6 +78,12 @@ impl Session { (m)(id, data).await } } + Message::Response { id, result } => { + s.tx.send((id, false, result)).unwrap(); + } + Message::ErrorResponse { id, error } => { + s.tx.send((id, true, error)).unwrap(); + } _ => {} } } @@ -116,27 +129,45 @@ impl Session { *id } - pub async fn request(&self, req: M::Request) -> crate::Result<()> { + pub async fn request( + &self, + req: M::Request, + ) -> crate::Result> { + let id = self.use_id().await; + self.send::(&Message::Request { - id: self.use_id().await, + id, method: M::NAME.to_string(), data: req, }) + .await?; + + let mut rx = self.tx.subscribe(); + + loop { + let r = rx.recv().await?; + + if r.0 == id { + break Ok(if r.1 { + Err(serde_json::from_value(r.2)?) + } else { + Ok(serde_json::from_value(r.2)?) + }); + } + } + } + + pub async fn respond(&self, to: u32, res: M::Response) -> crate::Result<()> { + self.send::(&Message::Response { + id: to, + result: res, + }) .await } - pub async fn respond( - &self, - to: u32, - error: bool, - res: M::Response, - ) -> crate::Result<()> { - self.send::(&Message::Response { - id: to, - error, - result: res, - }) - .await + pub async fn respond_error(&self, to: u32, err: M::Error) -> crate::Result<()> { + self.send::(&Message::ErrorResponse { id: to, error: err }) + .await } pub async fn notify(&self, data: M::Request) -> crate::Result<()> {