Skip to main content

roadster/service/http/middleware/
any.rs

1use crate::app::context::AppContext;
2use crate::error::RoadsterResult;
3use crate::service::http::middleware::Middleware;
4use axum::Router;
5use axum_core::extract::FromRef;
6
7type ApplyFn<S> = Box<dyn Fn(Router, &S) -> RoadsterResult<Router> + Send>;
8
9/// A [`Middleware`] that can be applied without creating a separate `struct`. Useful to easily
10/// apply a middleware that's based on a function, for example.
11///
12/// # Examples
13/// ```rust
14/// # use axum::response::Response;
15/// # use axum::middleware::Next;
16/// # use axum_core::extract::Request;
17/// # use tracing::info;
18/// # use roadster::service::http::middleware::any::AnyMiddleware;
19/// #
20/// pub(crate) async fn hello_world_middleware_fn(request: Request, next: Next) -> Response {
21///     info!("Running `hello-world` middleware");
22///
23///     next.run(request).await
24/// }
25///
26/// let middleware = AnyMiddleware::builder()
27///     .name("hello-world")
28///     .enabled(true)
29///     .apply(|router, _state| {
30///         Ok(router
31///             .layer(axum::middleware::from_fn(hello_world_middleware_fn)))
32///     })
33///     .build();
34/// ```
35#[derive(bon::Builder)]
36#[non_exhaustive]
37pub struct AnyMiddleware<S>
38where
39    S: Clone + Send + Sync + 'static,
40    AppContext: FromRef<S>,
41{
42    #[builder(into)]
43    name: String,
44    enabled: Option<bool>,
45    priority: Option<i32>,
46    // #[builder(setter(transform = |a: impl Fn(Router, &S) -> RoadsterResult<Router> + Send + 'static| to_box_fn(a) ))]
47    #[builder(setters(vis = "", name = apply_internal))]
48    apply: ApplyFn<S>,
49}
50
51impl<S, BS> AnyMiddlewareBuilder<S, BS>
52where
53    S: Clone + Send + Sync + 'static,
54    AppContext: FromRef<S>,
55    BS: any_middleware_builder::State,
56{
57    pub fn apply(
58        self,
59        apply_fn: impl Fn(Router, &S) -> RoadsterResult<Router> + Send + 'static,
60    ) -> AnyMiddlewareBuilder<S, any_middleware_builder::SetApply<BS>>
61    where
62        BS::Apply: any_middleware_builder::IsUnset,
63    {
64        self.apply_internal(Box::new(apply_fn))
65    }
66}
67
68impl<S> Middleware<S> for AnyMiddleware<S>
69where
70    S: Clone + Send + Sync + 'static,
71    AppContext: FromRef<S>,
72{
73    fn name(&self) -> String {
74        self.name.clone()
75    }
76
77    fn enabled(&self, state: &S) -> bool {
78        // If the field on `AnyMiddleware` is set, use that
79        if let Some(enabled) = self.enabled {
80            return enabled;
81        }
82
83        let context = AppContext::from_ref(state);
84        let custom_config = context
85            .config()
86            .service
87            .http
88            .custom
89            .middleware
90            .custom
91            .get(&self.name);
92
93        if let Some(custom_config) = custom_config {
94            custom_config.common.enabled(state)
95        } else {
96            context
97                .config()
98                .service
99                .http
100                .custom
101                .middleware
102                .default_enable
103        }
104    }
105
106    fn priority(&self, state: &S) -> i32 {
107        // If the field on `AnyMiddleware` is set, use that
108        if let Some(priority) = self.priority {
109            return priority;
110        }
111
112        AppContext::from_ref(state)
113            .config()
114            .service
115            .http
116            .custom
117            .middleware
118            .custom
119            .get(&self.name)
120            .map(|config| config.common.priority)
121            .unwrap_or_default()
122    }
123
124    fn install(&self, router: Router, state: &S) -> RoadsterResult<Router> {
125        (self.apply)(router, state)
126    }
127}
128
129#[cfg(test)]
130mod tests {
131    use crate::app::context::AppContext;
132    use crate::config::service::http::middleware::{CommonConfig, MiddlewareConfig};
133    use crate::config::{AppConfig, CustomConfig};
134    use crate::service::http::middleware::Middleware;
135    use crate::service::http::middleware::any::AnyMiddleware;
136    use crate::testing::snapshot::TestCase;
137    use rstest::{fixture, rstest};
138
139    const NAME: &str = "hello-world";
140
141    #[fixture]
142    fn case() -> TestCase {
143        Default::default()
144    }
145
146    #[test]
147    #[cfg_attr(coverage_nightly, coverage(off))]
148    fn name() {
149        let middleware = AnyMiddleware::builder()
150            .name(NAME)
151            .apply(|router, _state| Ok(router))
152            .build();
153
154        assert_eq!(middleware.name(), NAME);
155    }
156
157    #[rstest]
158    #[case(false, None, None, false)]
159    #[case(false, None, Some(false), false)]
160    #[case(false, Some(false), None, false)]
161    #[case(false, None, Some(true), true)]
162    #[case(false, Some(true), None, true)]
163    #[case(true, None, Some(false), false)]
164    #[case(true, Some(false), None, false)]
165    #[case(true, None, None, true)]
166    #[cfg_attr(coverage_nightly, coverage(off))]
167    fn enabled(
168        _case: TestCase,
169        #[case] default_enabled: bool,
170        #[case] enabled_config: Option<bool>,
171        #[case] enabled_field: Option<bool>,
172        #[case] expected: bool,
173    ) {
174        let mut config = AppConfig::test(None).unwrap();
175        config.service.http.custom.middleware.default_enable = default_enabled;
176        if let Some(enabled_config) = enabled_config {
177            let middleware_config: MiddlewareConfig<CustomConfig> = MiddlewareConfig {
178                common: CommonConfig {
179                    enable: Some(enabled_config),
180                    priority: 0,
181                },
182                custom: CustomConfig::default(),
183            };
184            config
185                .service
186                .http
187                .custom
188                .middleware
189                .custom
190                .insert(NAME.to_string(), middleware_config);
191        }
192        let context = AppContext::test(Some(config), None, None).unwrap();
193
194        let middleware = AnyMiddleware::builder()
195            .name(NAME)
196            .maybe_enabled(enabled_field)
197            .apply(|router, _state| Ok(router))
198            .build();
199
200        assert_eq!(middleware.enabled(&context), expected);
201    }
202
203    #[rstest]
204    #[case(None, None, 0)]
205    #[case(None, Some(10), 10)]
206    #[case(Some(20), None, 20)]
207    #[case(Some(20), Some(10), 10)]
208    #[cfg_attr(coverage_nightly, coverage(off))]
209    fn priority(
210        _case: TestCase,
211        #[case] config_priority: Option<i32>,
212        #[case] field_priority: Option<i32>,
213        #[case] expected: i32,
214    ) {
215        let mut config = AppConfig::test(None).unwrap();
216        if let Some(config_priority) = config_priority {
217            let middleware_config: MiddlewareConfig<CustomConfig> = MiddlewareConfig {
218                common: CommonConfig {
219                    enable: None,
220                    priority: config_priority,
221                },
222                custom: CustomConfig::default(),
223            };
224            config
225                .service
226                .http
227                .custom
228                .middleware
229                .custom
230                .insert(NAME.to_string(), middleware_config);
231        }
232        let context = AppContext::test(Some(config), None, None).unwrap();
233
234        let middleware = AnyMiddleware::builder()
235            .name(NAME)
236            .maybe_priority(field_priority)
237            .apply(|router, _state| Ok(router))
238            .build();
239
240        assert_eq!(middleware.priority(&context), expected);
241    }
242}