Skip to main content

fluxrpc_core/transport/
websocket.rs

1pub mod client {
2    use crate::codec::Codec;
3    use crate::session::{RpcSession, RpcSessionHandler, SessionState};
4    use crate::transport::{Transport, TransportMessage};
5    use async_trait::async_trait;
6    use ezsockets::{Bytes, Client, Error, Utf8Bytes};
7    use std::sync::Arc;
8    use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
9    use tokio::sync::{Mutex, oneshot};
10
11    pub use ezsockets::ClientConfig as WebsocketClientConfig;
12
13    pub struct ClientHandler {
14        tx: UnboundedSender<TransportMessage>,
15        on_connected: Option<oneshot::Sender<()>>,
16    }
17
18    pub struct ClientTransport {
19        rx: Mutex<UnboundedReceiver<TransportMessage>>,
20        handle: Client<ClientHandler>,
21    }
22
23    #[async_trait]
24    impl ezsockets::ClientExt for ClientHandler {
25        type Call = ();
26
27        async fn on_text(&mut self, text: Utf8Bytes) -> Result<(), Error> {
28            self.tx
29                .send(TransportMessage::Text(text.as_bytes().into()))?;
30            Ok(())
31        }
32
33        async fn on_binary(&mut self, bytes: Bytes) -> Result<(), Error> {
34            self.tx
35                .send(TransportMessage::Binary(bytes.to_vec().into()))?;
36            Ok(())
37        }
38
39        async fn on_call(&mut self, call: Self::Call) -> Result<(), Error> {
40            todo!()
41        }
42
43        async fn on_connect(&mut self) -> Result<(), ezsockets::Error> {
44            if let Some(on_connected) = self.on_connected.take() {
45                on_connected.send(()).ok();
46            }
47            Ok(())
48        }
49    }
50
51    #[async_trait]
52    impl Transport for ClientTransport {
53        async fn send(&self, data: &TransportMessage) -> anyhow::Result<()> {
54            let _ = match data {
55                TransportMessage::Text(data) => {
56                    self.handle.text(Utf8Bytes::try_from(data.clone())?)?
57                }
58                TransportMessage::Binary(data) => self.handle.binary(Bytes::from(data.clone()))?,
59            };
60            Ok(())
61        }
62
63        async fn receive(&self) -> anyhow::Result<TransportMessage> {
64            let mut rx = self.rx.lock().await;
65            rx.recv().await.ok_or(anyhow::anyhow!("socket closed"))
66        }
67    }
68
69    pub async fn connect_transport(
70        config: ezsockets::ClientConfig,
71    ) -> anyhow::Result<ClientTransport> {
72        let (tx_from_socket, rx_from_socket) = unbounded_channel();
73        let (tx_connected, rx_connected) = oneshot::channel::<()>();
74
75        let (handle, _) = ezsockets::connect(
76            move |_handle| ClientHandler {
77                tx: tx_from_socket.clone(),
78                on_connected: Some(tx_connected),
79            },
80            config,
81        )
82        .await;
83
84        // wait until connected
85        rx_connected.await?;
86
87        let transport = ClientTransport {
88            handle,
89            rx: Mutex::new(rx_from_socket),
90        };
91
92        Ok(transport)
93    }
94
95    /// Connect to a websocket server and return a started RPC session
96    pub async fn connect<C, S>(
97        config: ezsockets::ClientConfig,
98        codec: C,
99        handler: Arc<dyn RpcSessionHandler<State = S>>,
100        state: S,
101    ) -> anyhow::Result<Arc<RpcSession<C, ClientTransport, S>>>
102    where
103        C: Codec,
104        S: SessionState,
105    {
106        let transport = connect_transport(config).await?;
107        Ok(RpcSession::create(transport, codec, handler, state))
108    }
109}
110
111pub mod server {
112    use crate::codec::Codec;
113    use crate::session::{RpcSession, RpcSessionHandler, SessionState};
114    use crate::transport::{Transport, TransportMessage};
115    use async_trait::async_trait;
116    use ezsockets::{
117        Bytes, CloseFrame, Error, Request, Server, ServerExt, SessionExt, Socket, Utf8Bytes,
118    };
119    use std::net::SocketAddr;
120    use std::sync::Arc;
121    use tokio::sync::Mutex;
122    use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
123    use tracing::{debug, info};
124
125    type SessionID = u16;
126    type Session = ezsockets::Session<SessionID, ()>;
127
128    struct ServerHandler {
129        tx_accept: UnboundedSender<ServerSessionTransport>,
130        bearer_token: Option<String>,
131    }
132
133    struct ServerSession {
134        id: SessionID,
135        tx: UnboundedSender<TransportMessage>,
136        rx: Option<UnboundedReceiver<TransportMessage>>,
137        tx_accept: UnboundedSender<ServerSessionTransport>,
138        handle: Session,
139    }
140
141    #[async_trait]
142    impl SessionExt for ServerSession {
143        type ID = SessionID;
144        type Call = ();
145
146        fn id(&self) -> &Self::ID {
147            &self.id
148        }
149
150        async fn on_text(&mut self, text: Utf8Bytes) -> Result<(), Error> {
151            self.tx
152                .send(TransportMessage::Text(text.as_bytes().into()))?;
153            Ok(())
154        }
155
156        async fn on_binary(&mut self, bytes: Bytes) -> Result<(), Error> {
157            self.tx
158                .send(TransportMessage::Binary(bytes.to_vec().into()))?;
159            Ok(())
160        }
161
162        async fn on_call(&mut self, call: Self::Call) -> Result<(), Error> {
163            let t = ServerSessionTransport {
164                rx: Mutex::new(self.rx.take().unwrap()),
165                handle: self.handle.clone(),
166            };
167
168            self.tx_accept.send(t).unwrap();
169
170            Ok(())
171        }
172    }
173
174    #[async_trait]
175    impl ServerExt for ServerHandler {
176        type Session = ServerSession;
177        type Call = ();
178
179        async fn on_connect(
180            &mut self,
181            socket: Socket,
182            request: Request,
183            address: SocketAddr,
184        ) -> Result<Session, Option<CloseFrame>> {
185            info!(
186                "session connect uri={} port={}",
187                request.uri(),
188                address.port()
189            );
190
191            // TODO: validate socket
192            // TODO: validate request
193            // TODO: auth
194
195            let (tx, rx) = unbounded_channel();
196
197            let id = address.port();
198            let session = Session::create(
199                |handle| ServerSession {
200                    tx,
201                    id,
202                    rx: Some(rx),
203                    tx_accept: self.tx_accept.clone(),
204                    handle,
205                },
206                id,
207                socket,
208            );
209
210            session.call(()).unwrap();
211
212            Ok(session)
213        }
214
215        async fn on_disconnect(
216            &mut self,
217            id: SessionID,
218            reason: Result<Option<CloseFrame>, Error>,
219        ) -> Result<(), Error> {
220            debug!("session {} disconnected: {:?}", id, reason);
221            Ok(())
222        }
223
224        async fn on_call(&mut self, call: Self::Call) -> Result<(), Error> {
225            todo!()
226        }
227    }
228
229    pub struct ServerSessionTransport {
230        handle: Session,
231        rx: Mutex<UnboundedReceiver<TransportMessage>>,
232    }
233
234    #[async_trait]
235    impl Transport for ServerSessionTransport {
236        async fn send(&self, data: &TransportMessage) -> anyhow::Result<()> {
237            let _ = match data {
238                TransportMessage::Text(data) => {
239                    self.handle.text(Utf8Bytes::try_from(data.to_vec())?)?
240                }
241                TransportMessage::Binary(data) => self.handle.binary(Bytes::from(data.clone()))?,
242            };
243            Ok(())
244        }
245
246        async fn receive(&self) -> anyhow::Result<TransportMessage> {
247            let mut rx = self.rx.lock().await;
248            rx.recv().await.ok_or(anyhow::anyhow!("socket closed"))
249        }
250    }
251
252    pub async fn listen<C, S, F, Fut>(
253        addr: SocketAddr,
254        codec: C,
255        handler: Arc<dyn RpcSessionHandler<State = S>>,
256        state_factory: F,
257    ) -> anyhow::Result<()>
258    where
259        C: Codec,
260        S: SessionState,
261        F: Fn() -> Fut + Sync + Send + 'static,
262        Fut: Future<Output = anyhow::Result<S>> + Send + 'static,
263    {
264        let (tx_accept, mut rx_accept) = unbounded_channel();
265
266        tokio::spawn(async move {
267            let (server, _) = Server::create(|_server| ServerHandler {
268                tx_accept,
269                bearer_token: None,
270            });
271            ezsockets::tungstenite::run(server, addr).await.unwrap();
272        });
273
274        // run session per accepted client
275        tokio::spawn(async move {
276            while let Some(transport) = rx_accept.recv().await {
277                let _ = RpcSession::create(
278                    transport,
279                    codec.clone(),
280                    handler.clone(),
281                    state_factory().await.unwrap(),
282                );
283            }
284        });
285
286        // TODO: return a handle to stop the server
287
288        Ok(())
289    }
290}
291
292#[cfg(test)]
293mod tests {
294
295    use crate::codec::json::JsonCodec;
296    use crate::message::{ErrorBody, Request};
297    use crate::session::{RpcSessionHandler, SessionContext};
298    use crate::transport::websocket::client::connect;
299    use crate::transport::websocket::server::listen;
300    use async_trait::async_trait;
301    use ezsockets::ClientConfig;
302    use nanoid::nanoid;
303    use serde_json::Value;
304    use std::net::SocketAddr;
305    use std::sync::Arc;
306    use std::time::Duration;
307    use url::Url;
308
309    struct TestHandler {}
310
311    #[async_trait]
312    impl RpcSessionHandler for TestHandler {
313        type State = ();
314        async fn on_request(
315            &self,
316            s: Arc<dyn SessionContext<State = Self::State>>,
317            req: Request,
318        ) -> Result<Value, ErrorBody> {
319            match req.method.as_str() {
320                "ping" => Ok(Value::String("pong".to_string())),
321                _ => RpcSessionHandler::on_request(self, s, req).await,
322            }
323        }
324    }
325
326    #[tokio::test]
327    async fn client_server() {
328        let addr: SocketAddr = "127.0.0.1:8080".parse().unwrap();
329        let codec = JsonCodec::new();
330        let handler = Arc::new(TestHandler {});
331
332        // server
333        let server = listen(
334            addr,
335            codec.clone(),
336            handler.clone(),
337            || async move { Ok(()) },
338        )
339        .await
340        .unwrap();
341
342        // client
343        let client_url = Url::parse(format!("ws://{}", addr).as_str()).unwrap();
344        let client_config = ClientConfig::new(client_url);
345        let client = connect(client_config, codec, handler.clone(), ())
346            .await
347            .unwrap();
348
349        let response = client
350            .request(
351                &Request {
352                    id: nanoid!(),
353                    method: "ping".to_string(),
354                    params: None,
355                },
356                Duration::from_millis(100).into(),
357            )
358            .await
359            .unwrap();
360
361        assert_eq!(response.result, Value::String("pong".to_string()));
362    }
363}