actix-cloud 0.6.3

Actix Cloud is an all-in-one web framework based on Actix Web.
Documentation
//! Provide per-request extension data.
//!
//! Wrap the app with [`Middleware`] and handlers can read the [`Extension`] from request
//! extensions (`ReqData<Arc<Extension>>` or `req.extensions()`): request start time,
//! identified language, trace id and client real IP.
use std::{net::SocketAddr, rc::Rc, sync::Arc};

use actix_web::{
    dev::{forward_ready, Service, ServiceRequest, ServiceResponse, Transform},
    HttpMessage as _,
};
use chrono::{DateTime, Utc};
use futures::future::{ready, LocalBoxFuture, Ready};

#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[derive(Debug, Clone)]
pub struct Extension {
    /// Request start time.
    pub start_time: DateTime<Utc>,

    #[cfg(feature = "i18n")]
    /// Request language.
    pub lang: String,

    #[cfg(feature = "traceid")]
    /// Trace ID generated by `TracingLogger`.
    ///
    /// Wrap this middleware **before** `TracingLogger` so the id exists when it runs.
    pub trace_id: String,

    /// Client real IP, resolved through [`Middleware::real_ip`].
    pub real_ip: SocketAddr,
}

pub type RealIPFunc = Rc<dyn Fn(&ServiceRequest) -> SocketAddr>;
pub type LangFunc = Rc<dyn Fn(&ServiceRequest) -> Option<String>>;

/// Middleware inserting [`Extension`] into the request extensions.
pub struct Middleware {
    real_ip: RealIPFunc,
    #[cfg(feature = "traceid")]
    trace_header: Rc<Option<String>>,
    #[cfg(feature = "i18n")]
    lang: LangFunc,
}

impl Default for Middleware {
    fn default() -> Self {
        Self::new()
    }
}

impl Middleware {
    fn default_real_ip(req: &ServiceRequest) -> SocketAddr {
        // Peer address can be absent (e.g. unix sockets, tests); fall back to
        // the unspecified address instead of panicking.
        req.peer_addr()
            .unwrap_or_else(|| SocketAddr::new(std::net::Ipv4Addr::UNSPECIFIED.into(), 0))
    }

    #[cfg(feature = "i18n")]
    fn default_lang(_: &ServiceRequest) -> Option<String> {
        None
    }

    pub fn new() -> Self {
        Self {
            real_ip: Rc::new(Self::default_real_ip),
            #[cfg(feature = "traceid")]
            trace_header: Rc::new(None),
            #[cfg(feature = "i18n")]
            lang: Rc::new(Self::default_lang),
        }
    }

    #[cfg(feature = "traceid")]
    /// Set the response header name carrying the trace id (e.g. `X-Trace-Id`).
    pub fn trace_header<S>(mut self, s: S) -> Self
    where
        S: Into<String>,
    {
        self.trace_header = Rc::new(Some(s.into()));
        self
    }

    /// Set the callback to resolve the client real IP (e.g. reading `X-Real-IP`).
    ///
    /// Defaults to the peer address of the connection.
    pub fn real_ip<F>(mut self, f: F) -> Self
    where
        F: Fn(&ServiceRequest) -> SocketAddr + 'static,
    {
        self.real_ip = Rc::new(f);
        self
    }

    #[cfg(feature = "i18n")]
    /// Set the callback to identify the request language.
    ///
    /// Returning `None` falls back to `locale.default` in `GlobalState`.
    pub fn lang<F>(mut self, f: F) -> Self
    where
        F: Fn(&ServiceRequest) -> Option<String> + 'static,
    {
        self.lang = Rc::new(f);
        self
    }
}

impl<S, B> Transform<S, ServiceRequest> for Middleware
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error>,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = actix_web::Error;
    type InitError = ();
    type Transform = MiddlewareService<S>;
    type Future = Ready<Result<Self::Transform, Self::InitError>>;

    fn new_transform(&self, service: S) -> Self::Future {
        ready(Ok(MiddlewareService {
            service: Rc::new(service),
            real_ip: self.real_ip.clone(),
            #[cfg(feature = "traceid")]
            trace_header: self.trace_header.clone(),
            #[cfg(feature = "i18n")]
            lang: self.lang.clone(),
        }))
    }
}

pub struct MiddlewareService<S> {
    service: Rc<S>,
    real_ip: RealIPFunc,
    #[cfg(feature = "traceid")]
    trace_header: Rc<Option<String>>,
    #[cfg(feature = "i18n")]
    lang: LangFunc,
}

impl<S, B> Service<ServiceRequest> for MiddlewareService<S>
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error>,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = actix_web::Error;
    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;

    forward_ready!(service);

    fn call(&self, req: ServiceRequest) -> Self::Future {
        // Graceful fallbacks: no panic when `GlobalState` is not registered or
        // `TracingLogger` was not wrapped before this middleware.
        #[cfg(feature = "i18n")]
        let lang = (self.lang)(&req).unwrap_or_else(|| {
            req.app_data::<actix_web::web::Data<crate::state::GlobalState>>()
                .map(|state| state.locale.default.clone())
                .unwrap_or_else(|| String::from("en-US"))
        });
        #[cfg(feature = "traceid")]
        let trace_id = req
            .extensions()
            .get::<tracing_actix_web::RequestId>()
            .map(ToString::to_string)
            .unwrap_or_default();
        let ext = Extension {
            start_time: Utc::now(),
            #[cfg(feature = "i18n")]
            lang,
            #[cfg(feature = "traceid")]
            trace_id: trace_id.clone(),
            real_ip: (self.real_ip)(&req),
        };
        #[cfg(feature = "traceid")]
        let header = self.trace_header.clone();
        req.extensions_mut().insert(Arc::new(ext));

        #[cfg(not(feature = "traceid"))]
        return Box::pin(self.service.call(req));
        #[cfg(feature = "traceid")]
        {
            use futures::FutureExt;
            use std::str::FromStr;
            Box::pin(self.service.call(req).map(move |x| {
                if let Some(header) = header.as_ref() {
                    x.map(|mut x| {
                        x.headers_mut().insert(
                            actix_web::http::header::HeaderName::from_str(header).unwrap(),
                            actix_web::http::header::HeaderValue::from_str(&trace_id).unwrap(),
                        );
                        x
                    })
                } else {
                    x
                }
            }))
        }
    }
}