use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Mutex, MutexGuard};
use crate::storage::{ReadOptions, Storage, StorageChangeWatch, StorageError, WriteOptions};
static NEXT_SESSION_TOKEN: AtomicU64 = AtomicU64::new(1);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct StorageSessionToken(u64);
impl StorageSessionToken {
fn next() -> Self {
Self(NEXT_SESSION_TOKEN.fetch_add(1, Ordering::Relaxed))
}
pub fn to_decimal_string(self) -> String {
self.0.to_string()
}
pub fn from_decimal_string(value: &str) -> Option<Self> {
let parsed = value.parse::<u64>().ok()?;
(parsed.to_string() == value).then_some(Self(parsed))
}
}
#[derive(Clone, Debug)]
pub struct StorageSession<S> {
storage: S,
token: StorageSessionToken,
}
impl<S> StorageSession<S>
where
S: Storage,
{
pub async fn acquire(storage: S) -> Result<Self, StorageError> {
let token = storage.acquire_session().await?;
Ok(Self { storage, token })
}
pub fn token(&self) -> StorageSessionToken {
self.token
}
}
impl<S> Storage for StorageSession<S>
where
S: Storage,
{
type Read<'a>
= S::Read<'a>
where
Self: 'a;
type Write<'a>
= S::Write<'a>
where
Self: 'a;
async fn acquire_session(&self) -> Result<StorageSessionToken, StorageError> {
Ok(self.token)
}
async fn acquire_partial_replica_owner(
&self,
session: StorageSessionToken,
) -> Result<crate::storage::StorageOwnerLease, StorageError> {
if session != self.token {
return Err(StorageError::Fenced);
}
self.storage.acquire_partial_replica_owner(self.token).await
}
fn begin_read(
&self,
mut opts: ReadOptions,
) -> impl Future<Output = Result<Self::Read<'_>, StorageError>> + Send {
opts.session_token = Some(self.token);
self.storage.begin_read(opts)
}
fn begin_write(
&self,
mut opts: WriteOptions,
) -> impl Future<Output = Result<Self::Write<'_>, StorageError>> + Send {
opts.session_token = Some(self.token);
self.storage.begin_write(opts)
}
fn watch_for_changes(
&self,
) -> impl Future<Output = Result<StorageChangeWatch, StorageError>> + Send {
self.storage.watch_for_changes()
}
}
#[derive(Debug, Default)]
pub struct StorageSessionGate {
current: Mutex<Option<StorageSessionToken>>,
}
impl StorageSessionGate {
pub fn acquire(&self) -> Result<StorageSessionToken, StorageError> {
let mut current = self.lock()?;
Ok(*current.get_or_insert_with(StorageSessionToken::next))
}
pub fn validate(
&self,
token: Option<StorageSessionToken>,
) -> Result<StorageSessionPermit<'_>, StorageError> {
let current = self.lock()?;
match (*current, token) {
(None, None) | (Some(_), Some(_)) if *current == token => {
Ok(StorageSessionPermit { _barrier: current })
}
_ => Err(StorageError::Fenced),
}
}
fn lock(&self) -> Result<MutexGuard<'_, Option<StorageSessionToken>>, StorageError> {
self.current
.lock()
.map_err(|_| StorageError::Io("storage session gate lock poisoned".to_string()))
}
}
#[must_use = "dropping the permit releases the acquisition barrier"]
#[derive(Debug)]
pub struct StorageSessionPermit<'a> {
_barrier: MutexGuard<'a, Option<StorageSessionToken>>,
}
#[cfg(test)]
mod tests {
use super::StorageSessionToken;
#[test]
fn session_token_decimal_encoding_preserves_u64_precision() {
let encoded = u64::MAX.to_string();
let token = StorageSessionToken::from_decimal_string(&encoded).unwrap();
assert_eq!(token.to_decimal_string(), encoded);
assert!(StorageSessionToken::from_decimal_string("not-a-token").is_none());
assert!(StorageSessionToken::from_decimal_string("01").is_none());
assert!(StorageSessionToken::from_decimal_string("+1").is_none());
assert!(StorageSessionToken::from_decimal_string("18446744073709551616").is_none());
}
}