From 65f487ba260c37db69717ccd5750cb50d9faf016 Mon Sep 17 00:00:00 2001 From: fanyang Date: Thu, 18 Jun 2026 00:40:11 +0800 Subject: [PATCH] fix: make stats counters thread safe --- easytier/src/common/stats_manager.rs | 175 +++++++++++---------------- easytier/src/tunnel/stats.rs | 93 +++++++++----- 2 files changed, 133 insertions(+), 135 deletions(-) diff --git a/easytier/src/common/stats_manager.rs b/easytier/src/common/stats_manager.rs index ecbde860..0db8edbe 100644 --- a/easytier/src/common/stats_manager.rs +++ b/easytier/src/common/stats_manager.rs @@ -1,8 +1,9 @@ use dashmap::DashMap; +use parking_lot::Mutex; use serde::{Deserialize, Serialize}; -use std::cell::UnsafeCell; use std::fmt; use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; use std::time::{Duration, Instant}; use tokio::time::interval; use tokio_util::task::AbortOnDropHandle; @@ -374,10 +375,10 @@ impl Default for LabelSet { } } -/// UnsafeCounter provides a high-performance counter using UnsafeCell +/// UnsafeCounter provides a high-performance atomic counter #[derive(Debug)] pub struct UnsafeCounter { - value: UnsafeCell, + value: AtomicU64, } impl Default for UnsafeCounter { @@ -389,121 +390,79 @@ impl Default for UnsafeCounter { impl UnsafeCounter { pub fn new() -> Self { Self { - value: UnsafeCell::new(0), + value: AtomicU64::new(0), } } pub fn new_with_value(initial: u64) -> Self { Self { - value: UnsafeCell::new(initial), + value: AtomicU64::new(initial), } } /// Increment the counter by the given amount - /// # Safety - /// This method is unsafe because it uses UnsafeCell. The caller must ensure - /// that no other thread is accessing this counter simultaneously. - pub unsafe fn add(&self, delta: u64) { - let ptr = self.value.get(); - unsafe { - *ptr = (*ptr).saturating_add(delta); - } + pub fn add(&self, delta: u64) { + let _ = self + .value + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| { + Some(current.saturating_add(delta)) + }); } /// Increment the counter by 1 - /// # Safety - /// This method is unsafe because it uses UnsafeCell. The caller must ensure - /// that no other thread is accessing this counter simultaneously. - pub unsafe fn inc(&self) { - unsafe { - self.add(1); - } + pub fn inc(&self) { + self.add(1); } /// Get the current value of the counter - /// # Safety - /// This method is unsafe because it uses UnsafeCell. The caller must ensure - /// that no other thread is modifying this counter simultaneously. - pub unsafe fn get(&self) -> u64 { - let ptr = self.value.get(); - unsafe { *ptr } + pub fn get(&self) -> u64 { + self.value.load(Ordering::Relaxed) } /// Reset the counter to zero - /// # Safety - /// This method is unsafe because it uses UnsafeCell. The caller must ensure - /// that no other thread is accessing this counter simultaneously. - pub unsafe fn reset(&self) { - let ptr = self.value.get(); - unsafe { - *ptr = 0; - } + pub fn reset(&self) { + self.value.store(0, Ordering::Relaxed); } /// Set the counter to a specific value - /// # Safety - /// This method is unsafe because it uses UnsafeCell. The caller must ensure - /// that no other thread is accessing this counter simultaneously. - pub unsafe fn set(&self, value: u64) { - let ptr = self.value.get(); - unsafe { - *ptr = value; - } + pub fn set(&self, value: u64) { + self.value.store(value, Ordering::Relaxed); } } -// UnsafeCounter is Send + Sync because the safety is guaranteed by the caller -unsafe impl Send for UnsafeCounter {} -unsafe impl Sync for UnsafeCounter {} - /// MetricData contains both the counter and last update timestamp -/// Uses UnsafeCell for lock-free access #[derive(Debug)] struct MetricData { counter: UnsafeCounter, - last_updated: UnsafeCell, + last_updated: Mutex, } impl MetricData { fn new() -> Self { Self { counter: UnsafeCounter::new(), - last_updated: UnsafeCell::new(Instant::now()), + last_updated: Mutex::new(Instant::now()), } } fn new_with_value(initial: u64) -> Self { Self { counter: UnsafeCounter::new_with_value(initial), - last_updated: UnsafeCell::new(Instant::now()), + last_updated: Mutex::new(Instant::now()), } } /// Update the last_updated timestamp - /// # Safety - /// This method is unsafe because it uses UnsafeCell. The caller must ensure - /// that no other thread is accessing this timestamp simultaneously. - unsafe fn touch(&self) { - let ptr = self.last_updated.get(); - unsafe { - *ptr = Instant::now(); - } + fn touch(&self) { + *self.last_updated.lock() = Instant::now(); } /// Get the last updated timestamp - /// # Safety - /// This method is unsafe because it uses UnsafeCell. The caller must ensure - /// that no other thread is modifying this timestamp simultaneously. - unsafe fn get_last_updated(&self) -> Instant { - let ptr = self.last_updated.get(); - unsafe { *ptr } + fn get_last_updated(&self) -> Instant { + *self.last_updated.lock() } } -// MetricData is Send + Sync because the safety is guaranteed by the caller -unsafe impl Send for MetricData {} -unsafe impl Sync for MetricData {} - /// MetricKey uniquely identifies a metric with its name and labels #[derive(Debug, Clone, PartialEq, Eq, Hash)] struct MetricKey { @@ -546,39 +505,31 @@ impl CounterHandle { /// Increment the counter by the given amount pub fn add(&self, delta: u64) { - unsafe { - self.metric_data.counter.add(delta); - self.metric_data.touch(); - } + self.metric_data.counter.add(delta); + self.metric_data.touch(); } /// Increment the counter by 1 pub fn inc(&self) { - unsafe { - self.metric_data.counter.inc(); - self.metric_data.touch(); - } + self.metric_data.counter.inc(); + self.metric_data.touch(); } /// Get the current value of the counter pub fn get(&self) -> u64 { - unsafe { self.metric_data.counter.get() } + self.metric_data.counter.get() } /// Reset the counter to zero pub fn reset(&self) { - unsafe { - self.metric_data.counter.reset(); - self.metric_data.touch(); - } + self.metric_data.counter.reset(); + self.metric_data.touch(); } /// Set the counter to a specific value pub fn set(&self, value: u64) { - unsafe { - self.metric_data.counter.set(value); - self.metric_data.touch(); - } + self.metric_data.counter.set(value); + self.metric_data.touch(); } } @@ -624,7 +575,7 @@ impl StatsManager { counters.retain(|_, metric_data: &mut Arc| { Arc::strong_count(metric_data) > 1 - || unsafe { metric_data.get_last_updated() > cutoff_time } + || metric_data.get_last_updated() > cutoff_time }); counters.shrink_to_fit(); } @@ -662,7 +613,7 @@ impl StatsManager { let key = entry.key(); let metric_data = entry.value(); - let value = unsafe { metric_data.counter.get() }; + let value = metric_data.counter.get(); metrics.push(MetricSnapshot { name: key.name, @@ -695,7 +646,7 @@ impl StatsManager { let key = MetricKey::new(name, labels.clone()); if let Some(metric_data) = self.counters.get(&key) { - let value = unsafe { metric_data.counter.get() }; + let value = metric_data.counter.get(); Some(MetricSnapshot { name, labels: labels.clone(), @@ -796,17 +747,15 @@ mod tests { async fn test_unsafe_counter() { let counter = UnsafeCounter::new(); - unsafe { - assert_eq!(counter.get(), 0); - counter.inc(); - assert_eq!(counter.get(), 1); - counter.add(5); - assert_eq!(counter.get(), 6); - counter.set(10); - assert_eq!(counter.get(), 10); - counter.reset(); - assert_eq!(counter.get(), 0); - } + assert_eq!(counter.get(), 0); + counter.inc(); + assert_eq!(counter.get(), 1); + counter.add(5); + assert_eq!(counter.get(), 6); + counter.set(10); + assert_eq!(counter.get(), 10); + counter.reset(); + assert_eq!(counter.get(), 0); } #[tokio::test] @@ -951,8 +900,7 @@ mod tests { stats .counters .retain(|_, metric_data: &mut Arc| { - Arc::strong_count(metric_data) > 1 - || unsafe { metric_data.get_last_updated() > cutoff_time } + Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > cutoff_time }); assert_eq!(stats.metric_count(), 1); @@ -962,12 +910,33 @@ mod tests { stats .counters .retain(|_, metric_data: &mut Arc| { - Arc::strong_count(metric_data) > 1 - || unsafe { metric_data.get_last_updated() > cutoff_time } + Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > cutoff_time }); assert_eq!(stats.metric_count(), 0); } + #[tokio::test] + async fn test_counter_handle_concurrent_increment() { + const THREADS: usize = 8; + const INCREMENTS_PER_THREAD: usize = 10_000; + + let stats = StatsManager::new(); + let counter = stats.get_simple_counter(MetricName::TrafficPacketsForwarded); + + std::thread::scope(|scope| { + for _ in 0..THREADS { + let counter = counter.clone(); + scope.spawn(move || { + for _ in 0..INCREMENTS_PER_THREAD { + counter.inc(); + } + }); + } + }); + + assert_eq!(counter.get(), (THREADS * INCREMENTS_PER_THREAD) as u64); + } + #[tokio::test] async fn test_stats_rpc_data_structures() { // Test GetStatsRequest diff --git a/easytier/src/tunnel/stats.rs b/easytier/src/tunnel/stats.rs index f3b707a5..0edbc063 100644 --- a/easytier/src/tunnel/stats.rs +++ b/easytier/src/tunnel/stats.rs @@ -1,7 +1,4 @@ -use std::{ - cell::UnsafeCell, - sync::atomic::{AtomicU32, Ordering::Relaxed}, -}; +use std::sync::atomic::{AtomicU32, AtomicU64, Ordering::Relaxed}; pub struct WindowLatency { latency_us_window: Vec, @@ -63,34 +60,30 @@ impl WindowLatency { #[derive(Debug)] pub struct Throughput { - tx_bytes: UnsafeCell, - rx_bytes: UnsafeCell, - tx_packets: UnsafeCell, - rx_packets: UnsafeCell, + tx_bytes: AtomicU64, + rx_bytes: AtomicU64, + tx_packets: AtomicU64, + rx_packets: AtomicU64, } impl Clone for Throughput { fn clone(&self) -> Self { Self { - tx_bytes: UnsafeCell::new(unsafe { *self.tx_bytes.get() }), - rx_bytes: UnsafeCell::new(unsafe { *self.rx_bytes.get() }), - tx_packets: UnsafeCell::new(unsafe { *self.tx_packets.get() }), - rx_packets: UnsafeCell::new(unsafe { *self.rx_packets.get() }), + tx_bytes: AtomicU64::new(self.tx_bytes()), + rx_bytes: AtomicU64::new(self.rx_bytes()), + tx_packets: AtomicU64::new(self.tx_packets()), + rx_packets: AtomicU64::new(self.rx_packets()), } } } -// add sync::Send and sync::Sync traits to Throughput -unsafe impl Send for Throughput {} -unsafe impl Sync for Throughput {} - impl Default for Throughput { fn default() -> Self { Self { - tx_bytes: UnsafeCell::new(0), - rx_bytes: UnsafeCell::new(0), - tx_packets: UnsafeCell::new(0), - rx_packets: UnsafeCell::new(0), + tx_bytes: AtomicU64::new(0), + rx_bytes: AtomicU64::new(0), + tx_packets: AtomicU64::new(0), + rx_packets: AtomicU64::new(0), } } } @@ -101,32 +94,68 @@ impl Throughput { } pub fn tx_bytes(&self) -> u64 { - unsafe { *self.tx_bytes.get() } + self.tx_bytes.load(Relaxed) } pub fn rx_bytes(&self) -> u64 { - unsafe { *self.rx_bytes.get() } + self.rx_bytes.load(Relaxed) } pub fn tx_packets(&self) -> u64 { - unsafe { *self.tx_packets.get() } + self.tx_packets.load(Relaxed) } pub fn rx_packets(&self) -> u64 { - unsafe { *self.rx_packets.get() } + self.rx_packets.load(Relaxed) } pub fn record_tx_bytes(&self, bytes: u64) { - unsafe { - *self.tx_bytes.get() += bytes; - *self.tx_packets.get() += 1; - } + self.tx_bytes.fetch_add(bytes, Relaxed); + self.tx_packets.fetch_add(1, Relaxed); } pub fn record_rx_bytes(&self, bytes: u64) { - unsafe { - *self.rx_bytes.get() += bytes; - *self.rx_packets.get() += 1; - } + self.rx_bytes.fetch_add(bytes, Relaxed); + self.rx_packets.fetch_add(1, Relaxed); + } +} + +#[cfg(test)] +mod tests { + use super::Throughput; + use std::sync::Arc; + + #[test] + fn throughput_records_concurrent_tx_and_rx() { + const THREADS: usize = 8; + const RECORDS_PER_THREAD: usize = 10_000; + const TX_BYTES_PER_RECORD: u64 = 3; + const RX_BYTES_PER_RECORD: u64 = 7; + + let throughput = Arc::new(Throughput::new()); + + std::thread::scope(|scope| { + for _ in 0..THREADS { + let throughput = Arc::clone(&throughput); + scope.spawn(move || { + for _ in 0..RECORDS_PER_THREAD { + throughput.record_tx_bytes(TX_BYTES_PER_RECORD); + throughput.record_rx_bytes(RX_BYTES_PER_RECORD); + } + }); + } + }); + + let expected_packets = (THREADS * RECORDS_PER_THREAD) as u64; + assert_eq!(throughput.tx_packets(), expected_packets); + assert_eq!(throughput.rx_packets(), expected_packets); + assert_eq!( + throughput.tx_bytes(), + expected_packets * TX_BYTES_PER_RECORD + ); + assert_eq!( + throughput.rx_bytes(), + expected_packets * RX_BYTES_PER_RECORD + ); } }