rama_http/layer/forwarded/
set_forwarded.rs1use 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
17pub 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 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 #[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 #[must_use]
146 pub fn forwarded() -> Self {
147 Self::new()
148 }
149}
150
151impl SetForwardedHeaderLayer<Via> {
152 #[inline]
153 #[must_use]
155 pub fn via() -> Self {
156 Self::new()
157 }
158}
159
160impl SetForwardedHeaderLayer<XForwardedFor> {
161 #[inline]
162 #[must_use]
164 pub fn x_forwarded_for() -> Self {
165 Self::new()
166 }
167}
168
169impl SetForwardedHeaderLayer<XForwardedHost> {
170 #[inline]
171 #[must_use]
173 pub fn x_forwarded_host() -> Self {
174 Self::new()
175 }
176}
177
178impl SetForwardedHeaderLayer<XForwardedProto> {
179 #[inline]
180 #[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
207pub 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 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 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 pub fn forwarded(inner: S) -> Self {
267 Self::new(inner)
268 }
269}
270
271impl<S> SetForwardedHeaderService<S, Via> {
272 #[inline]
273 pub fn via(inner: S) -> Self {
275 Self::new(inner)
276 }
277}
278
279impl<S> SetForwardedHeaderService<S, XForwardedFor> {
280 #[inline]
281 pub fn x_forwarded_for(inner: S) -> Self {
283 Self::new(inner)
284 }
285}
286
287impl<S> SetForwardedHeaderService<S, XForwardedHost> {
288 #[inline]
289 pub fn x_forwarded_host(inner: S) -> Self {
291 Self::new(inner)
292 }
293}
294
295impl<S> SetForwardedHeaderService<S, XForwardedProto> {
296 #[inline]
297 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}