1use 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#[derive(Debug, Default)]
44pub struct Server {
45 services: HashMap<&'static str, mesh::Sender<(CancelContext, GenericRpc)>>,
46}
47
48#[derive(MeshPayload)]
52#[mesh(bound = "T: ServiceRpc")]
53pub struct RpcReceiver<T>(mesh::Receiver<(CancelContext, DecodedRpc<T>)>);
54
55impl<T: ServiceRpc> RpcReceiver<T> {
56 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 pub fn new() -> Self {
107 Self {
108 services: Default::default(),
109 }
110 }
111
112 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 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! { 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 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 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; 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 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! { 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 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 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 let mut buf = Vec::new();
585
586 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 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}