rskit_server/
middleware.rs1use std::sync::Arc;
2
3use axum::Router;
4
5pub const HTTP_INTERCEPTOR_ORDER: [&str; 5] =
10 ["tracing", "logging", "auth", "validation", "metrics"];
11
12pub const HTTP_BASELINE_LAYER_ORDER: [&str; 5] = [
16 "request_id",
17 "cors",
18 "security_headers",
19 "body_limit",
20 "timeout",
21];
22
23pub type RouterTransform = Arc<dyn Fn(Router) -> Router + Send + Sync + 'static>;
25
26#[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 #[must_use]
39 pub fn new() -> Self {
40 Self::default()
41 }
42
43 #[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 #[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 #[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 #[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 #[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 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}