use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use smallvec::SmallVec;
use std::fs;
use std::path::Path;
use std::sync::{
Mutex,
atomic::{AtomicU64, Ordering},
};
static SAVE_LOCK: Mutex<()> = Mutex::new(());
static SAVE_SEQ: AtomicU64 = AtomicU64::new(0);
const MAX_RETAINED_USER_METRICS: usize = 1024;
const STATS_FILE_VERSION: u32 = 1;
#[derive(Debug, Serialize, Deserialize)]
struct StatsFile {
version: u32,
saved_at: String, global: PersistedGlobal,
backends: SmallVec<[PersistedBackend; 8]>, users: SmallVec<[PersistedUser; 4]>, pipeline: PersistedPipeline,
}
#[derive(Debug, Serialize, Deserialize)]
pub(crate) struct PersistedGlobal {
total_connections: u64,
}
#[derive(Debug, Serialize, Deserialize)]
pub(crate) struct PersistedBackend {
name: String,
total_commands: u64,
bytes_sent: u64,
bytes_received: u64,
errors: u64,
errors_4xx: u64,
errors_5xx: u64,
article_bytes_total: u64,
article_count: u64,
ttfb_micros_total: u64,
ttfb_count: u64,
send_micros_total: u64,
recv_micros_total: u64,
connection_failures: u64,
}
#[derive(Debug, Serialize, Deserialize)]
pub(crate) struct PersistedUser {
username: String,
total_connections: u64,
bytes_sent: u64,
bytes_received: u64,
total_commands: u64,
errors: u64,
}
#[derive(Debug, Serialize, Deserialize)]
pub(crate) struct PersistedPipeline {
batches: u64,
commands: u64,
requests_queued: u64,
requests_completed: u64,
}
#[derive(Debug)]
pub struct BackendStore {
pub total_commands: AtomicU64,
pub bytes_sent: AtomicU64,
pub bytes_received: AtomicU64,
pub errors: AtomicU64,
pub errors_4xx: AtomicU64,
pub errors_5xx: AtomicU64,
pub article_bytes_total: AtomicU64,
pub article_count: AtomicU64,
pub ttfb_micros_total: AtomicU64,
pub ttfb_count: AtomicU64,
pub send_micros_total: AtomicU64,
pub recv_micros_total: AtomicU64,
pub connection_failures: AtomicU64,
}
impl Default for BackendStore {
fn default() -> Self {
Self {
total_commands: AtomicU64::new(0),
bytes_sent: AtomicU64::new(0),
bytes_received: AtomicU64::new(0),
errors: AtomicU64::new(0),
errors_4xx: AtomicU64::new(0),
errors_5xx: AtomicU64::new(0),
article_bytes_total: AtomicU64::new(0),
article_count: AtomicU64::new(0),
ttfb_micros_total: AtomicU64::new(0),
ttfb_count: AtomicU64::new(0),
send_micros_total: AtomicU64::new(0),
recv_micros_total: AtomicU64::new(0),
connection_failures: AtomicU64::new(0),
}
}
}
impl BackendStore {
pub(crate) fn to_persisted(&self, name: &str) -> PersistedBackend {
PersistedBackend {
name: name.to_string(),
total_commands: self.total_commands.load(Ordering::Relaxed),
bytes_sent: self.bytes_sent.load(Ordering::Relaxed),
bytes_received: self.bytes_received.load(Ordering::Relaxed),
errors: self.errors.load(Ordering::Relaxed),
errors_4xx: self.errors_4xx.load(Ordering::Relaxed),
errors_5xx: self.errors_5xx.load(Ordering::Relaxed),
article_bytes_total: self.article_bytes_total.load(Ordering::Relaxed),
article_count: self.article_count.load(Ordering::Relaxed),
ttfb_micros_total: self.ttfb_micros_total.load(Ordering::Relaxed),
ttfb_count: self.ttfb_count.load(Ordering::Relaxed),
send_micros_total: self.send_micros_total.load(Ordering::Relaxed),
recv_micros_total: self.recv_micros_total.load(Ordering::Relaxed),
connection_failures: self.connection_failures.load(Ordering::Relaxed),
}
}
pub(crate) fn restore_from(&self, persisted: &PersistedBackend) {
self.total_commands
.store(persisted.total_commands, Ordering::Relaxed);
self.bytes_sent
.store(persisted.bytes_sent, Ordering::Relaxed);
self.bytes_received
.store(persisted.bytes_received, Ordering::Relaxed);
self.errors.store(persisted.errors, Ordering::Relaxed);
self.errors_4xx
.store(persisted.errors_4xx, Ordering::Relaxed);
self.errors_5xx
.store(persisted.errors_5xx, Ordering::Relaxed);
self.article_bytes_total
.store(persisted.article_bytes_total, Ordering::Relaxed);
self.article_count
.store(persisted.article_count, Ordering::Relaxed);
self.ttfb_micros_total
.store(persisted.ttfb_micros_total, Ordering::Relaxed);
self.ttfb_count
.store(persisted.ttfb_count, Ordering::Relaxed);
self.send_micros_total
.store(persisted.send_micros_total, Ordering::Relaxed);
self.recv_micros_total
.store(persisted.recv_micros_total, Ordering::Relaxed);
self.connection_failures
.store(persisted.connection_failures, Ordering::Relaxed);
}
}
#[derive(Debug)]
pub struct MetricsStore {
pub total_connections: AtomicU64,
pub backend_stores: Vec<BackendStore>,
pub user_metrics: dashmap::DashMap<String, UserMetrics>,
pub pipeline_batches: AtomicU64,
pub pipeline_commands: AtomicU64,
pub pipeline_requests_queued: AtomicU64,
pub pipeline_requests_completed: AtomicU64,
}
#[derive(Debug, Clone, Default)]
pub struct UserMetrics {
pub username: String,
pub active_connections: usize, pub total_connections: u64,
pub bytes_sent: u64,
pub bytes_received: u64,
pub total_commands: u64,
pub errors: u64,
}
impl UserMetrics {
#[must_use]
pub const fn new(username: String) -> Self {
Self {
username,
active_connections: 0,
total_connections: 0,
bytes_sent: 0,
bytes_received: 0,
total_commands: 0,
errors: 0,
}
}
pub(crate) fn to_user_stats(&self) -> crate::metrics::UserStats {
use crate::metrics::types::{CommandCount, ErrorCount};
use crate::types::{BytesPerSecondRate, BytesReceived, BytesSent, TotalConnections};
crate::metrics::UserStats {
username: self.username.clone(),
active_connections: self.active_connections,
total_connections: TotalConnections::new(self.total_connections),
bytes_sent: BytesSent::new(self.bytes_sent),
bytes_received: BytesReceived::new(self.bytes_received),
total_commands: CommandCount::new(self.total_commands),
errors: ErrorCount::new(self.errors),
bytes_sent_per_sec: BytesPerSecondRate::ZERO,
bytes_received_per_sec: BytesPerSecondRate::ZERO,
}
}
fn to_persisted(&self) -> PersistedUser {
PersistedUser {
username: self.username.clone(),
total_connections: self.total_connections,
bytes_sent: self.bytes_sent,
bytes_received: self.bytes_received,
total_commands: self.total_commands,
errors: self.errors,
}
}
const fn restore_from(&mut self, persisted: &PersistedUser) {
self.total_connections = persisted.total_connections;
self.bytes_sent = persisted.bytes_sent;
self.bytes_received = persisted.bytes_received;
self.total_commands = persisted.total_commands;
self.errors = persisted.errors;
}
#[inline]
const fn total_bytes(&self) -> u64 {
self.bytes_sent.saturating_add(self.bytes_received)
}
}
impl MetricsStore {
#[must_use]
pub fn new(num_backends: usize) -> Self {
Self {
total_connections: AtomicU64::new(0),
backend_stores: (0..num_backends).map(|_| BackendStore::default()).collect(),
user_metrics: dashmap::DashMap::new(),
pipeline_batches: AtomicU64::new(0),
pipeline_commands: AtomicU64::new(0),
pipeline_requests_queued: AtomicU64::new(0),
pipeline_requests_completed: AtomicU64::new(0),
}
}
pub fn load(path: &Path, server_names: &[String]) -> Result<Option<Self>> {
if !path.exists() {
tracing::info!("No stats file found at {:?}, starting fresh", path);
return Ok(None);
}
let content = fs::read_to_string(path).context("Failed to read stats file")?;
let stats_file: StatsFile = match serde_json::from_str(&content) {
Ok(file) => file,
Err(e) => {
tracing::warn!("Corrupt stats file at {:?} ({}), starting fresh", path, e);
return Ok(None);
}
};
if stats_file.version != STATS_FILE_VERSION {
tracing::warn!(
"Unknown stats format version {} (expected {}), starting fresh",
stats_file.version,
STATS_FILE_VERSION
);
return Ok(None);
}
let store = Self::new(server_names.len());
store
.total_connections
.store(stats_file.global.total_connections, Ordering::Relaxed);
store
.pipeline_batches
.store(stats_file.pipeline.batches, Ordering::Relaxed);
store
.pipeline_commands
.store(stats_file.pipeline.commands, Ordering::Relaxed);
store
.pipeline_requests_queued
.store(stats_file.pipeline.requests_queued, Ordering::Relaxed);
store
.pipeline_requests_completed
.store(stats_file.pipeline.requests_completed, Ordering::Relaxed);
for persisted_backend in &stats_file.backends {
if let Some(index) = server_names
.iter()
.position(|name| name == &persisted_backend.name)
{
store.backend_stores[index].restore_from(persisted_backend);
tracing::debug!(
"Restored backend stats for '{}' at index {}",
persisted_backend.name,
index
);
} else {
tracing::debug!(
"Skipping backend '{}' (not in current config)",
persisted_backend.name
);
}
}
for persisted_user in &stats_file.users {
let mut user = UserMetrics::new(persisted_user.username.clone());
user.restore_from(persisted_user);
store
.user_metrics
.insert(persisted_user.username.clone(), user);
}
let pruned_users = store.prune_user_metrics();
if pruned_users > 0 {
tracing::warn!(
pruned_users,
retained_users = store.user_metrics.len(),
limit = MAX_RETAINED_USER_METRICS,
"Pruned persisted user metrics to cap memory usage"
);
}
tracing::info!(
"Loaded stats from {:?} (saved at {})",
path,
stats_file.saved_at
);
Ok(Some(store))
}
pub fn save(&self, path: &Path, server_names: &[String]) -> Result<()> {
let _save_guard = SAVE_LOCK
.lock()
.map_err(|_| anyhow::anyhow!("metrics save lock poisoned"))?;
let backends = server_names
.iter()
.enumerate()
.map(|(i, name)| {
self.backend_stores
.get(i)
.with_context(|| {
format!(
"Backend store index {} out of bounds (stores: {}, names: {})",
i,
self.backend_stores.len(),
server_names.len()
)
})
.map(|store| store.to_persisted(name))
})
.collect::<Result<Vec<_>>>()?
.into();
let stats_file = StatsFile {
version: STATS_FILE_VERSION,
saved_at: chrono::Utc::now().to_rfc3339(),
global: PersistedGlobal {
total_connections: self.total_connections.load(Ordering::Relaxed),
},
backends,
users: self
.user_metrics
.iter()
.map(|entry| entry.value().to_persisted())
.collect(),
pipeline: PersistedPipeline {
batches: self.pipeline_batches.load(Ordering::Relaxed),
commands: self.pipeline_commands.load(Ordering::Relaxed),
requests_queued: self.pipeline_requests_queued.load(Ordering::Relaxed),
requests_completed: self.pipeline_requests_completed.load(Ordering::Relaxed),
},
};
let json =
serde_json::to_string_pretty(&stats_file).context("Failed to serialize stats")?;
let seq = SAVE_SEQ.fetch_add(1, Ordering::Relaxed);
let tmp_filename = format!(
"{}.{}.tmp",
path.file_name().unwrap_or_default().to_string_lossy(),
seq
);
let tmp_path = path.with_file_name(tmp_filename);
fs::write(&tmp_path, json)
.with_context(|| format!("Failed to write stats to {}", tmp_path.display()))?;
if let Err(e) = atomic_replace_file(&tmp_path, path) {
let _ = fs::remove_file(&tmp_path);
return Err(e);
}
tracing::debug!("Saved stats to {}", path.display());
Ok(())
}
#[must_use]
pub fn prune_user_metrics(&self) -> usize {
self.prune_user_metrics_to(MAX_RETAINED_USER_METRICS)
}
fn prune_user_metrics_to(&self, limit: usize) -> usize {
if limit == 0 {
let keys = self
.user_metrics
.iter()
.map(|entry| entry.key().clone())
.collect::<Vec<_>>();
let removed = keys.len();
for username in keys {
self.user_metrics.remove(&username);
}
return removed;
}
let len = self.user_metrics.len();
if len <= limit {
return 0;
}
let mut ranked = self
.user_metrics
.iter()
.map(|entry| {
let user = entry.value();
(
entry.key().clone(),
user.active_connections > 0,
user.total_bytes(),
user.total_connections,
user.total_commands,
user.errors,
)
})
.collect::<Vec<_>>();
ranked.sort_by(|a, b| {
b.1.cmp(&a.1)
.then_with(|| b.2.cmp(&a.2))
.then_with(|| b.3.cmp(&a.3))
.then_with(|| b.4.cmp(&a.4))
.then_with(|| b.5.cmp(&a.5))
.then_with(|| a.0.cmp(&b.0))
});
let active_count = ranked.iter().take_while(|entry| entry.1).count();
let keep_count = limit.max(active_count);
let to_remove = ranked
.into_iter()
.skip(keep_count)
.map(|(username, _, _, _, _, _)| username)
.collect::<Vec<_>>();
let removed = to_remove.len();
for username in to_remove {
self.user_metrics.remove(&username);
}
removed
}
}
#[cfg(unix)]
fn atomic_replace_file(tmp_path: &Path, path: &Path) -> Result<()> {
fs::rename(tmp_path, path).with_context(|| {
format!(
"Failed to atomically replace {} with {}",
path.display(),
tmp_path.display()
)
})
}
#[cfg(windows)]
fn atomic_replace_file(tmp_path: &Path, path: &Path) -> Result<()> {
use std::os::windows::ffi::OsStrExt;
const MOVEFILE_REPLACE_EXISTING: u32 = 0x1;
const MOVEFILE_WRITE_THROUGH: u32 = 0x8;
#[link(name = "kernel32")]
unsafe extern "system" {
fn MoveFileExW(existing: *const u16, new: *const u16, flags: u32) -> i32;
}
fn wide_path(path: &Path) -> Vec<u16> {
path.as_os_str().encode_wide().chain(Some(0)).collect()
}
let tmp_wide = wide_path(tmp_path);
let path_wide = wide_path(path);
let ok = unsafe {
MoveFileExW(
tmp_wide.as_ptr(),
path_wide.as_ptr(),
MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH,
)
};
if ok == 0 {
Err(std::io::Error::last_os_error()).with_context(|| {
format!(
"Failed to atomically replace {} with {}",
path.display(),
tmp_path.display()
)
})
} else {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
#[test]
fn test_backend_store_roundtrip() {
let store = BackendStore::default();
store.total_commands.store(100, Ordering::Relaxed);
store.bytes_sent.store(50000, Ordering::Relaxed);
store.errors.store(5, Ordering::Relaxed);
let persisted = store.to_persisted("test-backend");
assert_eq!(persisted.name, "test-backend");
assert_eq!(persisted.total_commands, 100);
assert_eq!(persisted.bytes_sent, 50000);
assert_eq!(persisted.errors, 5);
let new_store = BackendStore::default();
new_store.restore_from(&persisted);
assert_eq!(new_store.total_commands.load(Ordering::Relaxed), 100);
assert_eq!(new_store.bytes_sent.load(Ordering::Relaxed), 50000);
assert_eq!(new_store.errors.load(Ordering::Relaxed), 5);
}
#[test]
fn test_user_metrics_roundtrip() {
let mut user = UserMetrics::new("alice".to_string());
user.total_connections = 10;
user.bytes_sent = 1000;
user.active_connections = 5;
let persisted = user.to_persisted();
assert_eq!(persisted.username, "alice");
assert_eq!(persisted.total_connections, 10);
assert_eq!(persisted.bytes_sent, 1000);
let mut new_user = UserMetrics::new("alice".to_string());
new_user.restore_from(&persisted);
assert_eq!(new_user.total_connections, 10);
assert_eq!(new_user.bytes_sent, 1000);
assert_eq!(new_user.active_connections, 0); }
#[test]
fn test_metrics_store_save_load_roundtrip() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("stats.json");
let server_names = vec!["backend1".to_string(), "backend2".to_string()];
let store = MetricsStore::new(2);
store.total_connections.store(42, Ordering::Relaxed);
store.backend_stores[0]
.total_commands
.store(100, Ordering::Relaxed);
store.backend_stores[1]
.total_commands
.store(200, Ordering::Relaxed);
store.pipeline_batches.store(10, Ordering::Relaxed);
let mut user = UserMetrics::new("bob".to_string());
user.total_connections = 5;
store.user_metrics.insert("bob".to_string(), user);
store.save(&stats_path, &server_names).unwrap();
let loaded = MetricsStore::load(&stats_path, &server_names)
.unwrap()
.expect("Should load");
assert_eq!(loaded.total_connections.load(Ordering::Relaxed), 42);
assert_eq!(
loaded.backend_stores[0]
.total_commands
.load(Ordering::Relaxed),
100
);
assert_eq!(
loaded.backend_stores[1]
.total_commands
.load(Ordering::Relaxed),
200
);
assert_eq!(loaded.pipeline_batches.load(Ordering::Relaxed), 10);
let bob = loaded.user_metrics.get("bob").unwrap();
assert_eq!(bob.total_connections, 5);
drop(bob);
}
#[test]
fn test_prune_user_metrics_to_keeps_most_active_and_heaviest_users() {
let store = MetricsStore::new(1);
let mut active = UserMetrics::new("active".to_string());
active.active_connections = 1;
active.bytes_sent = 1;
store.user_metrics.insert("active".to_string(), active);
let mut heavy = UserMetrics::new("heavy".to_string());
heavy.bytes_sent = 10_000;
store.user_metrics.insert("heavy".to_string(), heavy);
let mut medium = UserMetrics::new("medium".to_string());
medium.bytes_sent = 5_000;
store.user_metrics.insert("medium".to_string(), medium);
let mut light = UserMetrics::new("light".to_string());
light.bytes_sent = 100;
store.user_metrics.insert("light".to_string(), light);
let removed = store.prune_user_metrics_to(2);
assert_eq!(removed, 2);
assert!(store.user_metrics.contains_key("active"));
assert!(store.user_metrics.contains_key("heavy"));
assert!(!store.user_metrics.contains_key("medium"));
assert!(!store.user_metrics.contains_key("light"));
}
#[test]
fn test_metrics_store_load_prunes_excess_users() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("stats.json");
let server_names = vec!["backend1".to_string()];
let store = MetricsStore::new(1);
for i in 0..(MAX_RETAINED_USER_METRICS + 16) {
let mut user = UserMetrics::new(format!("user-{i:04}"));
user.bytes_sent = i as u64;
store.user_metrics.insert(user.username.clone(), user);
}
store.save(&stats_path, &server_names).unwrap();
let loaded = MetricsStore::load(&stats_path, &server_names)
.unwrap()
.expect("Should load");
assert_eq!(loaded.user_metrics.len(), MAX_RETAINED_USER_METRICS);
assert!(
loaded
.user_metrics
.contains_key(&format!("user-{:04}", MAX_RETAINED_USER_METRICS + 15))
);
assert!(!loaded.user_metrics.contains_key("user-0000"));
}
#[test]
fn test_metrics_store_concurrent_saves_leave_valid_file() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("stats.json");
let server_names = vec!["backend1".to_string()];
let handles = (0..8)
.map(|value| {
let path = stats_path.clone();
let names = server_names.clone();
std::thread::spawn(move || {
let store = MetricsStore::new(1);
store.total_connections.store(value, Ordering::Relaxed);
store.save(&path, &names).unwrap();
})
})
.collect::<Vec<_>>();
for handle in handles {
handle.join().unwrap();
}
let loaded = MetricsStore::load(&stats_path, &server_names)
.unwrap()
.unwrap();
assert!(loaded.total_connections.load(Ordering::Relaxed) < 8);
let tmp_files = std::fs::read_dir(temp_dir.path())
.unwrap()
.filter(|entry| {
entry
.as_ref()
.unwrap()
.file_name()
.to_string_lossy()
.ends_with(".tmp")
})
.count();
assert_eq!(tmp_files, 0);
}
#[test]
fn test_metrics_store_backend_name_mapping() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("stats.json");
let orig_names = vec![
"backend1".to_string(),
"backend2".to_string(),
"backend3".to_string(),
];
let store = MetricsStore::new(3);
store.backend_stores[0]
.total_commands
.store(100, Ordering::Relaxed);
store.backend_stores[1]
.total_commands
.store(200, Ordering::Relaxed);
store.backend_stores[2]
.total_commands
.store(300, Ordering::Relaxed);
store.save(&stats_path, &orig_names).unwrap();
let new_names = vec!["backend3".to_string(), "backend1".to_string()];
let loaded = MetricsStore::load(&stats_path, &new_names)
.unwrap()
.expect("Should load");
assert_eq!(
loaded.backend_stores[0]
.total_commands
.load(Ordering::Relaxed),
300
);
assert_eq!(
loaded.backend_stores[1]
.total_commands
.load(Ordering::Relaxed),
100
);
}
#[test]
fn test_metrics_store_missing_file_returns_none() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("nonexistent.json");
let server_names = vec!["backend1".to_string()];
let result = MetricsStore::load(&stats_path, &server_names).unwrap();
assert!(result.is_none());
}
#[test]
fn test_metrics_store_corrupt_file_returns_none() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("corrupt.json");
fs::write(&stats_path, "{ not valid json }").unwrap();
let server_names = vec!["backend1".to_string()];
let result = MetricsStore::load(&stats_path, &server_names).unwrap();
assert!(result.is_none());
}
#[test]
fn test_metrics_store_backend_mismatch() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("stats.json");
let orig_names = vec!["backend1".to_string(), "backend2".to_string()];
let store = MetricsStore::new(2);
store.backend_stores[0]
.total_commands
.store(100, Ordering::Relaxed);
store.backend_stores[1]
.total_commands
.store(200, Ordering::Relaxed);
store.save(&stats_path, &orig_names).unwrap();
let new_names = vec!["backend3".to_string(), "backend4".to_string()];
let loaded = MetricsStore::load(&stats_path, &new_names)
.unwrap()
.expect("Should load");
assert_eq!(
loaded.backend_stores[0]
.total_commands
.load(Ordering::Relaxed),
0
);
assert_eq!(
loaded.backend_stores[1]
.total_commands
.load(Ordering::Relaxed),
0
);
}
#[test]
fn test_metrics_store_save_with_too_many_names() {
let temp_dir = TempDir::new().unwrap();
let stats_path = temp_dir.path().join("stats.json");
let store = MetricsStore::new(1);
let server_names = vec!["backend1".to_string(), "backend2".to_string()];
let result = store.save(&stats_path, &server_names);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("out of bounds"), "error was: {err}");
}
}