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