Skip to main content

rskit_server/http/
component.rs

1use std::net::SocketAddr;
2use std::sync::Arc;
3
4use async_trait::async_trait;
5use axum::Router;
6use parking_lot::Mutex;
7use rskit_bootstrap::{Component, Health};
8use rskit_errors::{AppError, AppResult, ErrorCode};
9use tokio::net::TcpListener;
10use tokio_util::sync::CancellationToken;
11
12use super::serve::{ConnectionContext, serve_listener, serve_tls_listener};
13use super::tls::build_tls_acceptor;
14use crate::http_config::HttpServerConfig;
15
16/// HTTP server that implements the [`Component`] lifecycle.
17pub struct HttpServer {
18    pub(super) config: Arc<HttpServerConfig>,
19    pub(super) cancel: CancellationToken,
20    pub(super) router: Arc<tokio::sync::Mutex<Option<Router>>>,
21    pub(super) local_addr: Arc<Mutex<Option<SocketAddr>>>,
22}
23
24impl HttpServer {
25    /// Bind address as `host:port`.
26    #[must_use]
27    pub fn bind_addr(&self) -> String {
28        self.config.bind_addr()
29    }
30
31    /// Actual local socket address after [`start`](Component::start) binds.
32    #[must_use]
33    pub fn local_addr(&self) -> Option<SocketAddr> {
34        *self.local_addr.lock()
35    }
36}
37
38#[async_trait]
39impl Component for HttpServer {
40    fn name(&self) -> &str {
41        "http-server"
42    }
43
44    async fn start(&self) -> AppResult<()> {
45        let router = self
46            .router
47            .lock()
48            .await
49            .take()
50            .ok_or_else(|| AppError::new(ErrorCode::Internal, "HTTP server already started"))?;
51
52        let addr: SocketAddr = self.config.bind_addr().parse().map_err(|error| {
53            AppError::new(
54                ErrorCode::Internal,
55                format!("invalid bind address: {error}"),
56            )
57        })?;
58        let tls_acceptor = if let Some(tls) = &self.config.tls {
59            Some(build_tls_acceptor(tls)?)
60        } else {
61            None
62        };
63
64        let listener = TcpListener::bind(addr).await.map_err(|error| {
65            AppError::new(
66                ErrorCode::Internal,
67                format!("HTTP server bind failed for {addr}: {error}"),
68            )
69        })?;
70        let actual_addr = listener.local_addr().map_err(|error| {
71            AppError::new(
72                ErrorCode::Internal,
73                format!("failed to inspect HTTP server local address: {error}"),
74            )
75        })?;
76        *self.local_addr.lock() = Some(actual_addr);
77
78        let cancel = self.cancel.clone();
79        let config = Arc::clone(&self.config);
80        tokio::spawn(async move {
81            if let Some(acceptor) = tls_acceptor {
82                tracing::info!(addr = %actual_addr, "HTTPS server listening");
83                let context = ConnectionContext::new(router, Arc::clone(&config), true);
84                serve_tls_listener(listener, acceptor, context, cancel).await;
85            } else {
86                tracing::info!(addr = %actual_addr, "HTTP server listening");
87                let context =
88                    ConnectionContext::new(router, Arc::clone(&config), config.enable_h2c);
89                serve_listener(listener, context, cancel).await;
90            }
91        });
92
93        Ok(())
94    }
95
96    async fn stop(&self) -> AppResult<()> {
97        self.cancel.cancel();
98        Ok(())
99    }
100
101    fn health(&self) -> Health {
102        Health::healthy("http-server")
103    }
104}
105
106#[cfg(test)]
107mod tests {
108    use std::time::Duration;
109
110    use axum::Router;
111    use axum::routing::get;
112    use rskit_bootstrap::Component;
113    use rskit_errors::ErrorCode;
114    use rskit_security::TlsConfig;
115    use tokio::io::{AsyncReadExt, AsyncWriteExt};
116    use tokio_util::sync::CancellationToken;
117
118    use crate::http::HttpServerBuilder;
119    use crate::http::test_support::local_config;
120
121    fn testdata(name: &str) -> String {
122        format!("{}/testdata/{name}", env!("CARGO_MANIFEST_DIR"))
123    }
124
125    #[tokio::test]
126    async fn lifecycle_sets_local_address_and_cancels_shutdown() {
127        let server = HttpServerBuilder::new(local_config(), CancellationToken::new())
128            .build()
129            .expect("build server");
130
131        assert_eq!(server.bind_addr(), "127.0.0.1:0");
132        assert!(server.local_addr().is_none());
133        assert!(server.health().is_healthy());
134
135        server.start().await.expect("start http server");
136        assert!(server.local_addr().is_some());
137        server.stop().await.expect("stop http server");
138    }
139
140    #[tokio::test]
141    async fn local_http_listener_serves_requests_and_rejects_double_start() {
142        let server = HttpServerBuilder::new(local_config(), CancellationToken::new())
143            .with_router(Router::new().route("/ping", get(|| async { "pong" })))
144            .build()
145            .expect("build server");
146
147        server.start().await.expect("start http server");
148        let second_start = server.start().await.unwrap_err();
149        assert_eq!(second_start.code(), ErrorCode::Internal);
150        assert!(second_start.message().contains("already started"));
151
152        let addr = server.local_addr().expect("local address");
153        let mut stream = tokio::net::TcpStream::connect(addr)
154            .await
155            .expect("connect to local server");
156        stream
157            .write_all(b"GET /ping HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
158            .await
159            .expect("write request");
160        let mut response = String::new();
161        tokio::time::timeout(Duration::from_secs(2), stream.read_to_string(&mut response))
162            .await
163            .expect("response read timed out")
164            .expect("read response");
165
166        assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
167        assert!(response.contains("pong"), "{response}");
168        server.stop().await.expect("stop http server");
169    }
170
171    #[tokio::test]
172    async fn local_http_listener_serves_http1_when_h2c_disabled() {
173        let mut config = local_config();
174        config.enable_h2c = false;
175        let server = HttpServerBuilder::new(config, CancellationToken::new())
176            .with_router(Router::new().route("/http1", get(|| async { "ok" })))
177            .build()
178            .expect("build server");
179
180        server.start().await.expect("start http1 server");
181        let addr = server.local_addr().expect("local address");
182        let mut stream = tokio::net::TcpStream::connect(addr)
183            .await
184            .expect("connect to local server");
185        stream
186            .write_all(b"GET /http1 HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
187            .await
188            .expect("write request");
189        let mut response = String::new();
190        tokio::time::timeout(Duration::from_secs(2), stream.read_to_string(&mut response))
191            .await
192            .expect("response read timed out")
193            .expect("read response");
194
195        assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
196        assert!(response.contains("ok"), "{response}");
197        server.stop().await.expect("stop http1 server");
198    }
199
200    #[tokio::test]
201    async fn https_listener_completes_tls_handshake_and_serves_request() {
202        use std::sync::Arc;
203
204        use rustls::RootCertStore;
205        use rustls::pki_types::{CertificateDer, ServerName, pem::PemObject};
206        use tokio_rustls::TlsConnector;
207
208        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
209
210        let mut config = local_config();
211        config.tls = Some(TlsConfig {
212            cert_file: Some(testdata("cert.pem")),
213            key_file: Some(testdata("key.pem")),
214            ..Default::default()
215        });
216        let server = HttpServerBuilder::new(config, CancellationToken::new())
217            .with_router(Router::new().route("/secure", get(|| async { "encrypted" })))
218            .build()
219            .expect("build https server");
220
221        server.start().await.expect("start https server");
222        let addr = server.local_addr().expect("local address");
223
224        let mut roots = RootCertStore::empty();
225        roots
226            .add(CertificateDer::from_pem_file(testdata("cert.pem")).expect("load test cert"))
227            .expect("trust test cert");
228        let client_config = rustls::ClientConfig::builder()
229            .with_root_certificates(roots)
230            .with_no_client_auth();
231        let connector = TlsConnector::from(Arc::new(client_config));
232        let server_name = ServerName::try_from("localhost").expect("server name");
233
234        let tcp = tokio::net::TcpStream::connect(addr)
235            .await
236            .expect("connect to https server");
237        let mut stream = connector
238            .connect(server_name, tcp)
239            .await
240            .expect("tls handshake");
241        stream
242            .write_all(b"GET /secure HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
243            .await
244            .expect("write request");
245        let mut response = String::new();
246        tokio::time::timeout(Duration::from_secs(2), stream.read_to_string(&mut response))
247            .await
248            .expect("response read timed out")
249            .expect("read response");
250
251        assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
252        assert!(response.contains("encrypted"), "{response}");
253        server.stop().await.expect("stop https server");
254    }
255
256    #[tokio::test]
257    async fn start_reports_bind_failure_for_address_in_use() {
258        let occupied = tokio::net::TcpListener::bind("127.0.0.1:0")
259            .await
260            .expect("bind probe listener");
261        let taken = occupied.local_addr().expect("probe local address");
262
263        let mut config = local_config();
264        config.port = taken.port();
265        let server = HttpServerBuilder::new(config, CancellationToken::new())
266            .build()
267            .expect("build server");
268
269        let error = server.start().await.expect_err("bind should fail");
270        assert_eq!(error.code(), ErrorCode::Internal);
271        assert!(
272            error.message().contains("bind failed"),
273            "{}",
274            error.message()
275        );
276        assert!(server.local_addr().is_none());
277    }
278
279    #[tokio::test]
280    async fn https_listener_times_out_stalled_handshake() {
281        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
282
283        let mut config = local_config();
284        config.read_timeout = Duration::from_millis(50);
285        config.tls = Some(TlsConfig {
286            cert_file: Some(testdata("cert.pem")),
287            key_file: Some(testdata("key.pem")),
288            ..Default::default()
289        });
290        let server = HttpServerBuilder::new(config, CancellationToken::new())
291            .build()
292            .expect("build https server");
293
294        server.start().await.expect("start https server");
295        let addr = server.local_addr().expect("local address");
296        // Connect but never send a ClientHello so the handshake stalls and the
297        // server's read-timeout aborts it.
298        let _stream = tokio::net::TcpStream::connect(addr)
299            .await
300            .expect("connect to https server");
301        tokio::time::sleep(Duration::from_millis(150)).await;
302
303        server.stop().await.expect("stop https server");
304    }
305
306    async fn assert_serves_and_drains_on_shutdown(enable_h2c: bool) {
307        let mut config = local_config();
308        config.enable_h2c = enable_h2c;
309        let server = HttpServerBuilder::new(config, CancellationToken::new())
310            .with_router(Router::new().route("/keep", get(|| async { "ok" })))
311            .build()
312            .expect("build server");
313
314        server.start().await.expect("start server");
315        let addr = server.local_addr().expect("local address");
316        let mut stream = tokio::net::TcpStream::connect(addr)
317            .await
318            .expect("connect to local server");
319        // Keep-alive request (no `Connection: close`) leaves the server-side
320        // connection idle so that stopping the server exercises graceful drain.
321        stream
322            .write_all(b"GET /keep HTTP/1.1\r\nHost: localhost\r\n\r\n")
323            .await
324            .expect("write request");
325
326        let mut buf = vec![0u8; 1024];
327        let read = tokio::time::timeout(Duration::from_secs(2), stream.read(&mut buf))
328            .await
329            .expect("response read timed out")
330            .expect("read response");
331        let head = String::from_utf8_lossy(&buf[..read]);
332        assert!(head.starts_with("HTTP/1.1 200 OK"), "{head}");
333
334        server.stop().await.expect("stop server");
335
336        let mut rest = Vec::new();
337        let _ = tokio::time::timeout(Duration::from_secs(2), stream.read_to_end(&mut rest)).await;
338    }
339
340    #[tokio::test]
341    async fn h2c_connection_drains_on_graceful_shutdown() {
342        assert_serves_and_drains_on_shutdown(true).await;
343    }
344
345    #[tokio::test]
346    async fn http1_connection_drains_on_graceful_shutdown() {
347        assert_serves_and_drains_on_shutdown(false).await;
348    }
349
350    #[tokio::test]
351    async fn https_listener_survives_non_tls_client() {
352        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
353
354        let mut config = local_config();
355        config.tls = Some(TlsConfig {
356            cert_file: Some(testdata("cert.pem")),
357            key_file: Some(testdata("key.pem")),
358            ..Default::default()
359        });
360        let server = HttpServerBuilder::new(config, CancellationToken::new())
361            .build()
362            .expect("build https server");
363
364        server.start().await.expect("start https server");
365        let addr = server.local_addr().expect("local address");
366        let mut stream = tokio::net::TcpStream::connect(addr)
367            .await
368            .expect("connect to https server");
369        // Plain-text bytes cannot complete the TLS handshake; the server must
370        // log and drop the connection without terminating the accept loop.
371        stream
372            .write_all(b"GET / HTTP/1.1\r\nHost: localhost\r\n\r\n")
373            .await
374            .expect("write plaintext request");
375        let mut buf = vec![0u8; 64];
376        let _ = tokio::time::timeout(Duration::from_millis(200), stream.read(&mut buf)).await;
377
378        server.stop().await.expect("stop https server");
379    }
380}