Skip to main content

rskit_server/http/
builder.rs

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
20/// Builder for [`HttpServer`].
21pub struct HttpServerBuilder {
22    config: HttpServerConfig,
23    cancel: CancellationToken,
24    router: Router,
25    middleware: HttpMiddlewareStack,
26    security_headers: Option<SecurityHeadersConfig>,
27}
28
29impl HttpServerBuilder {
30    /// Create a new builder.
31    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    /// Merge an axum [`Router`] into the server.
42    #[must_use]
43    pub fn with_router(mut self, router: Router) -> Self {
44        self.router = self.router.merge(router);
45        self
46    }
47
48    /// Replace the ordered middleware stack.
49    #[must_use]
50    pub fn with_middleware_stack(mut self, middleware: HttpMiddlewareStack) -> Self {
51        self.middleware = middleware;
52        self
53    }
54
55    /// Append a logging-phase transform.
56    #[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    /// Append an auth-phase transform.
66    #[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    /// Append a validation-phase transform.
76    #[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    /// Append a metrics-phase transform.
86    #[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    /// Apply CORS from the server config (no-op if `cors` is `None`).
96    ///
97    /// # Errors
98    /// Returns an error when the configured CORS policy contains invalid origins, methods, headers,
99    /// or max-age values.
100    #[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    /// Add secure response headers using the default security policy.
109    ///
110    /// # Errors
111    /// Returns an error if the default security policy cannot be built (should never happen in practice — this is a programming error guard).
112    #[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    /// Add secure response headers using an explicit security policy.
118    ///
119    /// # Errors
120    /// Returns an error when the supplied policy is invalid.
121    #[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    /// Consume the builder and produce an [`HttpServer`].
132    /// # Errors
133    /// Returns an error when baseline transport middleware configuration is invalid.
134    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}