Skip to main content

mesh_rpc/
server.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! TTRPC server.
5
6use crate::message::MESSAGE_TYPE_REQUEST;
7use crate::message::MESSAGE_TYPE_RESPONSE;
8use crate::message::ReadResult;
9use crate::message::Request;
10use crate::message::Response;
11use crate::message::TooLongError;
12use crate::message::read_message;
13use crate::message::write_message;
14use crate::rpc::ProtocolError;
15use crate::rpc::status_from_err;
16use crate::service::Code;
17use crate::service::DecodedRpc;
18use crate::service::GenericRpc;
19use crate::service::ServiceRpc;
20use crate::service::ServiceRpcError;
21use crate::service::Status;
22use futures::FutureExt;
23use futures::Stream;
24use futures::StreamExt;
25use futures::stream::FusedStream;
26use futures_concurrency::future::TryJoin;
27use futures_concurrency::stream::Merge;
28use mesh::CancelContext;
29use mesh::MeshPayload;
30use mesh::local_node::Port;
31use pal_async::driver::Driver;
32use pal_async::socket::AsSockRef;
33use pal_async::socket::Listener;
34use pal_async::socket::PolledSocket;
35use std::collections::HashMap;
36use std::io::Read;
37use std::io::Write;
38use std::pin::Pin;
39use std::task::ready;
40use unicycle::FuturesUnordered;
41
42/// A ttrpc server.
43#[derive(Debug, Default)]
44pub struct Server {
45    services: HashMap<&'static str, mesh::Sender<(CancelContext, GenericRpc)>>,
46}
47
48/// A receiver for RPC requests for a given service.
49///
50/// Returned by [`Server::add_service`].
51#[derive(MeshPayload)]
52#[mesh(bound = "T: ServiceRpc")]
53pub struct RpcReceiver<T>(mesh::Receiver<(CancelContext, DecodedRpc<T>)>);
54
55impl<T: ServiceRpc> RpcReceiver<T> {
56    /// Returns a disconnected stream, useful for when a service is dynamically
57    /// not registered.
58    pub fn disconnected() -> Self {
59        let (_send, recv) = mesh::channel();
60        Self(recv)
61    }
62}
63
64impl<T: ServiceRpc> Stream for RpcReceiver<T> {
65    type Item = (CancelContext, T);
66
67    fn poll_next(
68        self: Pin<&mut Self>,
69        cx: &mut std::task::Context<'_>,
70    ) -> std::task::Poll<Option<Self::Item>> {
71        let this = self.get_mut();
72        while let Some((ctx, rpc)) = ready!(Pin::new(&mut this.0).poll_next(cx)) {
73            match rpc {
74                DecodedRpc::Rpc(rpc) => return Some((ctx, rpc)).into(),
75                DecodedRpc::Err { rpc, err } => {
76                    rpc.fail(err);
77                }
78            }
79        }
80        None.into()
81    }
82}
83
84impl<T: ServiceRpc> FusedStream for RpcReceiver<T> {
85    fn is_terminated(&self) -> bool {
86        self.0.is_terminated()
87    }
88}
89
90impl GenericRpc {
91    fn fail(self, err: ServiceRpcError) {
92        let status = match err {
93            ServiceRpcError::UnknownMethod => Status {
94                code: Code::Unimplemented.into(),
95                message: format!("unknown method {}", self.method),
96                details: Vec::new(),
97            },
98            ServiceRpcError::InvalidInput(error) => status_from_err(Code::InvalidArgument, error),
99        };
100        self.respond_status(status);
101    }
102}
103
104impl Server {
105    /// Creates a new ttrpc server.
106    pub fn new() -> Self {
107        Self {
108            services: Default::default(),
109        }
110    }
111
112    /// Adds or updates a channel for receiving service requests.
113    pub fn add_service<T: ServiceRpc>(&mut self) -> RpcReceiver<T> {
114        let (send, recv) = mesh::channel();
115        self.services.insert(T::NAME, Port::from(send).into());
116        RpcReceiver(recv)
117    }
118
119    /// Runs the server using the ttrpc transport, listening on `listener` and
120    /// servicing connections until `cancel`.
121    pub async fn run(
122        &mut self,
123        driver: &(impl Driver + ?Sized),
124        listener: impl Listener,
125        cancel: mesh::OneshotReceiver<()>,
126    ) -> anyhow::Result<()> {
127        let mut listener = PolledSocket::new(driver, listener)?;
128        let mut tasks = FuturesUnordered::new();
129        let mut cancel = cancel.fuse();
130        loop {
131            let conn = futures::select! { // merge semantics
132                r = listener.accept().fuse() => r,
133                _ = tasks.next() => continue,
134                _ = cancel => break,
135            };
136            if let Ok(conn) = conn.and_then(|(conn, _)| PolledSocket::new(driver, conn)) {
137                tasks.push(async {
138                    let _ = self.serve_connection(conn).await.map_err(|err| {
139                        tracing::error!(
140                            error = err.as_ref() as &dyn std::error::Error,
141                            "connection error"
142                        )
143                    });
144                });
145            }
146        }
147        Ok(())
148    }
149
150    /// Runs the server, servicing a single connection `conn`.
151    pub async fn run_single(
152        &mut self,
153        driver: &(impl Driver + ?Sized),
154        conn: impl AsSockRef + Read + Write,
155    ) -> anyhow::Result<()> {
156        self.serve_connection(PolledSocket::new(driver, conn)?)
157            .await
158    }
159
160    /// Services a single, already-accepted connection using the ttrpc protocol.
161    ///
162    /// This is useful when the caller owns the accept loop, for example to
163    /// dispatch a connection to ttrpc or another protocol based on a sniffed
164    /// prefix byte.
165    pub async fn serve_connection(
166        &self,
167        stream: PolledSocket<impl AsSockRef + Read + Write>,
168    ) -> anyhow::Result<()> {
169        let (mut reader, mut writer) = stream.split();
170        let (stream_send, mut stream_recv) = mesh::channel();
171        let ctx = CancelContext::new();
172        let recv_task = async {
173            let stream_send = stream_send; // move into this task
174            while let Some(message) = read_message(&mut reader).await? {
175                let (send, recv) = mesh::oneshot::<Result<Vec<u8>, Status>>();
176                stream_send.send((message.stream_id, recv));
177
178                let handle = handle_message(message).and_then(|request| {
179                    let service = self.services.get(request.service.as_str()).ok_or_else(|| {
180                        status_from_err(
181                            Code::Unimplemented,
182                            anyhow::anyhow!("unknown service {}", request.service),
183                        )
184                    })?;
185
186                    let ctx = if request.timeout_nano == 0 {
187                        ctx.clone()
188                    } else {
189                        ctx.with_timeout(std::time::Duration::from_nanos(request.timeout_nano))
190                    };
191
192                    Ok(move |port| {
193                        service.send((
194                            ctx,
195                            GenericRpc {
196                                method: request.method,
197                                data: request.payload,
198                                port,
199                            },
200                        ));
201                    })
202                });
203
204                match handle {
205                    Ok(handle) => handle(send.into()),
206                    Err(err) => send.send(Err(err)),
207                }
208            }
209            Ok(())
210        };
211        let send_task = async {
212            let mut responses = FuturesUnordered::new();
213            enum Event<T> {
214                Request((u32, mesh::OneshotReceiver<Result<Vec<u8>, Status>>)),
215                Response(T),
216            }
217            while let Some(event) = (
218                (&mut stream_recv).map(Event::Request),
219                (&mut responses).map(Event::Response),
220            )
221                .merge()
222                .next()
223                .await
224            {
225                match event {
226                    Event::Request((stream_id, recv)) => {
227                        let recv = recv.map(move |r| {
228                            (
229                                stream_id,
230                                match r {
231                                    Ok(Ok(payload)) => Response::Payload(payload),
232                                    Ok(Err(status)) => Response::Status(status),
233                                    Err(err) => {
234                                        Response::Status(status_from_err(Code::Internal, err))
235                                    }
236                                },
237                            )
238                        });
239                        responses.push(recv);
240                    }
241                    Event::Response((stream_id, payload)) => {
242                        write_message(
243                            &mut writer,
244                            stream_id,
245                            MESSAGE_TYPE_RESPONSE,
246                            &mesh::payload::encode(payload),
247                        )
248                        .await?;
249                    }
250                }
251            }
252            anyhow::Result::<_>::Ok(())
253        };
254        (recv_task, send_task).try_join().await?;
255        Ok(())
256    }
257}
258
259fn handle_message(message: ReadResult) -> Result<Request, Status> {
260    if message.stream_id % 2 != 1 {
261        return Err(status_from_err(
262            Code::InvalidArgument,
263            ProtocolError::EvenStreamId,
264        ));
265    }
266
267    match message.message_type {
268        MESSAGE_TYPE_REQUEST => {
269            let payload = message.payload.map_err(|err @ TooLongError { .. }| {
270                status_from_err(Code::ResourceExhausted, err)
271            })?;
272            let request = mesh::payload::decode::<Request>(&payload)
273                .map_err(|err| status_from_err(Code::InvalidArgument, err))?;
274
275            tracing::debug!(
276                stream_id = message.stream_id,
277                service = %request.service,
278                method = %request.method,
279                timeout = request.timeout_nano / 1000 / 1000,
280                "message",
281            );
282
283            Ok(request)
284        }
285        ty => Err(status_from_err(
286            Code::InvalidArgument,
287            ProtocolError::InvalidMessageType(ty),
288        )),
289    }
290}
291
292#[cfg(feature = "grpc")]
293mod grpc {
294    use super::Server;
295    use crate::rpc::status_from_err;
296    use crate::service::Code;
297    use crate::service::GenericRpc;
298    use crate::service::Status;
299    use anyhow::Context as _;
300    use futures::AsyncRead as _;
301    use futures::AsyncWrite;
302    use futures::FutureExt;
303    use futures::StreamExt;
304    use futures_concurrency::stream::Merge;
305    use h2::RecvStream;
306    use h2::server::SendResponse;
307    use http::HeaderMap;
308    use http::HeaderValue;
309    use mesh::CancelContext;
310    use pal_async::driver::Driver;
311    use pal_async::socket::AsSockRef;
312    use pal_async::socket::Listener;
313    use pal_async::socket::PolledSocket;
314    use prost::bytes::Bytes;
315    use std::io::Read;
316    use std::io::Write;
317    use std::pin::Pin;
318    use std::task::ready;
319    use thiserror::Error;
320    use unicycle::FuturesUnordered;
321
322    #[derive(Debug, Error)]
323    enum RequestError {
324        #[error("http error")]
325        Http(#[from] http::Error),
326        #[error("http2 error")]
327        H2(#[from] h2::Error),
328        #[error("unreachable")]
329        Status(http::StatusCode),
330        #[error("invalid message header")]
331        InvalidHeader,
332    }
333
334    impl From<http::StatusCode> for RequestError {
335        fn from(status: http::StatusCode) -> Self {
336            RequestError::Status(status)
337        }
338    }
339
340    impl Server {
341        /// Runs the server using the gRPC transport, listening on `listener` and servicing connections until
342        /// `cancel`.
343        pub async fn run_grpc(
344            &mut self,
345            driver: &(impl Driver + ?Sized),
346            listener: impl Listener,
347            cancel: mesh::OneshotReceiver<()>,
348        ) -> anyhow::Result<()> {
349            let mut listener = PolledSocket::new(driver, listener)?;
350            let mut tasks = FuturesUnordered::new();
351            let mut cancel = cancel.fuse();
352            loop {
353                let conn = futures::select! { // merge semantics
354                    r = listener.accept().fuse() => r,
355                    _ = tasks.next() => continue,
356                    _ = cancel => break,
357                };
358                if let Ok(conn) = conn.and_then(|(conn, _)| PolledSocket::new(driver, conn)) {
359                    tasks.push(async {
360                        let _ = self.serve_connection_grpc(conn).await.map_err(|err| {
361                            tracing::error!(
362                                error = err.as_ref() as &dyn std::error::Error,
363                                "connection error"
364                            )
365                        });
366                    });
367                }
368            }
369            Ok(())
370        }
371
372        /// Services a single, already-accepted connection using the gRPC
373        /// (HTTP/2) protocol.
374        ///
375        /// This is useful when the caller owns the accept loop, for example to
376        /// dispatch a connection to gRPC or another protocol based on a sniffed
377        /// prefix byte.
378        pub async fn serve_connection_grpc(
379            &self,
380            stream: PolledSocket<impl AsSockRef + Read + Write>,
381        ) -> anyhow::Result<()> {
382            struct Wrap<T>(T);
383
384            impl<T: AsSockRef + Read> tokio::io::AsyncRead for Wrap<PolledSocket<T>> {
385                fn poll_read(
386                    self: Pin<&mut Self>,
387                    cx: &mut std::task::Context<'_>,
388                    buf: &mut tokio::io::ReadBuf<'_>,
389                ) -> std::task::Poll<std::io::Result<()>> {
390                    let n = ready!(
391                        Pin::new(&mut self.get_mut().0).poll_read(cx, buf.initialize_unfilled())
392                    )?;
393                    buf.advance(n);
394                    std::task::Poll::Ready(Ok(()))
395                }
396            }
397
398            impl<T: AsSockRef + Write> tokio::io::AsyncWrite for Wrap<PolledSocket<T>> {
399                fn poll_write(
400                    self: Pin<&mut Self>,
401                    cx: &mut std::task::Context<'_>,
402                    buf: &[u8],
403                ) -> std::task::Poll<Result<usize, std::io::Error>> {
404                    Pin::new(&mut self.get_mut().0).poll_write(cx, buf)
405                }
406
407                fn poll_flush(
408                    self: Pin<&mut Self>,
409                    cx: &mut std::task::Context<'_>,
410                ) -> std::task::Poll<Result<(), std::io::Error>> {
411                    Pin::new(&mut self.get_mut().0).poll_flush(cx)
412                }
413
414                fn poll_shutdown(
415                    self: Pin<&mut Self>,
416                    cx: &mut std::task::Context<'_>,
417                ) -> std::task::Poll<Result<(), std::io::Error>> {
418                    Pin::new(&mut self.get_mut().0).poll_close(cx)
419                }
420            }
421
422            let mut conn = h2::server::handshake(Wrap(stream))
423                .await
424                .context("failed http2 handshake")?;
425
426            let mut tasks = FuturesUnordered::new();
427
428            loop {
429                enum Event<A, B> {
430                    Accept(A),
431                    Task(Result<(), B>),
432                }
433
434                let r = (
435                    (&mut conn).map(Event::Accept),
436                    (&mut tasks).map(Event::Task),
437                )
438                    .merge()
439                    .next()
440                    .await;
441
442                let (req, mut resp) = match r {
443                    None => break,
444                    Some(Event::Task(r)) => {
445                        r?;
446                        continue;
447                    }
448                    Some(Event::Accept(r)) => r.context("failed http2 stream accept")?,
449                };
450
451                let task = async move {
452                    match self.handle_request(req, &mut resp).await {
453                        Err(RequestError::Status(status)) => {
454                            tracing::debug!(status = status.as_u16(), "request error");
455                            resp.send_response(
456                                http::Response::builder().status(status).body(())?,
457                                true,
458                            )?;
459                            Ok(())
460                        }
461                        r => r,
462                    }
463                };
464                tasks.push(task);
465            }
466
467            std::future::poll_fn(|cx| conn.poll_closed(cx)).await?;
468            Ok(())
469        }
470
471        async fn handle_request(
472            &self,
473            req: http::Request<RecvStream>,
474            resp: &mut SendResponse<Bytes>,
475        ) -> Result<(), RequestError> {
476            tracing::debug!(url = %req.uri(), "rpc request");
477
478            if req.method() != http::Method::POST {
479                Err(http::StatusCode::METHOD_NOT_ALLOWED)?
480            }
481            let content_type = req.headers().get("content-type");
482            match content_type.map(|v| v.as_bytes()) {
483                Some(b"application/grpc" | b"application/grpc+proto") => {}
484                _ => Err(http::StatusCode::UNSUPPORTED_MEDIA_TYPE)?,
485            }
486
487            let response =
488                http::Response::builder().header("content-type", "application/grpc+proto");
489
490            let ctx = if let Some(timeout) = req.headers().get("grpc-timeout") {
491                let timeout = timeout
492                    .to_str()
493                    .map_err(|_| http::StatusCode::BAD_REQUEST)?;
494                let mul = match timeout
495                    .bytes()
496                    .last()
497                    .ok_or(http::StatusCode::BAD_REQUEST)?
498                {
499                    b'H' => std::time::Duration::from_secs(60 * 60),
500                    b'M' => std::time::Duration::from_secs(60),
501                    b'S' => std::time::Duration::from_secs(1),
502                    b'm' => std::time::Duration::from_millis(1),
503                    b'u' => std::time::Duration::from_micros(1),
504                    b'n' => std::time::Duration::from_nanos(1),
505                    _ => Err(http::StatusCode::BAD_REQUEST)?,
506                };
507                let timeout = timeout[..timeout.len() - 1]
508                    .parse::<u32>()
509                    .map_err(|_| http::StatusCode::BAD_REQUEST)?;
510                CancelContext::new().with_timeout(mul * timeout)
511            } else {
512                CancelContext::new()
513            };
514
515            let (head, body) = req.into_parts();
516            let path = head.uri.path();
517            let path = path.strip_prefix('/').ok_or(http::StatusCode::NOT_FOUND)?;
518            let (service, method) = path.split_once('/').ok_or(http::StatusCode::NOT_FOUND)?;
519
520            // No returning HTTP status code errors after this point.
521            let mut resp = resp.send_response(response.body(())?, false)?;
522
523            let result = self.invoke_rpc(service, method, body, ctx).await?;
524
525            let mut trailers = HeaderMap::new();
526            match result {
527                Ok(data) => {
528                    tracing::debug!(service, method, "rpc success");
529
530                    let mut buf = Vec::with_capacity(5 + data.len());
531                    buf.push(0);
532                    buf.extend(&(data.len() as u32).to_be_bytes());
533                    buf.extend(data);
534                    resp.send_data(buf.into(), false)?;
535                    trailers.insert("grpc-status", const { HeaderValue::from_static("0") });
536                }
537                Err(status) => {
538                    tracing::debug!(service, method, ?status, "rpc error");
539
540                    trailers.insert("grpc-status", status.code.into());
541                    trailers.insert(
542                        "grpc-message",
543                        urlencoding::encode(&status.message)
544                            .into_owned()
545                            .try_into()
546                            .unwrap(),
547                    );
548                    trailers.insert(
549                        "grpc-status-details-bin",
550                        base64::Engine::encode(
551                            &base64::engine::general_purpose::STANDARD,
552                            prost::Message::encode_to_vec(&status),
553                        )
554                        .try_into()
555                        .unwrap(),
556                    );
557                }
558            }
559            resp.send_trailers(trailers)?;
560            Ok(())
561        }
562
563        async fn invoke_rpc(
564            &self,
565            service: &str,
566            method: &str,
567            mut body: RecvStream,
568            ctx: CancelContext,
569        ) -> Result<Result<Vec<u8>, Status>, RequestError> {
570            let Some(service) = self.services.get(service) else {
571                return Ok(Err(Status {
572                    code: Code::Unimplemented.into(),
573                    message: format!("unknown service {}", service),
574                    details: Vec::new(),
575                }));
576            };
577
578            // For now, only non-stream RPCs are supported, so read the first
579            // message and ignore the rest.
580            //
581            // FUTURE: change the `GenericRpc` type to include channels for
582            // streams.
583
584            let mut buf = Vec::new();
585
586            // Read data frames until the header is complete.
587            while buf.len() < 5 {
588                let data = body.data().await.ok_or(RequestError::InvalidHeader)??;
589                buf.extend(&data);
590                body.flow_control().release_capacity(data.len()).unwrap();
591            }
592            let hdr = buf.get(0..5).ok_or(RequestError::InvalidHeader)?;
593            if hdr[0] != 0 {
594                // Compression was not advertised as supported, so the client
595                // should not send compressed messages.
596                return Err(RequestError::InvalidHeader);
597            }
598            let len = u32::from_be_bytes(hdr[1..5].try_into().unwrap()) as usize;
599
600            buf.drain(..5);
601            while buf.len() < len {
602                let data = body.data().await.ok_or(RequestError::InvalidHeader)??;
603                buf.extend(&data);
604                body.flow_control().release_capacity(data.len()).unwrap();
605            }
606
607            let (send, recv) = mesh::oneshot();
608
609            let rpc = GenericRpc {
610                method: method.to_owned(),
611                data: buf,
612                port: send.into(),
613            };
614
615            service.send((ctx, rpc));
616
617            Ok(recv
618                .await
619                .unwrap_or_else(|err| Err(status_from_err(Code::Internal, err))))
620        }
621    }
622}
623
624#[cfg(test)]
625mod tests {
626    use crate::Client;
627    use crate::Server;
628    use crate::client::ExistingConnection;
629    use crate::service::Code;
630    use crate::service::ServiceRpc;
631    use futures::StreamExt;
632    use futures::executor::block_on;
633    use pal_async::DefaultPool;
634    use pal_async::local::block_with_io;
635    use pal_async::socket::PolledSocket;
636    use test_with_tracing::test;
637
638    #[expect(clippy::allow_attributes)]
639    mod items {
640        include!(concat!(env!("OUT_DIR"), "/ttrpc.example.v1.rs"));
641    }
642
643    #[test]
644    fn client_server() {
645        let (c, s) = unix_socket::UnixStream::pair().unwrap();
646        let mut server = Server::new();
647        let mut recv = server.add_service::<items::Example>();
648        let server_thread = std::thread::spawn(move || {
649            block_with_io(async |driver| server.run_single(&driver, s).await)
650        });
651
652        let client_thread = std::thread::spawn(move || {
653            DefaultPool::run_with(async |driver| {
654                let client = Client::new(
655                    &driver,
656                    ExistingConnection::new(PolledSocket::new(&driver, c).unwrap()),
657                );
658                let response = client
659                    .call()
660                    .start(
661                        items::Example::Method1,
662                        items::Method1Request {
663                            foo: "abc".to_string(),
664                            bar: "def".to_string(),
665                        },
666                    )
667                    .await
668                    .unwrap();
669
670                assert_eq!(&response.foo, "abc123");
671                assert_eq!(&response.bar, "def456");
672
673                let status = client
674                    .call()
675                    .start_raw(items::Example::NAME, "unknown", Vec::new())
676                    .await
677                    .unwrap_err();
678
679                assert_eq!(status.code, Code::Unimplemented as i32);
680
681                client.shutdown().await;
682            })
683        });
684
685        block_on(async {
686            let (_, req) = recv.next().await.unwrap();
687            match req {
688                items::Example::Method1(input, resp) => {
689                    assert_eq!(&input.foo, "abc");
690                    assert_eq!(&input.bar, "def");
691                    resp.send(Ok(items::Method1Response {
692                        foo: input.foo + "123",
693                        bar: input.bar + "456",
694                    }));
695                }
696                _ => panic!("{:?}", req),
697            }
698
699            assert!(recv.next().await.is_none());
700        });
701
702        client_thread.join().unwrap();
703        server_thread.join().unwrap().unwrap();
704    }
705}