rama_http/layer/forwarded/
get_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::forwarded::ForwardHeader;
8use rama_core::{Layer, Service, extensions::ExtensionsRef};
9use rama_http_headers::HeaderMapExt;
10use rama_net::forwarded::Forwarded;
11use rama_net::forwarded::ForwardedElement;
12use rama_utils::macros::all_the_tuples_no_last_special_case;
13use std::fmt;
14use std::marker::PhantomData;
15
16pub struct GetForwardedHeadersLayer<T = Forwarded> {
48 _headers: PhantomData<fn() -> T>,
49}
50
51impl<T: fmt::Debug> fmt::Debug for GetForwardedHeadersLayer<T> {
52 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
53 f.debug_struct("GetForwardedHeadersLayer")
54 .field(
55 "_headers",
56 &format_args!("{}", std::any::type_name::<fn() -> T>()),
57 )
58 .finish()
59 }
60}
61
62impl<T: Clone> Clone for GetForwardedHeadersLayer<T> {
63 fn clone(&self) -> Self {
64 Self {
65 _headers: PhantomData,
66 }
67 }
68}
69
70impl<T> Default for GetForwardedHeadersLayer<T> {
71 #[inline]
72 fn default() -> Self {
73 Self::new()
74 }
75}
76
77impl<T> GetForwardedHeadersLayer<T> {
78 #[must_use]
80 pub const fn new() -> Self {
81 Self {
82 _headers: PhantomData,
83 }
84 }
85}
86
87impl<H, S> Layer<S> for GetForwardedHeadersLayer<H> {
88 type Service = GetForwardedHeadersService<S, H>;
89
90 fn layer(&self, inner: S) -> Self::Service {
91 Self::Service {
92 inner,
93 _headers: PhantomData,
94 }
95 }
96}
97
98pub struct GetForwardedHeadersService<S, T = Forwarded> {
102 inner: S,
103 _headers: PhantomData<fn() -> T>,
104}
105
106impl<S: fmt::Debug, T> fmt::Debug for GetForwardedHeadersService<S, T> {
107 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
108 f.debug_struct("GetForwardedHeadersService")
109 .field("inner", &self.inner)
110 .field("_headers", &format_args!("{}", std::any::type_name::<T>()))
111 .finish()
112 }
113}
114
115impl<S: Clone, T> Clone for GetForwardedHeadersService<S, T> {
116 fn clone(&self) -> Self {
117 Self {
118 inner: self.inner.clone(),
119 _headers: PhantomData,
120 }
121 }
122}
123
124impl<S, T> GetForwardedHeadersService<S, T> {
125 pub const fn new(inner: S) -> Self {
127 Self {
128 inner,
129 _headers: PhantomData,
130 }
131 }
132}
133
134macro_rules! get_forwarded_service_for_tuple {
135 ( $($ty:ident),* $(,)? ) => {
136 #[allow(non_snake_case)]
137 impl<$($ty,)* S, Body> Service<Request<Body>> for GetForwardedHeadersService<S, ($($ty,)*)>
138 where
139 $( $ty: ForwardHeader + Send + Sync + 'static, )*
140 S: Service<Request<Body>>,
141 Body: Send + 'static,
142 {
143 type Output = S::Output;
144 type Error = S::Error;
145
146 fn serve(
147 &self,
148 req: Request<Body>,
149 ) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send + '_ {
150 let mut forwarded_elements: Vec<ForwardedElement> = Vec::with_capacity(1);
151
152 $(
153 if let Some($ty) = req.headers().typed_get::<$ty>() {
154 let mut iter = $ty.into_iter();
155 for element in forwarded_elements.iter_mut() {
156 let other = iter.next();
157 match other {
158 Some(other) => {
159 element.merge(other);
160 }
161 None => break,
162 }
163 }
164 for other in iter {
165 forwarded_elements.push(other);
166 }
167 }
168 )*
169
170 if !forwarded_elements.is_empty() {
171 let forwarded = match req.extensions().get_ref::<Forwarded>().cloned() {
172 Some(mut forwarded) => {
173 forwarded.extend(forwarded_elements);
174 forwarded
175 }
176 None => {
177 let mut it = forwarded_elements.into_iter();
178 let mut forwarded = Forwarded::new(it.next().unwrap());
179 forwarded.extend(it);
180 forwarded
181 }
182 };
183 req.extensions().insert(forwarded);
184 }
185
186 self.inner.serve(req)
187 }
188 }
189 }
190}
191
192all_the_tuples_no_last_special_case!(get_forwarded_service_for_tuple);
193
194#[cfg(test)]
195mod tests {
196 use super::*;
197 use crate::{
198 Response, StatusCode,
199 headers::forwarded::{ClientIp, TrueClientIp, XClientIp},
200 service::web::response::IntoResponse,
201 };
202 use rama_core::{Layer, error::BoxError, extensions::ExtensionsRef, service::service_fn};
203 use rama_net::forwarded::ForwardedProtocol;
204 use std::{convert::Infallible, net::IpAddr};
205
206 fn assert_is_service<T: Service<Request<()>>>(_: T) {}
207
208 async fn dummy_service_fn() -> Result<Response, BoxError> {
209 Ok(StatusCode::OK.into_response())
210 }
211
212 #[test]
213 fn test_get_forwarded_service_is_service() {
214 assert_is_service(GetForwardedHeadersService::<_, (TrueClientIp,)>::new(
215 service_fn(dummy_service_fn),
216 ));
217 assert_is_service(
218 GetForwardedHeadersService::<_, (TrueClientIp, XClientIp)>::new(service_fn(
219 dummy_service_fn,
220 )),
221 );
222 assert_is_service(
223 GetForwardedHeadersLayer::<(ClientIp, TrueClientIp)>::new()
224 .into_layer(service_fn(dummy_service_fn)),
225 );
226 }
227
228 #[tokio::test]
229 async fn test_get_forwarded_headers() {
230 let service = GetForwardedHeadersLayer::<(rama_http_headers::forwarded::Forwarded,)>::new()
231 .into_layer(service_fn(async |req: Request<()>| {
232 let forwarded = req.extensions().get_ref::<Forwarded>().unwrap();
233 assert_eq!(forwarded.client_ip(), Some(IpAddr::from([12, 23, 34, 45])));
234 assert_eq!(forwarded.client_proto(), Some(ForwardedProtocol::HTTP));
235 Ok::<_, Infallible>(())
236 }));
237
238 let req = Request::builder()
239 .header("Forwarded", "for=\"12.23.34.45:5000\";proto=http")
240 .body(())
241 .unwrap();
242
243 service.serve(req).await.unwrap();
244 }
245}