Skip to main content

rama_http/layer/forwarded/
set_forwarded.rs

1use crate::Request;
2use crate::headers::HeaderMapExt;
3use crate::headers::forwarded::{
4    ForwardHeader, Via, XForwardedFor, XForwardedHost, XForwardedProto,
5};
6use rama_core::error::{BoxError, BoxErrorExt as _, ErrorContext as _};
7use rama_core::extensions::ExtensionsRef;
8use rama_core::{Layer, Service};
9use rama_http_headers::forwarded::Forwarded;
10use rama_net::address::Domain;
11use rama_net::forwarded::{ForwardedElement, NodeId};
12use rama_net::stream::SocketInfo;
13use rama_net::{AuthorityInputExt, Protocol, ProtocolInputExt};
14use std::fmt;
15use std::marker::PhantomData;
16
17/// Layer to write [`Forwarded`] information for this proxy,
18/// added to the end of the chain of forwarded information already known.
19///
20/// This layer can set any header as long as you have a [`ForwardHeader`] implementation
21/// for the header you want to set. You can pass it as the type to the layer when creating
22/// the layer using [`SetForwardedHeaderLayer::new`].
23///
24/// The following headers are supported out of the box with each their own constructor:
25///
26/// - [`SetForwardedHeaderLayer::forwarded`]: the standard [`Forwarded`] header [`RFC 7239`](https://github.com/plabayo/rama/blob/main/rama-http-headers/specifications/rfc7239.txt);
27/// - [`SetForwardedHeaderLayer::via`]: the canonical [`Via`] header (non-standard);
28/// - [`SetForwardedHeaderLayer::x_forwarded_for`]: the canonical [`X-Forwarded-For`][`XForwardedFor`] header (non-standard);
29/// - [`SetForwardedHeaderLayer::x_forwarded_host`]: the canonical [`X-Forwarded-Host`][`XForwardedHost`] header (non-standard);
30/// - [`SetForwardedHeaderLayer::x_forwarded_proto`]: the canonical [`X-Forwarded-Proto`][`XForwardedProto`] header (non-standard).
31///
32/// The "by" property is set to `rama` by default. Use [`SetForwardedHeaderLayer::set_forward_by`] to overwrite this,
33/// typically with the actual [`IPv4`]/[`IPv6`] address of your proxy.
34///
35/// [`IPv4`]: std::net::Ipv4Addr
36/// [`IPv6`]: std::net::Ipv6Addr
37///
38/// Rama also has the following headers already implemented for you to use:
39///
40/// > [`X-Real-Ip`], [`X-Client-Ip`], [`Client-Ip`], [`CF-Connecting-Ip`] and [`True-Client-Ip`].
41///
42/// There are no [`SetForwardedHeaderLayer`] constructors for these headers,
43/// but you can use the [`SetForwardedHeaderLayer::new`] constructor and pass the header type as a type parameter,
44/// alone or in a tuple with other headers.
45///
46/// [`X-Real-Ip`]: crate::headers::forwarded::XRealIp
47/// [`X-Client-Ip`]: crate::headers::forwarded::XClientIp
48/// [`Client-Ip`]: crate::headers::forwarded::ClientIp
49/// [`CF-Connecting-Ip`]: crate::headers::forwarded::CFConnectingIp
50/// [`True-Client-Ip`]: crate::headers::forwarded::TrueClientIp
51///
52/// ## Example
53///
54/// This example shows how you could expose the real Client IP using
55/// the [`X-Real-IP`][`crate::headers::forwarded::XRealIp`] header.
56///
57/// ```rust
58/// use rama_net::stream::SocketInfo;
59/// use rama_http::Request;
60/// use rama_core::service::service_fn;
61/// use rama_http::{headers::forwarded::XRealIp, layer::forwarded::SetForwardedHeaderLayer};
62/// use rama_core::{extensions::ExtensionsRef, Service, Layer};
63/// use std::convert::Infallible;
64///
65/// # type Body = ();
66/// # type State = ();
67///
68/// # #[tokio::main]
69/// # async fn main() {
70/// async fn svc(request: Request<Body>) -> Result<(), Infallible> {
71///     // ...
72///     # assert_eq!(
73///     #     request.headers().get("X-Real-Ip").unwrap(),
74///     #     "42.37.100.50:62345",
75///     # );
76///     # Ok(())
77/// }
78///
79/// let service = SetForwardedHeaderLayer::<XRealIp>::new()
80///     .into_layer(service_fn(svc));
81///
82/// # let mut req = Request::builder().uri("http://example.com").body(()).unwrap();
83/// # req.extensions().insert(SocketInfo::new(None, "42.37.100.50:62345".parse().unwrap()));
84/// service.serve(req).await.unwrap();
85/// # }
86/// ```
87pub struct SetForwardedHeaderLayer<T = Forwarded> {
88    by_node: NodeId,
89    _headers: PhantomData<fn() -> T>,
90}
91
92impl<T> fmt::Debug for SetForwardedHeaderLayer<T> {
93    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
94        f.debug_struct("SetForwardedHeaderLayer")
95            .field("by_node", &self.by_node)
96            .field(
97                "_headers",
98                &format_args!("{}", std::any::type_name::<fn() -> T>()),
99            )
100            .finish()
101    }
102}
103
104impl<T> Clone for SetForwardedHeaderLayer<T> {
105    fn clone(&self) -> Self {
106        Self {
107            by_node: self.by_node.clone(),
108            _headers: PhantomData,
109        }
110    }
111}
112
113impl<T> SetForwardedHeaderLayer<T> {
114    rama_utils::macros::generate_set_and_with! {
115        /// Set the given [`NodeId`] as the "by" property, identifying this proxy.
116        ///
117        /// Default of `None` will be set to `rama` otherwise.
118        pub fn forward_by(mut self, node_id: impl Into<NodeId>) -> Self {
119            self.by_node = node_id.into();
120            self
121        }
122    }
123}
124
125impl<T> SetForwardedHeaderLayer<T> {
126    /// Create a new `SetForwardedHeaderLayer` for the specified headers `T`.
127    #[must_use]
128    pub fn new() -> Self {
129        Self {
130            by_node: Domain::from_static("rama").into(),
131            _headers: PhantomData,
132        }
133    }
134}
135
136impl Default for SetForwardedHeaderLayer {
137    fn default() -> Self {
138        Self::forwarded()
139    }
140}
141
142impl SetForwardedHeaderLayer {
143    #[inline]
144    /// Create a new `SetForwardedHeaderLayer` for the standard [`Forwarded`] header.
145    #[must_use]
146    pub fn forwarded() -> Self {
147        Self::new()
148    }
149}
150
151impl SetForwardedHeaderLayer<Via> {
152    #[inline]
153    /// Create a new `SetForwardedHeaderLayer` for the canonical [`Via`] header.
154    #[must_use]
155    pub fn via() -> Self {
156        Self::new()
157    }
158}
159
160impl SetForwardedHeaderLayer<XForwardedFor> {
161    #[inline]
162    /// Create a new `SetForwardedHeaderLayer` for the canonical [`X-Forwarded-For`][XForwardedFor] header.
163    #[must_use]
164    pub fn x_forwarded_for() -> Self {
165        Self::new()
166    }
167}
168
169impl SetForwardedHeaderLayer<XForwardedHost> {
170    #[inline]
171    /// Create a new `SetForwardedHeaderLayer` for the canonical [`X-Forwarded-Host`][XForwardedHost] header.
172    #[must_use]
173    pub fn x_forwarded_host() -> Self {
174        Self::new()
175    }
176}
177
178impl SetForwardedHeaderLayer<XForwardedProto> {
179    #[inline]
180    /// Create a new `SetForwardedHeaderLayer` for the canonical [`X-Forwarded-Proto`][XForwardedProto] header.
181    #[must_use]
182    pub fn x_forwarded_proto() -> Self {
183        Self::new()
184    }
185}
186
187impl<H, S> Layer<S> for SetForwardedHeaderLayer<H> {
188    type Service = SetForwardedHeaderService<S, H>;
189
190    fn layer(&self, inner: S) -> Self::Service {
191        Self::Service {
192            inner,
193            by_node: self.by_node.clone(),
194            _headers: PhantomData,
195        }
196    }
197
198    fn into_layer(self, inner: S) -> Self::Service {
199        Self::Service {
200            inner,
201            by_node: self.by_node,
202            _headers: PhantomData,
203        }
204    }
205}
206
207/// Middleware [`Service`] to write [`Forwarded`] information for this proxy,
208/// added to the end of the chain of forwarded information already known.
209///
210/// See [`SetForwardedHeaderLayer`] for more information.
211pub struct SetForwardedHeaderService<S, T = Forwarded> {
212    inner: S,
213    by_node: NodeId,
214    _headers: PhantomData<fn() -> T>,
215}
216
217impl<S: fmt::Debug, T> fmt::Debug for SetForwardedHeaderService<S, T> {
218    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
219        f.debug_struct("SetForwardedHeaderService")
220            .field("inner", &self.inner)
221            .field("by_node", &self.by_node)
222            .field(
223                "_headers",
224                &format_args!("{}", std::any::type_name::<fn() -> T>()),
225            )
226            .finish()
227    }
228}
229
230impl<S: Clone, T> Clone for SetForwardedHeaderService<S, T> {
231    fn clone(&self) -> Self {
232        Self {
233            inner: self.inner.clone(),
234            by_node: self.by_node.clone(),
235            _headers: PhantomData,
236        }
237    }
238}
239
240impl<S, T> SetForwardedHeaderService<S, T> {
241    rama_utils::macros::generate_set_and_with! {
242        /// Set the given [`NodeId`] as the "by" property, identifying this proxy.
243        ///
244        /// Default of `None` will be set to `rama` otherwise.
245        pub fn forward_by(mut self, node_id: impl Into<NodeId>) -> Self {
246            self.by_node = node_id.into();
247            self
248        }
249    }
250}
251
252impl<S, T> SetForwardedHeaderService<S, T> {
253    /// Create a new `SetForwardedHeaderService` for the specified headers `T`.
254    pub fn new(inner: S) -> Self {
255        Self {
256            inner,
257            by_node: Domain::from_static("rama").into(),
258            _headers: PhantomData,
259        }
260    }
261}
262
263impl<S> SetForwardedHeaderService<S> {
264    #[inline]
265    /// Create a new `SetForwardedHeaderService` for the standard [`Forwarded`] header.
266    pub fn forwarded(inner: S) -> Self {
267        Self::new(inner)
268    }
269}
270
271impl<S> SetForwardedHeaderService<S, Via> {
272    #[inline]
273    /// Create a new `SetForwardedHeaderService` for the canonical [`Via`] header.
274    pub fn via(inner: S) -> Self {
275        Self::new(inner)
276    }
277}
278
279impl<S> SetForwardedHeaderService<S, XForwardedFor> {
280    #[inline]
281    /// Create a new `SetForwardedHeaderService` for the canonical [`X-Forwarded-For`][XForwardedFor] header.
282    pub fn x_forwarded_for(inner: S) -> Self {
283        Self::new(inner)
284    }
285}
286
287impl<S> SetForwardedHeaderService<S, XForwardedHost> {
288    #[inline]
289    /// Create a new `SetForwardedHeaderService` for the canonical [`X-Forwarded-Host`][XForwardedHost] header.
290    pub fn x_forwarded_host(inner: S) -> Self {
291        Self::new(inner)
292    }
293}
294
295impl<S> SetForwardedHeaderService<S, XForwardedProto> {
296    #[inline]
297    /// Create a new `SetForwardedHeaderService` for the canonical [`X-Forwarded-Proto`][XForwardedProto] header.
298    pub fn x_forwarded_proto(inner: S) -> Self {
299        Self::new(inner)
300    }
301}
302
303impl<S, H, Body> Service<Request<Body>> for SetForwardedHeaderService<S, H>
304where
305    S: Service<Request<Body>, Error: Into<BoxError>>,
306    H: ForwardHeader + Send + Sync + 'static,
307    Body: Send + 'static,
308{
309    type Output = S::Output;
310    type Error = BoxError;
311
312    async fn serve(&self, mut req: Request<Body>) -> Result<Self::Output, Self::Error> {
313        let forwarded: Option<rama_net::forwarded::Forwarded> = req.extensions().get_ref().cloned();
314
315        let mut forwarded_element = ForwardedElement::new_forwarded_by(self.by_node.clone());
316
317        if let Some(peer_addr) = req
318            .extensions()
319            .get_ref::<SocketInfo>()
320            .map(|socket| socket.peer_addr())
321        {
322            forwarded_element.set_forwarded_for(peer_addr);
323        }
324        let authority = req
325            .authority()
326            .ok_or_else(|| BoxError::from_static_str("set forwarded: no authority"))?;
327
328        forwarded_element.set_forwarded_host(authority);
329
330        let protocol = req.protocol().unwrap_or(&Protocol::HTTP);
331        if let Ok(forwarded_proto) = protocol.try_into() {
332            forwarded_element.set_forwarded_proto(forwarded_proto);
333        }
334
335        let forwarded = match forwarded {
336            None => Some(rama_net::forwarded::Forwarded::new(forwarded_element)),
337            Some(mut forwarded) => {
338                forwarded.append(forwarded_element);
339                Some(forwarded)
340            }
341        };
342
343        if let Some(forwarded) = forwarded
344            && let Some(header) = H::try_from_forwarded(forwarded.iter())
345        {
346            req.headers_mut().typed_insert(header);
347        }
348
349        self.inner.serve(req).await.into_box_error()
350    }
351}
352
353#[cfg(test)]
354mod tests {
355    use super::*;
356    use crate::{
357        Response, StatusCode,
358        headers::forwarded::{TrueClientIp, XRealIp},
359        service::web::response::IntoResponse,
360    };
361    use rama_core::{Layer, error::BoxError, extensions::ExtensionsRef, service::service_fn};
362    use std::{convert::Infallible, net::IpAddr};
363
364    fn assert_is_service<T: Service<Request<()>>>(_: T) {}
365
366    async fn dummy_service_fn() -> Result<Response, BoxError> {
367        Ok(StatusCode::OK.into_response())
368    }
369
370    #[test]
371    fn test_set_forwarded_service_is_service() {
372        assert_is_service(SetForwardedHeaderService::forwarded(service_fn(
373            dummy_service_fn,
374        )));
375        assert_is_service(SetForwardedHeaderService::via(service_fn(dummy_service_fn)));
376        assert_is_service(SetForwardedHeaderService::x_forwarded_for(service_fn(
377            dummy_service_fn,
378        )));
379        assert_is_service(SetForwardedHeaderService::x_forwarded_proto(service_fn(
380            dummy_service_fn,
381        )));
382        assert_is_service(SetForwardedHeaderService::x_forwarded_host(service_fn(
383            dummy_service_fn,
384        )));
385        assert_is_service(SetForwardedHeaderService::<_, TrueClientIp>::new(
386            service_fn(dummy_service_fn),
387        ));
388        assert_is_service(SetForwardedHeaderLayer::via().into_layer(service_fn(dummy_service_fn)));
389        assert_is_service(
390            SetForwardedHeaderLayer::<XRealIp>::new().into_layer(service_fn(dummy_service_fn)),
391        );
392    }
393
394    #[tokio::test]
395    async fn test_set_forwarded_service_forwarded() {
396        async fn svc(request: Request<()>) -> Result<(), Infallible> {
397            assert_eq!(
398                request.headers().get("Forwarded").unwrap(),
399                "by=rama;host=\"example.com:80\";proto=http"
400            );
401            Ok(())
402        }
403
404        let service = SetForwardedHeaderService::forwarded(service_fn(svc));
405        let req = Request::builder()
406            .uri("http://example.com")
407            .body(())
408            .unwrap();
409        service.serve(req).await.unwrap();
410    }
411
412    #[tokio::test]
413    async fn test_set_forwarded_service_forwarded_with_chain() {
414        async fn svc(request: Request<()>) -> Result<(), Infallible> {
415            assert_eq!(
416                request.headers().get("Forwarded").unwrap(),
417                "for=12.23.34.45,by=rama;for=\"127.0.0.1:62345\";host=\"www.example.com:443\";proto=https",
418            );
419            Ok(())
420        }
421
422        let service = SetForwardedHeaderService::forwarded(service_fn(svc));
423        let req = Request::builder()
424            .uri("https://www.example.com")
425            .body(())
426            .unwrap();
427        req.extensions().insert(rama_net::forwarded::Forwarded::new(
428            ForwardedElement::new_forwarded_for(IpAddr::from([12, 23, 34, 45])),
429        ));
430        req.extensions()
431            .insert(SocketInfo::new(None, "127.0.0.1:62345".parse().unwrap()));
432        service.serve(req).await.unwrap();
433    }
434
435    #[tokio::test]
436    async fn test_set_forwarded_service_x_forwarded_for_with_chain() {
437        async fn svc(request: Request<()>) -> Result<(), Infallible> {
438            assert_eq!(
439                request.headers().get("X-Forwarded-For").unwrap(),
440                "12.23.34.45, 127.0.0.1",
441            );
442            Ok(())
443        }
444
445        let service = SetForwardedHeaderService::x_forwarded_for(service_fn(svc));
446        let req = Request::builder()
447            .uri("https://www.example.com")
448            .body(())
449            .unwrap();
450        req.extensions().insert(rama_net::forwarded::Forwarded::new(
451            ForwardedElement::new_forwarded_for(IpAddr::from([12, 23, 34, 45])),
452        ));
453        req.extensions()
454            .insert(SocketInfo::new(None, "127.0.0.1:62345".parse().unwrap()));
455        service.serve(req).await.unwrap();
456    }
457
458    #[tokio::test]
459    async fn test_set_forwarded_service_forwarded_fully_defined() {
460        async fn svc(request: Request<()>) -> Result<(), Infallible> {
461            assert_eq!(
462                request.headers().get("Forwarded").unwrap(),
463                "by=12.23.34.45;for=\"127.0.0.1:62345\";host=\"www.example.com:443\";proto=https",
464            );
465            Ok(())
466        }
467
468        let service = SetForwardedHeaderService::forwarded(service_fn(svc))
469            .with_forward_by(IpAddr::from([12, 23, 34, 45]));
470        let req = Request::builder()
471            .uri("https://www.example.com")
472            .body(())
473            .unwrap();
474        req.extensions()
475            .insert(SocketInfo::new(None, "127.0.0.1:62345".parse().unwrap()));
476        service.serve(req).await.unwrap();
477    }
478
479    #[tokio::test]
480    async fn test_set_forwarded_service_forwarded_fully_defined_with_chain() {
481        async fn svc(request: Request<()>) -> Result<(), Infallible> {
482            assert_eq!(
483                request.headers().get("Forwarded").unwrap(),
484                "by=rama;for=\"127.0.0.1:62345\";host=\"www.example.com:443\";proto=https",
485            );
486            Ok(())
487        }
488
489        let service = SetForwardedHeaderService::forwarded(service_fn(svc));
490        let req = Request::builder()
491            .uri("https://www.example.com")
492            .body(())
493            .unwrap();
494        req.extensions()
495            .insert(SocketInfo::new(None, "127.0.0.1:62345".parse().unwrap()));
496        service.serve(req).await.unwrap();
497    }
498}