use dashmap::DashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
const MIN_SHARDS: usize = 1;
const MAX_SHARDS: usize = 1024;
type ShardMap = DashMap<String, AtomicU64>;
pub struct HotKeyTracker {
shards: Vec<Arc<ShardMap>>,
}
impl HotKeyTracker {
pub fn new(shards: usize) -> Self {
let shards = shards.clamp(MIN_SHARDS, MAX_SHARDS);
Self {
shards: (0..shards).map(|_| Arc::new(DashMap::new())).collect(),
}
}
fn shard_index(&self, key: &str) -> usize {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
key.hash(&mut hasher);
(hasher.finish() as usize) % self.shards.len()
}
pub fn record(&self, key: &str) {
let idx = self.shard_index(key);
self.shards[idx]
.entry(key.to_string())
.and_modify(|c| {
c.fetch_add(1, Ordering::Relaxed);
})
.or_insert(AtomicU64::new(1));
}
pub fn peek(&self, key: &str) -> u64 {
let idx = self.shard_index(key);
self.shards[idx]
.get(key)
.map(|c| c.load(Ordering::Relaxed))
.unwrap_or(0)
}
pub fn snapshot_top(&self, k: usize) -> Vec<(String, u64)> {
let mut all: Vec<(String, u64)> = Vec::new();
for shard in &self.shards {
for entry in shard.iter() {
let count = entry.value().load(Ordering::Relaxed);
if count > 0 {
all.push((entry.key().clone(), count));
}
}
}
all.sort_unstable_by(|a, b| b.1.cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
all.truncate(k);
for shard in &self.shards {
for entry in shard.iter() {
let count = entry.value().load(Ordering::Relaxed);
entry.value().store(count >> 1, Ordering::Relaxed);
}
shard.retain(|_, c| c.load(Ordering::Relaxed) > 0);
}
all
}
pub fn reset(&self) {
for shard in &self.shards {
shard.clear();
}
}
}
pub type SharedHotKeyTracker = Arc<HotKeyTracker>;
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Barrier;
use std::thread;
#[test]
fn concurrent_record_then_top_k_sorted() {
let tracker = Arc::new(HotKeyTracker::new(16));
let barrier = Arc::new(Barrier::new(8));
let mut handles = Vec::new();
for t in 0..8u64 {
let tracker = tracker.clone();
let barrier = barrier.clone();
handles.push(thread::spawn(move || {
barrier.wait();
for _ in 0..100 {
tracker.record(&format!("key-{}", t % 2));
}
}));
}
for h in handles {
h.join().unwrap();
}
let top = tracker.snapshot_top(10);
assert_eq!(top.len(), 2);
assert!(top[0].1 >= top[1].1);
let total: u64 = top.iter().map(|(_, c)| c).sum();
assert!(total >= 800, "并发采样不应大量丢失,实际 {total}");
}
#[test]
fn snapshot_halves_counts() {
let tracker = HotKeyTracker::new(4);
for _ in 0..100 {
tracker.record("old-hot");
}
let top = tracker.snapshot_top(10);
assert_eq!(top[0].0, "old-hot");
assert_eq!(top[0].1, 100);
assert_eq!(tracker.peek("old-hot"), 50);
for _ in 0..60 {
tracker.record("new-hot");
}
let top = tracker.snapshot_top(10);
assert_eq!(top[0].0, "new-hot", "半衰应允许新热点顶替旧热点");
}
#[test]
fn reset_clears_all() {
let tracker = HotKeyTracker::new(8);
tracker.record("k");
tracker.reset();
assert!(tracker.snapshot_top(10).is_empty());
}
}