vacuna 0.4.1

Simple web server for static files
Documentation
use std::{
    env, fmt,
    fmt::Display,
    future::{Ready, ready},
    ops::RangeBounds,
    sync::Arc,
};

use crate::config::{Log, LogLevel};
use actix_web::{
    Error, HttpRequest,
    dev::{self, Service, ServiceRequest, ServiceResponse, Transform},
    http::header::HeaderName,
};
use futures_util::future::LocalBoxFuture;
use log::{debug, error, info, trace, warn};
use regex::Regex;

use time::{OffsetDateTime, format_description::well_known::Rfc3339};

pub struct CustomLog {
    log: Arc<Vec<Log>>,
}
impl CustomLog {
    pub fn new(log: &Vec<Log>) -> Self {
        Self {
            log: Arc::new(log.clone()),
        }
    }
}

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

    fn new_transform(&self, service: S) -> Self::Future {
        ready(Ok(CustomLogMiddleware::new(service, &self.log)))
    }
}

pub struct CustomLogMiddleware<S> {
    service: S,
    log: Arc<Vec<Log>>,
}

impl<S> CustomLogMiddleware<S> {
    fn new(service: S, log: &Arc<Vec<Log>>) -> Self {
        Self {
            service,
            log: log.clone(),
        }
    }
}

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

    dev::forward_ready!(service);

    fn call(&self, req: ServiceRequest) -> Self::Future {
        let fut = self.service.call(req);

        let lg = self.log.clone();
        Box::pin(async move {
            let res = fut.await?;
            for l in lg.iter() {
                if l.status_range().contains(&res.status().as_u16()) {
                    let now = OffsetDateTime::now_utc();
                    let mut fmt = Format::new(l.format());
                    for item in &mut fmt.0 {
                        item.render_request(now, res.request());
                    }
                    for item in &mut fmt.0 {
                        item.render_response(&res);
                    }
                    let formatted = fmt;
                    match l.level() {
                        LogLevel::Trace => trace!("{}", formatted),
                        LogLevel::Debug => debug!("{}", formatted),
                        LogLevel::Info => info!("{}", formatted),
                        LogLevel::Warn => warn!("{}", formatted),
                        LogLevel::Error => error!("{}", formatted),
                    }
                }
            }
            Ok(res)
        })
    }
}

impl Display for Format {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        for item in &self.0 {
            item.render(f)?;
        }
        Ok(())
    }
}

/// A formatting style for the `Logger` consisting of multiple concatenated `FormatText` items.
#[derive(Debug, Clone)]
struct Format(Vec<FormatText>);

impl Default for Format {
    /// Return the default formatting style for the `Logger`:
    fn default() -> Format {
        Format::new(r#"%a "%r" %s %b "%{Referer}i" "%{User-Agent}i" %T"#)
    }
}

impl Format {
    /// Create a `Format` from a format string.
    ///
    /// Returns `None` if the format string syntax is incorrect.
    pub fn new(s: &str) -> Format {
        log::trace!("Access log format: {}", s);
        let fmt = Regex::new(r"%(\{([A-Za-z0-9\-_]+)\}([aioe]|x[io])|[%atPrUsbTD]?)").unwrap();

        let mut idx = 0;
        let mut results = Vec::new();
        for cap in fmt.captures_iter(s) {
            let m = cap.get(0).unwrap();
            let pos = m.start();
            if idx != pos {
                results.push(FormatText::Str(s[idx..pos].to_owned()));
            }
            idx = m.end();

            if let Some(key) = cap.get(2) {
                results.push(match cap.get(3).unwrap().as_str() {
                    "a" => {
                        if key.as_str() == "r" {
                            FormatText::RealIpRemoteAddr
                        } else {
                            unreachable!("regex and code mismatch")
                        }
                    }
                    "i" => FormatText::RequestHeader(HeaderName::try_from(key.as_str()).unwrap()),
                    "o" => FormatText::ResponseHeader(HeaderName::try_from(key.as_str()).unwrap()),
                    "e" => FormatText::EnvironHeader(key.as_str().to_owned()),
                    _ => unreachable!(),
                })
            } else {
                let m = cap.get(1).unwrap();
                results.push(match m.as_str() {
                    "%" => FormatText::Percent,
                    "a" => FormatText::RemoteAddr,
                    "t" => FormatText::RequestTime,
                    "r" => FormatText::RequestLine,
                    "s" => FormatText::ResponseStatus,
                    "b" => FormatText::ResponseSize,
                    "U" => FormatText::UrlPath,
                    "T" => FormatText::Time,
                    "D" => FormatText::TimeMillis,
                    _ => FormatText::Str(m.as_str().to_owned()),
                });
            }
        }
        if idx != s.len() {
            results.push(FormatText::Str(s[idx..].to_owned()));
        }

        Format(results)
    }
}

/// A string of text to be logged.
///
/// This is either one of the data fields supported by the `Logger`, or a custom `String`.
#[non_exhaustive]
#[derive(Debug, Clone)]
enum FormatText {
    Str(String),
    Percent,
    RequestLine,
    RequestTime,
    ResponseStatus,
    ResponseSize,
    Time,
    TimeMillis,
    RemoteAddr,
    RealIpRemoteAddr,
    UrlPath,
    RequestHeader(HeaderName),
    ResponseHeader(HeaderName),
    EnvironHeader(String),
}

impl FormatText {
    fn render(&self, fmt: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
        match self {
            FormatText::Str(string) => fmt.write_str(string),
            FormatText::Percent => "%".fmt(fmt),
            FormatText::EnvironHeader(name) => {
                if let Ok(val) = env::var(name) {
                    fmt.write_fmt(format_args!("{}", val))
                } else {
                    "-".fmt(fmt)
                }
            }
            _ => Ok(()),
        }
    }

    fn render_response(&mut self, res: &ServiceResponse) {
        match self {
            FormatText::ResponseStatus => *self = FormatText::Str(format!("{}", res.status().as_u16())),

            FormatText::ResponseHeader(name) => {
                let s = res.headers().get(&*name).and_then(|x| x.to_str().ok()).unwrap_or("-");
                *self = FormatText::Str(s.to_string())
            }

            _ => {}
        }
    }

    fn render_request(&mut self, now: OffsetDateTime, req: &HttpRequest) {
        match self {
            FormatText::RequestLine => {
                *self = if req.query_string().is_empty() {
                    FormatText::Str(format!("{} {} {:?}", req.method(), req.path(), req.version()))
                } else {
                    FormatText::Str(format!(
                        "{} {}?{} {:?}",
                        req.method(),
                        req.path(),
                        req.query_string(),
                        req.version()
                    ))
                };
            }
            FormatText::UrlPath => *self = FormatText::Str(req.path().to_string()),
            FormatText::RequestTime => *self = FormatText::Str(now.format(&Rfc3339).unwrap()),
            FormatText::RequestHeader(name) => {
                let s = req
                    .headers()
                    .get(&*name)
                    .and_then(|val| val.to_str().ok())
                    .unwrap_or("-");
                *self = FormatText::Str(s.to_string());
            }
            FormatText::RemoteAddr => {
                let s = if let Some(peer) = req.connection_info().peer_addr() {
                    FormatText::Str((*peer).to_string())
                } else {
                    FormatText::Str("-".to_string())
                };
                *self = s;
            }
            FormatText::RealIpRemoteAddr => {
                let s = if let Some(remote) = req.connection_info().realip_remote_addr() {
                    FormatText::Str(remote.to_string())
                } else {
                    FormatText::Str("-".to_string())
                };
                *self = s;
            }
            _ => {}
        }
    }
}