use std::hash::Hash;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Instant;
use dashmap::DashMap;
use flume::{unbounded, Receiver, Sender};
use crate::error::TransportError;
pub struct LockFreeHashMap<K, V>
where
K: Hash + Eq + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
map: DashMap<K, V>,
stats: Arc<LockFreeStats>,
}
#[derive(Debug)]
pub struct LockFreeStats {
pub reads: AtomicU64,
pub writes: AtomicU64,
pub cas_retries: AtomicU64,
pub avg_read_latency_ns: AtomicU64,
}
impl<K, V> LockFreeHashMap<K, V>
where
K: Hash + Eq + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
pub fn new() -> Self {
Self::with_capacity(16) }
pub fn with_capacity(shard_count: usize) -> Self {
let shard_count = shard_count.max(1).next_power_of_two();
Self {
map: DashMap::with_shard_amount(shard_count),
stats: Arc::new(LockFreeStats::new()),
}
}
pub fn get(&self, key: &K) -> Option<V> {
let start = Instant::now();
let read_count = self.stats.reads.fetch_add(1, Ordering::Relaxed) + 1;
let result = self.map.get(key).map(|v| v.clone());
let latency = start.elapsed().as_nanos() as u64;
let mut current = self.stats.avg_read_latency_ns.load(Ordering::Relaxed);
loop {
let next = if read_count <= 1 {
latency
} else if latency >= current {
current + (latency - current) / read_count
} else {
current - (current - latency) / read_count
};
match self.stats.avg_read_latency_ns.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
result
}
pub fn insert(&self, key: K, value: V) -> Result<Option<V>, TransportError> {
self.stats.writes.fetch_add(1, Ordering::Relaxed);
Ok(self.map.insert(key, value))
}
pub fn remove(&self, key: &K) -> Result<Option<V>, TransportError> {
self.stats.writes.fetch_add(1, Ordering::Relaxed);
Ok(self.map.remove(key).map(|(_, v)| v))
}
pub fn snapshot(&self) -> Result<Vec<(K, V)>, String> {
Ok(self
.map
.iter()
.map(|entry| (entry.key().clone(), entry.value().clone()))
.collect())
}
pub fn len(&self) -> usize {
self.map.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn keys(&self) -> Result<Vec<K>, String> {
Ok(self.map.iter().map(|entry| entry.key().clone()).collect())
}
pub fn for_each<F>(&self, mut f: F) -> Result<(), String>
where
F: FnMut(&K, &V),
{
for entry in self.map.iter() {
f(entry.key(), entry.value());
}
Ok(())
}
pub fn stats(&self) -> LockFreeStats {
LockFreeStats {
reads: AtomicU64::new(self.stats.reads.load(Ordering::Relaxed)),
writes: AtomicU64::new(self.stats.writes.load(Ordering::Relaxed)),
cas_retries: AtomicU64::new(self.stats.cas_retries.load(Ordering::Relaxed)),
avg_read_latency_ns: AtomicU64::new(
self.stats.avg_read_latency_ns.load(Ordering::Relaxed),
),
}
}
}
impl LockFreeStats {
fn new() -> Self {
Self {
reads: AtomicU64::new(0),
writes: AtomicU64::new(0),
cas_retries: AtomicU64::new(0),
avg_read_latency_ns: AtomicU64::new(0),
}
}
pub fn cas_success_rate(&self) -> f64 {
let writes = self.writes.load(Ordering::Relaxed) as f64;
let retries = self.cas_retries.load(Ordering::Relaxed) as f64;
if writes == 0.0 {
1.0
} else {
writes / (writes + retries)
}
}
}
pub struct LockFreeQueue<T>
where
T: Send + Sync + 'static,
{
sender: Sender<T>,
receiver: Receiver<T>,
stats: Arc<QueueStats>,
}
#[derive(Debug)]
pub struct QueueStats {
pub enqueued: AtomicU64,
pub dequeued: AtomicU64,
pub current_size: AtomicUsize,
}
impl<T> LockFreeQueue<T>
where
T: Send + Sync + 'static,
{
pub fn new() -> Self {
let (sender, receiver) = unbounded();
Self {
sender,
receiver,
stats: Arc::new(QueueStats {
enqueued: AtomicU64::new(0),
dequeued: AtomicU64::new(0),
current_size: AtomicUsize::new(0),
}),
}
}
pub fn push(&self, item: T) -> Result<(), TransportError> {
match self.sender.send(item) {
Ok(_) => {
self.stats.enqueued.fetch_add(1, Ordering::Relaxed);
self.stats.current_size.fetch_add(1, Ordering::Relaxed);
Ok(())
}
Err(_) => Err(TransportError::resource_error("queue_push", 1, 0)),
}
}
pub fn pop(&self) -> Option<T> {
match self.receiver.try_recv() {
Ok(item) => {
self.stats.dequeued.fetch_add(1, Ordering::Relaxed);
self.stats.current_size.fetch_sub(1, Ordering::Relaxed);
Some(item)
}
Err(_) => None,
}
}
pub fn len(&self) -> usize {
self.stats.current_size.load(Ordering::Relaxed)
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn stats(&self) -> &QueueStats {
&self.stats
}
}
pub struct LockFreeCounter {
value: AtomicUsize,
stats: Arc<CounterStats>,
}
#[derive(Debug)]
pub struct CounterStats {
pub increments: AtomicU64,
pub decrements: AtomicU64,
pub reads: AtomicU64,
}
impl LockFreeCounter {
pub fn new(initial: usize) -> Self {
Self {
value: AtomicUsize::new(initial),
stats: Arc::new(CounterStats {
increments: AtomicU64::new(0),
decrements: AtomicU64::new(0),
reads: AtomicU64::new(0),
}),
}
}
pub fn increment(&self) -> usize {
self.stats.increments.fetch_add(1, Ordering::Relaxed);
self.value.fetch_add(1, Ordering::Relaxed) + 1
}
pub fn decrement(&self) -> usize {
self.stats.decrements.fetch_add(1, Ordering::Relaxed);
self.value.fetch_sub(1, Ordering::Relaxed).saturating_sub(1)
}
pub fn get(&self) -> usize {
self.stats.reads.fetch_add(1, Ordering::Relaxed);
self.value.load(Ordering::Relaxed)
}
pub fn set(&self, value: usize) {
self.value.store(value, Ordering::Relaxed);
}
pub fn swap(&self, value: usize) -> usize {
self.value.swap(value, Ordering::Relaxed)
}
pub fn stats(&self) -> &CounterStats {
&self.stats
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::thread;
#[test]
fn test_lockfree_hashmap_basic() {
let map = LockFreeHashMap::new();
assert!(map.insert("key1".to_string(), "value1".to_string()).is_ok());
assert_eq!(map.get(&"key1".to_string()), Some("value1".to_string()));
assert!(map.insert("key1".to_string(), "value2".to_string()).is_ok());
assert_eq!(map.get(&"key1".to_string()), Some("value2".to_string()));
assert!(map.remove(&"key1".to_string()).is_ok());
assert_eq!(map.get(&"key1".to_string()), None);
}
#[test]
fn test_lockfree_hashmap_concurrent() {
let map = Arc::new(LockFreeHashMap::new());
let mut handles = vec![];
for i in 0..10 {
let map_clone = Arc::clone(&map);
let handle = thread::spawn(move || {
for j in 0..100 {
let key = format!("key_{}", i * 100 + j);
let value = format!("value_{}", i * 100 + j);
map_clone.insert(key, value).unwrap();
}
});
handles.push(handle);
}
for i in 0..5 {
let map_clone = Arc::clone(&map);
let handle = thread::spawn(move || {
for _ in 0..1000 {
let key = format!("key_{}", i);
let _value = map_clone.get(&key);
}
});
handles.push(handle);
}
for handle in handles {
handle.join().unwrap();
}
assert_eq!(map.len(), 1000);
let stats = map.stats();
println!("CAS success rate: {:.2}%", stats.cas_success_rate() * 100.0);
}
#[test]
fn test_lockfree_queue() {
let queue = LockFreeQueue::new();
assert!(queue.push(1).is_ok());
assert!(queue.push(2).is_ok());
assert_eq!(queue.len(), 2);
assert_eq!(queue.pop(), Some(1));
assert_eq!(queue.pop(), Some(2));
assert_eq!(queue.pop(), None);
assert!(queue.is_empty());
}
#[test]
fn test_lockfree_counter() {
let counter = LockFreeCounter::new(0);
assert_eq!(counter.increment(), 1);
assert_eq!(counter.increment(), 2);
assert_eq!(counter.get(), 2);
assert_eq!(counter.decrement(), 1);
assert_eq!(counter.get(), 1);
}
}