rama_http/layer/forwarded/
get_forwarded.rs1use 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
13pub 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 #[must_use]
106 pub const fn new() -> Self {
107 Self {
108 _headers: PhantomData,
109 }
110 }
111}
112
113impl GetForwardedHeaderLayer {
114 #[inline]
115 #[must_use]
117 pub fn forwarded() -> Self {
118 Self::new()
119 }
120}
121
122impl GetForwardedHeaderLayer<Via> {
123 #[inline]
124 #[must_use]
126 pub fn via() -> Self {
127 Self::new()
128 }
129}
130
131impl GetForwardedHeaderLayer<XForwardedFor> {
132 #[inline]
133 #[must_use]
135 pub fn x_forwarded_for() -> Self {
136 Self::new()
137 }
138}
139
140impl GetForwardedHeaderLayer<XForwardedHost> {
141 #[inline]
142 #[must_use]
144 pub fn x_forwarded_host() -> Self {
145 Self::new()
146 }
147}
148
149impl GetForwardedHeaderLayer<XForwardedProto> {
150 #[inline]
151 #[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
169pub 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 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 pub fn forwarded(inner: S) -> Self {
209 Self::new(inner)
210 }
211}
212
213impl<S> GetForwardedHeaderService<S, Via> {
214 #[inline]
215 pub fn via(inner: S) -> Self {
217 Self::new(inner)
218 }
219}
220
221impl<S> GetForwardedHeaderService<S, XForwardedFor> {
222 #[inline]
223 pub fn x_forwarded_for(inner: S) -> Self {
225 Self::new(inner)
226 }
227}
228
229impl<S> GetForwardedHeaderService<S, XForwardedHost> {
230 #[inline]
231 pub fn x_forwarded_host(inner: S) -> Self {
233 Self::new(inner)
234 }
235}
236
237impl<S> GetForwardedHeaderService<S, XForwardedProto> {
238 #[inline]
239 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}