1use std::path::Path;
34use std::sync::Arc;
35
36use axum::serve::Listener;
37use axum::Router;
38use tokio::net::TcpListener;
39use tokio_rustls::rustls::pki_types::{pem::PemObject, CertificateDer, PrivateKeyDer};
40use tokio_rustls::rustls::ServerConfig;
41use tokio_rustls::server::TlsStream;
42use tokio_rustls::TlsAcceptor;
43
44#[derive(Debug, thiserror::Error)]
46pub enum TlsError {
47 #[error("failed to bind TCP listener: {0}")]
49 Bind(#[source] std::io::Error),
50
51 #[error("failed to read certificate file: {0}")]
53 ReadCertFile(#[source] std::io::Error),
54
55 #[error("failed to read private key file: {0}")]
57 ReadKeyFile(#[source] std::io::Error),
58
59 #[error("failed to parse certificate PEM: {0}")]
61 ParseCert(#[source] std::io::Error),
62
63 #[error("failed to parse private key PEM: {0}")]
65 ParseKey(#[source] std::io::Error),
66
67 #[error("no private key found in PEM file")]
69 NoPrivateKey,
70
71 #[error("failed to build rustls ServerConfig: {0}")]
73 BuildServerConfig(#[source] tokio_rustls::rustls::Error),
74
75 #[error("server error: {0}")]
77 Server(#[source] std::io::Error),
78}
79
80pub async fn load_tls_config(
90 cert_path: impl AsRef<Path>,
91 key_path: impl AsRef<Path>,
92) -> Result<ServerConfig, TlsError> {
93 let certs = load_certs(cert_path).await?;
94 let key = load_private_key(key_path).await?;
95
96 let mut config = ServerConfig::builder()
97 .with_no_client_auth()
98 .with_single_cert(certs, key)
99 .map_err(TlsError::BuildServerConfig)?;
100 config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
102 Ok(config)
103}
104
105async fn load_certs(path: impl AsRef<Path>) -> Result<Vec<CertificateDer<'static>>, TlsError> {
107 let cert_data = tokio::fs::read(path.as_ref())
108 .await
109 .map_err(TlsError::ReadCertFile)?;
110 let certs: Vec<CertificateDer<'static>> = CertificateDer::pem_slice_iter(&cert_data)
111 .collect::<Result<Vec<_>, _>>()
112 .map_err(|e| TlsError::ParseCert(std::io::Error::other(e)))?;
113 if certs.is_empty() {
114 return Err(TlsError::ParseCert(std::io::Error::new(
115 std::io::ErrorKind::InvalidData,
116 "no certificate found in PEM file",
117 )));
118 }
119 Ok(certs)
120}
121
122async fn load_private_key(path: impl AsRef<Path>) -> Result<PrivateKeyDer<'static>, TlsError> {
124 let key_data = tokio::fs::read(path.as_ref())
125 .await
126 .map_err(TlsError::ReadKeyFile)?;
127 let keys: Vec<PrivateKeyDer<'static>> = PrivateKeyDer::pem_slice_iter(&key_data)
128 .collect::<Result<Vec<_>, _>>()
129 .map_err(|e| TlsError::ParseKey(std::io::Error::other(e)))?;
130 keys.into_iter().next().ok_or(TlsError::NoPrivateKey)
131}
132
133pub fn tls_acceptor(config: ServerConfig) -> TlsAcceptor {
135 TlsAcceptor::from(Arc::new(config))
136}
137
138pub struct TlsListener {
151 tcp: TcpListener,
152 acceptor: TlsAcceptor,
153 pending_rx: tokio::sync::mpsc::UnboundedReceiver<(
155 TlsStream<tokio::net::TcpStream>,
156 std::net::SocketAddr,
157 )>,
158 pending_tx: tokio::sync::mpsc::UnboundedSender<(
160 TlsStream<tokio::net::TcpStream>,
161 std::net::SocketAddr,
162 )>,
163}
164
165impl std::fmt::Debug for TlsListener {
166 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
167 f.debug_struct("TlsListener")
168 .field("tcp", &self.tcp)
169 .finish_non_exhaustive()
170 }
171}
172
173impl TlsListener {
174 pub fn new(tcp: TcpListener, acceptor: TlsAcceptor) -> Self {
176 let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
177 Self {
178 tcp,
179 acceptor,
180 pending_rx: rx,
181 pending_tx: tx,
182 }
183 }
184
185 pub fn from_config(tcp: TcpListener, config: ServerConfig) -> Self {
187 Self::new(tcp, tls_acceptor(config))
188 }
189}
190
191impl Listener for TlsListener {
192 type Io = TlsStream<tokio::net::TcpStream>;
193 type Addr = std::net::SocketAddr;
194
195 async fn accept(&mut self) -> (Self::Io, Self::Addr) {
196 let mut backoff_ms: u64 = 100;
198 const MAX_BACKOFF_MS: u64 = 5_000;
199
200 loop {
201 tokio::select! {
202 Some((tls, addr)) = self.pending_rx.recv() => {
204 return (tls, addr);
205 }
206 accept_result = self.tcp.accept() => {
208 match accept_result {
209 Ok((tcp, addr)) => {
210 let acceptor = self.acceptor.clone();
211 let tx = self.pending_tx.clone();
212 tokio::spawn(async move {
213 match acceptor.accept(tcp).await {
214 Ok(tls) => {
215 if tx.send((tls, addr)).is_err() {
216 tracing::warn!(
217 "pending channel closed, dropping TLS connection from {addr}"
218 );
219 }
220 }
221 Err(e) => {
222 tracing::warn!("TLS handshake failed from {addr}: {e}");
223 }
224 }
225 });
226 backoff_ms = 100;
228 }
229 Err(e) => {
230 if matches!(
231 e.kind(),
232 std::io::ErrorKind::ConnectionRefused
233 | std::io::ErrorKind::ConnectionAborted
234 | std::io::ErrorKind::ConnectionReset
235 ) {
236 continue;
238 }
239 tracing::error!("accept error: {e}");
240 tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
242 backoff_ms = (backoff_ms * 2).min(MAX_BACKOFF_MS);
243 }
244 }
245 }
246 }
247 }
248 }
249
250 fn local_addr(&self) -> std::io::Result<Self::Addr> {
251 self.tcp.local_addr()
252 }
253}
254
255pub async fn serve_h2(
266 router: Router,
267 addr: &str,
268 cert_path: impl AsRef<Path>,
269 key_path: impl AsRef<Path>,
270) -> Result<(), TlsError> {
271 let listener = TcpListener::bind(addr).await.map_err(TlsError::Bind)?;
272 let config = load_tls_config(cert_path, key_path).await?;
273 serve_h2_with_listener(router, listener, config).await
274}
275
276pub async fn serve_h2_with_graceful_shutdown(
280 router: Router,
281 addr: &str,
282 cert_path: impl AsRef<Path>,
283 key_path: impl AsRef<Path>,
284) -> Result<(), TlsError> {
285 let listener = TcpListener::bind(addr).await.map_err(TlsError::Bind)?;
286 let config = load_tls_config(cert_path, key_path).await?;
287 let tls_listener = TlsListener::from_config(listener, config);
288 axum::serve(tls_listener, router.into_make_service())
289 .with_graceful_shutdown(shutdown_signal())
290 .await
291 .map_err(TlsError::Server)?;
292 Ok(())
293}
294
295pub async fn serve_h2_with_listener(
299 router: Router,
300 listener: TcpListener,
301 config: ServerConfig,
302) -> Result<(), TlsError> {
303 let tls_listener = TlsListener::from_config(listener, config);
304 axum::serve(tls_listener, router.into_make_service())
305 .await
306 .map_err(TlsError::Server)?;
307 Ok(())
308}
309
310async fn shutdown_signal() {
312 let ctrl_c = async {
313 tokio::signal::ctrl_c()
314 .await
315 .expect("failed to install Ctrl+C handler");
316 };
317
318 #[cfg(unix)]
319 let terminate = async {
320 tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
321 .expect("failed to install signal handler")
322 .recv()
323 .await;
324 };
325
326 #[cfg(not(unix))]
327 let terminate = std::future::pending::<()>();
328
329 tokio::select! {
330 _ = ctrl_c => {},
331 _ = terminate => {},
332 }
333}
334
335#[cfg(test)]
336mod tests {
337 use super::*;
338 use std::sync::Arc;
339
340 fn generate_self_signed_cert() -> (Vec<u8>, Vec<u8>) {
344 let params = rcgen::CertificateParams::new(vec!["localhost".to_string()]).unwrap();
346 let key_pair = rcgen::KeyPair::generate().unwrap();
347 let cert = params.self_signed(&key_pair).unwrap();
348 let cert_pem = cert.pem().into_bytes();
349 let key_pem = key_pair.serialize_pem().into_bytes();
350 (cert_pem, key_pem)
351 }
352
353 fn write_temp_pem(name: &str, data: &[u8]) -> std::path::PathBuf {
355 let dir = std::env::temp_dir().join("sz-rust-h2-tests");
356 std::fs::create_dir_all(&dir).unwrap();
357 let path = dir.join(name);
358 std::fs::write(&path, data).unwrap();
359 path
360 }
361
362 #[tokio::test]
367 async fn test_load_tls_config_valid_pem() {
368 let (cert, key) = generate_self_signed_cert();
369 let cert_path = write_temp_pem("valid_cert.pem", &cert);
370 let key_path = write_temp_pem("valid_key.pem", &key);
371
372 let config = load_tls_config(&cert_path, &key_path).await;
373 assert!(
374 config.is_ok(),
375 "failed to load TLS config: {:?}",
376 config.err()
377 );
378
379 let config = config.unwrap();
380 assert!(config.alpn_protocols.contains(&b"h2".to_vec()));
382 assert!(config.alpn_protocols.contains(&b"http/1.1".to_vec()));
383 }
384
385 #[tokio::test]
386 async fn test_load_tls_config_missing_cert_file() {
387 let result = load_tls_config("nonexistent_cert.pem", "nonexistent_key.pem").await;
388 assert!(matches!(result, Err(TlsError::ReadCertFile(_))));
389 }
390
391 #[tokio::test]
392 async fn test_load_tls_config_missing_key_file() {
393 let (cert, _key) = generate_self_signed_cert();
394 let cert_path = write_temp_pem("valid_cert_for_missing_key.pem", &cert);
395
396 let result = load_tls_config(&cert_path, "nonexistent_key.pem").await;
397 assert!(matches!(result, Err(TlsError::ReadKeyFile(_))));
398 }
399
400 #[tokio::test]
401 async fn test_load_tls_config_invalid_pem_content() {
402 let cert_path = write_temp_pem("invalid_cert.pem", b"not a valid PEM");
403 let key_path = write_temp_pem("invalid_key.pem", b"not a valid PEM");
404
405 let result = load_tls_config(&cert_path, &key_path).await;
406 assert!(result.is_err());
408 }
409
410 #[tokio::test]
411 async fn test_load_tls_config_empty_pem_files() {
412 let cert_path = write_temp_pem("empty_cert.pem", b"");
413 let key_path = write_temp_pem("empty_key.pem", b"");
414
415 let result = load_tls_config(&cert_path, &key_path).await;
417 assert!(matches!(result, Err(TlsError::ParseCert(_))));
418 }
419
420 #[tokio::test]
421 async fn test_load_tls_config_key_without_cert() {
422 let (_cert, key) = generate_self_signed_cert();
424 let key_path = write_temp_pem("only_key.pem", &key);
425 let empty_cert_path = write_temp_pem("empty_for_key_only.pem", b"");
426
427 let result = load_tls_config(&empty_cert_path, &key_path).await;
428 assert!(matches!(result, Err(TlsError::ParseCert(_))));
429 }
430
431 #[tokio::test]
436 async fn test_tls_acceptor_constructible() {
437 let (cert, key) = generate_self_signed_cert();
438 let cert_path = write_temp_pem("acceptor_cert.pem", &cert);
439 let key_path = write_temp_pem("acceptor_key.pem", &key);
440
441 let config = load_tls_config(&cert_path, &key_path).await.unwrap();
442 let _acceptor = tls_acceptor(config);
443 }
444
445 #[tokio::test]
446 async fn test_tls_listener_constructible() {
447 let (cert, key) = generate_self_signed_cert();
448 let cert_path = write_temp_pem("listener_cert.pem", &cert);
449 let key_path = write_temp_pem("listener_key.pem", &key);
450
451 let config = load_tls_config(&cert_path, &key_path).await.unwrap();
452 let (tcp, _addr) = crate::server::build_tcp_listener("127.0.0.1:0")
453 .await
454 .unwrap();
455 let tls_listener = TlsListener::from_config(tcp, config);
456 assert!(tls_listener.local_addr().is_ok());
458 }
459
460 #[tokio::test]
461 async fn test_tls_listener_new() {
462 let (cert, key) = generate_self_signed_cert();
463 let cert_path = write_temp_pem("new_cert.pem", &cert);
464 let key_path = write_temp_pem("new_key.pem", &key);
465
466 let config = load_tls_config(&cert_path, &key_path).await.unwrap();
467 let _arc: Arc<ServerConfig> = Arc::new(config);
468 }
469
470 #[tokio::test]
475 async fn test_serve_h2_with_listener_starts_and_accepts_connections() {
476 use tokio::net::TcpStream;
477
478 let router = Router::new().route("/", axum::routing::get(|| async { "hello h2" }));
480
481 let (cert, key) = generate_self_signed_cert();
483 let cert_path = write_temp_pem("serve_cert.pem", &cert);
484 let key_path = write_temp_pem("serve_key.pem", &key);
485 let config = load_tls_config(&cert_path, &key_path).await.unwrap();
486 let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
487 .await
488 .unwrap();
489
490 tokio::spawn(async move {
492 let _ = serve_h2_with_listener(router, listener, config).await;
493 });
494
495 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
497
498 let _stream = TcpStream::connect(addr).await.expect("TCP connect failed");
500 }
501
502 #[tokio::test]
503 async fn test_serve_h2_with_invalid_cert_returns_error() {
504 let router = Router::new();
505 let result = serve_h2(router, "127.0.0.1:0", "nonexistent.pem", "nonexistent.pem").await;
506 assert!(result.is_err());
507 match result {
510 Err(TlsError::ReadCertFile(_)) => {}
511 other => panic!("expected ReadCertFile error, got: {:?}", other),
512 }
513 }
514
515 #[tokio::test]
516 async fn test_serve_h2_bind_failure() {
517 let (cert, key) = generate_self_signed_cert();
518 let cert_path = write_temp_pem("bind_fail_cert.pem", &cert);
519 let key_path = write_temp_pem("bind_fail_key.pem", &key);
520
521 let router = Router::new();
523 let result = serve_h2(router, "127.0.0.1:99999", &cert_path, &key_path).await;
524 assert!(matches!(result, Err(TlsError::Bind(_))));
525 }
526
527 #[tokio::test]
532 async fn test_tls_handshake_completes_with_valid_client() {
533 use tokio::io::AsyncReadExt;
534 use tokio::io::AsyncWriteExt;
535 use tokio::net::TcpStream;
536
537 let router = Router::new().route("/", axum::routing::get(|| async { "tls ok" }));
539
540 let (cert, key) = generate_self_signed_cert();
542 let cert_path = write_temp_pem("handshake_cert.pem", &cert);
543 let key_path = write_temp_pem("handshake_key.pem", &key);
544 let config = load_tls_config(&cert_path, &key_path).await.unwrap();
545 let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
546 .await
547 .unwrap();
548
549 tokio::spawn(async move {
551 let _ = serve_h2_with_listener(router, listener, config).await;
552 });
553
554 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
555
556 let mut stream = TcpStream::connect(addr).await.unwrap();
559
560 let _ = stream.write_all(b"GET / HTTP/1.1\r\n\r\n").await;
562
563 let mut buf = [0u8; 64];
565 let _ = stream.read(&mut buf).await;
566 }
567
568 #[tokio::test]
569 async fn test_serve_h2_full_tls_request_with_rustls_client() {
570 use tokio::io::{AsyncReadExt, AsyncWriteExt};
572 use tokio::net::TcpStream;
573 use tokio_rustls::rustls::pki_types::ServerName;
574 use tokio_rustls::rustls::{ClientConfig, RootCertStore};
575 use tokio_rustls::TlsConnector;
576
577 let router = Router::new().route("/ping", axum::routing::get(|| async { "pong" }));
579
580 let (cert, key) = generate_self_signed_cert();
582 let cert_path = write_temp_pem("full_cert.pem", &cert);
583 let key_path = write_temp_pem("full_key.pem", &key);
584
585 let server_config = load_tls_config(&cert_path, &key_path).await.unwrap();
587 let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
588 .await
589 .unwrap();
590
591 tokio::spawn(async move {
593 let _ = serve_h2_with_listener(router, listener, server_config).await;
594 });
595 tokio::time::sleep(std::time::Duration::from_millis(150)).await;
596
597 let mut root_store = RootCertStore::empty();
599 let cert_der = CertificateDer::pem_slice_iter(&cert[..])
600 .collect::<Result<Vec<_>, _>>()
601 .unwrap();
602 for c in cert_der {
603 root_store.add(c).unwrap();
604 }
605 let client_config = ClientConfig::builder()
606 .with_root_certificates(root_store)
607 .with_no_client_auth();
608
609 let connector = TlsConnector::from(Arc::new(client_config));
611 let tcp_stream = TcpStream::connect(addr).await.unwrap();
612 let server_name = ServerName::try_from("localhost").unwrap();
613 let mut tls_stream = connector.connect(server_name, tcp_stream).await.unwrap();
614
615 tls_stream
617 .write_all(b"GET /ping HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
618 .await
619 .unwrap();
620
621 let mut response = Vec::new();
623 tls_stream.read_to_end(&mut response).await.unwrap();
624 let response_str = String::from_utf8_lossy(&response);
625
626 assert!(
628 response_str.contains("pong"),
629 "expected response to contain 'pong', got: {}",
630 response_str
631 );
632 assert!(
634 response_str.starts_with("HTTP/1.1") || response_str.starts_with("HTTP/2"),
635 "expected HTTP response, got: {}",
636 response_str.lines().next().unwrap_or("")
637 );
638 }
639}