1use crate::HeaderName;
55use crate::{
56 Request,
57 utils::{HeaderValueErr, HeaderValueGetter},
58};
59use rama_core::error::ErrorContext as _;
60use rama_core::extensions::{Extension, ExtensionsRef};
61use rama_core::{Layer, Service, error::BoxError};
62use rama_utils::macros::define_inner_service_accessors;
63use std::iter::FromIterator;
64use std::str::FromStr;
65use std::{fmt, marker::PhantomData};
66
67pub struct HeaderFromStrConfigService<T, S, C = Vec<T>> {
72 inner: S,
73 header_name: HeaderName,
74 optional: bool,
75 repeat: bool,
76 _marker: PhantomData<fn() -> (T, C)>,
77}
78
79impl<T, S, C> HeaderFromStrConfigService<T, S, C> {
80 define_inner_service_accessors!();
81
82 pub const fn required(inner: S, header_name: HeaderName) -> Self {
86 Self {
87 inner,
88 header_name,
89 optional: false,
90 repeat: false,
91 _marker: PhantomData,
92 }
93 }
94
95 pub const fn optional(inner: S, header_name: HeaderName) -> Self {
99 Self {
100 inner,
101 header_name,
102 optional: true,
103 repeat: false,
104 _marker: PhantomData,
105 }
106 }
107
108 rama_utils::macros::generate_set_and_with! {
109 pub fn repeat(mut self, repeat: bool) -> Self {
112 self.repeat = repeat;
113 self
114 }
115 }
116}
117
118impl<T, S, C> fmt::Debug for HeaderFromStrConfigService<T, S, C>
119where
120 S: fmt::Debug,
121{
122 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
123 f.debug_struct("HeaderFromStrConfigService")
124 .field("inner", &self.inner)
125 .field("header_name", &self.header_name)
126 .field("optional", &self.optional)
127 .field("repeat", &self.repeat)
128 .field(
129 "_marker",
130 &format_args!("{}", std::any::type_name::<fn() -> (T, C)>()),
131 )
132 .finish()
133 }
134}
135
136impl<T, S, C> Clone for HeaderFromStrConfigService<T, S, C>
137where
138 S: Clone,
139{
140 fn clone(&self) -> Self {
141 Self {
142 inner: self.inner.clone(),
143 header_name: self.header_name.clone(),
144 optional: self.optional,
145 repeat: self.repeat,
146 _marker: PhantomData,
147 }
148 }
149}
150
151impl<T, S, Body, E, C> Service<Request<Body>> for HeaderFromStrConfigService<T, S, C>
152where
153 S: Service<Request<Body>, Error = E>,
154 T: FromStr<Err: Into<BoxError> + Send + Sync + 'static>
155 + Send
156 + Sync
157 + Clone
158 + std::fmt::Debug
159 + Extension
160 + 'static,
161 C: FromIterator<T> + Send + Sync + Clone + std::fmt::Debug + Extension + 'static,
162 Body: Send + Sync + 'static,
163 E: Into<BoxError> + Send + Sync + 'static,
164{
165 type Output = S::Output;
166 type Error = BoxError;
167
168 async fn serve(&self, request: Request<Body>) -> Result<Self::Output, Self::Error> {
169 if self.repeat {
170 let headers = request.headers().get_all(&self.header_name);
171 let mut parsed_values = headers
172 .into_iter()
173 .flat_map(|value| {
174 value.to_str().into_iter().flat_map(|string| {
175 string
176 .split(',')
177 .filter_map(|x| match x.trim() {
178 "" => None,
179 y => Some(y),
180 })
181 .map(|x| x.parse::<T>().into_box_error())
182 })
183 })
184 .peekable();
185
186 if parsed_values.peek().is_none() {
187 if !self.optional {
188 return Err(HeaderValueErr::HeaderMissing(self.header_name.to_string()).into());
189 }
190 } else {
191 let values = parsed_values.collect::<Result<C, _>>()?;
192 request.extensions().insert(values);
193 }
194 } else {
195 match request.header_str(&self.header_name) {
196 Ok(s) => {
197 let cfg: T = s.parse().into_box_error()?;
198 request.extensions().insert(cfg);
199 }
200 Err(HeaderValueErr::HeaderMissing(_)) if self.optional => (),
201 Err(err) => {
202 return Err(err.into());
203 }
204 }
205 }
206
207 self.inner.serve(request).await.into_box_error()
208 }
209}
210
211pub struct HeaderFromStrConfigLayer<T, C = Vec<T>> {
216 header_name: HeaderName,
217 optional: bool,
218 repeat: bool,
219 _marker: PhantomData<fn() -> (T, C)>,
220}
221
222impl<T, C: fmt::Debug> fmt::Debug for HeaderFromStrConfigLayer<T, C> {
223 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
224 f.debug_struct("HeaderFromStrConfigLayer")
225 .field("header_name", &self.header_name)
226 .field("optional", &self.optional)
227 .field("repeat", &self.repeat)
228 .field(
229 "_marker",
230 &format_args!("{}", std::any::type_name::<fn() -> (T, C)>()),
231 )
232 .finish()
233 }
234}
235
236impl<T, C> Clone for HeaderFromStrConfigLayer<T, C> {
237 fn clone(&self) -> Self {
238 Self {
239 header_name: self.header_name.clone(),
240 optional: self.optional,
241 repeat: self.repeat,
242 _marker: PhantomData,
243 }
244 }
245}
246
247impl<T, C> HeaderFromStrConfigLayer<T, C> {
248 pub fn required(header_name: HeaderName) -> Self {
252 Self {
253 header_name,
254 optional: false,
255 repeat: false,
256 _marker: PhantomData,
257 }
258 }
259
260 pub fn optional(header_name: HeaderName) -> Self {
264 Self {
265 header_name,
266 optional: true,
267 repeat: false,
268 _marker: PhantomData,
269 }
270 }
271
272 rama_utils::macros::generate_set_and_with! {
273 pub fn repeat(mut self, repeat: bool) -> Self {
276 self.repeat = repeat;
277 self
278 }
279 }
280}
281
282impl<T, S, C> Layer<S> for HeaderFromStrConfigLayer<T, C> {
283 type Service = HeaderFromStrConfigService<T, S, C>;
284
285 fn layer(&self, inner: S) -> Self::Service {
286 HeaderFromStrConfigService {
287 inner,
288 header_name: self.header_name.clone(),
289 optional: self.optional,
290 repeat: self.repeat,
291 _marker: PhantomData,
292 }
293 }
294
295 fn into_layer(self, inner: S) -> Self::Service {
296 HeaderFromStrConfigService {
297 inner,
298 header_name: self.header_name,
299 optional: self.optional,
300 repeat: self.repeat,
301 _marker: PhantomData,
302 }
303 }
304}
305
306#[cfg(test)]
307mod test {
308 use rama_core::extensions::{Extension, ExtensionsRef};
309
310 use super::*;
311 use crate::Method;
312 use ahash::HashSet;
313 use std::{collections::VecDeque, convert::Infallible, num::ParseIntError, str::FromStr};
314
315 #[derive(Debug, Clone, Copy, PartialEq, Eq, Extension)]
316 struct ProxyId(usize);
317
318 impl FromStr for ProxyId {
319 type Err = ParseIntError;
320
321 fn from_str(s: &str) -> Result<Self, Self::Err> {
322 s.parse::<usize>().map(Self)
323 }
324 }
325
326 #[derive(Debug, Clone, PartialEq, Eq, Hash, Extension)]
327 struct ProxyLabel(String);
328
329 impl From<&str> for ProxyLabel {
330 fn from(value: &str) -> Self {
331 Self(value.to_owned())
332 }
333 }
334
335 impl FromStr for ProxyLabel {
336 type Err = Infallible;
337
338 fn from_str(s: &str) -> Result<Self, Self::Err> {
339 Ok(Self(s.to_owned()))
340 }
341 }
342
343 #[derive(Debug, Clone, Default, Extension)]
344 struct ProxyLabelList(Vec<ProxyLabel>);
345
346 impl FromIterator<ProxyLabel> for ProxyLabelList {
347 fn from_iter<T: IntoIterator<Item = ProxyLabel>>(iter: T) -> Self {
348 Self(iter.into_iter().collect())
349 }
350 }
351
352 impl ProxyLabelList {
353 fn join_with_plus(&self) -> String {
354 self.0
355 .iter()
356 .map(|label| label.0.as_str())
357 .collect::<Vec<_>>()
358 .join("+")
359 }
360 }
361
362 #[derive(Debug, Clone, Default, Extension)]
363 struct ProxyLabelSet(HashSet<ProxyLabel>);
364
365 impl FromIterator<ProxyLabel> for ProxyLabelSet {
366 fn from_iter<T: IntoIterator<Item = ProxyLabel>>(iter: T) -> Self {
367 Self(iter.into_iter().collect())
368 }
369 }
370
371 impl ProxyLabelSet {
372 fn contains_value(&self, value: &str) -> bool {
373 self.0.contains(&ProxyLabel::from(value))
374 }
375 }
376
377 #[derive(Debug, Clone, Default, Extension)]
378 struct ProxyLabelQueue(VecDeque<ProxyLabel>);
379
380 impl FromIterator<ProxyLabel> for ProxyLabelQueue {
381 fn from_iter<T: IntoIterator<Item = ProxyLabel>>(iter: T) -> Self {
382 Self(iter.into_iter().collect())
383 }
384 }
385
386 #[derive(Debug, Clone, Default, Extension)]
387 struct ProxyIdList {
388 _ids: Vec<ProxyId>,
389 }
390
391 impl FromIterator<ProxyId> for ProxyIdList {
392 fn from_iter<T: IntoIterator<Item = ProxyId>>(iter: T) -> Self {
393 Self {
394 _ids: iter.into_iter().collect(),
395 }
396 }
397 }
398
399 #[tokio::test]
400 async fn test_header_config_required_happy_path() {
401 let request = Request::builder()
402 .method(Method::GET)
403 .uri("https://www.example.com")
404 .header("x-proxy-id", "42")
405 .body(())
406 .unwrap();
407
408 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
409 let id: &ProxyId = req.extensions().get_ref().unwrap();
410 assert_eq!(id.0, 42);
411
412 Ok::<_, Infallible>(())
413 });
414
415 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::required(
416 inner_service,
417 HeaderName::from_static("x-proxy-id"),
418 );
419
420 service.serve(request).await.unwrap();
421 }
422
423 #[tokio::test]
424 async fn test_header_config_required_repeat_happy_path() {
425 let request = Request::builder()
426 .method(Method::GET)
427 .uri("https://www.example.com")
428 .header("x-proxy-labels", "foo,bar ,baz, fin ")
429 .body(())
430 .unwrap();
431
432 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
433 let labels: &ProxyLabelList = req.extensions().get_ref().unwrap();
434 assert_eq!("foo+bar+baz+fin", labels.join_with_plus());
435
436 Ok::<_, Infallible>(())
437 });
438
439 let service = HeaderFromStrConfigService::<ProxyLabel, _, ProxyLabelList>::required(
440 inner_service,
441 HeaderName::from_static("x-proxy-labels"),
442 )
443 .with_repeat(true);
444
445 service.serve(request).await.unwrap();
446 }
447
448 #[tokio::test]
449 async fn test_header_config_required_repeat_custom_container() {
450 let request = Request::builder()
451 .method(Method::GET)
452 .uri("https://www.example.com")
453 .header("x-proxy-labels", "foo,bar,baz,foo")
454 .body(())
455 .unwrap();
456
457 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
458 let labels: &ProxyLabelSet = req.extensions().get_ref().unwrap();
459 assert_eq!(3, labels.0.len());
460 assert!(labels.contains_value("foo"));
461 assert!(labels.contains_value("bar"));
462 assert!(labels.contains_value("baz"));
463
464 Ok::<_, Infallible>(())
465 });
466
467 let service = HeaderFromStrConfigService::<ProxyLabel, _, ProxyLabelSet>::required(
468 inner_service,
469 HeaderName::from_static("x-proxy-labels"),
470 )
471 .with_repeat(true);
472
473 service.serve(request).await.unwrap();
474 }
475
476 #[tokio::test]
477 async fn test_header_config_required_repeat_linked_list() {
478 let request = Request::builder()
479 .method(Method::GET)
480 .uri("https://www.example.com")
481 .header("x-proxy-labels", "foo,bar,baz")
482 .body(())
483 .unwrap();
484
485 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
486 let labels: &ProxyLabelQueue = req.extensions().get_ref().unwrap();
487 let mut iter = labels.0.iter();
488 assert_eq!(Some("foo"), iter.next().map(|x| x.0.as_str()));
489 assert_eq!(Some("bar"), iter.next().map(|x| x.0.as_str()));
490 assert_eq!(Some("baz"), iter.next().map(|x| x.0.as_str()));
491 assert_eq!(None, iter.next());
492
493 Ok::<_, Infallible>(())
494 });
495
496 let service = HeaderFromStrConfigService::<ProxyLabel, _, ProxyLabelQueue>::required(
497 inner_service,
498 HeaderName::from_static("x-proxy-labels"),
499 )
500 .with_repeat(true);
501
502 service.serve(request).await.unwrap();
503 }
504
505 #[tokio::test]
506 async fn test_header_config_required_repeat_happy_path_multi_header() {
507 let request = Request::builder()
508 .method(Method::GET)
509 .uri("https://www.example.com")
510 .header("x-proxy-labels", "foo,bar ")
511 .header("x-Proxy-Labels", "baz ")
512 .header("X-PROXY-LABELS", " fin")
513 .body(())
514 .unwrap();
515
516 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
517 let labels: &ProxyLabelList = req.extensions().get_ref().unwrap();
518 assert_eq!("foo+bar+baz+fin", labels.join_with_plus());
519
520 Ok::<_, Infallible>(())
521 });
522
523 let service = HeaderFromStrConfigService::<ProxyLabel, _, ProxyLabelList>::required(
524 inner_service,
525 HeaderName::from_static("x-proxy-labels"),
526 )
527 .with_repeat(true);
528
529 service.serve(request).await.unwrap();
530 }
531
532 #[tokio::test]
533 async fn test_header_config_optional_found() {
534 let request = Request::builder()
535 .method(Method::GET)
536 .uri("https://www.example.com")
537 .header("x-proxy-id", "42")
538 .body(())
539 .unwrap();
540
541 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
542 let id: &ProxyId = req.extensions().get_ref().unwrap();
543 assert_eq!(id.0, 42);
544
545 Ok::<_, Infallible>(())
546 });
547
548 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::optional(
549 inner_service,
550 HeaderName::from_static("x-proxy-id"),
551 );
552
553 service.serve(request).await.unwrap();
554 }
555
556 #[tokio::test]
557 async fn test_header_config_repeat_optional_found() {
558 let request = Request::builder()
559 .method(Method::GET)
560 .uri("https://www.example.com")
561 .header("x-proxy-labels", "foo,bar ,baz, fin ")
562 .body(())
563 .unwrap();
564
565 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
566 let labels: &ProxyLabelList = req.extensions().get_ref().unwrap();
567 assert_eq!("foo+bar+baz+fin", labels.join_with_plus());
568
569 Ok::<_, Infallible>(())
570 });
571
572 let service = HeaderFromStrConfigService::<ProxyLabel, _, ProxyLabelList>::optional(
573 inner_service,
574 HeaderName::from_static("x-proxy-labels"),
575 )
576 .with_repeat(true);
577
578 service.serve(request).await.unwrap();
579 }
580
581 #[tokio::test]
582 async fn test_header_config_optional_missing() {
583 let request = Request::builder()
584 .method(Method::GET)
585 .uri("https://www.example.com")
586 .body(())
587 .unwrap();
588
589 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
590 assert!(req.extensions().get_ref::<ProxyId>().is_none());
591 Ok::<_, Infallible>(())
592 });
593
594 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::optional(
595 inner_service,
596 HeaderName::from_static("x-proxy-id"),
597 );
598
599 service.serve(request).await.unwrap();
600 }
601
602 #[tokio::test]
603 async fn test_header_config_repeat_optional_missing() {
604 let request = Request::builder()
605 .method(Method::GET)
606 .uri("https://www.example.com")
607 .body(())
608 .unwrap();
609
610 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
611 assert!(req.extensions().get_ref::<ProxyLabelList>().is_none());
612
613 Ok::<_, Infallible>(())
614 });
615
616 let service = HeaderFromStrConfigService::<ProxyLabel, _, ProxyLabelList>::optional(
617 inner_service,
618 HeaderName::from_static("x-proxy-labels"),
619 )
620 .with_repeat(true);
621
622 service.serve(request).await.unwrap();
623 }
624
625 #[tokio::test]
626 async fn test_header_config_required_missing_header() {
627 let request = Request::builder()
628 .method(Method::GET)
629 .uri("https://www.example.com")
630 .body(())
631 .unwrap();
632
633 let inner_service =
634 rama_core::service::service_fn(async |_req: Request<()>| Ok::<_, Infallible>(()));
635
636 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::required(
637 inner_service,
638 HeaderName::from_static("x-proxy-id"),
639 );
640
641 let result = service.serve(request).await;
642 assert!(result.is_err());
643 }
644
645 #[tokio::test]
646 async fn test_header_config_repeat_required_missing() {
647 let request = Request::builder()
648 .method(Method::GET)
649 .uri("https://www.example.com")
650 .body(())
651 .unwrap();
652
653 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
654 assert!(req.extensions().get_ref::<ProxyLabelList>().is_none());
655
656 Ok::<_, Infallible>(())
657 });
658
659 let service = HeaderFromStrConfigService::<ProxyLabel, _, ProxyLabelList>::required(
660 inner_service,
661 HeaderName::from_static("x-proxy-labels"),
662 )
663 .with_repeat(true);
664
665 let result = service.serve(request).await;
666 assert!(result.is_err());
667 }
668
669 #[tokio::test]
670 async fn test_header_config_required_invalid_config() {
671 let request = Request::builder()
672 .method(Method::GET)
673 .uri("https://www.example.com")
674 .header("x-proxy-id", "foo")
675 .body(())
676 .unwrap();
677
678 let inner_service =
679 rama_core::service::service_fn(async |_req: Request<()>| Ok::<_, Infallible>(()));
680
681 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::required(
682 inner_service,
683 HeaderName::from_static("x-proxy-id"),
684 );
685
686 let result = service.serve(request).await;
687 assert!(result.is_err());
688 }
689
690 #[tokio::test]
691 async fn test_header_config_repeat_required_invalid_config() {
692 let request = Request::builder()
693 .method(Method::GET)
694 .uri("https://www.example.com")
695 .header("x-proxy-labels", "42,foo")
696 .body(())
697 .unwrap();
698
699 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
700 assert!(req.extensions().get_ref::<ProxyIdList>().is_none());
701
702 Ok::<_, Infallible>(())
703 });
704
705 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::required(
706 inner_service,
707 HeaderName::from_static("x-proxy-labels"),
708 )
709 .with_repeat(true);
710
711 let result = service.serve(request).await;
712 assert!(result.is_err());
713 }
714
715 #[tokio::test]
716 async fn test_header_config_optional_invalid_config() {
717 let request = Request::builder()
718 .method(Method::GET)
719 .uri("https://www.example.com")
720 .header("x-proxy-id", "foo")
721 .body(())
722 .unwrap();
723
724 let inner_service =
725 rama_core::service::service_fn(async |_req: Request<()>| Ok::<_, Infallible>(()));
726
727 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::optional(
728 inner_service,
729 HeaderName::from_static("x-proxy-id"),
730 );
731
732 let result = service.serve(request).await;
733 assert!(result.is_err());
734 }
735
736 #[tokio::test]
737 async fn test_header_config_repeat_optional_invalid_config() {
738 let request = Request::builder()
739 .method(Method::GET)
740 .uri("https://www.example.com")
741 .header("x-proxy-labels", "42,foo")
742 .body(())
743 .unwrap();
744
745 let inner_service = rama_core::service::service_fn(async |req: Request<()>| {
746 assert!(req.extensions().get_ref::<ProxyIdList>().is_none());
747
748 Ok::<_, Infallible>(())
749 });
750
751 let service = HeaderFromStrConfigService::<ProxyId, _, ProxyIdList>::optional(
752 inner_service,
753 HeaderName::from_static("x-proxy-labels"),
754 )
755 .with_repeat(true);
756
757 let result = service.serve(request).await;
758 assert!(result.is_err());
759 }
760}