sz_rust_http_facade/
tls.rs1use std::net::SocketAddr;
9use std::path::PathBuf;
10use std::sync::Arc;
11
12use axum::Router;
13use thiserror::Error;
14use tokio::net::TcpListener;
15use tokio_rustls::TlsAcceptor;
16use tower::ServiceExt;
17
18#[derive(Debug, Error)]
20pub enum TlsError {
21 #[error("IO error: {0}")]
23 Io(#[from] std::io::Error),
24 #[error("No valid certificate found")]
26 NoCertificate,
27 #[error("No valid private key found")]
29 NoPrivateKey,
30 #[error("TLS error: {0}")]
32 Tls(String),
33}
34
35#[derive(Debug, Clone)]
37pub struct TlsConfig {
38 pub cert_path: PathBuf,
40 pub key_path: PathBuf,
42 pub alpn: Vec<Vec<u8>>,
44}
45
46impl TlsConfig {
47 pub fn new(cert_path: impl Into<PathBuf>, key_path: impl Into<PathBuf>) -> Self {
49 Self {
50 cert_path: cert_path.into(),
51 key_path: key_path.into(),
52 alpn: vec![b"h2".to_vec(), b"http/1.1".to_vec()],
53 }
54 }
55
56 pub fn h2_only(mut self) -> Self {
58 self.alpn = vec![b"h2".to_vec()];
59 self
60 }
61
62 pub fn http1_only(mut self) -> Self {
64 self.alpn = vec![b"http/1.1".to_vec()];
65 self
66 }
67
68 pub async fn build_acceptor(&self) -> Result<TlsAcceptor, TlsError> {
72 let cert = load_certs(&self.cert_path).await?;
73 let key = load_private_key(&self.key_path).await?;
74
75 let mut config = rustls::ServerConfig::builder()
76 .with_no_client_auth()
77 .with_single_cert(cert, key)
78 .map_err(|e| TlsError::Tls(e.to_string()))?;
79
80 config.alpn_protocols = self.alpn.clone();
81
82 Ok(TlsAcceptor::from(Arc::new(config)))
83 }
84}
85
86async fn load_certs(
88 path: &std::path::Path,
89) -> Result<Vec<rustls::pki_types::CertificateDer<'static>>, TlsError> {
90 let file = tokio::fs::read_to_string(path).await?;
91 let mut reader = std::io::BufReader::new(file.as_bytes());
92 let certs: Vec<_> = rustls_pemfile::certs(&mut reader).collect::<Result<_, _>>()?;
93 if certs.is_empty() {
94 return Err(TlsError::NoCertificate);
95 }
96 Ok(certs)
97}
98
99async fn load_private_key(
101 path: &std::path::Path,
102) -> Result<rustls::pki_types::PrivateKeyDer<'static>, TlsError> {
103 let file = tokio::fs::read_to_string(path).await?;
104 let mut reader = std::io::BufReader::new(file.as_bytes());
105 rustls_pemfile::private_key(&mut reader)
106 .map_err(TlsError::Io)?
107 .ok_or(TlsError::NoPrivateKey)
108}
109
110pub async fn serve_http2(router: Router, addr: SocketAddr, tls: TlsConfig) -> Result<(), TlsError> {
116 let acceptor = tls.build_acceptor().await?;
117 let listener = TcpListener::bind(addr).await?;
118 tracing::info!("HTTPS/HTTP2 server listening on {}", addr);
119
120 loop {
121 let (tcp_stream, _remote) = listener.accept().await?;
122 let acceptor = acceptor.clone();
123 let router = router.clone();
124
125 tokio::spawn(async move {
126 let tls_stream = match acceptor.accept(tcp_stream).await {
127 Ok(s) => s,
128 Err(e) => {
129 tracing::warn!("TLS handshake failed: {}", e);
130 return;
131 }
132 };
133
134 let io = hyper_util::rt::TokioIo::new(tls_stream);
135 let svc = hyper::service::service_fn(move |req| {
136 let router = router.clone();
137 async move { router.oneshot(req).await }
138 });
139
140 if let Err(e) =
141 hyper::server::conn::http2::Builder::new(hyper_util::rt::TokioExecutor::new())
142 .serve_connection(io, svc)
143 .await
144 {
145 tracing::warn!("HTTP/2 connection error: {}", e);
146 }
147 });
148 }
149}
150
151#[cfg(test)]
152mod tests {
153 use super::*;
154
155 #[test]
156 fn test_tls_config_new() {
157 let config = TlsConfig::new("/path/to/cert.pem", "/path/to/key.pem");
158 assert_eq!(config.cert_path, PathBuf::from("/path/to/cert.pem"));
159 assert_eq!(config.key_path, PathBuf::from("/path/to/key.pem"));
160 assert_eq!(config.alpn, vec![b"h2".to_vec(), b"http/1.1".to_vec()]);
161 }
162
163 #[test]
164 fn test_tls_config_h2_only() {
165 let config = TlsConfig::new("/cert.pem", "/key.pem").h2_only();
166 assert_eq!(config.alpn, vec![b"h2".to_vec()]);
167 }
168
169 #[test]
170 fn test_tls_config_http1_only() {
171 let config = TlsConfig::new("/cert.pem", "/key.pem").http1_only();
172 assert_eq!(config.alpn, vec![b"http/1.1".to_vec()]);
173 }
174
175 #[test]
176 fn test_tls_error_display() {
177 let err = TlsError::NoCertificate;
178 assert_eq!(err.to_string(), "No valid certificate found");
179
180 let err = TlsError::NoPrivateKey;
181 assert_eq!(err.to_string(), "No valid private key found");
182 }
183}