use std::sync::atomic::{AtomicU8, AtomicU64, Ordering::Relaxed};
use gxhash::HashMap as GxHashMap;
use parking_lot::Mutex;
use wdev::Device;
use wkv::{BatchStoreSession, TtlOpt};
type WatchVersion = u64;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum StoreType {
None = 0,
Main = 1,
Object = 2,
All = 3,
}
pub struct StorageSession<'a, D: Device> {
pub batch: BatchStoreSession<'a, D>,
pub session_found: AtomicU64,
pub session_notfound: AtomicU64,
pub session_pending: AtomicU64,
pub(super) pending_start_ms: AtomicU64,
pub pending_total_ms: AtomicU64,
resp_version: AtomicU8,
watch_versions: Mutex<GxHashMap<Vec<u8>, WatchVersion>>,
}
impl<'a, D: Device> StorageSession<'a, D> {
pub fn new(batch: BatchStoreSession<'a, D>) -> Self {
Self {
batch,
session_found: AtomicU64::new(0),
session_notfound: AtomicU64::new(0),
session_pending: AtomicU64::new(0),
pending_start_ms: AtomicU64::new(0),
pending_total_ms: AtomicU64::new(0),
resp_version: AtomicU8::new(2),
watch_versions: Mutex::new(GxHashMap::default()),
}
}
#[inline]
pub fn resp_protocol_version(&self) -> u8 {
self.resp_version.load(Relaxed)
}
pub fn watch_key(&self, key: &[u8]) {
let version = self.batch.store.tail_address();
self.watch_versions.lock().insert(key.to_vec(), version);
}
pub fn watched_version(&self, key: &[u8]) -> Option<WatchVersion> {
self.watch_versions.lock().get(key).copied()
}
pub fn validate_watch_version(&self) -> bool {
let now = self.batch.store.tail_address();
self.watch_versions.lock().values().all(|&v| v == now)
}
pub fn clear_watches(&self) {
self.watch_versions.lock().clear();
}
pub async fn read_string_with<R>(
&self,
key: &[u8],
f: impl Fn(&[u8]) -> R,
) -> wkv::Result<Option<R>> {
match self.batch.try_read_sync(key, &f)? {
Some(r) => {
self.record_read_outcome(r.is_some());
Ok(r)
}
None => {
self.session_pending.fetch_add(1, Relaxed);
let r = self.batch.read_with(key, f).await?;
self.record_read_outcome(r.is_some());
Ok(r)
}
}
}
pub async fn read_string(&self, key: &[u8]) -> wkv::Result<Option<Vec<u8>>> {
self.read_string_with(key, |v| v.to_vec()).await
}
pub async fn upsert_string(&self, key: &[u8], val: &[u8]) -> wkv::Result<()> {
match self.batch.try_upsert_sync(key, val)? {
Ok(_) => Ok(()),
Err(_) => self.batch.upsert(key, val).await.map(|_| ()),
}
}
pub async fn delete_string(&self, key: &[u8]) -> wkv::Result<bool> {
match self.batch.try_delete_sync(key)? {
Ok(deleted) => Ok(deleted),
Err(_) => self.batch.delete(key).await,
}
}
pub async fn expire_at_ms(&self, key: &[u8], expire_at_ms: u64) -> wkv::Result<i32> {
self.batch.expire_at(key, expire_at_ms, TtlOpt::NONE).await
}
pub async fn expire_in_ms(&self, key: &[u8], ttl_ms: u64) -> wkv::Result<i32> {
let now = coarsetime::Clock::now_since_epoch().as_millis();
self.expire_at_ms(key, now.saturating_add(ttl_ms)).await
}
pub async fn persist_key(&self, key: &[u8]) -> wkv::Result<i32> {
self.batch.persist(key).await
}
pub async fn pttl_ms(&self, key: &[u8]) -> wkv::Result<i64> {
self.batch.pttl_ms(key).await
}
pub async fn expiretime_ms(&self, key: &[u8]) -> wkv::Result<i64> {
self.batch.expiretime_ms(key).await
}
#[inline]
fn record_read_outcome(&self, found: bool) {
if found {
self.session_found.fetch_add(1, Relaxed);
} else {
self.session_notfound.fetch_add(1, Relaxed);
}
}
}
impl<'a, D: Device> StorageSession<'a, D> {
pub fn update_resp_protocol_version(&self, resp_protocol_version: u8) {
self.resp_version.store(resp_protocol_version, Relaxed);
}
}