use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
#[derive(Clone, Debug)]
pub struct CountMinSketch {
counters: Vec<Vec<u32>>,
depth: usize,
width: usize,
total_count: u64,
}
impl CountMinSketch {
pub fn new(width: usize, depth: usize) -> Self {
Self {
counters: vec![vec![0u32; width]; depth],
depth,
width,
total_count: 0,
}
}
pub fn increment<K: Hash>(&mut self, key: &K) {
for i in 0..self.depth {
let index = self.hash(key, i);
self.counters[i][index] = self.counters[i][index].saturating_add(1);
}
self.total_count = self.total_count.saturating_add(1);
}
pub fn estimate<K: Hash>(&self, key: &K) -> u32 {
let mut min_count = u32::MAX;
for i in 0..self.depth {
let index = self.hash(key, i);
min_count = min_count.min(self.counters[i][index]);
}
min_count
}
pub fn decay(&mut self) {
for row in &mut self.counters {
for counter in row.iter_mut() {
*counter /= 2;
}
}
self.total_count /= 2;
}
pub fn reset(&mut self) {
for row in &mut self.counters {
for counter in row.iter_mut() {
*counter = 0;
}
}
self.total_count = 0;
}
pub fn total_count(&self) -> u64 {
self.total_count
}
fn hash<K: Hash>(&self, key: &K, hash_index: usize) -> usize {
let mut hasher = DefaultHasher::new();
hash_index.hash(&mut hasher);
key.hash(&mut hasher);
(hasher.finish() as usize) % self.width
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_increment_and_estimate() {
let mut sketch = CountMinSketch::new(1024, 4);
sketch.increment(&"key1");
assert_eq!(sketch.estimate(&"key1"), 1);
sketch.increment(&"key1");
assert_eq!(sketch.estimate(&"key1"), 2);
sketch.increment(&"key2");
assert_eq!(sketch.estimate(&"key2"), 1);
assert_eq!(sketch.estimate(&"key1"), 2);
}
#[test]
fn test_estimate_nonexistent_key() {
let sketch = CountMinSketch::new(1024, 4);
assert_eq!(sketch.estimate(&"nonexistent"), 0);
}
#[test]
fn test_decay() {
let mut sketch = CountMinSketch::new(1024, 4);
sketch.increment(&"key1");
sketch.increment(&"key1");
sketch.increment(&"key1");
sketch.increment(&"key1");
assert_eq!(sketch.estimate(&"key1"), 4);
sketch.decay();
assert_eq!(sketch.estimate(&"key1"), 2);
sketch.decay();
assert_eq!(sketch.estimate(&"key1"), 1);
sketch.decay();
assert_eq!(sketch.estimate(&"key1"), 0);
}
#[test]
fn test_reset() {
let mut sketch = CountMinSketch::new(1024, 4);
sketch.increment(&"key1");
sketch.increment(&"key2");
sketch.increment(&"key3");
assert_eq!(sketch.total_count(), 3);
assert!(sketch.estimate(&"key1") > 0);
sketch.reset();
assert_eq!(sketch.total_count(), 0);
assert_eq!(sketch.estimate(&"key1"), 0);
assert_eq!(sketch.estimate(&"key2"), 0);
assert_eq!(sketch.estimate(&"key3"), 0);
}
#[test]
fn test_total_count() {
let mut sketch = CountMinSketch::new(1024, 4);
assert_eq!(sketch.total_count(), 0);
sketch.increment(&"key1");
assert_eq!(sketch.total_count(), 1);
sketch.increment(&"key2");
sketch.increment(&"key1");
assert_eq!(sketch.total_count(), 3);
sketch.decay();
assert_eq!(sketch.total_count(), 1); }
#[test]
fn test_multiple_keys() {
let mut sketch = CountMinSketch::new(2048, 4);
for i in 0..100 {
sketch.increment(&format!("key{}", i));
}
for i in 0..100 {
assert_eq!(sketch.estimate(&format!("key{}", i)), 1);
}
for _ in 0..10 {
sketch.increment(&"key5");
}
assert_eq!(sketch.estimate(&"key5"), 11);
assert_eq!(sketch.estimate(&"key10"), 1);
}
#[test]
fn test_saturation() {
let mut sketch = CountMinSketch::new(1024, 4);
let test_key = "key1";
for i in 0..sketch.depth {
let index = sketch.hash(&test_key, i);
sketch.counters[i][index] = u32::MAX - 1;
}
sketch.increment(&test_key);
sketch.increment(&test_key);
sketch.increment(&test_key);
let estimate = sketch.estimate(&test_key);
assert_eq!(estimate, u32::MAX);
}
#[test]
fn test_different_types() {
let mut sketch = CountMinSketch::new(1024, 4);
sketch.increment(&42u32);
sketch.increment(&"string_key");
sketch.increment(&(1, 2, 3));
assert_eq!(sketch.estimate(&42u32), 1);
assert_eq!(sketch.estimate(&"string_key"), 1);
assert_eq!(sketch.estimate(&(1, 2, 3)), 1);
}
}