Skip to main content

fluxrpc_core/
session.rs

1use crate::codec::Codec;
2use crate::message::{
3    ErrorBody, Event, Message, Request, RequestError, RequestResult, Response, StandardErrorCode,
4};
5use crate::transport::{Transport, TransportMessage};
6use async_trait::async_trait;
7use dashmap::DashMap;
8use serde_json::Value;
9use std::fmt;
10use std::sync::Arc;
11use std::time::Duration;
12use tokio::sync::{mpsc, oneshot};
13use tokio::time;
14use tracing::{debug, error};
15
16pub trait SessionState: Send + Sync + 'static {}
17
18impl SessionState for () {}
19
20#[derive(Debug)]
21pub enum RpcSessionError {
22    Transport(anyhow::Error),
23    Request(RequestError),
24}
25
26impl fmt::Display for RpcSessionError {
27    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
28        match self {
29            RpcSessionError::Transport(e) => write!(f, "transport error: {}", e),
30            RpcSessionError::Request(e) => write!(f, "request error: {}", e),
31        }
32    }
33}
34
35impl std::error::Error for RpcSessionError {}
36
37pub enum HandlerError {
38    Unimplemented { method: String },
39}
40
41impl Into<ErrorBody> for HandlerError {
42    fn into(self) -> ErrorBody {
43        match self {
44            Self::Unimplemented { method } => ErrorBody {
45                message: format!("Method [{}] is not implemented", method),
46                code: StandardErrorCode::NotImplemented.into(),
47                data: None,
48            },
49        }
50    }
51}
52
53#[async_trait]
54pub trait SessionContext: Sync + Send {
55    type State: SessionState;
56
57    fn state(&self) -> &Self::State;
58
59    async fn send_binary(&self, data: Vec<u8>) -> anyhow::Result<()>;
60    async fn notify(&self, event: &Event) -> anyhow::Result<()>;
61    async fn request(
62        &self,
63        request: &Request,
64        timeout: Option<Duration>,
65    ) -> Result<RequestResult, RpcSessionError>;
66}
67
68#[async_trait]
69pub trait RpcSessionHandler: Send + Sync + 'static {
70    type State: SessionState;
71
72    async fn on_open(&self, s: Arc<dyn SessionContext<State = Self::State>>) -> anyhow::Result<()> {
73        Ok(())
74    }
75    async fn on_close(
76        &self,
77        s: Arc<dyn SessionContext<State = Self::State>>,
78    ) -> anyhow::Result<()> {
79        Ok(())
80    }
81    async fn on_data(
82        &self,
83        s: Arc<dyn SessionContext<State = Self::State>>,
84        data: Vec<u8>,
85    ) -> anyhow::Result<()> {
86        Ok(())
87    }
88    async fn on_event(
89        &self,
90        s: Arc<dyn SessionContext<State = Self::State>>,
91        evt: Event,
92    ) -> anyhow::Result<()> {
93        Ok(())
94    }
95    async fn on_request(
96        &self,
97        s: Arc<dyn SessionContext<State = Self::State>>,
98        req: Request,
99    ) -> Result<Value, ErrorBody> {
100        Err(HandlerError::Unimplemented { method: req.method }.into())
101    }
102}
103
104pub struct RpcSession<C, T, S>
105where
106    C: Codec,
107    T: Transport,
108    S: SessionState,
109{
110    state: Arc<S>,
111    transport: Arc<T>,
112    codec: C,
113    pending_requests: Arc<DashMap<String, oneshot::Sender<Response>>>,
114    handler: Arc<dyn RpcSessionHandler<State = S>>,
115    _foo: std::marker::PhantomData<S>,
116}
117
118impl<C, T, S> RpcSession<C, T, S>
119where
120    T: Transport,
121    C: Codec,
122    S: SessionState,
123{
124    pub fn create(
125        transport: T,
126        codec: C,
127        handler: Arc<dyn RpcSessionHandler<State = S>>,
128        state: S,
129    ) -> Arc<Self> {
130        let s = Arc::new(Self {
131            codec,
132            transport: Arc::new(transport),
133            pending_requests: Arc::new(DashMap::new()),
134            handler,
135            state: Arc::new(state),
136            _foo: std::marker::PhantomData,
137        });
138
139        let s1 = s.clone();
140        tokio::spawn(async move {
141            s1.start().await;
142        });
143
144        s
145    }
146
147    pub async fn start(self: Arc<Self>) {
148        let session = self.clone();
149        tokio::spawn(async move {
150            session.run().await;
151        });
152    }
153
154    pub async fn notify(&self, event: &Event) -> anyhow::Result<()> {
155        let msg = Message::Event(event.clone());
156        let data = self.codec.encode(&msg)?;
157        self.transport.send(&TransportMessage::Text(data)).await
158    }
159
160    pub async fn send_binary(&self, data: Vec<u8>) -> anyhow::Result<()> {
161        self.transport.send(&TransportMessage::Binary(data)).await
162    }
163
164    pub async fn request(
165        &self,
166        request: &Request,
167        timeout: Option<Duration>,
168    ) -> Result<RequestResult, RpcSessionError> {
169        let started_at = time::Instant::now();
170
171        debug!("Sending request {request:?}");
172        let msg = Message::Request(request.clone());
173        let data = self
174            .codec
175            .encode(&msg)
176            .map_err(|err| RpcSessionError::Transport(err))?;
177
178        let (tx, rx) = oneshot::channel();
179        let id = request.id.clone();
180        {
181            self.pending_requests.insert(id.clone(), tx);
182        }
183
184        self.transport
185            .send(&TransportMessage::Text(data))
186            .await
187            .map_err(RpcSessionError::Transport)?;
188
189        let result = match timeout {
190            Some(dur) => time::timeout(dur, rx).await.map_err(|_| {
191                RpcSessionError::Request(RequestError {
192                    id: id.clone(),
193                    error: ErrorBody::timeout(),
194                })
195            })?,
196            None => rx.await,
197        };
198
199        let took = started_at.elapsed().as_micros();
200        debug!("Request {request:?} took {took} microseconds");
201
202        match result {
203            Ok(Response::Ok(r)) => Ok(r),
204            Ok(Response::Error(e)) => Err(RpcSessionError::Request(e)),
205            Err(err) => Err(RpcSessionError::Request(RequestError {
206                id: id.clone(),
207                error: ErrorBody::internal_error(err.to_string()),
208            })),
209        }
210    }
211
212    async fn handle_msg(
213        codec: C,
214        handler: Arc<dyn RpcSessionHandler<State = S>>,
215        handle: Arc<dyn SessionContext<State = S>>,
216        transport: Arc<T>,
217        pending: Arc<DashMap<String, oneshot::Sender<Response>>>,
218        msg: TransportMessage,
219    ) -> anyhow::Result<()> {
220        match msg {
221            TransportMessage::Binary(data) => handler.on_data(handle.clone(), data).await,
222            TransportMessage::Text(data) => {
223                let msg: Message = codec.decode(&data)?;
224
225                match &msg {
226                    Message::Response(res) => match pending.remove(res.id()) {
227                        Some((_, tx)) => {
228                            tx.send(res.clone())
229                                .map_err(|_| anyhow::Error::msg("failed to send response"))?;
230                            Ok(())
231                        }
232                        None => Err(anyhow::Error::msg("received response for unknown request"))?,
233                    },
234                    Message::Event(evt) => handler.on_event(handle.clone(), evt.clone()).await,
235                    Message::Request(req) => {
236                        let req = req.clone();
237                        let request_id = req.id.clone();
238                        let res: Response =
239                            match handler.on_request(handle.clone(), req.clone()).await {
240                                Ok(v) => Response::Ok(RequestResult {
241                                    id: request_id,
242                                    result: v,
243                                }),
244                                Err(err) => Response::Error(RequestError {
245                                    id: request_id,
246                                    error: err.into(),
247                                }),
248                            };
249                        let msg = Message::Response(res);
250                        let data = codec.encode(&msg).expect("failed to encode response");
251                        transport.send(&TransportMessage::Text(data)).await.unwrap();
252
253                        Ok(())
254                    }
255                }
256            }
257        }
258    }
259
260    async fn run(self: Arc<Self>) {
261        let ctx: Arc<dyn SessionContext<State = S>> = self.clone();
262
263        self.handler
264            .on_open(ctx.clone())
265            .await
266            .expect("TODO: panic message");
267
268        let (tx, mut rx) = mpsc::channel::<TransportMessage>(100);
269
270        tokio::spawn({
271            let transport = self.transport.clone();
272            async move {
273                while let Ok(msg) = transport.receive().await {
274                    if tx.send(msg).await.is_err() {
275                        break;
276                    }
277                }
278            }
279        });
280
281        tokio::spawn({
282            let codec = self.codec.clone();
283            let handler = self.handler.clone();
284            let ctx: Arc<dyn SessionContext<State = S>> = self.clone();
285            let transport = self.transport.clone();
286            let pending = self.pending_requests.clone();
287
288            async move {
289                while let Some(msg) = rx.recv().await {
290                    debug!("Received message: {:?}", msg);
291
292                    let codec = codec.clone();
293                    let handler = handler.clone();
294                    let ctx = ctx.clone();
295                    let transport = transport.clone();
296                    let pending = pending.clone();
297
298                    tokio::spawn(async move {
299                        if let Err(err) =
300                            Self::handle_msg(codec, handler, ctx, transport, pending, msg.clone())
301                                .await
302                        {
303                            error!("Error handling message: {msg:?} {err}");
304                        }
305                    });
306                }
307            }
308        });
309    }
310}
311
312#[async_trait]
313impl<C, T, S> SessionContext for RpcSession<C, T, S>
314where
315    C: Codec,
316    T: Transport,
317    S: SessionState,
318{
319    type State = S;
320
321    fn state(&self) -> &Self::State {
322        self.state.as_ref()
323    }
324
325    async fn send_binary(&self, data: Vec<u8>) -> anyhow::Result<()> {
326        self.send_binary(data).await
327    }
328
329    async fn notify(&self, event: &Event) -> anyhow::Result<()> {
330        self.notify(event).await
331    }
332
333    async fn request(
334        &self,
335        request: &Request,
336        timeout: Option<Duration>,
337    ) -> Result<RequestResult, RpcSessionError> {
338        self.request(request, timeout).await
339    }
340}
341
342#[cfg(test)]
343mod tests {
344    use super::*;
345    use crate::codec::json::JsonCodec;
346    use crate::transport::channel::channel_transport_pair;
347    use async_trait::async_trait;
348    use serde_json::{Value, json};
349
350    struct MyHandler;
351
352    #[async_trait]
353    impl RpcSessionHandler for MyHandler {
354        type State = ();
355        async fn on_request(
356            &self,
357            s: Arc<dyn SessionContext<State = Self::State>>,
358            req: Request,
359        ) -> Result<Value, ErrorBody> {
360            assert_eq!(req.method, "ping");
361            Ok(json!("pong"))
362        }
363    }
364
365    #[tokio::test]
366    async fn test_request_response() {
367        let handler = Arc::new(MyHandler);
368        let (a, b) = channel_transport_pair(10);
369        let session_a = RpcSession::create(a, JsonCodec::new(), handler.clone(), ());
370        let session_b = RpcSession::create(b, JsonCodec::new(), handler.clone(), ());
371
372        let req = Request::new("ping", None);
373
374        let res = session_a
375            .request(&req, Some(Duration::from_millis(100)))
376            .await
377            .expect("request failed");
378
379        assert_eq!(res.id, req.id);
380        assert_eq!(res.result, json!("pong"));
381    }
382}