Skip to main content

sz_rust_http_facade/
tls.rs

1//! TLS 配置 + HTTP/2 服务
2//!
3//! 提供 `TlsConfig` 结构体和 `serve_http2` 函数,支持:
4//! - ALPN 协商 h2 + http/1.1 自动回退
5//! - 从 PEM 文件加载证书
6//! - axum::serve 集成
7
8use 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/// TLS 错误
19#[derive(Debug, Error)]
20pub enum TlsError {
21    /// IO 错误
22    #[error("IO error: {0}")]
23    Io(#[from] std::io::Error),
24    /// 未找到有效证书
25    #[error("No valid certificate found")]
26    NoCertificate,
27    /// 未找到有效私钥
28    #[error("No valid private key found")]
29    NoPrivateKey,
30    /// TLS 协议错误
31    #[error("TLS error: {0}")]
32    Tls(String),
33}
34
35/// TLS 配置
36#[derive(Debug, Clone)]
37pub struct TlsConfig {
38    /// 证书文件路径(PEM 格式)
39    pub cert_path: PathBuf,
40    /// 私钥文件路径(PEM 格式)
41    pub key_path: PathBuf,
42    /// ALPN 协议列表(默认 h2 + http/1.1)
43    pub alpn: Vec<Vec<u8>>,
44}
45
46impl TlsConfig {
47    /// 创建 TLS 配置
48    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    /// 仅 HTTP/2(不回退 http/1.1)
57    pub fn h2_only(mut self) -> Self {
58        self.alpn = vec![b"h2".to_vec()];
59        self
60    }
61
62    /// 仅 HTTP/1.1
63    pub fn http1_only(mut self) -> Self {
64        self.alpn = vec![b"http/1.1".to_vec()];
65        self
66    }
67
68    /// 构建 TlsAcceptor
69    ///
70    /// 异步读取证书和私钥文件(遵循 tokio::fs 铁律)。
71    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
86/// 从 PEM 文件加载证书
87async 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
99/// 从 PEM 文件加载私钥
100async 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
110/// 启动 HTTP/2 + HTTPS 服务
111///
112/// ALPN 协商:客户端支持 h2 时用 HTTP/2,否则回退 http/1.1。
113///
114/// 实现方式:手动接受 TCP 连接 → TLS 握手 → hyper HTTP/2 服务。
115pub 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}