use serde::{Deserialize, Serialize};
use std::fmt;
use std::time::Instant;
use axum::extract::Request;
use axum::middleware::Next;
use axum::response::Response;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
pub enum AuditOutputTarget {
Tracing,
File,
Stdout,
}
impl Default for AuditOutputTarget {
fn default() -> Self {
Self::Tracing
}
}
impl fmt::Display for AuditOutputTarget {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Tracing => write!(f, "tracing"),
Self::File => write!(f, "file"),
Self::Stdout => write!(f, "stdout"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
pub enum AuditEventType {
HttpRequest,
AuthSuccess,
AuthFailure,
AccessDenied,
IpRejected,
BodyTooLarge,
}
impl AuditEventType {
pub fn is_security_event(self) -> bool {
matches!(
self,
Self::AuthFailure | Self::AccessDenied | Self::IpRejected | Self::BodyTooLarge
)
}
}
#[derive(Debug, Clone, Serialize)]
pub struct AuditEvent {
pub timestamp: chrono::DateTime<chrono::Utc>,
pub event_type: AuditEventType,
pub user_id: Option<i64>,
pub client_ip: String,
pub method: String,
pub path: String,
pub status: u16,
pub duration_ms: u64,
pub headers: Option<serde_json::Value>,
pub body: Option<serde_json::Value>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct AuditLogConfig {
#[serde(default)]
pub enabled: bool,
#[serde(default = "default_sample_rate")]
pub sample_rate: f64,
#[serde(default)]
pub exclude_paths: Vec<String>,
#[serde(default = "default_sensitive_fields")]
pub sensitive_fields: Vec<String>,
#[serde(default)]
pub log_headers: bool,
#[serde(default)]
pub log_body: bool,
#[serde(default)]
pub output_target: AuditOutputTarget,
}
fn default_sample_rate() -> f64 {
1.0
}
fn default_sensitive_fields() -> Vec<String> {
vec![
"password".to_string(),
"token".to_string(),
"authorization".to_string(),
"credit_card".to_string(),
]
}
impl Default for AuditLogConfig {
fn default() -> Self {
Self {
enabled: false,
sample_rate: 1.0,
exclude_paths: Vec::new(),
sensitive_fields: default_sensitive_fields(),
log_headers: false,
log_body: false,
output_target: AuditOutputTarget::Tracing,
}
}
}
pub fn redact_sensitive_fields(value: &mut serde_json::Value, sensitive_fields: &[String]) {
match value {
serde_json::Value::Object(map) => {
for (key, val) in map.iter_mut() {
if sensitive_fields.contains(key) {
*val = serde_json::Value::String("[REDACTED]".to_string());
} else {
redact_sensitive_fields(val, sensitive_fields);
}
}
}
serde_json::Value::Array(arr) => {
for item in arr.iter_mut() {
redact_sensitive_fields(item, sensitive_fields);
}
}
_ => {}
}
}
pub fn should_sample(sample_rate: f64, is_security_event: bool) -> bool {
if is_security_event {
return true;
}
if sample_rate <= 0.0 {
return false;
}
if sample_rate >= 1.0 {
return true;
}
use rand::Rng;
rand::rngs::OsRng.gen::<f64>() < sample_rate
}
async fn write_audit_log(event: &AuditEvent, target: AuditOutputTarget) -> Result<(), std::io::Error> {
let json = serde_json::to_string(event).unwrap_or_default();
match target {
AuditOutputTarget::Tracing => {
tracing::info!("audit: {json}");
}
AuditOutputTarget::Stdout => {
println!("{json}");
}
AuditOutputTarget::File => {
use tokio::io::AsyncWriteExt;
let mut file = tokio::fs::OpenOptions::new()
.append(true)
.create(true)
.open("audit.log")
.await?;
file.write_all(json.as_bytes()).await?;
file.write_all(b"\n").await?;
}
}
Ok(())
}
fn infer_event_type(status: u16) -> AuditEventType {
match status {
401 => AuditEventType::AuthFailure,
403 => AuditEventType::AccessDenied,
413 => AuditEventType::BodyTooLarge,
_ => AuditEventType::HttpRequest,
}
}
pub async fn audit_log_middleware(
axum::extract::State(config): axum::extract::State<AuditLogConfig>,
req: Request,
next: Next,
) -> Response {
if !config.enabled {
return next.run(req).await;
}
let path = req.uri().path().to_string();
if config.exclude_paths.contains(&path) {
return next.run(req).await;
}
let method = req.method().to_string();
let client_ip = extract_client_ip_simple(req.headers());
let start = Instant::now();
let response = next.run(req).await;
let duration_ms = start.elapsed().as_millis() as u64;
let status = response.status().as_u16();
let event_type = infer_event_type(status);
if !should_sample(config.sample_rate, event_type.is_security_event()) {
return response;
}
let event = AuditEvent {
timestamp: chrono::Utc::now(),
event_type,
user_id: None,
client_ip,
method,
path,
status,
duration_ms,
headers: None,
body: None,
};
let target = config.output_target;
tokio::spawn(async move {
if let Err(e) = write_audit_log(&event, target).await {
tracing::error!("审计日志写入失败: {e}");
}
});
response
}
fn extract_client_ip_simple(headers: &axum::http::HeaderMap) -> String {
if let Some(forwarded) = headers.get("x-forwarded-for") {
if let Ok(value) = forwarded.to_str() {
if let Some(first) = value.split(',').next() {
let trimmed = first.trim();
if !trimmed.is_empty() {
return trimmed.to_string();
}
}
}
}
if let Some(real_ip) = headers.get("x-real-ip") {
if let Ok(value) = real_ip.to_str() {
let trimmed = value.trim();
if !trimmed.is_empty() {
return trimmed.to_string();
}
}
}
"unknown".to_string()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config_disabled() {
let cfg = AuditLogConfig::default();
assert!(!cfg.enabled);
assert_eq!(cfg.sample_rate, 1.0);
assert!(cfg.sensitive_fields.contains(&"password".to_string()));
assert_eq!(cfg.output_target, AuditOutputTarget::Tracing);
}
#[test]
fn test_audit_output_target_display() {
assert_eq!(AuditOutputTarget::Tracing.to_string(), "tracing");
assert_eq!(AuditOutputTarget::File.to_string(), "file");
assert_eq!(AuditOutputTarget::Stdout.to_string(), "stdout");
}
#[test]
fn test_event_type_is_security_event() {
assert!(!AuditEventType::HttpRequest.is_security_event());
assert!(!AuditEventType::AuthSuccess.is_security_event());
assert!(AuditEventType::AuthFailure.is_security_event());
assert!(AuditEventType::AccessDenied.is_security_event());
assert!(AuditEventType::IpRejected.is_security_event());
assert!(AuditEventType::BodyTooLarge.is_security_event());
}
#[test]
fn test_redact_simple() {
let mut value = serde_json::json!({"password": "secret123", "name": "alice"});
redact_sensitive_fields(&mut value, &["password".to_string()]);
assert_eq!(value["password"], "[REDACTED]");
assert_eq!(value["name"], "alice");
}
#[test]
fn test_redact_nested() {
let mut value = serde_json::json!({"user": {"token": "xxx", "name": "bob"}});
redact_sensitive_fields(&mut value, &["token".to_string()]);
assert_eq!(value["user"]["token"], "[REDACTED]");
assert_eq!(value["user"]["name"], "bob");
}
#[test]
fn test_redact_array() {
let mut value = serde_json::json!({"items": [{"password": "a"}, {"password": "b"}]});
redact_sensitive_fields(&mut value, &["password".to_string()]);
assert_eq!(value["items"][0]["password"], "[REDACTED]");
assert_eq!(value["items"][1]["password"], "[REDACTED]");
}
#[test]
fn test_should_sample_security_event_always() {
assert!(should_sample(0.0, true));
assert!(should_sample(1.0, true));
}
#[test]
fn test_should_sample_zero_rate() {
assert!(!should_sample(0.0, false));
}
#[test]
fn test_should_sample_full_rate() {
assert!(should_sample(1.0, false));
}
#[tokio::test]
async fn test_middleware_disabled_passes_through() {
use axum::routing::get;
use tower::ServiceExt;
let config = AuditLogConfig::default();
let app = axum::Router::new()
.route("/", get(|| async { "ok" }))
.layer(axum::middleware::from_fn_with_state(
config,
audit_log_middleware,
));
let resp = app
.oneshot(
axum::http::Request::builder()
.uri("/")
.body(axum::body::Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
}
}