Respond awaiter

This commit is contained in:
2026-02-19 04:46:08 +01:00
parent bad32bfc36
commit e09623e8b9
3 changed files with 57 additions and 16 deletions
+2 -1
View File
@@ -8,6 +8,7 @@ impl Method for Data {
const NAME: &'static str = "data"; const NAME: &'static str = "data";
type Request = (); type Request = ();
type Response = (); type Response = ();
type Error = ();
} }
#[tokio::main(flavor = "current_thread")] #[tokio::main(flavor = "current_thread")]
@@ -16,7 +17,7 @@ async fn main() -> session_rs::Result<()> {
session.start_receiver(); session.start_receiver();
session.request::<Data>(()).await?; println!("{:?}", session.request::<Data>(()).await?);
session session
.on::<Data, _>(async |i, d| println!("Ok {i} {d:?}")) .on::<Data, _>(async |i, d| println!("Ok {i} {d:?}"))
+9
View File
@@ -14,6 +14,7 @@ pub trait Method {
const NAME: &'static str; const NAME: &'static str;
type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync; type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync;
type Response: Serialize + for<'de> Deserialize<'de>; type Response: Serialize + for<'de> Deserialize<'de>;
type Error: Serialize + for<'de> Deserialize<'de>;
} }
pub struct GenericMethod; pub struct GenericMethod;
@@ -22,6 +23,7 @@ impl Method for GenericMethod {
const NAME: &'static str = "generic_do_not_use"; const NAME: &'static str = "generic_do_not_use";
type Request = serde_json::Value; type Request = serde_json::Value;
type Response = serde_json::Value; type Response = serde_json::Value;
type Error = serde_json::Value;
} }
#[derive(Debug)] #[derive(Debug)]
@@ -29,6 +31,7 @@ pub enum Error {
WebSocket(ws::Error), WebSocket(ws::Error),
Json(serde_json::Error), Json(serde_json::Error),
Io(std::io::Error), Io(std::io::Error),
RecvError(tokio::sync::broadcast::error::RecvError),
} }
impl From<ws::Error> for Error { impl From<ws::Error> for Error {
@@ -48,3 +51,9 @@ impl From<serde_json::Error> for Error {
Self::Json(value) Self::Json(value)
} }
} }
impl From<tokio::sync::broadcast::error::RecvError> for Error {
fn from(value: tokio::sync::broadcast::error::RecvError) -> Self {
Self::RecvError(value)
}
}
+45 -14
View File
@@ -2,6 +2,7 @@ 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 crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket};
@@ -15,9 +16,12 @@ pub enum Message<M: Method> {
}, },
Response { Response {
id: u32, id: u32,
error: bool,
result: M::Response, result: M::Response,
}, },
ErrorResponse {
id: u32,
error: M::Error,
},
Notification { Notification {
method: String, method: String,
data: M::Request, data: M::Request,
@@ -28,6 +32,7 @@ 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>>>,
tx: broadcast::Sender<(u32, bool, serde_json::Value)>,
} }
impl Session { impl Session {
@@ -36,6 +41,7 @@ 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(),
tx: self.tx.clone(),
} }
} }
} }
@@ -46,6 +52,7 @@ impl Session {
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,
} }
} }
@@ -71,6 +78,12 @@ impl Session {
(m)(id, data).await (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,26 +129,44 @@ impl Session {
*id *id
} }
pub async fn request<M: Method>(&self, req: M::Request) -> crate::Result<()> { pub async fn request<M: Method>(
&self,
req: M::Request,
) -> crate::Result<std::result::Result<M::Response, M::Error>> {
let id = self.use_id().await;
self.send::<M>(&Message::Request { self.send::<M>(&Message::Request {
id: self.use_id().await, id,
method: M::NAME.to_string(), method: M::NAME.to_string(),
data: req, 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<M: Method>(&self, to: u32, res: M::Response) -> crate::Result<()> {
self.send::<M>(&Message::Response {
id: to,
result: res,
})
.await .await
} }
pub async fn respond<M: Method>( pub async fn respond_error<M: Method>(&self, to: u32, err: M::Error) -> crate::Result<()> {
&self, self.send::<M>(&Message::ErrorResponse { id: to, error: err })
to: u32,
error: bool,
res: M::Response,
) -> crate::Result<()> {
self.send::<M>(&Message::Response {
id: to,
error,
result: res,
})
.await .await
} }