Skip to main content

actix_cloud/
request.rs

1//! Provide per-request extension data.
2//!
3//! Wrap the app with [`Middleware`] and handlers can read the [`Extension`] from request
4//! extensions (`ReqData<Arc<Extension>>` or `req.extensions()`): request start time,
5//! identified language, trace id and client real IP.
6use 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    /// Request start time.
19    pub start_time: DateTime<Utc>,
20
21    #[cfg(feature = "i18n")]
22    /// Request language.
23    pub lang: String,
24
25    #[cfg(feature = "traceid")]
26    /// Trace ID generated by `TracingLogger`.
27    ///
28    /// Wrap this middleware **before** `TracingLogger` so the id exists when it runs.
29    pub trace_id: String,
30
31    /// Client real IP, resolved through [`Middleware::real_ip`].
32    pub real_ip: SocketAddr,
33}
34
35pub type RealIPFunc = Rc<dyn Fn(&ServiceRequest) -> SocketAddr>;
36pub type LangFunc = Rc<dyn Fn(&ServiceRequest) -> Option<String>>;
37
38/// Middleware inserting [`Extension`] into the request extensions.
39pub 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        // Peer address can be absent (e.g. unix sockets, tests); fall back to
56        // the unspecified address instead of panicking.
57        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    /// Set the response header name carrying the trace id (e.g. `X-Trace-Id`).
78    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    /// Set the callback to resolve the client real IP (e.g. reading `X-Real-IP`).
87    ///
88    /// Defaults to the peer address of the connection.
89    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    /// Set the callback to identify the request language.
99    ///
100    /// Returning `None` falls back to `locale.default` in `GlobalState`.
101    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        // Graceful fallbacks: no panic when `GlobalState` is not registered or
157        // `TracingLogger` was not wrapped before this middleware.
158        #[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}