1use std::io::BufReader;
34use std::path::Path;
35use std::sync::Arc;
36
37use axum::serve::Listener;
38use axum::Router;
39use tokio::net::TcpListener;
40use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer};
41use tokio_rustls::rustls::ServerConfig;
42use tokio_rustls::server::TlsStream;
43use tokio_rustls::TlsAcceptor;
44
45#[derive(Debug, thiserror::Error)]
47pub enum TlsError {
48 #[error("failed to bind TCP listener: {0}")]
50 Bind(#[source] std::io::Error),
51
52 #[error("failed to read certificate file: {0}")]
54 ReadCertFile(#[source] std::io::Error),
55
56 #[error("failed to read private key file: {0}")]
58 ReadKeyFile(#[source] std::io::Error),
59
60 #[error("failed to parse certificate PEM: {0}")]
62 ParseCert(#[source] std::io::Error),
63
64 #[error("failed to parse private key PEM: {0}")]
66 ParseKey(#[source] std::io::Error),
67
68 #[error("no private key found in PEM file")]
70 NoPrivateKey,
71
72 #[error("failed to build rustls ServerConfig: {0}")]
74 BuildServerConfig(#[source] tokio_rustls::rustls::Error),
75
76 #[error("server error: {0}")]
78 Server(#[source] std::io::Error),
79}
80
81pub fn load_tls_config(
91 cert_path: impl AsRef<Path>,
92 key_path: impl AsRef<Path>,
93) -> Result<ServerConfig, TlsError> {
94 let certs = load_certs(cert_path)?;
95 let key = load_private_key(key_path)?;
96
97 let mut config = ServerConfig::builder()
98 .with_no_client_auth()
99 .with_single_cert(certs, key)
100 .map_err(TlsError::BuildServerConfig)?;
101 config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
103 Ok(config)
104}
105
106fn load_certs(path: impl AsRef<Path>) -> Result<Vec<CertificateDer<'static>>, TlsError> {
108 let file = std::fs::File::open(path).map_err(TlsError::ReadCertFile)?;
109 let mut reader = BufReader::new(file);
110 let certs: Vec<CertificateDer<'static>> = rustls_pemfile::certs(&mut reader)
111 .collect::<Result<Vec<_>, _>>()
112 .map_err(TlsError::ParseCert)?;
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
122fn load_private_key(path: impl AsRef<Path>) -> Result<PrivateKeyDer<'static>, TlsError> {
124 let file = std::fs::File::open(path).map_err(TlsError::ReadKeyFile)?;
125 let mut reader = BufReader::new(file);
126 let mut keys = Vec::new();
127 for item in rustls_pemfile::read_all(&mut reader) {
128 match item.map_err(TlsError::ParseKey)? {
129 rustls_pemfile::Item::Pkcs8Key(k) => keys.push(PrivateKeyDer::Pkcs8(k)),
130 rustls_pemfile::Item::Pkcs1Key(k) => keys.push(PrivateKeyDer::Pkcs1(k)),
131 rustls_pemfile::Item::Sec1Key(k) => keys.push(PrivateKeyDer::Sec1(k)),
132 _ => {}
133 }
134 }
135 keys.into_iter().next().ok_or(TlsError::NoPrivateKey)
136}
137
138pub fn tls_acceptor(config: ServerConfig) -> TlsAcceptor {
140 TlsAcceptor::from(Arc::new(config))
141}
142
143pub struct TlsListener {
156 tcp: TcpListener,
157 acceptor: TlsAcceptor,
158 pending_rx: tokio::sync::mpsc::UnboundedReceiver<(
160 TlsStream<tokio::net::TcpStream>,
161 std::net::SocketAddr,
162 )>,
163 pending_tx: tokio::sync::mpsc::UnboundedSender<(
165 TlsStream<tokio::net::TcpStream>,
166 std::net::SocketAddr,
167 )>,
168}
169
170impl std::fmt::Debug for TlsListener {
171 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
172 f.debug_struct("TlsListener")
173 .field("tcp", &self.tcp)
174 .finish_non_exhaustive()
175 }
176}
177
178impl TlsListener {
179 pub fn new(tcp: TcpListener, acceptor: TlsAcceptor) -> Self {
181 let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
182 Self {
183 tcp,
184 acceptor,
185 pending_rx: rx,
186 pending_tx: tx,
187 }
188 }
189
190 pub fn from_config(tcp: TcpListener, config: ServerConfig) -> Self {
192 Self::new(tcp, tls_acceptor(config))
193 }
194}
195
196impl Listener for TlsListener {
197 type Io = TlsStream<tokio::net::TcpStream>;
198 type Addr = std::net::SocketAddr;
199
200 async fn accept(&mut self) -> (Self::Io, Self::Addr) {
201 let mut backoff_ms: u64 = 100;
203 const MAX_BACKOFF_MS: u64 = 5_000;
204
205 loop {
206 tokio::select! {
207 Some((tls, addr)) = self.pending_rx.recv() => {
209 return (tls, addr);
210 }
211 accept_result = self.tcp.accept() => {
213 match accept_result {
214 Ok((tcp, addr)) => {
215 let acceptor = self.acceptor.clone();
216 let tx = self.pending_tx.clone();
217 tokio::spawn(async move {
218 match acceptor.accept(tcp).await {
219 Ok(tls) => {
220 if tx.send((tls, addr)).is_err() {
221 tracing::warn!(
222 "pending channel closed, dropping TLS connection from {addr}"
223 );
224 }
225 }
226 Err(e) => {
227 tracing::warn!("TLS handshake failed from {addr}: {e}");
228 }
229 }
230 });
231 backoff_ms = 100;
233 }
234 Err(e) => {
235 if matches!(
236 e.kind(),
237 std::io::ErrorKind::ConnectionRefused
238 | std::io::ErrorKind::ConnectionAborted
239 | std::io::ErrorKind::ConnectionReset
240 ) {
241 continue;
243 }
244 tracing::error!("accept error: {e}");
245 tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
247 backoff_ms = (backoff_ms * 2).min(MAX_BACKOFF_MS);
248 }
249 }
250 }
251 }
252 }
253 }
254
255 fn local_addr(&self) -> std::io::Result<Self::Addr> {
256 self.tcp.local_addr()
257 }
258}
259
260pub async fn serve_h2(
271 router: Router,
272 addr: &str,
273 cert_path: impl AsRef<Path>,
274 key_path: impl AsRef<Path>,
275) -> Result<(), TlsError> {
276 let listener = TcpListener::bind(addr).await.map_err(TlsError::Bind)?;
277 let config = load_tls_config(cert_path, key_path)?;
278 serve_h2_with_listener(router, listener, config).await
279}
280
281pub async fn serve_h2_with_graceful_shutdown(
285 router: Router,
286 addr: &str,
287 cert_path: impl AsRef<Path>,
288 key_path: impl AsRef<Path>,
289) -> Result<(), TlsError> {
290 let listener = TcpListener::bind(addr).await.map_err(TlsError::Bind)?;
291 let config = load_tls_config(cert_path, key_path)?;
292 let tls_listener = TlsListener::from_config(listener, config);
293 axum::serve(tls_listener, router.into_make_service())
294 .with_graceful_shutdown(shutdown_signal())
295 .await
296 .map_err(TlsError::Server)?;
297 Ok(())
298}
299
300pub async fn serve_h2_with_listener(
304 router: Router,
305 listener: TcpListener,
306 config: ServerConfig,
307) -> Result<(), TlsError> {
308 let tls_listener = TlsListener::from_config(listener, config);
309 axum::serve(tls_listener, router.into_make_service())
310 .await
311 .map_err(TlsError::Server)?;
312 Ok(())
313}
314
315async fn shutdown_signal() {
317 let ctrl_c = async {
318 tokio::signal::ctrl_c()
319 .await
320 .expect("failed to install Ctrl+C handler");
321 };
322
323 #[cfg(unix)]
324 let terminate = async {
325 tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
326 .expect("failed to install signal handler")
327 .recv()
328 .await;
329 };
330
331 #[cfg(not(unix))]
332 let terminate = std::future::pending::<()>();
333
334 tokio::select! {
335 _ = ctrl_c => {},
336 _ = terminate => {},
337 }
338}
339
340#[cfg(test)]
341mod tests {
342 use super::*;
343 use std::sync::Arc;
344
345 fn generate_self_signed_cert() -> (Vec<u8>, Vec<u8>) {
349 let params = rcgen::CertificateParams::new(vec!["localhost".to_string()]).unwrap();
351 let key_pair = rcgen::KeyPair::generate().unwrap();
352 let cert = params.self_signed(&key_pair).unwrap();
353 let cert_pem = cert.pem().into_bytes();
354 let key_pem = key_pair.serialize_pem().into_bytes();
355 (cert_pem, key_pem)
356 }
357
358 fn write_temp_pem(name: &str, data: &[u8]) -> std::path::PathBuf {
360 let dir = std::env::temp_dir().join("sz-rust-h2-tests");
361 std::fs::create_dir_all(&dir).unwrap();
362 let path = dir.join(name);
363 std::fs::write(&path, data).unwrap();
364 path
365 }
366
367 #[test]
372 fn test_load_tls_config_valid_pem() {
373 let (cert, key) = generate_self_signed_cert();
374 let cert_path = write_temp_pem("valid_cert.pem", &cert);
375 let key_path = write_temp_pem("valid_key.pem", &key);
376
377 let config = load_tls_config(&cert_path, &key_path);
378 assert!(
379 config.is_ok(),
380 "failed to load TLS config: {:?}",
381 config.err()
382 );
383
384 let config = config.unwrap();
385 assert!(config.alpn_protocols.contains(&b"h2".to_vec()));
387 assert!(config.alpn_protocols.contains(&b"http/1.1".to_vec()));
388 }
389
390 #[test]
391 fn test_load_tls_config_missing_cert_file() {
392 let result = load_tls_config("nonexistent_cert.pem", "nonexistent_key.pem");
393 assert!(matches!(result, Err(TlsError::ReadCertFile(_))));
394 }
395
396 #[test]
397 fn test_load_tls_config_missing_key_file() {
398 let (cert, _key) = generate_self_signed_cert();
399 let cert_path = write_temp_pem("valid_cert_for_missing_key.pem", &cert);
400
401 let result = load_tls_config(&cert_path, "nonexistent_key.pem");
402 assert!(matches!(result, Err(TlsError::ReadKeyFile(_))));
403 }
404
405 #[test]
406 fn test_load_tls_config_invalid_pem_content() {
407 let cert_path = write_temp_pem("invalid_cert.pem", b"not a valid PEM");
408 let key_path = write_temp_pem("invalid_key.pem", b"not a valid PEM");
409
410 let result = load_tls_config(&cert_path, &key_path);
411 assert!(result.is_err());
413 }
414
415 #[test]
416 fn test_load_tls_config_empty_pem_files() {
417 let cert_path = write_temp_pem("empty_cert.pem", b"");
418 let key_path = write_temp_pem("empty_key.pem", b"");
419
420 let result = load_tls_config(&cert_path, &key_path);
422 assert!(matches!(result, Err(TlsError::ParseCert(_))));
423 }
424
425 #[test]
426 fn test_load_tls_config_key_without_cert() {
427 let (_cert, key) = generate_self_signed_cert();
429 let key_path = write_temp_pem("only_key.pem", &key);
430 let empty_cert_path = write_temp_pem("empty_for_key_only.pem", b"");
431
432 let result = load_tls_config(&empty_cert_path, &key_path);
433 assert!(matches!(result, Err(TlsError::ParseCert(_))));
434 }
435
436 #[test]
441 fn test_tls_acceptor_constructible() {
442 let (cert, key) = generate_self_signed_cert();
443 let cert_path = write_temp_pem("acceptor_cert.pem", &cert);
444 let key_path = write_temp_pem("acceptor_key.pem", &key);
445
446 let config = load_tls_config(&cert_path, &key_path).unwrap();
447 let _acceptor = tls_acceptor(config);
448 }
449
450 #[tokio::test]
451 async fn test_tls_listener_constructible() {
452 let (cert, key) = generate_self_signed_cert();
453 let cert_path = write_temp_pem("listener_cert.pem", &cert);
454 let key_path = write_temp_pem("listener_key.pem", &key);
455
456 let config = load_tls_config(&cert_path, &key_path).unwrap();
457 let (tcp, _addr) = crate::server::build_tcp_listener("127.0.0.1:0")
458 .await
459 .unwrap();
460 let tls_listener = TlsListener::from_config(tcp, config);
461 assert!(tls_listener.local_addr().is_ok());
463 }
464
465 #[test]
466 fn test_tls_listener_new() {
467 let (cert, key) = generate_self_signed_cert();
468 let cert_path = write_temp_pem("new_cert.pem", &cert);
469 let key_path = write_temp_pem("new_key.pem", &key);
470
471 let config = load_tls_config(&cert_path, &key_path).unwrap();
472 let _arc: Arc<ServerConfig> = Arc::new(config);
473 }
474
475 #[tokio::test]
480 async fn test_serve_h2_with_listener_starts_and_accepts_connections() {
481 use tokio::net::TcpStream;
482
483 let router = Router::new().route("/", axum::routing::get(|| async { "hello h2" }));
485
486 let (cert, key) = generate_self_signed_cert();
488 let cert_path = write_temp_pem("serve_cert.pem", &cert);
489 let key_path = write_temp_pem("serve_key.pem", &key);
490 let config = load_tls_config(&cert_path, &key_path).unwrap();
491 let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
492 .await
493 .unwrap();
494
495 tokio::spawn(async move {
497 let _ = serve_h2_with_listener(router, listener, config).await;
498 });
499
500 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
502
503 let _stream = TcpStream::connect(addr).await.expect("TCP connect failed");
505 }
506
507 #[tokio::test]
508 async fn test_serve_h2_with_invalid_cert_returns_error() {
509 let router = Router::new();
510 let result = serve_h2(router, "127.0.0.1:0", "nonexistent.pem", "nonexistent.pem").await;
511 assert!(result.is_err());
512 match result {
515 Err(TlsError::ReadCertFile(_)) => {}
516 other => panic!("expected ReadCertFile error, got: {:?}", other),
517 }
518 }
519
520 #[tokio::test]
521 async fn test_serve_h2_bind_failure() {
522 let (cert, key) = generate_self_signed_cert();
523 let cert_path = write_temp_pem("bind_fail_cert.pem", &cert);
524 let key_path = write_temp_pem("bind_fail_key.pem", &key);
525
526 let router = Router::new();
528 let result = serve_h2(router, "127.0.0.1:99999", &cert_path, &key_path).await;
529 assert!(matches!(result, Err(TlsError::Bind(_))));
530 }
531
532 #[tokio::test]
537 async fn test_tls_handshake_completes_with_valid_client() {
538 use tokio::io::AsyncReadExt;
539 use tokio::io::AsyncWriteExt;
540 use tokio::net::TcpStream;
541
542 let router = Router::new().route("/", axum::routing::get(|| async { "tls ok" }));
544
545 let (cert, key) = generate_self_signed_cert();
547 let cert_path = write_temp_pem("handshake_cert.pem", &cert);
548 let key_path = write_temp_pem("handshake_key.pem", &key);
549 let config = load_tls_config(&cert_path, &key_path).unwrap();
550 let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
551 .await
552 .unwrap();
553
554 tokio::spawn(async move {
556 let _ = serve_h2_with_listener(router, listener, config).await;
557 });
558
559 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
560
561 let mut stream = TcpStream::connect(addr).await.unwrap();
564
565 let _ = stream.write_all(b"GET / HTTP/1.1\r\n\r\n").await;
567
568 let mut buf = [0u8; 64];
570 let _ = stream.read(&mut buf).await;
571 }
572
573 #[tokio::test]
574 async fn test_serve_h2_full_tls_request_with_rustls_client() {
575 use tokio::io::{AsyncReadExt, AsyncWriteExt};
577 use tokio::net::TcpStream;
578 use tokio_rustls::rustls::pki_types::ServerName;
579 use tokio_rustls::rustls::{ClientConfig, RootCertStore};
580 use tokio_rustls::TlsConnector;
581
582 let router = Router::new().route("/ping", axum::routing::get(|| async { "pong" }));
584
585 let (cert, key) = generate_self_signed_cert();
587 let cert_path = write_temp_pem("full_cert.pem", &cert);
588 let key_path = write_temp_pem("full_key.pem", &key);
589
590 let server_config = load_tls_config(&cert_path, &key_path).unwrap();
592 let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
593 .await
594 .unwrap();
595
596 tokio::spawn(async move {
598 let _ = serve_h2_with_listener(router, listener, server_config).await;
599 });
600 tokio::time::sleep(std::time::Duration::from_millis(150)).await;
601
602 let mut root_store = RootCertStore::empty();
604 let cert_der = rustls_pemfile::certs(&mut BufReader::new(cert.as_slice()))
605 .collect::<Result<Vec<_>, _>>()
606 .unwrap();
607 for c in cert_der {
608 root_store.add(c).unwrap();
609 }
610 let client_config = ClientConfig::builder()
611 .with_root_certificates(root_store)
612 .with_no_client_auth();
613
614 let connector = TlsConnector::from(Arc::new(client_config));
616 let tcp_stream = TcpStream::connect(addr).await.unwrap();
617 let server_name = ServerName::try_from("localhost").unwrap();
618 let mut tls_stream = connector.connect(server_name, tcp_stream).await.unwrap();
619
620 tls_stream
622 .write_all(b"GET /ping HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
623 .await
624 .unwrap();
625
626 let mut response = Vec::new();
628 tls_stream.read_to_end(&mut response).await.unwrap();
629 let response_str = String::from_utf8_lossy(&response);
630
631 assert!(
633 response_str.contains("pong"),
634 "expected response to contain 'pong', got: {}",
635 response_str
636 );
637 assert!(
639 response_str.starts_with("HTTP/1.1") || response_str.starts_with("HTTP/2"),
640 "expected HTTP response, got: {}",
641 response_str.lines().next().unwrap_or("")
642 );
643 }
644}