use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::time::Instant;
use bytes::Bytes;
use dashmap::DashMap;
use crate::rendezvous::protocol::DataMetadata;
pub const DEFAULT_CHUNK_SIZE: u32 = 512 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StageMode {
InMemory,
Pinned,
}
#[allow(dead_code)]
pub(crate) struct DataSlot {
pub data: Bytes,
pub mode: StageMode,
pub refcount: AtomicU32,
pub read_lock_count: AtomicU32,
pub total_len: u64,
pub created_at: Instant,
pub ttl: Option<std::time::Duration>,
}
#[allow(dead_code)]
pub(crate) struct TransferState {
pub slot_local_id: u64,
pub lease_id: u64,
pub chunk_size: u32,
pub chunk_count: u32,
pub created_at: Instant,
}
#[derive(Debug, Clone)]
pub struct RegisterOptions {
pub ttl: Option<std::time::Duration>,
}
impl RegisterOptions {
pub fn new() -> Self {
Self { ttl: None }
}
pub fn ttl(mut self, ttl: std::time::Duration) -> Self {
self.ttl = Some(ttl);
self
}
}
impl Default for RegisterOptions {
fn default() -> Self {
Self::new()
}
}
pub struct DataStore {
next_id: AtomicU64,
pub(crate) slots: DashMap<u64, DataSlot>,
pub(crate) transfers: DashMap<u64, TransferState>,
next_transfer_id: AtomicU64,
next_lease_id: AtomicU64,
active_leases: DashMap<u64, u64>,
}
impl DataStore {
pub fn new() -> Self {
Self {
next_id: AtomicU64::new(1),
slots: DashMap::new(),
transfers: DashMap::new(),
next_transfer_id: AtomicU64::new(1),
next_lease_id: AtomicU64::new(1),
active_leases: DashMap::new(),
}
}
pub fn register(&self, data: Bytes, opts: Option<RegisterOptions>) -> u64 {
let local_id = self.next_id.fetch_add(1, Ordering::Relaxed);
let total_len = data.len() as u64;
let ttl = opts.as_ref().and_then(|o| o.ttl);
self.slots.insert(
local_id,
DataSlot {
data,
mode: StageMode::InMemory,
refcount: AtomicU32::new(1),
read_lock_count: AtomicU32::new(0),
total_len,
created_at: Instant::now(),
ttl,
},
);
local_id
}
pub fn metadata(&self, local_id: u64) -> Option<DataMetadata> {
self.slots.get(&local_id).map(|slot| DataMetadata {
total_len: slot.total_len,
refcount: slot.refcount.load(Ordering::Relaxed),
pinned: slot.mode == StageMode::Pinned,
})
}
pub fn acquire_read_lock(&self, local_id: u64) -> Option<u64> {
let slot = self.slots.get(&local_id)?;
slot.read_lock_count.fetch_add(1, Ordering::Relaxed);
let lease_id = self.next_lease_id.fetch_add(1, Ordering::Relaxed);
self.active_leases.insert(lease_id, local_id);
Some(lease_id)
}
pub fn consume_lease(&self, lease_id: u64) -> Option<u64> {
self.active_leases
.remove(&lease_id)
.map(|(_, local_id)| local_id)
}
pub fn release_read_lock(&self, local_id: u64) -> bool {
if let Some(slot) = self.slots.get(&local_id) {
let result =
slot.read_lock_count
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |v| {
if v > 0 { Some(v - 1) } else { None }
});
match result {
Ok(prev) => {
let read_locks = prev - 1;
let refcount = slot.refcount.load(Ordering::Relaxed);
read_locks == 0 && refcount == 0
}
Err(_) => {
tracing::warn!(
"release_read_lock: read_lock_count already 0 for slot {local_id}"
);
false
}
}
} else {
false
}
}
pub fn ref_increment(&self, local_id: u64) -> bool {
if let Some(slot) = self.slots.get(&local_id) {
slot.refcount.fetch_add(1, Ordering::Relaxed);
true
} else {
false
}
}
pub fn ref_decrement(&self, local_id: u64) -> bool {
if let Some(slot) = self.slots.get(&local_id) {
let result = slot
.refcount
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |v| {
if v > 0 { Some(v - 1) } else { None }
});
match result {
Ok(prev) => {
let refcount = prev - 1;
let read_locks = slot.read_lock_count.load(Ordering::Relaxed);
refcount == 0 && read_locks == 0
}
Err(_) => {
tracing::warn!("ref_decrement: refcount already 0 for slot {local_id}");
false
}
}
} else {
false
}
}
pub fn remove(&self, local_id: u64) -> Option<Bytes> {
self.slots.remove(&local_id).map(|(_, slot)| slot.data)
}
pub fn try_free(&self, local_id: u64) {
self.slots.remove_if(&local_id, |_, slot| {
slot.refcount.load(Ordering::Relaxed) == 0
&& slot.read_lock_count.load(Ordering::Relaxed) == 0
});
}
pub fn get_data(&self, local_id: u64) -> Option<Bytes> {
self.slots.get(&local_id).map(|slot| slot.data.clone())
}
pub fn get_total_len(&self, local_id: u64) -> Option<u64> {
self.slots.get(&local_id).map(|slot| slot.total_len)
}
pub fn create_transfer(
&self,
local_id: u64,
lease_id: u64,
max_chunk_size: u32,
) -> Option<(u64, u32, u32)> {
let slot = self.slots.get(&local_id)?;
let total_len = slot.total_len;
let chunk_size = max_chunk_size.min(DEFAULT_CHUNK_SIZE);
let chunk_count = total_len.div_ceil(chunk_size as u64) as u32;
let transfer_id = self.next_transfer_id.fetch_add(1, Ordering::Relaxed);
self.transfers.insert(
transfer_id,
TransferState {
slot_local_id: local_id,
lease_id,
chunk_size,
chunk_count,
created_at: Instant::now(),
},
);
Some((transfer_id, chunk_size, chunk_count))
}
pub fn get_chunk(&self, transfer_id: u64, chunk_index: u32) -> Option<Bytes> {
let transfer = self.transfers.get(&transfer_id)?;
let slot = self.slots.get(&transfer.slot_local_id)?;
let offset = chunk_index as u64 * transfer.chunk_size as u64;
let end = (offset + transfer.chunk_size as u64).min(slot.total_len);
if offset >= slot.total_len {
return None;
}
Some(slot.data.slice(offset as usize..end as usize))
}
pub fn remove_transfer(&self, transfer_id: u64) {
self.transfers.remove(&transfer_id);
}
pub fn remove_transfers_by_lease(&self, lease_id: u64) {
self.transfers.retain(|_, state| state.lease_id != lease_id);
}
}
impl Default for DataStore {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_register_and_get() {
let store = DataStore::new();
let data = Bytes::from(vec![1u8, 2, 3, 4]);
let id = store.register(data.clone(), None);
assert_eq!(id, 1);
assert_eq!(store.get_data(id).unwrap(), data);
}
#[test]
fn test_metadata() {
let store = DataStore::new();
let data = Bytes::from(vec![0u8; 1024]);
let id = store.register(data, None);
let meta = store.metadata(id).unwrap();
assert_eq!(meta.total_len, 1024);
assert_eq!(meta.refcount, 1);
assert!(!meta.pinned);
}
#[test]
fn test_ref_counting() {
let store = DataStore::new();
let id = store.register(Bytes::from("hello"), None);
assert_eq!(store.metadata(id).unwrap().refcount, 1);
assert!(store.ref_increment(id));
assert_eq!(store.metadata(id).unwrap().refcount, 2);
assert!(!store.ref_decrement(id));
assert!(store.ref_decrement(id));
store.try_free(id);
assert!(store.metadata(id).is_none());
}
#[test]
fn test_read_lock_prevents_free() {
let store = DataStore::new();
let id = store.register(Bytes::from("data"), None);
let _lease = store.acquire_read_lock(id).unwrap();
let should_free = store.ref_decrement(id);
assert!(!should_free);
let should_free = store.release_read_lock(id);
assert!(should_free);
}
#[test]
fn test_chunked_transfer() {
let store = DataStore::new();
let data = Bytes::from(vec![0xAA; 2000]);
let id = store.register(data, None);
let lease_id = store.acquire_read_lock(id).unwrap();
let (transfer_id, chunk_size, chunk_count) =
store.create_transfer(id, lease_id, 1024).unwrap();
assert_eq!(chunk_size, 1024);
assert_eq!(chunk_count, 2);
let chunk0 = store.get_chunk(transfer_id, 0).unwrap();
assert_eq!(chunk0.len(), 1024);
assert!(chunk0.iter().all(|&b| b == 0xAA));
let chunk1 = store.get_chunk(transfer_id, 1).unwrap();
assert_eq!(chunk1.len(), 976); assert!(chunk1.iter().all(|&b| b == 0xAA));
assert!(store.get_chunk(transfer_id, 2).is_none());
store.remove_transfer(transfer_id);
assert!(store.transfers.get(&transfer_id).is_none());
}
}