use std::sync::atomic::{AtomicI64, Ordering};
use std::sync::OnceLock;
use dashmap::DashMap;
#[derive(Default)]
pub struct Counter {
v: AtomicI64,
}
impl Counter {
#[inline]
pub fn inc(&self, n: i64) {
self.v.fetch_add(n, Ordering::Relaxed);
}
#[inline]
pub fn get(&self) -> i64 {
self.v.load(Ordering::Relaxed)
}
}
#[derive(Default)]
pub struct Gauge {
v: AtomicI64,
}
impl Gauge {
#[inline]
pub fn set(&self, val: i64) {
self.v.store(val, Ordering::Relaxed);
}
#[inline]
pub fn get(&self) -> i64 {
self.v.load(Ordering::Relaxed)
}
}
pub(crate) struct Registry {
pub(crate) counters: DashMap<String, std::sync::Arc<Counter>>,
pub(crate) gauges: DashMap<String, std::sync::Arc<Gauge>>,
}
impl Default for Registry {
fn default() -> Self {
Self {
counters: DashMap::new(),
gauges: DashMap::new(),
}
}
}
pub(crate) static REGISTRY: OnceLock<Registry> = OnceLock::new();
pub(crate) fn get_registry() -> &'static Registry {
REGISTRY.get_or_init(Registry::default)
}
pub fn counter(name: &str) -> std::sync::Arc<Counter> {
let registry = get_registry();
if let Some(c) = registry.counters.get(name) {
return c.value().clone();
}
let c = std::sync::Arc::new(Counter::default());
registry
.counters
.entry(name.to_string())
.or_insert(c.clone());
registry
.counters
.get(name)
.map(|entry| entry.value().clone())
.unwrap_or(c)
}
pub fn gauge(name: &str) -> std::sync::Arc<Gauge> {
let registry = get_registry();
if let Some(g) = registry.gauges.get(name) {
return g.value().clone();
}
let g = std::sync::Arc::new(Gauge::default());
registry.gauges.entry(name.to_string()).or_insert(g.clone());
registry
.gauges
.get(name)
.map(|entry| entry.value().clone())
.unwrap_or(g)
}
pub mod name {
pub const CLIENT_BYTES_READ_LOCAL: &str = "Client.BytesReadLocal";
pub const CLIENT_BYTES_WRITTEN_LOCAL: &str = "Client.BytesWrittenLocal";
pub const CLIENT_BYTES_WRITTEN_UFS: &str = "Client.BytesWrittenUfs";
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn counter_inc_and_get() {
let c = Counter::default();
assert_eq!(c.get(), 0);
c.inc(42);
assert_eq!(c.get(), 42);
c.inc(8);
assert_eq!(c.get(), 50);
c.inc(-10);
assert_eq!(c.get(), 40);
}
#[test]
fn counter_negative_increment() {
let c = Counter::default();
c.inc(-5);
assert_eq!(c.get(), -5);
}
#[test]
fn gauge_set_and_get() {
let g = Gauge::default();
assert_eq!(g.get(), 0);
g.set(99);
assert_eq!(g.get(), 99);
g.set(-10);
assert_eq!(g.get(), -10);
}
#[test]
fn registry_counter_factory() {
let c1 = counter("my_counter");
assert_eq!(c1.get(), 0);
c1.inc(10);
let c2 = counter("my_counter");
assert_eq!(c2.get(), 10);
let c3 = counter("other_counter");
assert_eq!(c3.get(), 0);
}
#[test]
fn registry_gauge_factory() {
let g1 = gauge("my_gauge");
assert_eq!(g1.get(), 0);
g1.set(55);
let g2 = gauge("my_gauge");
assert_eq!(g2.get(), 55);
let g3 = gauge("other_gauge");
assert_eq!(g3.get(), 0);
}
#[test]
fn registry_counter_concurrent() {
use std::thread;
let c = counter("concurrent_counter");
let mut handles = vec![];
for _ in 0..10 {
let c = c.clone();
let handle = thread::spawn(move || {
for _ in 0..1000 {
c.inc(1);
}
});
handles.push(handle);
}
for h in handles {
h.join().unwrap();
}
assert_eq!(c.get(), 10_000);
}
#[test]
fn name_constants() {
assert_eq!(name::CLIENT_BYTES_READ_LOCAL, "Client.BytesReadLocal");
assert_eq!(name::CLIENT_BYTES_WRITTEN_LOCAL, "Client.BytesWrittenLocal");
assert_eq!(name::CLIENT_BYTES_WRITTEN_UFS, "Client.BytesWrittenUfs");
}
}