Skip to main content

rama_http/layer/forwarded/
get_forwarded_multi.rs

1#![expect(
2    clippy::allow_attributes,
3    reason = "macro-generated `#[allow]` attributes whose underlying lints fire only for some expansions"
4)]
5
6use crate::Request;
7use crate::headers::forwarded::ForwardHeader;
8use rama_core::{Layer, Service, extensions::ExtensionsRef};
9use rama_http_headers::HeaderMapExt;
10use rama_net::forwarded::Forwarded;
11use rama_net::forwarded::ForwardedElement;
12use rama_utils::macros::all_the_tuples_no_last_special_case;
13use std::fmt;
14use std::marker::PhantomData;
15
16/// Layer to extract [`Forwarded`] information from the specified `T` headers.
17///
18/// Use [`GetForwardedHeaderLayer`] if you only need a single a header.
19///
20/// [`GetForwardedHeaderLayer`]: super::GetForwardedHeaderLayer
21///
22/// This layer can be used to extract the [`Forwarded`] information from any specified header `T`,
23/// as long as the header implements the [`ForwardHeader`] trait. Multiple headers can be specified
24/// as a tuple, and the layer will extract information from them all, and combine the information.
25///
26/// Please take into consideration the following when combining headers:
27///
28/// - The last header in the tuple will take precedence over the previous headers,
29///   if the same information is present in multiple headers.
30/// - Headers that can contain multiple elements, (e.g. X-Forwarded-For, Via)
31///   will combine their elements in the order as specified. That does however mean that in
32///   case one header has less elements then the other, that the combination down the line
33///   will not be accurate.
34///
35/// Rama also has the following headers already implemented for you to use:
36///
37/// > [`X-Real-Ip`], [`X-Client-Ip`], [`Client-Ip`], [`Cf-Connecting-Ip`] and [`True-Client-Ip`].
38///
39/// There are no [`GetForwardedHeadersLayer`] constructors for these headers,
40/// but you can use the [`GetForwardedHeadersLayer::new`] constructor and pass the header type as a type parameter in a tuple with other headers.
41///
42/// [`X-Real-Ip`]: crate::headers::forwarded::XRealIp
43/// [`X-Client-Ip`]: crate::headers::forwarded::XClientIp
44/// [`Client-Ip`]: crate::headers::forwarded::ClientIp
45/// [`CF-Connecting-Ip`]: crate::headers::forwarded::CFConnectingIp
46/// [`True-Client-Ip`]: crate::headers::forwarded::TrueClientIp
47pub struct GetForwardedHeadersLayer<T = Forwarded> {
48    _headers: PhantomData<fn() -> T>,
49}
50
51impl<T: fmt::Debug> fmt::Debug for GetForwardedHeadersLayer<T> {
52    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
53        f.debug_struct("GetForwardedHeadersLayer")
54            .field(
55                "_headers",
56                &format_args!("{}", std::any::type_name::<fn() -> T>()),
57            )
58            .finish()
59    }
60}
61
62impl<T: Clone> Clone for GetForwardedHeadersLayer<T> {
63    fn clone(&self) -> Self {
64        Self {
65            _headers: PhantomData,
66        }
67    }
68}
69
70impl<T> Default for GetForwardedHeadersLayer<T> {
71    #[inline]
72    fn default() -> Self {
73        Self::new()
74    }
75}
76
77impl<T> GetForwardedHeadersLayer<T> {
78    /// Create a new `GetForwardedHeadersLayer` for the specified headers `T`.
79    #[must_use]
80    pub const fn new() -> Self {
81        Self {
82            _headers: PhantomData,
83        }
84    }
85}
86
87impl<H, S> Layer<S> for GetForwardedHeadersLayer<H> {
88    type Service = GetForwardedHeadersService<S, H>;
89
90    fn layer(&self, inner: S) -> Self::Service {
91        Self::Service {
92            inner,
93            _headers: PhantomData,
94        }
95    }
96}
97
98/// Middleware service to extract [`Forwarded`] information from the specified `T` headers.
99///
100/// See [`GetForwardedHeadersLayer`] for more information.
101pub struct GetForwardedHeadersService<S, T = Forwarded> {
102    inner: S,
103    _headers: PhantomData<fn() -> T>,
104}
105
106impl<S: fmt::Debug, T> fmt::Debug for GetForwardedHeadersService<S, T> {
107    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
108        f.debug_struct("GetForwardedHeadersService")
109            .field("inner", &self.inner)
110            .field("_headers", &format_args!("{}", std::any::type_name::<T>()))
111            .finish()
112    }
113}
114
115impl<S: Clone, T> Clone for GetForwardedHeadersService<S, T> {
116    fn clone(&self) -> Self {
117        Self {
118            inner: self.inner.clone(),
119            _headers: PhantomData,
120        }
121    }
122}
123
124impl<S, T> GetForwardedHeadersService<S, T> {
125    /// Create a new `GetForwardedHeadersService` for the specified headers `T`.
126    pub const fn new(inner: S) -> Self {
127        Self {
128            inner,
129            _headers: PhantomData,
130        }
131    }
132}
133
134macro_rules! get_forwarded_service_for_tuple {
135    ( $($ty:ident),* $(,)? ) => {
136        #[allow(non_snake_case)]
137        impl<$($ty,)* S, Body> Service<Request<Body>> for GetForwardedHeadersService<S, ($($ty,)*)>
138        where
139            $( $ty: ForwardHeader + Send + Sync + 'static, )*
140            S: Service<Request<Body>>,
141            Body: Send + 'static,
142        {
143            type Output = S::Output;
144            type Error = S::Error;
145
146            fn serve(
147                &self,
148                req: Request<Body>,
149            ) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send + '_ {
150                let mut forwarded_elements: Vec<ForwardedElement> = Vec::with_capacity(1);
151
152                $(
153                    if let Some($ty) = req.headers().typed_get::<$ty>() {
154                        let mut iter = $ty.into_iter();
155                        for element in forwarded_elements.iter_mut() {
156                            let other = iter.next();
157                            match other {
158                                Some(other) => {
159                                    element.merge(other);
160                                }
161                                None => break,
162                            }
163                        }
164                        for other in iter {
165                            forwarded_elements.push(other);
166                        }
167                    }
168                )*
169
170                if !forwarded_elements.is_empty() {
171                    let forwarded = match req.extensions().get_ref::<Forwarded>().cloned() {
172                        Some(mut forwarded) => {
173                            forwarded.extend(forwarded_elements);
174                            forwarded
175                        }
176                        None => {
177                            let mut it = forwarded_elements.into_iter();
178                            let mut forwarded = Forwarded::new(it.next().unwrap());
179                            forwarded.extend(it);
180                            forwarded
181                        }
182                    };
183                    req.extensions().insert(forwarded);
184                }
185
186                self.inner.serve(req)
187            }
188        }
189    }
190}
191
192all_the_tuples_no_last_special_case!(get_forwarded_service_for_tuple);
193
194#[cfg(test)]
195mod tests {
196    use super::*;
197    use crate::{
198        Response, StatusCode,
199        headers::forwarded::{ClientIp, TrueClientIp, XClientIp},
200        service::web::response::IntoResponse,
201    };
202    use rama_core::{Layer, error::BoxError, extensions::ExtensionsRef, service::service_fn};
203    use rama_net::forwarded::ForwardedProtocol;
204    use std::{convert::Infallible, net::IpAddr};
205
206    fn assert_is_service<T: Service<Request<()>>>(_: T) {}
207
208    async fn dummy_service_fn() -> Result<Response, BoxError> {
209        Ok(StatusCode::OK.into_response())
210    }
211
212    #[test]
213    fn test_get_forwarded_service_is_service() {
214        assert_is_service(GetForwardedHeadersService::<_, (TrueClientIp,)>::new(
215            service_fn(dummy_service_fn),
216        ));
217        assert_is_service(
218            GetForwardedHeadersService::<_, (TrueClientIp, XClientIp)>::new(service_fn(
219                dummy_service_fn,
220            )),
221        );
222        assert_is_service(
223            GetForwardedHeadersLayer::<(ClientIp, TrueClientIp)>::new()
224                .into_layer(service_fn(dummy_service_fn)),
225        );
226    }
227
228    #[tokio::test]
229    async fn test_get_forwarded_headers() {
230        let service = GetForwardedHeadersLayer::<(rama_http_headers::forwarded::Forwarded,)>::new()
231            .into_layer(service_fn(async |req: Request<()>| {
232                let forwarded = req.extensions().get_ref::<Forwarded>().unwrap();
233                assert_eq!(forwarded.client_ip(), Some(IpAddr::from([12, 23, 34, 45])));
234                assert_eq!(forwarded.client_proto(), Some(ForwardedProtocol::HTTP));
235                Ok::<_, Infallible>(())
236            }));
237
238        let req = Request::builder()
239            .header("Forwarded", "for=\"12.23.34.45:5000\";proto=http")
240            .body(())
241            .unwrap();
242
243        service.serve(req).await.unwrap();
244    }
245}