1use std::sync::Arc;
2use std::time::Duration;
3
4use axum::extract::DefaultBodyLimit;
5use axum::{Router, http::StatusCode};
6use parking_lot::Mutex;
7use rskit_errors::AppResult;
8use rskit_http::{SecurityHeadersConfig, SecurityHeadersLayer};
9use tokio_util::sync::CancellationToken;
10use tower_http::{
11 request_id::{MakeRequestUuid, SetRequestIdLayer},
12 timeout::TimeoutLayer,
13 trace::TraceLayer,
14};
15
16use super::component::HttpServer;
17use crate::http_config::HttpServerConfig;
18use crate::middleware::HttpMiddlewareStack;
19
20pub struct HttpServerBuilder {
22 config: HttpServerConfig,
23 cancel: CancellationToken,
24 router: Router,
25 middleware: HttpMiddlewareStack,
26 security_headers: Option<SecurityHeadersConfig>,
27}
28
29impl HttpServerBuilder {
30 pub fn new(config: HttpServerConfig, cancel: CancellationToken) -> Self {
32 Self {
33 config,
34 cancel,
35 router: Router::new(),
36 middleware: HttpMiddlewareStack::new(),
37 security_headers: None,
38 }
39 }
40
41 #[must_use]
43 pub fn with_router(mut self, router: Router) -> Self {
44 self.router = self.router.merge(router);
45 self
46 }
47
48 #[must_use]
50 pub fn with_middleware_stack(mut self, middleware: HttpMiddlewareStack) -> Self {
51 self.middleware = middleware;
52 self
53 }
54
55 #[must_use]
57 pub fn with_logging_transform<F>(mut self, transform: F) -> Self
58 where
59 F: Fn(Router) -> Router + Send + Sync + 'static,
60 {
61 self.middleware = self.middleware.with_logging_transform(transform);
62 self
63 }
64
65 #[must_use]
67 pub fn with_auth_transform<F>(mut self, transform: F) -> Self
68 where
69 F: Fn(Router) -> Router + Send + Sync + 'static,
70 {
71 self.middleware = self.middleware.with_auth_transform(transform);
72 self
73 }
74
75 #[must_use]
77 pub fn with_validation_transform<F>(mut self, transform: F) -> Self
78 where
79 F: Fn(Router) -> Router + Send + Sync + 'static,
80 {
81 self.middleware = self.middleware.with_validation_transform(transform);
82 self
83 }
84
85 #[must_use]
87 pub fn with_metrics_transform<F>(mut self, transform: F) -> Self
88 where
89 F: Fn(Router) -> Router + Send + Sync + 'static,
90 {
91 self.middleware = self.middleware.with_metrics_transform(transform);
92 self
93 }
94
95 #[must_use = "builder methods return a new builder; use the returned value"]
101 pub fn with_cors(self) -> AppResult<Self> {
102 if let Some(cors_cfg) = self.config.cors.as_ref() {
103 let _ = cors_cfg.layer()?;
104 }
105 Ok(self)
106 }
107
108 #[must_use = "builder methods return a new builder; use the returned value"]
113 pub fn with_security_headers(self) -> AppResult<Self> {
114 self.with_security_headers_config(SecurityHeadersConfig::default())
115 }
116
117 #[must_use = "builder methods return a new builder; use the returned value"]
122 pub fn with_security_headers_config(
123 mut self,
124 config: SecurityHeadersConfig,
125 ) -> AppResult<Self> {
126 let _ = SecurityHeadersLayer::new(&config)?;
127 self.security_headers = Some(config);
128 Ok(self)
129 }
130
131 pub fn build(self) -> AppResult<HttpServer> {
135 let builder = self;
136 let security_headers = builder.security_headers.clone();
137 let request_timeout = builder.config.request_timeout;
138 let max_body_bytes = builder.config.max_body_bytes;
139 let cors = builder.config.cors.clone();
140 let router = builder.middleware.apply(builder.router);
141 let router = apply_canonical_tracing(router);
142 let router = apply_baseline_layers(
143 router,
144 security_headers,
145 request_timeout,
146 max_body_bytes,
147 cors,
148 )?;
149 Ok(HttpServer {
150 config: Arc::new(builder.config),
151 cancel: builder.cancel,
152 router: Arc::new(tokio::sync::Mutex::new(Some(router))),
153 local_addr: Arc::new(Mutex::new(None)),
154 })
155 }
156}
157
158fn apply_canonical_tracing(router: Router) -> Router {
159 use http::Request;
160 use tower_http::trace::DefaultOnResponse;
161 use tracing::Level;
162
163 let trace_layer = TraceLayer::new_for_http()
164 .make_span_with(|request: &Request<_>| {
165 let path = request.uri().path();
166 tracing::info_span!(
167 "http_request",
168 method = %request.method(),
169 "http.target" = path,
170 status_code = tracing::field::Empty,
171 )
172 })
173 .on_response(DefaultOnResponse::new().level(Level::INFO));
174
175 router.layer(trace_layer)
176}
177
178fn apply_baseline_layers(
179 router: Router,
180 security_headers: Option<SecurityHeadersConfig>,
181 request_timeout: Duration,
182 max_body_bytes: usize,
183 cors: Option<crate::CorsPolicy>,
184) -> AppResult<Router> {
185 let security_headers = SecurityHeadersLayer::new(&security_headers.unwrap_or_default())?;
186 let router = router
187 .layer(TimeoutLayer::with_status_code(
188 StatusCode::REQUEST_TIMEOUT,
189 request_timeout,
190 ))
191 .layer(DefaultBodyLimit::max(max_body_bytes))
192 .layer(security_headers);
193 let router = if let Some(cors_cfg) = cors.as_ref() {
194 router.layer(cors_cfg.layer()?)
195 } else {
196 router
197 };
198 Ok(router.layer(SetRequestIdLayer::x_request_id(MakeRequestUuid)))
199}
200
201#[cfg(test)]
202mod tests {
203 use std::sync::Arc;
204 use std::sync::atomic::{AtomicUsize, Ordering};
205
206 use axum::body::Body;
207 use axum::{http::Request, routing::get};
208 use rskit_bootstrap::Component;
209 use rskit_errors::ErrorCode;
210 use rskit_http::CorsPolicy;
211 use rskit_security::TransportSecurity;
212 use tower::ServiceExt;
213
214 use super::*;
215 use crate::http::test_support::local_config;
216
217 #[tokio::test]
218 async fn builder_applies_baseline_layers_to_application_routes() {
219 let server = HttpServerBuilder::new(local_config(), CancellationToken::new())
220 .with_router(Router::new().route("/", get(|| async { "ok" })))
221 .with_security_headers_config(
222 SecurityHeadersConfig::default()
223 .with_transport_security(TransportSecurity::AllowInsecureLocal),
224 )
225 .expect("configure security headers")
226 .build()
227 .expect("build server");
228 let router = server
229 .router
230 .lock()
231 .await
232 .take()
233 .expect("router is present");
234
235 let response = router
236 .oneshot(
237 Request::builder()
238 .uri("/")
239 .body(Body::empty())
240 .expect("request"),
241 )
242 .await
243 .expect("route response");
244
245 assert_eq!(response.status(), StatusCode::OK);
246 assert_eq!(
247 response
248 .headers()
249 .get(http::header::X_CONTENT_TYPE_OPTIONS)
250 .expect("x-content-type-options header"),
251 "nosniff"
252 );
253 }
254
255 #[tokio::test]
256 async fn builder_applies_ordered_transform_helpers_and_cors() {
257 let calls = Arc::new(AtomicUsize::new(0));
258 let transform = |calls: Arc<AtomicUsize>| {
259 move |router: Router| {
260 calls.fetch_add(1, Ordering::SeqCst);
261 router
262 }
263 };
264
265 let server = HttpServerBuilder::new(local_config(), CancellationToken::new())
266 .with_router(Router::new().route("/", get(|| async { "ok" })))
267 .with_logging_transform(transform(Arc::clone(&calls)))
268 .with_auth_transform(transform(Arc::clone(&calls)))
269 .with_validation_transform(transform(Arc::clone(&calls)))
270 .with_metrics_transform(transform(Arc::clone(&calls)))
271 .with_cors()
272 .expect("default cors state is valid")
273 .build()
274 .expect("build server");
275
276 let router = server
277 .router
278 .lock()
279 .await
280 .take()
281 .expect("router is present");
282 let response = router
283 .oneshot(
284 Request::builder()
285 .uri("/")
286 .body(Body::empty())
287 .expect("request"),
288 )
289 .await
290 .expect("route response");
291
292 assert_eq!(response.status(), StatusCode::OK);
293 assert_eq!(calls.load(Ordering::SeqCst), 4);
294 }
295
296 #[test]
297 fn builder_helpers_validate_cors_security_headers_and_middleware_stack() {
298 let invalid_cors = HttpServerConfig {
299 cors: Some(CorsPolicy {
300 allow_credentials: true,
301 ..Default::default()
302 }),
303 ..local_config()
304 };
305 assert_eq!(
306 HttpServerBuilder::new(invalid_cors, CancellationToken::new())
307 .with_cors()
308 .err()
309 .unwrap()
310 .code(),
311 ErrorCode::InvalidInput
312 );
313 let invalid_cors = HttpServerConfig {
314 cors: Some(CorsPolicy {
315 allowed_origins: vec!["*".to_string()],
316 ..Default::default()
317 }),
318 ..local_config()
319 };
320 assert_eq!(
321 HttpServerBuilder::new(invalid_cors, CancellationToken::new())
322 .build()
323 .err()
324 .unwrap()
325 .code(),
326 ErrorCode::InvalidInput
327 );
328
329 let builder = HttpServerBuilder::new(local_config(), CancellationToken::new())
330 .with_middleware_stack(HttpMiddlewareStack::new())
331 .with_security_headers()
332 .expect("default security headers are valid");
333 assert!(builder.security_headers.is_some());
334 }
335
336 #[test]
337 fn builder_builds_with_valid_cors_and_exposes_server_accessors() {
338 let config = HttpServerConfig {
339 cors: Some(CorsPolicy {
340 allowed_origins: vec!["https://example.com".to_string()],
341 allowed_methods: vec!["GET".to_string()],
342 allowed_headers: vec!["x-test".to_string()],
343 ..Default::default()
344 }),
345 ..local_config()
346 };
347 let server = HttpServerBuilder::new(config, CancellationToken::new())
348 .with_router(Router::new())
349 .build()
350 .expect("valid CORS policy builds");
351
352 assert_eq!(server.name(), "http-server");
353 assert_eq!(server.bind_addr(), "127.0.0.1:0");
354 assert_eq!(server.local_addr(), None);
355 }
356
357 #[tokio::test]
358 async fn builder_applies_ordered_transform_shortcuts() {
359 let server = HttpServerBuilder::new(local_config(), CancellationToken::new())
360 .with_logging_transform(|router| router.route("/logging", get(|| async { "logging" })))
361 .with_auth_transform(|router| router.route("/auth", get(|| async { "auth" })))
362 .with_validation_transform(|router| {
363 router.route("/validation", get(|| async { "validation" }))
364 })
365 .with_metrics_transform(|router| router.route("/metrics", get(|| async { "metrics" })))
366 .with_cors()
367 .expect("empty cors config is accepted")
368 .build()
369 .expect("build server");
370 let router = server
371 .router
372 .lock()
373 .await
374 .take()
375 .expect("router is present");
376
377 for (path, expected) in [
378 ("/logging", "logging"),
379 ("/auth", "auth"),
380 ("/validation", "validation"),
381 ("/metrics", "metrics"),
382 ] {
383 let response = router
384 .clone()
385 .oneshot(
386 Request::builder()
387 .uri(path)
388 .body(Body::empty())
389 .expect("request"),
390 )
391 .await
392 .expect("route response");
393 let body = axum::body::to_bytes(response.into_body(), usize::MAX)
394 .await
395 .expect("body bytes");
396 assert_eq!(&body[..], expected.as_bytes());
397 }
398 }
399
400 #[test]
401 fn security_header_configuration_is_validated_at_build_time() {
402 let config = SecurityHeadersConfig::default()
403 .with_transport_security(TransportSecurity::AllowInsecureLocal)
404 .with_content_security_policy(None)
405 .with_permissions_policy(None)
406 .with_referrer_policy(None)
407 .with_frame_options(None)
408 .with_content_type_options(None);
409
410 let error = match HttpServerBuilder::new(local_config(), CancellationToken::new())
411 .with_security_headers_config(config)
412 {
413 Ok(_) => panic!("invalid security header policy should be rejected"),
414 Err(error) => error,
415 };
416
417 assert_eq!(error.code(), ErrorCode::InvalidInput);
418 }
419}