use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use serde_json::json;
use crate::StatsMergeMode;
#[derive(Debug)]
pub struct RateCounter {
total: Arc<AtomicU64>,
speed: Arc<AtomicU64>,
last_total: Arc<AtomicU64>,
current: Arc<AtomicI64>,
max: Arc<AtomicI64>,
mode: StatsMergeMode,
}
impl Clone for RateCounter {
fn clone(&self) -> Self {
Self {
total: Arc::clone(&self.total),
speed: Arc::clone(&self.speed),
last_total: Arc::clone(&self.last_total),
current: Arc::clone(&self.current),
max: Arc::clone(&self.max),
mode: self.mode.clone(),
}
}
}
impl Default for RateCounter {
fn default() -> Self {
Self::new()
}
}
impl RateCounter {
#[inline]
pub fn new() -> Self {
Self::new_with_mode(StatsMergeMode::None)
}
#[inline]
pub fn new_with_mode(mode: StatsMergeMode) -> Self {
Self {
total: Arc::new(AtomicU64::new(0)),
speed: Arc::new(AtomicU64::new(0)),
last_total: Arc::new(AtomicU64::new(0)),
current: Arc::new(AtomicI64::new(0)),
max: Arc::new(AtomicI64::new(0)),
mode,
}
}
#[inline]
pub fn snapshot(&self) -> Self {
Self {
total: Arc::new(AtomicU64::new(self.total())),
speed: Arc::new(AtomicU64::new(f64::to_bits(self.speed()))),
last_total: Arc::new(AtomicU64::new(self.total())),
current: Arc::new(AtomicI64::new(self.current())),
max: Arc::new(AtomicI64::new(self.max())),
mode: self.mode.clone(),
}
}
#[inline]
pub fn inc(&self) {
self.total.fetch_add(1, Ordering::Relaxed);
let old = self.current.fetch_add(1, Ordering::Relaxed);
self.max.fetch_max(old + 1, Ordering::Relaxed);
}
#[inline]
pub fn incs(&self, n: u64) {
self.total.fetch_add(n, Ordering::Relaxed);
let old = self.current.fetch_add(n as i64, Ordering::Relaxed);
self.max.fetch_max(old + n as i64, Ordering::Relaxed);
}
#[inline]
pub fn total(&self) -> u64 {
self.total.load(Ordering::Relaxed)
}
#[inline]
pub fn speed(&self) -> f64 {
f64::from_bits(self.speed.load(Ordering::Relaxed))
}
#[inline]
pub fn tick(&self, interval: Duration) {
let curr = self.total.load(Ordering::Relaxed);
let prev = self.last_total.swap(curr, Ordering::Relaxed);
let delta = curr.wrapping_sub(prev) as f64;
let rate = delta / interval.as_secs_f64();
self.speed.store(f64::to_bits(rate), Ordering::Relaxed);
}
#[inline]
pub fn reset(&self) {
self.total.store(0, Ordering::Relaxed);
self.speed.store(f64::to_bits(0.0), Ordering::Relaxed);
self.last_total.store(0, Ordering::Relaxed);
self.current.store(0, Ordering::Relaxed);
self.max.store(0, Ordering::Relaxed);
}
#[inline]
pub fn current(&self) -> i64 {
self.current.load(Ordering::Relaxed)
}
#[inline]
pub fn max(&self) -> i64 {
self.max.load(Ordering::Relaxed)
}
#[inline]
pub fn to_json(&self) -> serde_json::Value {
json!({
"total": self.total(),
"speed": self.speed(),
"current": self.current(),
"max": self.max(),
})
}
#[inline]
pub fn add(&self, other: &Self) {
self.total.fetch_add(other.total(), Ordering::Relaxed);
self.current.fetch_add(other.current(), Ordering::Relaxed);
self.max.fetch_add(other.max(), Ordering::Relaxed);
}
#[inline]
pub fn set(&self, other: &Self) {
self.total.store(other.total(), Ordering::Relaxed);
self.speed.store(f64::to_bits(other.speed()), Ordering::Relaxed);
self.last_total.store(other.total(), Ordering::Relaxed);
self.current.store(other.current(), Ordering::Relaxed);
self.max.store(other.max(), Ordering::Relaxed);
}
#[inline]
pub fn merge(&self, other: &Self) {
match self.mode {
StatsMergeMode::None => {}
StatsMergeMode::Sum => {
self.add(other);
}
StatsMergeMode::Max => {
self.current.fetch_max(other.current(), Ordering::Relaxed);
self.max.fetch_max(other.max(), Ordering::Relaxed);
self.total.fetch_max(other.total(), Ordering::Relaxed);
}
StatsMergeMode::Min => {
self.current.fetch_min(other.current(), Ordering::Relaxed);
self.max.fetch_min(other.max(), Ordering::Relaxed);
self.total.fetch_min(other.total(), Ordering::Relaxed);
}
_ => {}
}
}
#[inline]
pub fn dec(&self) {
self.current.fetch_sub(1, Ordering::Relaxed);
}
#[inline]
pub fn decs(&self, n: i64) {
self.current.fetch_sub(n, Ordering::Relaxed);
}
}
#[derive(Serialize, Deserialize)]
struct RateCounterData {
total: u64,
speed: f64,
current: i64,
max: i64,
mode: StatsMergeMode,
}
impl Serialize for RateCounter {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
RateCounterData {
total: self.total(),
speed: self.speed(),
current: self.current(),
max: self.max(),
mode: self.mode.clone(),
}
.serialize(serializer)
}
}
impl<'de> Deserialize<'de> for RateCounter {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let data = RateCounterData::deserialize(deserializer)?;
Ok(RateCounter {
total: Arc::new(AtomicU64::new(data.total)),
speed: Arc::new(AtomicU64::new(f64::to_bits(data.speed))),
last_total: Arc::new(AtomicU64::new(data.total)),
current: Arc::new(AtomicI64::new(data.current)),
max: Arc::new(AtomicI64::new(data.max)),
mode: data.mode,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn new_is_zero() {
let rc = RateCounter::new();
assert_eq!(rc.total(), 0);
assert_eq!(rc.speed(), 0.0);
assert_eq!(rc.current(), 0);
assert_eq!(rc.max(), 0);
}
#[test]
fn inc_updates_current_and_max() {
let rc = RateCounter::new();
rc.inc();
assert_eq!(rc.total(), 1);
assert_eq!(rc.current(), 1);
assert_eq!(rc.max(), 1);
}
#[test]
fn incs_updates_current_and_max() {
let rc = RateCounter::new();
rc.incs(42);
assert_eq!(rc.total(), 42);
assert_eq!(rc.current(), 42);
assert_eq!(rc.max(), 42);
}
#[test]
fn max_tracks_peak() {
let rc = RateCounter::new();
rc.incs(10);
assert_eq!(rc.max(), 10);
rc.decs(5);
assert_eq!(rc.current(), 5);
assert_eq!(rc.max(), 10, "decrease should not lower max");
rc.incs(20);
assert_eq!(rc.current(), 25);
assert_eq!(rc.max(), 25, "new higher peak should update max");
}
#[test]
fn dec_does_not_affect_total_nor_max() {
let rc = RateCounter::new();
rc.incs(10);
assert_eq!(rc.total(), 10);
assert_eq!(rc.max(), 10);
rc.dec();
assert_eq!(rc.total(), 10, "dec should not change total");
assert_eq!(rc.current(), 9);
assert_eq!(rc.max(), 10, "dec should not change max");
rc.decs(5);
assert_eq!(rc.total(), 10, "decs should not change total");
assert_eq!(rc.current(), 4);
assert_eq!(rc.max(), 10, "decs should not change max");
}
#[test]
fn tick_does_not_affect_current_nor_max() {
let rc = RateCounter::new();
rc.incs(100);
assert_eq!(rc.current(), 100);
assert_eq!(rc.max(), 100);
rc.tick(Duration::from_secs(1));
assert!((rc.speed() - 100.0).abs() < f64::EPSILON);
assert_eq!(rc.total(), 100);
assert_eq!(rc.current(), 100, "tick should not affect current");
assert_eq!(rc.max(), 100, "tick should not affect max");
}
#[test]
fn reset_clears_all_including_max() {
let rc = RateCounter::new();
rc.incs(100);
rc.decs(20);
rc.tick(Duration::from_secs(1));
rc.reset();
assert_eq!(rc.total(), 0);
assert_eq!(rc.speed(), 0.0);
assert_eq!(rc.current(), 0);
assert_eq!(rc.max(), 0);
rc.tick(Duration::from_secs(1));
assert_eq!(rc.speed(), 0.0);
}
#[test]
fn concurrent_incs_max() {
use std::thread;
let rc = Arc::new(RateCounter::new());
let mut handles = Vec::new();
for _ in 0..8 {
let rc = rc.clone();
handles.push(thread::spawn(move || {
for _ in 0..1_000 {
rc.inc();
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(rc.total(), 8_000);
assert_eq!(rc.current(), 8_000);
assert_eq!(rc.max(), 8_000);
}
#[test]
fn clone_shares_max() {
let a = RateCounter::new();
let b = a.clone();
a.incs(10);
assert_eq!(b.total(), 10, "clone should see the same total");
assert_eq!(b.current(), 10, "clone should see the same current");
assert_eq!(b.max(), 10, "clone should see the same max");
b.dec();
assert_eq!(a.current(), 9, "original should see b's dec");
a.incs(20);
assert_eq!(b.current(), 29, "clone should see a's new current");
assert_eq!(b.max(), 29, "clone should see a's new max");
a.tick(Duration::from_secs(1));
assert!((b.speed() - 30.0).abs() < f64::EPSILON, "speed computed on a should also be visible on b");
}
#[test]
fn concurrent_inc_dec_pair_max() {
use std::thread;
let rc = Arc::new(RateCounter::new());
let mut handles = Vec::new();
for _ in 0..8 {
let rc = rc.clone();
handles.push(thread::spawn(move || {
for _ in 0..1_000 {
rc.inc();
rc.dec();
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(rc.total(), 8_000);
assert_eq!(rc.current(), 0);
assert!(rc.max() > 0, "concurrent inc/dec should have produced a peak");
}
#[test]
fn to_json_includes_all_fields() {
let rc = RateCounter::new();
rc.incs(100);
rc.decs(20);
rc.tick(Duration::from_secs(5));
let json = rc.to_json();
let obj = json.as_object().expect("to_json should return an object");
assert_eq!(obj["total"].as_u64(), Some(100));
assert!((obj["speed"].as_f64().unwrap() - 20.0).abs() < f64::EPSILON);
assert_eq!(obj["current"].as_i64(), Some(80));
assert_eq!(obj["max"].as_i64(), Some(100));
}
}