xref: /tonic/examples/src/tls_rustls/server.rs (revision 555a8bcf)
1 pub mod pb {
2     tonic::include_proto!("/grpc.examples.unaryecho");
3 }
4 
5 use hyper::server::conn::Http;
6 use pb::{EchoRequest, EchoResponse};
7 use std::sync::Arc;
8 use tokio::net::TcpListener;
9 use tokio_rustls::{
10     rustls::{Certificate, PrivateKey, ServerConfig},
11     TlsAcceptor,
12 };
13 use tonic::{transport::Server, Request, Response, Status};
14 use tower_http::ServiceBuilderExt;
15 
16 #[tokio::main]
17 async fn main() -> Result<(), Box<dyn std::error::Error>> {
18     let data_dir = std::path::PathBuf::from_iter([std::env!("CARGO_MANIFEST_DIR"), "data"]);
19     let certs = {
20         let fd = std::fs::File::open(data_dir.join("tls/server.pem"))?;
21         let mut buf = std::io::BufReader::new(&fd);
22         rustls_pemfile::certs(&mut buf)?
23             .into_iter()
24             .map(Certificate)
25             .collect()
26     };
27     let key = {
28         let fd = std::fs::File::open(data_dir.join("tls/server.key"))?;
29         let mut buf = std::io::BufReader::new(&fd);
30         rustls_pemfile::pkcs8_private_keys(&mut buf)?
31             .into_iter()
32             .map(PrivateKey)
33             .next()
34             .unwrap()
35 
36         // let key = std::fs::read(data_dir.join("tls/server.key"))?;
37         // PrivateKey(key)
38     };
39 
40     let mut tls = ServerConfig::builder()
41         .with_safe_defaults()
42         .with_no_client_auth()
43         .with_single_cert(certs, key)?;
44     tls.alpn_protocols = vec![b"h2".to_vec()];
45 
46     let server = EchoServer::default();
47 
48     let svc = Server::builder()
49         .add_service(pb::echo_server::EchoServer::new(server))
50         .into_service();
51 
52     let mut http = Http::new();
53     http.http2_only(true);
54 
55     let listener = TcpListener::bind("[::1]:50051").await?;
56     let tls_acceptor = TlsAcceptor::from(Arc::new(tls));
57 
58     loop {
59         let (conn, addr) = match listener.accept().await {
60             Ok(incoming) => incoming,
61             Err(e) => {
62                 eprintln!("Error accepting connection: {}", e);
63                 continue;
64             }
65         };
66 
67         let http = http.clone();
68         let tls_acceptor = tls_acceptor.clone();
69         let svc = svc.clone();
70 
71         tokio::spawn(async move {
72             let mut certificates = Vec::new();
73 
74             let conn = tls_acceptor
75                 .accept_with(conn, |info| {
76                     if let Some(certs) = info.peer_certificates() {
77                         for cert in certs {
78                             certificates.push(cert.clone());
79                         }
80                     }
81                 })
82                 .await
83                 .unwrap();
84 
85             let svc = tower::ServiceBuilder::new()
86                 .add_extension(Arc::new(ConnInfo { addr, certificates }))
87                 .service(svc);
88 
89             http.serve_connection(conn, svc).await.unwrap();
90         });
91     }
92 }
93 
94 #[derive(Debug)]
95 struct ConnInfo {
96     addr: std::net::SocketAddr,
97     certificates: Vec<Certificate>,
98 }
99 
100 type EchoResult<T> = Result<Response<T>, Status>;
101 
102 #[derive(Default)]
103 pub struct EchoServer;
104 
105 #[tonic::async_trait]
106 impl pb::echo_server::Echo for EchoServer {
107     async fn unary_echo(&self, request: Request<EchoRequest>) -> EchoResult<EchoResponse> {
108         let conn_info = request.extensions().get::<Arc<ConnInfo>>().unwrap();
109         println!(
110             "Got a request from: {:?} with certs: {:?}",
111             conn_info.addr, conn_info.certificates
112         );
113 
114         let message = request.into_inner().message;
115         Ok(Response::new(EchoResponse { message }))
116     }
117 }
118