Skip to main content

rama_http/layer/forwarded/
set_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::HeaderMapExt;
8use crate::headers::forwarded::ForwardHeader;
9use rama_core::error::{BoxError, BoxErrorExt as _, ErrorContext as _};
10use rama_core::{Layer, Service, extensions::ExtensionsRef};
11use rama_net::address::Domain;
12use rama_net::forwarded::{Forwarded, ForwardedElement, NodeId};
13use rama_net::stream::SocketInfo;
14use rama_net::{AuthorityInputExt, Protocol, ProtocolInputExt};
15use rama_utils::macros::all_the_tuples_no_last_special_case;
16use std::fmt;
17use std::marker::PhantomData;
18
19/// Layer to write [`Forwarded`] information for this proxy,
20/// added to the end of the chain of forwarded information already known.
21///
22/// Use [`super::SetForwardedHeaderLayer`] if you only need a single a header.
23///
24/// This layer can set any headers as long as you have a [`ForwardHeader`] implementation
25/// for the headers you want to set. You can pass it as the type to the layer when creating
26/// the layer using [`SetForwardedHeadersLayer::new`], with the headers in a single tuple.
27pub struct SetForwardedHeadersLayer<T = Forwarded> {
28    by_node: NodeId,
29    _headers: PhantomData<fn() -> T>,
30}
31
32impl<T: fmt::Debug> fmt::Debug for SetForwardedHeadersLayer<T> {
33    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
34        f.debug_struct("SetForwardedHeadersLayer")
35            .field("by_node", &self.by_node)
36            .field(
37                "_headers",
38                &format_args!("{}", std::any::type_name::<fn() -> T>()),
39            )
40            .finish()
41    }
42}
43
44impl<T: Clone> Clone for SetForwardedHeadersLayer<T> {
45    fn clone(&self) -> Self {
46        Self {
47            by_node: self.by_node.clone(),
48            _headers: PhantomData,
49        }
50    }
51}
52
53impl<T> Default for SetForwardedHeadersLayer<T> {
54    #[inline]
55    fn default() -> Self {
56        Self::new()
57    }
58}
59
60impl<T> SetForwardedHeadersLayer<T> {
61    /// Create a new `SetForwardedHeadersLayer` for the specified headers `T`.
62    #[must_use]
63    pub fn new() -> Self {
64        Self {
65            by_node: Domain::from_static("rama").into(),
66            _headers: PhantomData,
67        }
68    }
69}
70
71impl<H, S> Layer<S> for SetForwardedHeadersLayer<H> {
72    type Service = SetForwardedHeadersService<S, H>;
73
74    fn layer(&self, inner: S) -> Self::Service {
75        Self::Service {
76            inner,
77            by_node: self.by_node.clone(),
78            _headers: PhantomData,
79        }
80    }
81
82    fn into_layer(self, inner: S) -> Self::Service {
83        Self::Service {
84            inner,
85            by_node: self.by_node,
86            _headers: PhantomData,
87        }
88    }
89}
90
91/// Middleware [`Service`] to write [`Forwarded`] information for this proxy,
92/// added to the end of the chain of forwarded information already known.
93///
94/// See [`SetForwardedHeadersLayer`] for more information.
95pub struct SetForwardedHeadersService<S, T = Forwarded> {
96    inner: S,
97    by_node: NodeId,
98    _headers: PhantomData<fn() -> T>,
99}
100
101impl<S: fmt::Debug, T> fmt::Debug for SetForwardedHeadersService<S, T> {
102    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
103        f.debug_struct("SetForwardedHeadersService")
104            .field("inner", &self.inner)
105            .field("by_node", &self.by_node)
106            .field(
107                "_headers",
108                &format_args!("{}", std::any::type_name::<fn() -> T>()),
109            )
110            .finish()
111    }
112}
113
114impl<S: Clone, T> Clone for SetForwardedHeadersService<S, T> {
115    fn clone(&self) -> Self {
116        Self {
117            inner: self.inner.clone(),
118            by_node: self.by_node.clone(),
119            _headers: PhantomData,
120        }
121    }
122}
123
124impl<S, T> SetForwardedHeadersService<S, T> {
125    /// Create a new `SetForwardedHeadersService` for the specified headers `T`.
126    pub fn new(inner: S) -> Self {
127        Self {
128            inner,
129            by_node: Domain::from_static("rama").into(),
130            _headers: PhantomData,
131        }
132    }
133}
134
135macro_rules! set_forwarded_service_for_tuple {
136    ( $($ty:ident),* $(,)? ) => {
137        #[allow(non_snake_case)]
138        impl<S, $($ty),* , Body> Service<Request<Body>> for SetForwardedHeadersService<S, ($($ty,)*)>
139        where
140            $( $ty: ForwardHeader + Send + Sync + 'static, )*
141            S: Service<Request<Body>, Error: Into<BoxError>>,
142            Body: Send + 'static,
143        {
144            type Output = S::Output;
145            type Error = BoxError;
146
147            async fn serve(
148                &self,
149                mut req: Request<Body>,
150            ) -> Result<Self::Output, Self::Error> {
151                let forwarded: Option<Forwarded> = req.extensions().get_ref().cloned();
152
153                let mut forwarded_element = ForwardedElement::new_forwarded_by(self.by_node.clone());
154
155                if let Some(peer_addr) = req
156                    .extensions()
157                    .ingress()
158                    .and_then(|ext|ext.get_ref::<SocketInfo>())
159                    .map(|socket| socket.peer_addr())
160                {
161
162                    forwarded_element.set_forwarded_for(peer_addr);
163                }
164
165                let authority = req
166                    .authority()
167                    .ok_or_else(|| BoxError::from_static_str("set forwarded: no authority"))?;
168
169                forwarded_element.set_forwarded_host(authority);
170
171                let protocol = req.protocol().unwrap_or(&Protocol::HTTP);
172                if let Ok(forwarded_proto) = protocol.try_into() {
173                    forwarded_element.set_forwarded_proto(forwarded_proto);
174                }
175
176                let forwarded = match forwarded {
177                    None => Some(Forwarded::new(forwarded_element)),
178                    Some(mut forwarded) => {
179                        forwarded.append(forwarded_element);
180                        Some(forwarded)
181                    }
182                };
183
184                if let Some(forwarded) = forwarded {
185                    $(
186                        if let Some(header) = $ty::try_from_forwarded(forwarded.iter()) {
187                            req.headers_mut().typed_insert(header);
188                        }
189                    )*
190                }
191
192                self.inner.serve(req).await.into_box_error()
193            }
194        }
195    };
196}
197all_the_tuples_no_last_special_case!(set_forwarded_service_for_tuple);
198
199#[cfg(test)]
200mod tests {
201    use super::*;
202    use crate::{
203        Response, StatusCode,
204        headers::forwarded::{TrueClientIp, XClientIp, XRealIp},
205        service::web::response::IntoResponse,
206    };
207    use rama_core::{Layer, error::BoxError, service::service_fn};
208    use rama_http_headers::forwarded::XForwardedProto;
209    use std::convert::Infallible;
210
211    fn assert_is_service<T: Service<Request<()>>>(_: T) {}
212
213    async fn dummy_service_fn() -> Result<Response, BoxError> {
214        Ok(StatusCode::OK.into_response())
215    }
216
217    #[test]
218    fn test_set_forwarded_service_is_service() {
219        assert_is_service(SetForwardedHeadersService::<_, (TrueClientIp,)>::new(
220            service_fn(dummy_service_fn),
221        ));
222        assert_is_service(
223            SetForwardedHeadersService::<_, (TrueClientIp, XClientIp)>::new(service_fn(
224                dummy_service_fn,
225            )),
226        );
227        assert_is_service(
228            SetForwardedHeadersLayer::<(XRealIp, XForwardedProto)>::new()
229                .into_layer(service_fn(dummy_service_fn)),
230        );
231    }
232
233    #[tokio::test]
234    async fn test_set_forwarded_service_forwarded() {
235        async fn svc(request: Request<()>) -> Result<(), Infallible> {
236            assert_eq!(
237                request.headers().get("Forwarded").unwrap(),
238                "by=rama;host=\"example.com:80\";proto=http"
239            );
240            Ok(())
241        }
242
243        let service =
244            SetForwardedHeadersService::<_, (rama_http_headers::forwarded::Forwarded,)>::new(
245                service_fn(svc),
246            );
247        let req = Request::builder()
248            .uri("http://example.com")
249            .body(())
250            .unwrap();
251        service.serve(req).await.unwrap();
252    }
253}