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}