rskit_server/http/
component.rs1use 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
16pub 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 #[must_use]
27 pub fn bind_addr(&self) -> String {
28 self.config.bind_addr()
29 }
30
31 #[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 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 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 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}