1use std::{net::SocketAddr, rc::Rc, sync::Arc};
7
8use actix_web::{
9 dev::{forward_ready, Service, ServiceRequest, ServiceResponse, Transform},
10 HttpMessage as _,
11};
12use chrono::{DateTime, Utc};
13use futures::future::{ready, LocalBoxFuture, Ready};
14
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16#[derive(Debug, Clone)]
17pub struct Extension {
18 pub start_time: DateTime<Utc>,
20
21 #[cfg(feature = "i18n")]
22 pub lang: String,
24
25 #[cfg(feature = "traceid")]
26 pub trace_id: String,
30
31 pub real_ip: SocketAddr,
33}
34
35pub type RealIPFunc = Rc<dyn Fn(&ServiceRequest) -> SocketAddr>;
36pub type LangFunc = Rc<dyn Fn(&ServiceRequest) -> Option<String>>;
37
38pub struct Middleware {
40 real_ip: RealIPFunc,
41 #[cfg(feature = "traceid")]
42 trace_header: Rc<Option<String>>,
43 #[cfg(feature = "i18n")]
44 lang: LangFunc,
45}
46
47impl Default for Middleware {
48 fn default() -> Self {
49 Self::new()
50 }
51}
52
53impl Middleware {
54 fn default_real_ip(req: &ServiceRequest) -> SocketAddr {
55 req.peer_addr()
58 .unwrap_or_else(|| SocketAddr::new(std::net::Ipv4Addr::UNSPECIFIED.into(), 0))
59 }
60
61 #[cfg(feature = "i18n")]
62 fn default_lang(_: &ServiceRequest) -> Option<String> {
63 None
64 }
65
66 pub fn new() -> Self {
67 Self {
68 real_ip: Rc::new(Self::default_real_ip),
69 #[cfg(feature = "traceid")]
70 trace_header: Rc::new(None),
71 #[cfg(feature = "i18n")]
72 lang: Rc::new(Self::default_lang),
73 }
74 }
75
76 #[cfg(feature = "traceid")]
77 pub fn trace_header<S>(mut self, s: S) -> Self
79 where
80 S: Into<String>,
81 {
82 self.trace_header = Rc::new(Some(s.into()));
83 self
84 }
85
86 pub fn real_ip<F>(mut self, f: F) -> Self
90 where
91 F: Fn(&ServiceRequest) -> SocketAddr + 'static,
92 {
93 self.real_ip = Rc::new(f);
94 self
95 }
96
97 #[cfg(feature = "i18n")]
98 pub fn lang<F>(mut self, f: F) -> Self
102 where
103 F: Fn(&ServiceRequest) -> Option<String> + 'static,
104 {
105 self.lang = Rc::new(f);
106 self
107 }
108}
109
110impl<S, B> Transform<S, ServiceRequest> for Middleware
111where
112 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error>,
113 S::Future: 'static,
114 B: 'static,
115{
116 type Response = ServiceResponse<B>;
117 type Error = actix_web::Error;
118 type InitError = ();
119 type Transform = MiddlewareService<S>;
120 type Future = Ready<Result<Self::Transform, Self::InitError>>;
121
122 fn new_transform(&self, service: S) -> Self::Future {
123 ready(Ok(MiddlewareService {
124 service: Rc::new(service),
125 real_ip: self.real_ip.clone(),
126 #[cfg(feature = "traceid")]
127 trace_header: self.trace_header.clone(),
128 #[cfg(feature = "i18n")]
129 lang: self.lang.clone(),
130 }))
131 }
132}
133
134pub struct MiddlewareService<S> {
135 service: Rc<S>,
136 real_ip: RealIPFunc,
137 #[cfg(feature = "traceid")]
138 trace_header: Rc<Option<String>>,
139 #[cfg(feature = "i18n")]
140 lang: LangFunc,
141}
142
143impl<S, B> Service<ServiceRequest> for MiddlewareService<S>
144where
145 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error>,
146 S::Future: 'static,
147 B: 'static,
148{
149 type Response = ServiceResponse<B>;
150 type Error = actix_web::Error;
151 type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
152
153 forward_ready!(service);
154
155 fn call(&self, req: ServiceRequest) -> Self::Future {
156 #[cfg(feature = "i18n")]
159 let lang = (self.lang)(&req).unwrap_or_else(|| {
160 req.app_data::<actix_web::web::Data<crate::state::GlobalState>>()
161 .map(|state| state.locale.default.clone())
162 .unwrap_or_else(|| String::from("en-US"))
163 });
164 #[cfg(feature = "traceid")]
165 let trace_id = req
166 .extensions()
167 .get::<tracing_actix_web::RequestId>()
168 .map(ToString::to_string)
169 .unwrap_or_default();
170 let ext = Extension {
171 start_time: Utc::now(),
172 #[cfg(feature = "i18n")]
173 lang,
174 #[cfg(feature = "traceid")]
175 trace_id: trace_id.clone(),
176 real_ip: (self.real_ip)(&req),
177 };
178 #[cfg(feature = "traceid")]
179 let header = self.trace_header.clone();
180 req.extensions_mut().insert(Arc::new(ext));
181
182 #[cfg(not(feature = "traceid"))]
183 return Box::pin(self.service.call(req));
184 #[cfg(feature = "traceid")]
185 {
186 use futures::FutureExt;
187 use std::str::FromStr;
188 Box::pin(self.service.call(req).map(move |x| {
189 if let Some(header) = header.as_ref() {
190 x.map(|mut x| {
191 x.headers_mut().insert(
192 actix_web::http::header::HeaderName::from_str(header).unwrap(),
193 actix_web::http::header::HeaderValue::from_str(&trace_id).unwrap(),
194 );
195 x
196 })
197 } else {
198 x
199 }
200 }))
201 }
202 }
203}