use serde::{Deserialize, Serialize};
use sz_orm_flamegraph::QueryPhaseTiming;
use sz_orm_masking::{DataMasker, MaskingRule};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SerializablePhaseTiming {
pub phase: String,
pub start_ms: u64,
pub duration_ms: u64,
}
impl From<&QueryPhaseTiming> for SerializablePhaseTiming {
fn from(t: &QueryPhaseTiming) -> Self {
Self {
phase: t.phase.as_str().to_string(),
start_ms: t.start_ms,
duration_ms: t.duration_ms,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum LogLevel {
Debug,
Info,
Warn,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct QueryLogEntry {
pub query_key: String,
pub sql: String,
pub params: Vec<String>,
pub total_elapsed_ms: u64,
pub phase_breakdown: Vec<SerializablePhaseTiming>,
pub slow: bool,
pub from_cache: bool,
pub timestamp: String,
}
impl QueryLogEntry {
pub fn to_json(&self) -> String {
serde_json::to_string(self).unwrap_or_else(|e| format!("{{\"error\":\"{e}\"}}"))
}
}
pub struct QueryLogger {
sample_rate: f64,
level: LogLevel,
counter: std::sync::atomic::AtomicU64,
masking_rules: Vec<MaskingRule>,
}
impl Default for QueryLogger {
fn default() -> Self {
Self::new()
}
}
impl QueryLogger {
pub fn new() -> Self {
Self {
sample_rate: 0.01,
level: LogLevel::Info,
counter: std::sync::atomic::AtomicU64::new(0),
masking_rules: vec![
MaskingRule::Phone,
MaskingRule::Email,
MaskingRule::IdCard,
MaskingRule::BankCard,
MaskingRule::Password,
MaskingRule::ApiKey,
],
}
}
pub fn with_sample_rate(mut self, rate: f64) -> Self {
self.sample_rate = rate.clamp(0.0, 1.0);
self
}
pub fn with_level(mut self, level: LogLevel) -> Self {
self.level = level;
self
}
pub fn with_masking_rules(mut self, rules: Vec<MaskingRule>) -> Self {
self.masking_rules = rules;
self
}
pub fn log(&self, mut entry: QueryLogEntry) -> Option<String> {
if self.level == LogLevel::Warn && !entry.slow {
return None;
}
if !entry.slow {
let count = self
.counter
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let threshold = (self.sample_rate * 1000.0) as u64;
let rand_val = count % 1000;
if rand_val >= threshold {
return None;
}
}
entry.params = mask_params(&entry.params, &self.masking_rules);
if self.level == LogLevel::Info {
entry.sql.clear();
entry.params.clear();
}
Some(entry.to_json())
}
}
pub fn mask_params(params: &[String], rules: &[MaskingRule]) -> Vec<String> {
params
.iter()
.map(|p| {
for rule in rules {
let masked = DataMasker::apply(rule, p);
if masked != *p && masked != "***" {
return masked;
}
}
p.clone()
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use sz_orm_flamegraph::Phase;
fn timing(phase: Phase, ms: u64) -> SerializablePhaseTiming {
SerializablePhaseTiming {
phase: phase.as_str().to_string(),
start_ms: 0,
duration_ms: ms,
}
}
fn entry(slow: bool) -> QueryLogEntry {
QueryLogEntry {
query_key: "q1".into(),
sql: "SELECT * FROM users WHERE phone = ?".into(),
params: vec!["13800138000".into()],
total_elapsed_ms: 150,
phase_breakdown: vec![timing(Phase::SqlExecute, 150)],
slow,
from_cache: false,
timestamp: "2026-08-12T10:00:00Z".into(),
}
}
#[test]
fn log_entry_json_contains_all_fields() {
let e = entry(false);
let json = e.to_json();
assert!(json.contains("query_key"));
assert!(json.contains("sql"));
assert!(json.contains("params"));
assert!(json.contains("total_elapsed_ms"));
assert!(json.contains("phase_breakdown"));
assert!(json.contains("slow"));
assert!(json.contains("from_cache"));
assert!(json.contains("timestamp"));
}
#[test]
fn log_entry_json_roundtrip() {
let e = entry(true);
let json = e.to_json();
let back: QueryLogEntry = serde_json::from_str(&json).unwrap();
assert_eq!(e, back);
}
#[test]
fn slow_query_always_sampled() {
let logger = QueryLogger::new().with_sample_rate(0.0);
let e = entry(true);
assert!(logger.log(e).is_some());
}
#[test]
fn sample_rate_zero_no_non_slow() {
let logger = QueryLogger::new().with_sample_rate(0.0);
for _ in 0..100 {
assert!(logger.log(entry(false)).is_none());
}
}
#[test]
fn sample_rate_one_all_sampled() {
let logger = QueryLogger::new().with_sample_rate(1.0);
for _ in 0..100 {
assert!(logger.log(entry(false)).is_some());
}
}
#[test]
fn warn_level_only_slow_queries() {
let logger = QueryLogger::new().with_level(LogLevel::Warn);
assert!(logger.log(entry(false)).is_none());
assert!(logger.log(entry(true)).is_some());
}
#[test]
fn info_level_strips_sql_and_params() {
let logger = QueryLogger::new()
.with_level(LogLevel::Info)
.with_sample_rate(1.0);
let json = logger.log(entry(false)).unwrap();
let parsed: QueryLogEntry = serde_json::from_str(&json).unwrap();
assert!(parsed.sql.is_empty());
assert!(parsed.params.is_empty());
}
#[test]
fn debug_level_keeps_sql_and_params() {
let logger = QueryLogger::new()
.with_level(LogLevel::Debug)
.with_sample_rate(1.0);
let json = logger.log(entry(false)).unwrap();
let parsed: QueryLogEntry = serde_json::from_str(&json).unwrap();
assert!(!parsed.sql.is_empty());
}
#[test]
fn phone_param_masked() {
let params = vec!["13800138000".to_string()];
let rules = vec![MaskingRule::Phone];
let masked = mask_params(¶ms, &rules);
assert!(masked[0].contains('*'));
assert!(!masked[0].contains("13800138000"));
}
#[test]
fn email_param_masked() {
let params = vec!["user@example.com".to_string()];
let rules = vec![MaskingRule::Email];
let masked = mask_params(¶ms, &rules);
assert!(masked[0].contains('*'));
}
#[test]
fn non_sensitive_param_not_masked() {
let params = vec!["42".to_string()];
let rules = vec![MaskingRule::Phone, MaskingRule::Email];
let masked = mask_params(¶ms, &rules);
assert_eq!(masked[0], "42");
}
#[test]
fn empty_phase_breakdown_still_logs() {
let mut e = entry(true);
e.phase_breakdown.clear();
let logger = QueryLogger::new();
assert!(logger.log(e).is_some());
}
}