Skip to main content

rskit_server/
middleware.rs

1use std::sync::Arc;
2
3use axum::Router;
4
5/// Canonical interceptor ordering shared across rskit transports.
6///
7/// Request processing flows through these phases in order. Metrics wrap the handler completion
8/// so response observations happen after the handler returns.
9pub const HTTP_INTERCEPTOR_ORDER: [&str; 5] =
10    ["tracing", "logging", "auth", "validation", "metrics"];
11
12/// Baseline HTTP transport layers applied by [`HttpServerBuilder`](crate::HttpServerBuilder).
13///
14/// These layers wrap every HTTP server before application middleware executes.
15pub const HTTP_BASELINE_LAYER_ORDER: [&str; 5] = [
16    "request_id",
17    "cors",
18    "security_headers",
19    "body_limit",
20    "timeout",
21];
22
23/// Boxed router transform used to inject transport middleware without exposing axum's concrete layer stack in the public API.
24pub type RouterTransform = Arc<dyn Fn(Router) -> Router + Send + Sync + 'static>;
25
26/// Ordered HTTP middleware phases for service-facing servers.
27#[derive(Clone, Default)]
28pub struct HttpMiddlewareStack {
29    tracing: Vec<RouterTransform>,
30    logging: Vec<RouterTransform>,
31    auth: Vec<RouterTransform>,
32    validation: Vec<RouterTransform>,
33    metrics: Vec<RouterTransform>,
34}
35
36impl HttpMiddlewareStack {
37    /// Create an empty middleware stack.
38    #[must_use]
39    pub fn new() -> Self {
40        Self::default()
41    }
42
43    /// Append a tracing-phase transform.
44    #[must_use]
45    pub fn with_tracing_transform<F>(mut self, transform: F) -> Self
46    where
47        F: Fn(Router) -> Router + Send + Sync + 'static,
48    {
49        self.tracing.push(Arc::new(transform));
50        self
51    }
52
53    /// Append a logging-phase transform.
54    #[must_use]
55    pub fn with_logging_transform<F>(mut self, transform: F) -> Self
56    where
57        F: Fn(Router) -> Router + Send + Sync + 'static,
58    {
59        self.logging.push(Arc::new(transform));
60        self
61    }
62
63    /// Append an auth-phase transform.
64    #[must_use]
65    pub fn with_auth_transform<F>(mut self, transform: F) -> Self
66    where
67        F: Fn(Router) -> Router + Send + Sync + 'static,
68    {
69        self.auth.push(Arc::new(transform));
70        self
71    }
72
73    /// Append a validation-phase transform.
74    #[must_use]
75    pub fn with_validation_transform<F>(mut self, transform: F) -> Self
76    where
77        F: Fn(Router) -> Router + Send + Sync + 'static,
78    {
79        self.validation.push(Arc::new(transform));
80        self
81    }
82
83    /// Append a metrics-phase transform.
84    #[must_use]
85    pub fn with_metrics_transform<F>(mut self, transform: F) -> Self
86    where
87        F: Fn(Router) -> Router + Send + Sync + 'static,
88    {
89        self.metrics.push(Arc::new(transform));
90        self
91    }
92
93    /// Apply the configured phases around a router.
94    pub fn apply(&self, router: Router) -> Router {
95        let router = apply_phase(router, &self.metrics);
96        let router = apply_phase(router, &self.validation);
97        let router = apply_phase(router, &self.auth);
98        let router = apply_phase(router, &self.logging);
99        apply_phase(router, &self.tracing)
100    }
101}
102
103fn apply_phase(router: Router, transforms: &[RouterTransform]) -> Router {
104    transforms
105        .iter()
106        .fold(router, |router, transform| transform(router))
107}
108
109#[cfg(test)]
110mod tests {
111    use std::convert::Infallible;
112    use std::sync::Arc;
113
114    use axum::{Router, body::Body, http::Request, response::Response, routing::get};
115    use parking_lot::Mutex;
116    use tower::{Layer, Service, ServiceExt};
117
118    use super::HttpMiddlewareStack;
119
120    #[derive(Clone)]
121    struct RecordLayer {
122        name: &'static str,
123        events: Arc<Mutex<Vec<&'static str>>>,
124    }
125
126    impl RecordLayer {
127        fn new(name: &'static str, events: Arc<Mutex<Vec<&'static str>>>) -> Self {
128            Self { name, events }
129        }
130    }
131
132    impl<S> Layer<S> for RecordLayer {
133        type Service = RecordService<S>;
134
135        fn layer(&self, inner: S) -> Self::Service {
136            RecordService {
137                inner,
138                name: self.name,
139                events: Arc::clone(&self.events),
140            }
141        }
142    }
143
144    #[derive(Clone)]
145    struct RecordService<S> {
146        inner: S,
147        name: &'static str,
148        events: Arc<Mutex<Vec<&'static str>>>,
149    }
150
151    impl<S> Service<Request<Body>> for RecordService<S>
152    where
153        S: Service<Request<Body>, Response = Response, Error = Infallible> + Clone + Send + 'static,
154        S::Future: Send + 'static,
155    {
156        type Response = Response;
157        type Error = Infallible;
158        type Future = futures::future::BoxFuture<'static, Result<Response, Infallible>>;
159
160        fn poll_ready(
161            &mut self,
162            cx: &mut std::task::Context<'_>,
163        ) -> std::task::Poll<Result<(), Self::Error>> {
164            self.inner.poll_ready(cx)
165        }
166
167        fn call(&mut self, req: Request<Body>) -> Self::Future {
168            let mut inner = self.inner.clone();
169            let events = Arc::clone(&self.events);
170            let name = self.name;
171            Box::pin(async move {
172                events.lock().push(name);
173                inner.call(req).await
174            })
175        }
176    }
177
178    #[tokio::test]
179    async fn stack_applies_request_phases_in_locked_order() {
180        let events = Arc::new(Mutex::new(Vec::new()));
181        let app = Router::new().route("/", get(|| async { "ok" }));
182        let app = HttpMiddlewareStack::new()
183            .with_tracing_transform({
184                let events = Arc::clone(&events);
185                move |router| router.layer(RecordLayer::new("tracing", Arc::clone(&events)))
186            })
187            .with_logging_transform({
188                let events = Arc::clone(&events);
189                move |router| router.layer(RecordLayer::new("logging", Arc::clone(&events)))
190            })
191            .with_auth_transform({
192                let events = Arc::clone(&events);
193                move |router| router.layer(RecordLayer::new("auth", Arc::clone(&events)))
194            })
195            .with_validation_transform({
196                let events = Arc::clone(&events);
197                move |router| router.layer(RecordLayer::new("validation", Arc::clone(&events)))
198            })
199            .with_metrics_transform({
200                let events = Arc::clone(&events);
201                move |router| router.layer(RecordLayer::new("metrics", Arc::clone(&events)))
202            })
203            .apply(app);
204
205        let response = app
206            .oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
207            .await
208            .unwrap();
209
210        assert_eq!(response.status(), axum::http::StatusCode::OK);
211        assert_eq!(
212            *events.lock(),
213            vec!["tracing", "logging", "auth", "validation", "metrics"]
214        );
215    }
216}