Skip to main content

actix_cloud/
request.rs

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    /// Request start time.
14    pub start_time: DateTime<Utc>,
15
16    #[cfg(feature = "i18n")]
17    /// Request language.
18    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}