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(())
}
}
#[derive(Debug, Clone)]
struct Format(Vec<FormatText>);
impl Default for Format {
fn default() -> Format {
Format::new(r#"%a "%r" %s %b "%{Referer}i" "%{User-Agent}i" %T"#)
}
}
impl Format {
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)
}
}
#[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;
}
_ => {}
}
}
}