Skip to main content

rama_http/layer/
header_config.rs

1//! Extract a header config from a request or response and insert it into [`Extensions`].
2//!
3//! [`Extensions`]: rama_core::extensions::Extensions
4//!
5//! # Example
6//!
7//! ```rust
8//! use rama_http::layer::header_config::{HeaderConfigLayer, HeaderConfigService};
9//! use rama_http::service::web::{WebService};
10//! use rama_http::{Body, Request, StatusCode, HeaderName};
11//! use rama_core::{extensions::Extensions, Service, Layer};
12//! use serde::Deserialize;
13//!
14//! use rama_core::extensions::Extension;
15//!
16//! #[derive(Debug, Deserialize, Clone, Extension)]
17//! struct Config {
18//!     s: String,
19//!     n: i32,
20//!     m: Option<i32>,
21//!     b: bool,
22//! }
23//!
24//! #[tokio::main]
25//! async fn main() {
26//!     let service = HeaderConfigLayer::<Config>::required(HeaderName::from_static("x-proxy-config"))
27//!         .into_layer(WebService::default()
28//!             .with_get("/", async |ext: Extensions| {
29//!                 let cfg = ext.get_ref::<Config>().unwrap();
30//!                 assert_eq!(cfg.s, "E&G");
31//!                 assert_eq!(cfg.n, 1);
32//!                 assert!(cfg.m.is_none());
33//!                 assert!(cfg.b);
34//!             }),
35//!         );
36//!
37//!     let request = Request::builder()
38//!         .header("x-proxy-config", "s=E%26G&n=1&b=true")
39//!         .body(Body::empty())
40//!         .unwrap();
41//!
42//!     let resp = service.serve(request).await.unwrap();
43//!     assert_eq!(resp.status(), StatusCode::OK);
44//! }
45//! ```
46
47use crate::{HeaderName, header::AsHeaderName};
48use crate::{
49    Request,
50    utils::{HeaderValueErr, HeaderValueGetter},
51};
52use rama_core::error::ErrorContext as _;
53use rama_core::extensions::{Extension, ExtensionsRef};
54use rama_core::telemetry::tracing;
55use rama_core::{Layer, Service, error::BoxError};
56use rama_utils::macros::define_inner_service_accessors;
57use serde::de::DeserializeOwned;
58use std::{fmt, marker::PhantomData};
59
60/// Extract a header config from a request or response without consuming it.
61pub fn extract_header_config<H, T, G>(request: &G, header_name: H) -> Result<T, HeaderValueErr>
62where
63    H: AsHeaderName + Copy,
64    T: DeserializeOwned + Send + Sync + 'static,
65    G: HeaderValueGetter,
66{
67    let value = request.header_str(header_name)?;
68    let config = serde_html_form::from_str::<T>(value)
69        .map_err(|_e| HeaderValueErr::HeaderInvalid(header_name.as_str().to_owned()))?;
70    Ok(config)
71}
72
73/// A [`Service`] which extracts a header config from a request or response
74/// and inserts it into the [`Extensions`] of that object.
75///
76/// [`Extensions`]: rama_core::extensions::Extensions
77pub struct HeaderConfigService<T, S> {
78    inner: S,
79    header_name: HeaderName,
80    optional: bool,
81    _marker: PhantomData<fn() -> T>,
82}
83
84impl<T, S> HeaderConfigService<T, S> {
85    /// Create a new [`HeaderConfigService`].
86    ///
87    /// Alias for [`HeaderConfigService::required`] if `!optional`
88    /// and [`HeaderConfigService::optional`] if `optional`.
89    pub const fn new(inner: S, header_name: HeaderName, optional: bool) -> Self {
90        Self {
91            inner,
92            header_name,
93            optional,
94            _marker: PhantomData,
95        }
96    }
97
98    define_inner_service_accessors!();
99
100    /// Create a new [`HeaderConfigService`] with the given inner service
101    /// and header name, on which to extract the config,
102    /// and which will fail if the header is missing.
103    pub const fn required(inner: S, header_name: HeaderName) -> Self {
104        Self::new(inner, header_name, false)
105    }
106
107    /// Create a new [`HeaderConfigService`] with the given inner service
108    /// and header name, on which to extract the config,
109    /// and which will gracefully accept if the header is missing.
110    pub const fn optional(inner: S, header_name: HeaderName) -> Self {
111        Self::new(inner, header_name, true)
112    }
113}
114
115impl<T, S: fmt::Debug> fmt::Debug for HeaderConfigService<T, S> {
116    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
117        f.debug_struct("HeaderConfigService")
118            .field("inner", &self.inner)
119            .field("header_name", &self.header_name)
120            .field("optional", &self.optional)
121            .field(
122                "_marker",
123                &format_args!("{}", std::any::type_name::<fn() -> T>()),
124            )
125            .finish()
126    }
127}
128
129impl<T, S> Clone for HeaderConfigService<T, S>
130where
131    S: Clone,
132{
133    fn clone(&self) -> Self {
134        Self {
135            inner: self.inner.clone(),
136            header_name: self.header_name.clone(),
137            optional: self.optional,
138            _marker: PhantomData,
139        }
140    }
141}
142
143impl<T, S, Body, E> Service<Request<Body>> for HeaderConfigService<T, S>
144where
145    S: Service<Request<Body>, Error = E>,
146    T: DeserializeOwned + Extension,
147    Body: Send + Sync + 'static,
148    E: Into<BoxError> + Send + Sync + 'static,
149{
150    type Output = S::Output;
151    type Error = BoxError;
152
153    async fn serve(&self, request: Request<Body>) -> Result<Self::Output, Self::Error> {
154        let config = match extract_header_config::<_, T, _>(&request, &self.header_name) {
155            Ok(config) => config,
156            Err(err) => {
157                if self.optional && matches!(err, crate::utils::HeaderValueErr::HeaderMissing(_)) {
158                    tracing::debug!("failed to extract header config: {err:?}");
159                    return self.inner.serve(request).await.into_box_error();
160                } else {
161                    return Err(err.into());
162                }
163            }
164        };
165        request.extensions().insert(config);
166        self.inner.serve(request).await.into_box_error()
167    }
168}
169
170/// Layer which extracts a header config for the given HeaderName
171/// from a request or response and inserts it into the [`Extensions`] of that object.
172///
173/// [`Extensions`]: rama_core::extensions::Extensions
174pub struct HeaderConfigLayer<T> {
175    header_name: HeaderName,
176    optional: bool,
177    _marker: PhantomData<fn() -> T>,
178}
179
180impl<T> fmt::Debug for HeaderConfigLayer<T> {
181    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
182        f.debug_struct("HeaderConfigLayer")
183            .field("header_name", &self.header_name)
184            .field("optional", &self.optional)
185            .field(
186                "_marker",
187                &format_args!("{}", std::any::type_name::<fn() -> T>()),
188            )
189            .finish()
190    }
191}
192
193impl<T> Clone for HeaderConfigLayer<T> {
194    fn clone(&self) -> Self {
195        Self {
196            header_name: self.header_name.clone(),
197            optional: self.optional,
198            _marker: PhantomData,
199        }
200    }
201}
202
203impl<T> HeaderConfigLayer<T> {
204    /// Create a new [`HeaderConfigLayer`] with the given header name,
205    /// on which to extract the config,
206    /// and which will fail if the header is missing.
207    pub fn required(header_name: HeaderName) -> Self {
208        Self {
209            header_name,
210            optional: false,
211            _marker: PhantomData,
212        }
213    }
214
215    /// Create a new [`HeaderConfigLayer`] with the given header name,
216    /// on which to extract the config,
217    /// and which will gracefully accept if the header is missing.
218    pub fn optional(header_name: HeaderName) -> Self {
219        Self {
220            header_name,
221            optional: true,
222            _marker: PhantomData,
223        }
224    }
225}
226
227impl<T, S> Layer<S> for HeaderConfigLayer<T> {
228    type Service = HeaderConfigService<T, S>;
229
230    fn layer(&self, inner: S) -> Self::Service {
231        HeaderConfigService::new(inner, self.header_name.clone(), self.optional)
232    }
233
234    fn into_layer(self, inner: S) -> Self::Service {
235        HeaderConfigService::new(inner, self.header_name, self.optional)
236    }
237}
238
239#[cfg(test)]
240mod test {
241    use rama_core::extensions::{Extension, ExtensionsRef};
242    use serde::Deserialize;
243
244    use crate::Method;
245
246    use super::*;
247
248    #[tokio::test]
249    async fn test_header_config_required_happy_path() {
250        let request = Request::builder()
251            .method(Method::GET)
252            .uri("https://www.example.com")
253            .header("x-proxy-config", "s=E%26G&n=1&b=true")
254            .body(())
255            .unwrap();
256
257        let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
258            let cfg: &Config = req.extensions().get_ref().unwrap();
259            assert_eq!(cfg.s, "E&G");
260            assert_eq!(cfg.n, 1);
261            assert!(cfg.m.is_none());
262            assert!(cfg.b);
263
264            Ok::<_, std::convert::Infallible>(())
265        });
266
267        let service = HeaderConfigService::<Config, _>::required(
268            inner_service,
269            HeaderName::from_static("x-proxy-config"),
270        );
271
272        service.serve(request).await.unwrap();
273    }
274
275    #[tokio::test]
276    async fn test_header_config_optional_found() {
277        let request = Request::builder()
278            .method(Method::GET)
279            .uri("https://www.example.com")
280            .header("x-proxy-config", "s=E%26G&n=1&b=true")
281            .body(())
282            .unwrap();
283
284        let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
285            let cfg: &Config = req.extensions().get_ref().unwrap();
286            assert_eq!(cfg.s, "E&G");
287            assert_eq!(cfg.n, 1);
288            assert!(cfg.m.is_none());
289            assert!(cfg.b);
290
291            Ok::<_, std::convert::Infallible>(())
292        });
293
294        let service = HeaderConfigService::<Config, _>::optional(
295            inner_service,
296            HeaderName::from_static("x-proxy-config"),
297        );
298
299        service.serve(request).await.unwrap();
300    }
301
302    #[tokio::test]
303    async fn test_header_config_optional_missing() {
304        let request = Request::builder()
305            .method(Method::GET)
306            .uri("https://www.example.com")
307            .body(())
308            .unwrap();
309
310        let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
311            assert!(req.extensions().get_ref::<Config>().is_none());
312
313            Ok::<_, std::convert::Infallible>(())
314        });
315
316        let service = HeaderConfigService::<Config, _>::optional(
317            inner_service,
318            HeaderName::from_static("x-proxy-config"),
319        );
320
321        service.serve(request).await.unwrap();
322    }
323
324    #[tokio::test]
325    async fn test_header_config_required_missing_header() {
326        let request = Request::builder()
327            .method(Method::GET)
328            .uri("https://www.example.com")
329            .body(())
330            .unwrap();
331
332        let inner_service = rama_core::service::service_fn(async |_: Request<()>| {
333            Ok::<_, std::convert::Infallible>(())
334        });
335
336        let service = HeaderConfigService::<Config, _>::required(
337            inner_service,
338            HeaderName::from_static("x-proxy-config"),
339        );
340
341        let result = service.serve(request).await;
342        assert!(result.is_err());
343    }
344
345    #[tokio::test]
346    async fn test_header_config_required_invalid_config() {
347        let request = Request::builder()
348            .method(Method::GET)
349            .uri("https://www.example.com")
350            .header("x-proxy-config", "s=bar&n=1&b=invalid")
351            .body(())
352            .unwrap();
353
354        let inner_service = rama_core::service::service_fn(async |_: Request<()>| {
355            Ok::<_, std::convert::Infallible>(())
356        });
357
358        let service = HeaderConfigService::<Config, _>::required(
359            inner_service,
360            HeaderName::from_static("x-proxy-config"),
361        );
362
363        let result = service.serve(request).await;
364        assert!(result.is_err());
365    }
366
367    #[tokio::test]
368    async fn test_header_config_optional_invalid_config() {
369        let request = Request::builder()
370            .method(Method::GET)
371            .uri("https://www.example.com")
372            .header("x-proxy-config", "s=bar&n=1&b=invalid")
373            .body(())
374            .unwrap();
375
376        let inner_service = rama_core::service::service_fn(async |_: Request<()>| {
377            Ok::<_, std::convert::Infallible>(())
378        });
379
380        let service = HeaderConfigService::<Config, _>::optional(
381            inner_service,
382            HeaderName::from_static("x-proxy-config"),
383        );
384
385        let result = service.serve(request).await;
386        assert!(result.is_err());
387    }
388
389    #[derive(Debug, Deserialize, Clone, Extension)]
390    struct Config {
391        s: String,
392        n: i32,
393        m: Option<i32>,
394        b: bool,
395    }
396}