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 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 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 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 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 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 let server = listen(
334 addr,
335 codec.clone(),
336 handler.clone(),
337 || async move { Ok(()) },
338 )
339 .await
340 .unwrap();
341
342 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}