1use 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
60pub 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
73pub 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 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 pub const fn required(inner: S, header_name: HeaderName) -> Self {
104 Self::new(inner, header_name, false)
105 }
106
107 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
170pub 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 pub fn required(header_name: HeaderName) -> Self {
208 Self {
209 header_name,
210 optional: false,
211 _marker: PhantomData,
212 }
213 }
214
215 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}