rama_http/layer/forwarded/
set_forwarded_multi.rs1#![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
19pub 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 #[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
91pub 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 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}