Skip to main content

rama_http/layer/forwarded/
get_forwarded.rs

1use crate::Request;
2use crate::headers::forwarded::{
3    ForwardHeader, Via, XForwardedFor, XForwardedHost, XForwardedProto,
4};
5use rama_core::extensions::ExtensionsRef;
6use rama_core::{Layer, Service};
7use rama_http_headers::HeaderMapExt;
8use rama_http_headers::forwarded::Forwarded;
9use rama_net::forwarded::ForwardedElement;
10use std::fmt;
11use std::marker::PhantomData;
12
13/// Layer to extract [`Forwarded`] information from the specified `T` headers.
14///
15/// This layer can be used to extract the [`Forwarded`] information from any specified header `T`,
16/// as long as the header implements the [`ForwardHeader`] trait.
17///
18/// The following headers are supported by default:
19///
20/// - [`GetForwardedHeaderLayer::forwarded`]: The standard [`Forwarded`] header [`RFC 7239`](https://github.com/plabayo/rama/blob/main/rama-http-headers/specifications/rfc7239.txt).
21/// - [`GetForwardedHeaderLayer::via`]: The canonical [`Via`] header [`RFC 9110`](https://github.com/plabayo/rama/blob/main/rama-http-core/specifications/rfc9110.txt#section-7.6.3).
22/// - [`GetForwardedHeaderLayer::x_forwarded_for`]: The canonical [`X-Forwarded-For`][XForwardedFor] header [`RFC 7239`](https://github.com/plabayo/rama/blob/main/rama-http-headers/specifications/rfc7239.txt#section-5.2).
23/// - [`GetForwardedHeaderLayer::x_forwarded_host`]: The canonical [`X-Forwarded-Host`][XForwardedHost] header [`RFC 7239`](https://github.com/plabayo/rama/blob/main/rama-http-headers/specifications/rfc7239.txt#section-5.4).
24/// - [`GetForwardedHeaderLayer::x_forwarded_proto`]: The canonical [`X-Forwarded-Proto`][XForwardedProto] header [`RFC 7239`](https://github.com/plabayo/rama/blob/main/rama-http-headers/specifications/rfc7239.txt#section-5.3).
25///
26/// Rama also has the following headers already implemented for you to use:
27///
28/// > [`X-Real-Ip`], [`X-Client-Ip`], [`Client-Ip`], [`Cf-Connecting-Ip`] and [`True-Client-Ip`].
29///
30/// There are no [`GetForwardedHeaderLayer`] constructors for these headers,
31/// but you can use the [`GetForwardedHeaderLayer::new`] constructor and pass the header type as a type parameter,
32/// alone or in a tuple with other headers.
33///
34/// [`X-Real-Ip`]: crate::headers::forwarded::XRealIp
35/// [`X-Client-Ip`]: crate::headers::forwarded::XClientIp
36/// [`Client-Ip`]: crate::headers::forwarded::ClientIp
37/// [`CF-Connecting-Ip`]: crate::headers::forwarded::CFConnectingIp
38/// [`True-Client-Ip`]: crate::headers::forwarded::TrueClientIp
39///
40/// ## Example
41///
42/// This example shows you can extract the client IP from the `X-Forwarded-For`
43/// header in case your application is behind a proxy which sets this header.
44///
45/// ```rust
46/// use rama_core::{
47///     service::service_fn,
48///     extensions::ExtensionsRef, Service, Layer,
49/// };
50/// use rama_http::{headers::forwarded::Forwarded, layer::forwarded::GetForwardedHeaderLayer, Request};
51/// use std::{convert::Infallible, net::IpAddr};
52///
53/// #[tokio::main]
54/// async fn main() {
55///     let service = GetForwardedHeaderLayer::x_forwarded_for()
56///         .into_layer(service_fn(async |req: Request<()>| {
57///             let forwarded = req.extensions().get_ref::<rama_net::forwarded::Forwarded>().unwrap();
58///             assert_eq!(forwarded.client_ip(), Some(IpAddr::from([12, 23, 34, 45])));
59///             assert!(forwarded.client_proto().is_none());
60///
61///             // ...
62///
63///             Ok::<_, Infallible>(())
64///         }));
65///
66///     let req = Request::builder()
67///         .header("X-Forwarded-For", "12.23.34.45")
68///         .body(())
69///         .unwrap();
70///
71///     service.serve(req).await.unwrap();
72/// }
73/// ```
74pub struct GetForwardedHeaderLayer<T = rama_http_headers::forwarded::Forwarded> {
75    _headers: PhantomData<fn() -> T>,
76}
77
78impl<T> fmt::Debug for GetForwardedHeaderLayer<T> {
79    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
80        f.debug_struct("GetForwardedHeaderLayer")
81            .field(
82                "_headers",
83                &format_args!("{}", std::any::type_name::<fn() -> T>()),
84            )
85            .finish()
86    }
87}
88
89impl<T> Clone for GetForwardedHeaderLayer<T> {
90    fn clone(&self) -> Self {
91        Self {
92            _headers: PhantomData,
93        }
94    }
95}
96
97impl Default for GetForwardedHeaderLayer {
98    fn default() -> Self {
99        Self::forwarded()
100    }
101}
102
103impl<T> GetForwardedHeaderLayer<T> {
104    /// Create a new `GetForwardedHeaderLayer` for the specified headers `T`.
105    #[must_use]
106    pub const fn new() -> Self {
107        Self {
108            _headers: PhantomData,
109        }
110    }
111}
112
113impl GetForwardedHeaderLayer {
114    #[inline]
115    /// Create a new `GetForwardedHeaderLayer` for the standard [`Forwarded`] header.
116    #[must_use]
117    pub fn forwarded() -> Self {
118        Self::new()
119    }
120}
121
122impl GetForwardedHeaderLayer<Via> {
123    #[inline]
124    /// Create a new `GetForwardedHeaderLayer` for the canonical [`Via`] header.
125    #[must_use]
126    pub fn via() -> Self {
127        Self::new()
128    }
129}
130
131impl GetForwardedHeaderLayer<XForwardedFor> {
132    #[inline]
133    /// Create a new `GetForwardedHeaderLayer` for the canonical [`X-Forwarded-For`][XForwardedFor] header.
134    #[must_use]
135    pub fn x_forwarded_for() -> Self {
136        Self::new()
137    }
138}
139
140impl GetForwardedHeaderLayer<XForwardedHost> {
141    #[inline]
142    /// Create a new `GetForwardedHeaderLayer` for the canonical [`X-Forwarded-Host`][XForwardedHost] header.
143    #[must_use]
144    pub fn x_forwarded_host() -> Self {
145        Self::new()
146    }
147}
148
149impl GetForwardedHeaderLayer<XForwardedProto> {
150    #[inline]
151    /// Create a new `GetForwardedHeaderLayer` for the canonical [`X-Forwarded-Proto`][XForwardedProto] header.
152    #[must_use]
153    pub fn x_forwarded_proto() -> Self {
154        Self::new()
155    }
156}
157
158impl<H, S> Layer<S> for GetForwardedHeaderLayer<H> {
159    type Service = GetForwardedHeaderService<S, H>;
160
161    fn layer(&self, inner: S) -> Self::Service {
162        Self::Service {
163            inner,
164            _headers: PhantomData,
165        }
166    }
167}
168
169/// Middleware service to extract [`Forwarded`] information from the specified `T` headers.
170///
171/// See [`GetForwardedHeaderLayer`] for more information.
172pub struct GetForwardedHeaderService<S, T = Forwarded> {
173    inner: S,
174    _headers: PhantomData<fn() -> T>,
175}
176
177impl<S: fmt::Debug, T> fmt::Debug for GetForwardedHeaderService<S, T> {
178    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
179        f.debug_struct("GetForwardedHeaderService")
180            .field("inner", &self.inner)
181            .field("_headers", &format_args!("{}", std::any::type_name::<T>()))
182            .finish()
183    }
184}
185
186impl<S: Clone, T> Clone for GetForwardedHeaderService<S, T> {
187    fn clone(&self) -> Self {
188        Self {
189            inner: self.inner.clone(),
190            _headers: PhantomData,
191        }
192    }
193}
194
195impl<S, T> GetForwardedHeaderService<S, T> {
196    /// Create a new `GetForwardedHeaderService` for the specified headers `T`.
197    pub const fn new(inner: S) -> Self {
198        Self {
199            inner,
200            _headers: PhantomData,
201        }
202    }
203}
204
205impl<S> GetForwardedHeaderService<S> {
206    #[inline]
207    /// Create a new `GetForwardedHeaderService` for the standard [`Forwarded`] header.
208    pub fn forwarded(inner: S) -> Self {
209        Self::new(inner)
210    }
211}
212
213impl<S> GetForwardedHeaderService<S, Via> {
214    #[inline]
215    /// Create a new `GetForwardedHeaderService` for the canonical [`Via`] header.
216    pub fn via(inner: S) -> Self {
217        Self::new(inner)
218    }
219}
220
221impl<S> GetForwardedHeaderService<S, XForwardedFor> {
222    #[inline]
223    /// Create a new `GetForwardedHeaderService` for the canonical [`X-Forwarded-For`](XForwardedFor) header.
224    pub fn x_forwarded_for(inner: S) -> Self {
225        Self::new(inner)
226    }
227}
228
229impl<S> GetForwardedHeaderService<S, XForwardedHost> {
230    #[inline]
231    /// Create a new `GetForwardedHeaderService` for the canonical [`X-Forwarded-Host`][XForwardedHost] header.
232    pub fn x_forwarded_host(inner: S) -> Self {
233        Self::new(inner)
234    }
235}
236
237impl<S> GetForwardedHeaderService<S, XForwardedProto> {
238    #[inline]
239    /// Create a new `GetForwardedHeaderService` for the canonical [`X-Forwarded-Proto`][XForwardedProto] header.
240    pub fn x_forwarded_proto(inner: S) -> Self {
241        Self::new(inner)
242    }
243}
244
245impl<H, S, Body> Service<Request<Body>> for GetForwardedHeaderService<S, H>
246where
247    H: ForwardHeader + Send + Sync + 'static,
248    S: Service<Request<Body>>,
249    Body: Send + 'static,
250{
251    type Output = S::Output;
252    type Error = S::Error;
253
254    fn serve(
255        &self,
256        req: Request<Body>,
257    ) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send + '_ {
258        let mut forwarded_elements: Vec<ForwardedElement> = Vec::with_capacity(1);
259
260        if let Some(header) = req.headers().typed_get::<H>() {
261            forwarded_elements.extend(header);
262        }
263
264        if !forwarded_elements.is_empty() {
265            let forwarded = if let Some(mut forwarded) = req
266                .extensions()
267                .get_ref::<rama_net::forwarded::Forwarded>()
268                .cloned()
269            {
270                forwarded.extend(forwarded_elements);
271                Some(forwarded)
272            } else {
273                let mut it = forwarded_elements.into_iter();
274                if let Some(first) = it.next() {
275                    let mut forwarded = rama_net::forwarded::Forwarded::new(first);
276                    forwarded.extend(it);
277                    Some(forwarded)
278                } else {
279                    None
280                }
281            };
282
283            if let Some(forwarded) = forwarded {
284                req.extensions().insert(forwarded);
285            }
286        }
287
288        self.inner.serve(req)
289    }
290}
291
292#[cfg(test)]
293mod tests {
294    use super::*;
295    use crate::{Response, StatusCode, service::web::response::IntoResponse};
296    use rama_core::{Layer, error::BoxError, extensions::ExtensionsRef, service::service_fn};
297    use rama_http_headers::forwarded::{TrueClientIp, XRealIp};
298    use rama_net::forwarded::{ForwardedProtocol, ForwardedVersion};
299    use std::{convert::Infallible, net::IpAddr};
300
301    fn assert_is_service<T: Service<Request<()>>>(_: T) {}
302
303    async fn dummy_service_fn() -> Result<Response, BoxError> {
304        Ok(StatusCode::OK.into_response())
305    }
306
307    #[test]
308    fn test_get_forwarded_service_is_service() {
309        assert_is_service(GetForwardedHeaderService::forwarded(service_fn(
310            dummy_service_fn,
311        )));
312        assert_is_service(GetForwardedHeaderService::via(service_fn(dummy_service_fn)));
313        assert_is_service(GetForwardedHeaderService::x_forwarded_for(service_fn(
314            dummy_service_fn,
315        )));
316        assert_is_service(GetForwardedHeaderService::x_forwarded_proto(service_fn(
317            dummy_service_fn,
318        )));
319        assert_is_service(GetForwardedHeaderService::x_forwarded_host(service_fn(
320            dummy_service_fn,
321        )));
322        assert_is_service(GetForwardedHeaderService::<_, TrueClientIp>::new(
323            service_fn(dummy_service_fn),
324        ));
325        assert_is_service(
326            GetForwardedHeaderLayer::forwarded().into_layer(service_fn(dummy_service_fn)),
327        );
328        assert_is_service(GetForwardedHeaderLayer::via().into_layer(service_fn(dummy_service_fn)));
329        assert_is_service(
330            GetForwardedHeaderLayer::<XRealIp>::new().into_layer(service_fn(dummy_service_fn)),
331        );
332    }
333
334    #[tokio::test]
335    async fn test_get_forwarded_header_forwarded() {
336        let service =
337            GetForwardedHeaderLayer::forwarded().into_layer(service_fn(async |req: Request<()>| {
338                let forwarded = req
339                    .extensions()
340                    .get_ref::<rama_net::forwarded::Forwarded>()
341                    .unwrap();
342                assert_eq!(forwarded.client_ip(), Some(IpAddr::from([12, 23, 34, 45])));
343                assert_eq!(forwarded.client_proto(), Some(ForwardedProtocol::HTTP));
344                Ok::<_, Infallible>(())
345            }));
346
347        let req = Request::builder()
348            .header("Forwarded", "for=\"12.23.34.45:5000\";proto=http")
349            .body(())
350            .unwrap();
351
352        service.serve(req).await.unwrap();
353    }
354
355    #[tokio::test]
356    async fn test_get_forwarded_header_via() {
357        let service =
358            GetForwardedHeaderLayer::via().into_layer(service_fn(async |req: Request<()>| {
359                let forwarded = req
360                    .extensions()
361                    .get_ref::<rama_net::forwarded::Forwarded>()
362                    .unwrap();
363                assert!(forwarded.client_ip().is_none());
364                assert_eq!(
365                    forwarded.iter().next().unwrap().forwarded_by(),
366                    Some(&(IpAddr::from([12, 23, 34, 45]), 5000).into())
367                );
368                assert!(forwarded.client_proto().is_none());
369                assert_eq!(forwarded.client_version(), Some(ForwardedVersion::HTTP_11));
370                Ok::<_, Infallible>(())
371            }));
372
373        let req = Request::builder()
374            .header("Via", "1.1 12.23.34.45:5000")
375            .body(())
376            .unwrap();
377
378        service.serve(req).await.unwrap();
379    }
380
381    #[tokio::test]
382    async fn test_get_forwarded_header_x_forwarded_for() {
383        let service = GetForwardedHeaderLayer::x_forwarded_for().into_layer(service_fn(
384            async |req: Request<()>| {
385                let forwarded = req
386                    .extensions()
387                    .get_ref::<rama_net::forwarded::Forwarded>()
388                    .unwrap();
389                assert_eq!(forwarded.client_ip(), Some(IpAddr::from([12, 23, 34, 45])));
390                assert!(forwarded.client_proto().is_none());
391                Ok::<_, Infallible>(())
392            },
393        ));
394
395        let req = Request::builder()
396            .header("X-Forwarded-For", "12.23.34.45, 127.0.0.1")
397            .body(())
398            .unwrap();
399
400        service.serve(req).await.unwrap();
401    }
402}