use std::path::Path;
use std::sync::Arc;
use axum::serve::Listener;
use axum::Router;
use tokio::net::TcpListener;
use tokio_rustls::rustls::pki_types::{pem::PemObject, CertificateDer, PrivateKeyDer};
use tokio_rustls::rustls::ServerConfig;
use tokio_rustls::server::TlsStream;
use tokio_rustls::TlsAcceptor;
#[derive(Debug, thiserror::Error)]
pub enum TlsError {
#[error("failed to bind TCP listener: {0}")]
Bind(#[source] std::io::Error),
#[error("failed to read certificate file: {0}")]
ReadCertFile(#[source] std::io::Error),
#[error("failed to read private key file: {0}")]
ReadKeyFile(#[source] std::io::Error),
#[error("failed to parse certificate PEM: {0}")]
ParseCert(#[source] std::io::Error),
#[error("failed to parse private key PEM: {0}")]
ParseKey(#[source] std::io::Error),
#[error("no private key found in PEM file")]
NoPrivateKey,
#[error("failed to build rustls ServerConfig: {0}")]
BuildServerConfig(#[source] tokio_rustls::rustls::Error),
#[error("server error: {0}")]
Server(#[source] std::io::Error),
}
pub async fn load_tls_config(
cert_path: impl AsRef<Path>,
key_path: impl AsRef<Path>,
) -> Result<ServerConfig, TlsError> {
let certs = load_certs(cert_path).await?;
let key = load_private_key(key_path).await?;
let mut config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.map_err(TlsError::BuildServerConfig)?;
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
Ok(config)
}
async fn load_certs(path: impl AsRef<Path>) -> Result<Vec<CertificateDer<'static>>, TlsError> {
let cert_data = tokio::fs::read(path.as_ref())
.await
.map_err(TlsError::ReadCertFile)?;
let certs: Vec<CertificateDer<'static>> = CertificateDer::pem_slice_iter(&cert_data)
.collect::<Result<Vec<_>, _>>()
.map_err(|e| TlsError::ParseCert(std::io::Error::other(e)))?;
if certs.is_empty() {
return Err(TlsError::ParseCert(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"no certificate found in PEM file",
)));
}
Ok(certs)
}
async fn load_private_key(path: impl AsRef<Path>) -> Result<PrivateKeyDer<'static>, TlsError> {
let key_data = tokio::fs::read(path.as_ref())
.await
.map_err(TlsError::ReadKeyFile)?;
let keys: Vec<PrivateKeyDer<'static>> = PrivateKeyDer::pem_slice_iter(&key_data)
.collect::<Result<Vec<_>, _>>()
.map_err(|e| TlsError::ParseKey(std::io::Error::other(e)))?;
keys.into_iter().next().ok_or(TlsError::NoPrivateKey)
}
pub fn tls_acceptor(config: ServerConfig) -> TlsAcceptor {
TlsAcceptor::from(Arc::new(config))
}
pub struct TlsListener {
tcp: TcpListener,
acceptor: TlsAcceptor,
pending_rx: tokio::sync::mpsc::UnboundedReceiver<(
TlsStream<tokio::net::TcpStream>,
std::net::SocketAddr,
)>,
pending_tx: tokio::sync::mpsc::UnboundedSender<(
TlsStream<tokio::net::TcpStream>,
std::net::SocketAddr,
)>,
}
impl std::fmt::Debug for TlsListener {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TlsListener")
.field("tcp", &self.tcp)
.finish_non_exhaustive()
}
}
impl TlsListener {
pub fn new(tcp: TcpListener, acceptor: TlsAcceptor) -> Self {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
Self {
tcp,
acceptor,
pending_rx: rx,
pending_tx: tx,
}
}
pub fn from_config(tcp: TcpListener, config: ServerConfig) -> Self {
Self::new(tcp, tls_acceptor(config))
}
}
impl Listener for TlsListener {
type Io = TlsStream<tokio::net::TcpStream>;
type Addr = std::net::SocketAddr;
async fn accept(&mut self) -> (Self::Io, Self::Addr) {
let mut backoff_ms: u64 = 100;
const MAX_BACKOFF_MS: u64 = 5_000;
loop {
tokio::select! {
Some((tls, addr)) = self.pending_rx.recv() => {
return (tls, addr);
}
accept_result = self.tcp.accept() => {
match accept_result {
Ok((tcp, addr)) => {
let acceptor = self.acceptor.clone();
let tx = self.pending_tx.clone();
tokio::spawn(async move {
match acceptor.accept(tcp).await {
Ok(tls) => {
if tx.send((tls, addr)).is_err() {
tracing::warn!(
"pending channel closed, dropping TLS connection from {addr}"
);
}
}
Err(e) => {
tracing::warn!("TLS handshake failed from {addr}: {e}");
}
}
});
backoff_ms = 100;
}
Err(e) => {
if matches!(
e.kind(),
std::io::ErrorKind::ConnectionRefused
| std::io::ErrorKind::ConnectionAborted
| std::io::ErrorKind::ConnectionReset
) {
continue;
}
tracing::error!("accept error: {e}");
tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
backoff_ms = (backoff_ms * 2).min(MAX_BACKOFF_MS);
}
}
}
}
}
}
fn local_addr(&self) -> std::io::Result<Self::Addr> {
self.tcp.local_addr()
}
}
pub async fn serve_h2(
router: Router,
addr: &str,
cert_path: impl AsRef<Path>,
key_path: impl AsRef<Path>,
) -> Result<(), TlsError> {
let listener = TcpListener::bind(addr).await.map_err(TlsError::Bind)?;
let config = load_tls_config(cert_path, key_path).await?;
serve_h2_with_listener(router, listener, config).await
}
pub async fn serve_h2_with_graceful_shutdown(
router: Router,
addr: &str,
cert_path: impl AsRef<Path>,
key_path: impl AsRef<Path>,
) -> Result<(), TlsError> {
let listener = TcpListener::bind(addr).await.map_err(TlsError::Bind)?;
let config = load_tls_config(cert_path, key_path).await?;
let tls_listener = TlsListener::from_config(listener, config);
axum::serve(tls_listener, router.into_make_service())
.with_graceful_shutdown(shutdown_signal())
.await
.map_err(TlsError::Server)?;
Ok(())
}
pub async fn serve_h2_with_listener(
router: Router,
listener: TcpListener,
config: ServerConfig,
) -> Result<(), TlsError> {
let tls_listener = TlsListener::from_config(listener, config);
axum::serve(tls_listener, router.into_make_service())
.await
.map_err(TlsError::Server)?;
Ok(())
}
async fn shutdown_signal() {
let ctrl_c = async {
tokio::signal::ctrl_c()
.await
.expect("failed to install Ctrl+C handler");
};
#[cfg(unix)]
let terminate = async {
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
.expect("failed to install signal handler")
.recv()
.await;
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {},
_ = terminate => {},
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
fn generate_self_signed_cert() -> (Vec<u8>, Vec<u8>) {
let params = rcgen::CertificateParams::new(vec!["localhost".to_string()]).unwrap();
let key_pair = rcgen::KeyPair::generate().unwrap();
let cert = params.self_signed(&key_pair).unwrap();
let cert_pem = cert.pem().into_bytes();
let key_pem = key_pair.serialize_pem().into_bytes();
(cert_pem, key_pem)
}
fn write_temp_pem(name: &str, data: &[u8]) -> std::path::PathBuf {
let dir = std::env::temp_dir().join("sz-rust-h2-tests");
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join(name);
std::fs::write(&path, data).unwrap();
path
}
#[tokio::test]
async fn test_load_tls_config_valid_pem() {
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("valid_cert.pem", &cert);
let key_path = write_temp_pem("valid_key.pem", &key);
let config = load_tls_config(&cert_path, &key_path).await;
assert!(
config.is_ok(),
"failed to load TLS config: {:?}",
config.err()
);
let config = config.unwrap();
assert!(config.alpn_protocols.contains(&b"h2".to_vec()));
assert!(config.alpn_protocols.contains(&b"http/1.1".to_vec()));
}
#[tokio::test]
async fn test_load_tls_config_missing_cert_file() {
let result = load_tls_config("nonexistent_cert.pem", "nonexistent_key.pem").await;
assert!(matches!(result, Err(TlsError::ReadCertFile(_))));
}
#[tokio::test]
async fn test_load_tls_config_missing_key_file() {
let (cert, _key) = generate_self_signed_cert();
let cert_path = write_temp_pem("valid_cert_for_missing_key.pem", &cert);
let result = load_tls_config(&cert_path, "nonexistent_key.pem").await;
assert!(matches!(result, Err(TlsError::ReadKeyFile(_))));
}
#[tokio::test]
async fn test_load_tls_config_invalid_pem_content() {
let cert_path = write_temp_pem("invalid_cert.pem", b"not a valid PEM");
let key_path = write_temp_pem("invalid_key.pem", b"not a valid PEM");
let result = load_tls_config(&cert_path, &key_path).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_load_tls_config_empty_pem_files() {
let cert_path = write_temp_pem("empty_cert.pem", b"");
let key_path = write_temp_pem("empty_key.pem", b"");
let result = load_tls_config(&cert_path, &key_path).await;
assert!(matches!(result, Err(TlsError::ParseCert(_))));
}
#[tokio::test]
async fn test_load_tls_config_key_without_cert() {
let (_cert, key) = generate_self_signed_cert();
let key_path = write_temp_pem("only_key.pem", &key);
let empty_cert_path = write_temp_pem("empty_for_key_only.pem", b"");
let result = load_tls_config(&empty_cert_path, &key_path).await;
assert!(matches!(result, Err(TlsError::ParseCert(_))));
}
#[tokio::test]
async fn test_tls_acceptor_constructible() {
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("acceptor_cert.pem", &cert);
let key_path = write_temp_pem("acceptor_key.pem", &key);
let config = load_tls_config(&cert_path, &key_path).await.unwrap();
let _acceptor = tls_acceptor(config);
}
#[tokio::test]
async fn test_tls_listener_constructible() {
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("listener_cert.pem", &cert);
let key_path = write_temp_pem("listener_key.pem", &key);
let config = load_tls_config(&cert_path, &key_path).await.unwrap();
let (tcp, _addr) = crate::server::build_tcp_listener("127.0.0.1:0")
.await
.unwrap();
let tls_listener = TlsListener::from_config(tcp, config);
assert!(tls_listener.local_addr().is_ok());
}
#[tokio::test]
async fn test_tls_listener_new() {
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("new_cert.pem", &cert);
let key_path = write_temp_pem("new_key.pem", &key);
let config = load_tls_config(&cert_path, &key_path).await.unwrap();
let _arc: Arc<ServerConfig> = Arc::new(config);
}
#[tokio::test]
async fn test_serve_h2_with_listener_starts_and_accepts_connections() {
use tokio::net::TcpStream;
let router = Router::new().route("/", axum::routing::get(|| async { "hello h2" }));
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("serve_cert.pem", &cert);
let key_path = write_temp_pem("serve_key.pem", &key);
let config = load_tls_config(&cert_path, &key_path).await.unwrap();
let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
.await
.unwrap();
tokio::spawn(async move {
let _ = serve_h2_with_listener(router, listener, config).await;
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let _stream = TcpStream::connect(addr).await.expect("TCP connect failed");
}
#[tokio::test]
async fn test_serve_h2_with_invalid_cert_returns_error() {
let router = Router::new();
let result = serve_h2(router, "127.0.0.1:0", "nonexistent.pem", "nonexistent.pem").await;
assert!(result.is_err());
match result {
Err(TlsError::ReadCertFile(_)) => {}
other => panic!("expected ReadCertFile error, got: {:?}", other),
}
}
#[tokio::test]
async fn test_serve_h2_bind_failure() {
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("bind_fail_cert.pem", &cert);
let key_path = write_temp_pem("bind_fail_key.pem", &key);
let router = Router::new();
let result = serve_h2(router, "127.0.0.1:99999", &cert_path, &key_path).await;
assert!(matches!(result, Err(TlsError::Bind(_))));
}
#[tokio::test]
async fn test_tls_handshake_completes_with_valid_client() {
use tokio::io::AsyncReadExt;
use tokio::io::AsyncWriteExt;
use tokio::net::TcpStream;
let router = Router::new().route("/", axum::routing::get(|| async { "tls ok" }));
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("handshake_cert.pem", &cert);
let key_path = write_temp_pem("handshake_key.pem", &key);
let config = load_tls_config(&cert_path, &key_path).await.unwrap();
let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
.await
.unwrap();
tokio::spawn(async move {
let _ = serve_h2_with_listener(router, listener, config).await;
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let mut stream = TcpStream::connect(addr).await.unwrap();
let _ = stream.write_all(b"GET / HTTP/1.1\r\n\r\n").await;
let mut buf = [0u8; 64];
let _ = stream.read(&mut buf).await;
}
#[tokio::test]
async fn test_serve_h2_full_tls_request_with_rustls_client() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tokio_rustls::rustls::pki_types::ServerName;
use tokio_rustls::rustls::{ClientConfig, RootCertStore};
use tokio_rustls::TlsConnector;
let router = Router::new().route("/ping", axum::routing::get(|| async { "pong" }));
let (cert, key) = generate_self_signed_cert();
let cert_path = write_temp_pem("full_cert.pem", &cert);
let key_path = write_temp_pem("full_key.pem", &key);
let server_config = load_tls_config(&cert_path, &key_path).await.unwrap();
let (listener, addr) = crate::server::build_tcp_listener("127.0.0.1:0")
.await
.unwrap();
tokio::spawn(async move {
let _ = serve_h2_with_listener(router, listener, server_config).await;
});
tokio::time::sleep(std::time::Duration::from_millis(150)).await;
let mut root_store = RootCertStore::empty();
let cert_der = CertificateDer::pem_slice_iter(&cert[..])
.collect::<Result<Vec<_>, _>>()
.unwrap();
for c in cert_der {
root_store.add(c).unwrap();
}
let client_config = ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
let connector = TlsConnector::from(Arc::new(client_config));
let tcp_stream = TcpStream::connect(addr).await.unwrap();
let server_name = ServerName::try_from("localhost").unwrap();
let mut tls_stream = connector.connect(server_name, tcp_stream).await.unwrap();
tls_stream
.write_all(b"GET /ping HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut response = Vec::new();
tls_stream.read_to_end(&mut response).await.unwrap();
let response_str = String::from_utf8_lossy(&response);
assert!(
response_str.contains("pong"),
"expected response to contain 'pong', got: {}",
response_str
);
assert!(
response_str.starts_with("HTTP/1.1") || response_str.starts_with("HTTP/2"),
"expected HTTP response, got: {}",
response_str.lines().next().unwrap_or("")
);
}
}