diff --git a/examples/client.rs b/examples/client.rs index 0a7e647..da7847b 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -20,7 +20,11 @@ async fn main() -> session_rs::Result<()> { println!("{:?}", session.request::(()).await?); session - .on::(async |i, d| println!("Ok {i} {d:?}")) + .on::(async |i, d| { + println!("Ok {i} {d:?}"); + + Ok(()) + }) .await; tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await; diff --git a/examples/server.rs b/examples/server.rs index d515a16..770f39d 100644 --- a/examples/server.rs +++ b/examples/server.rs @@ -1,7 +1,17 @@ -use std::sync::Arc; +use serde::{Deserialize, Serialize}; use tokio::net::TcpListener; -use session_rs::ws::{Frame, WebSocket}; +use session_rs::{Method, session::Session, ws::WebSocket}; + +#[derive(Debug, Serialize, Deserialize)] +struct Data; + +impl Method for Data { + const NAME: &'static str = "data"; + type Request = (); + type Response = (); + type Error = (); +} #[tokio::main(flavor = "current_thread")] async fn main() -> session_rs::Result<()> { @@ -14,35 +24,15 @@ async fn main() -> session_rs::Result<()> { tokio::spawn(async move { // Wrap session in Arc so tasks can share it - let session = match WebSocket::handshake(stream).await { - Ok(s) => Arc::new(s), - Err(e) => { - eprintln!("Handshake failed: {:?}", e); - return; - } - }; + let session = Session::from_ws( + WebSocket::handshake(stream) + .await + .expect("Failed to initialize websocket"), + ); - session.start_ping_loop(); + session.start_receiver(); - // Read loop - loop { - match session.read().await { - Ok(Frame::Text(text)) => { - println!("Received text: {}", text); - - // Echo back - if let Err(e) = session.send(&text).await { - eprintln!("Send error: {:?}", e); - break; - } - } - Ok(_) => {} - Err(e) => { - eprintln!("{e:?}"); - break; - } - } - } + session.on::(async |_, _| Ok(())).await; }); } } diff --git a/src/lib.rs b/src/lib.rs index 86365df..310676f 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -7,7 +7,7 @@ pub mod session; pub mod ws; pub type Result = std::result::Result; -pub type BoxFuture<'a> = Pin + Send + 'a>>; +pub type BoxFuture<'a> = 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 7aec8d1..ef8e1a3 100644 --- a/src/session.rs +++ b/src/session.rs @@ -75,7 +75,14 @@ impl Session { match msg { Message::Request { id, method, data } => { if let Some(m) = s.methods.lock().await.get(&method) { - (m)(id, data).await + (m)(id, data).await; + // let s = s.clone(); + // if let Ok(req) = serde_json::from_value(value) { + // match handler(id, req).await { + // Ok(res) => s.respond::(id, res).await, + // Err(res) => s.respond_error::(id, res).await, + // }; + // } } } Message::Response { id, result } => { @@ -94,7 +101,10 @@ impl Session { }); } - pub async fn on + Send + 'static>( + pub async fn on< + M: Method, + Fut: Future> + Send + 'static, + >( &self, handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static, ) { @@ -106,9 +116,12 @@ impl Session { let handler = Arc::clone(&handler); Box::pin(async move { - if let Ok(req) = serde_json::from_value(value) { - handler(id, req).await; - } + Some( + match handler(id, serde_json::from_value(value).ok()?).await { + Ok(v) => (false, serde_json::to_value(v).ok()?), + Err(v) => (true, serde_json::to_value(v).ok()?), + }, + ) }) }), );