use std::collections::HashMap;
use bytes::Bytes;
use tokio::sync::RwLock;
use tracing::debug;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub(crate) struct StateId {
pub raw: [u8; 16],
}
impl StateId {
pub fn anonymous() -> Self {
Self { raw: [0u8; 16] }
}
pub fn from_bytes(bytes: &[u8; 16]) -> Self {
Self { raw: *bytes }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum AccessMode {
Read,
Write,
Both,
}
impl AccessMode {
pub fn share_access(&self) -> u32 {
match self {
AccessMode::Read => 0x00000001, AccessMode::Write => 0x00000002, AccessMode::Both => 0x00000003, }
}
pub fn covers(&self, requested: AccessMode) -> bool {
matches!(
(self, requested),
(AccessMode::Both, _)
| (AccessMode::Read, AccessMode::Read)
| (AccessMode::Write, AccessMode::Write)
)
}
pub fn upgrade(&self, requested: AccessMode) -> AccessMode {
match (self, requested) {
(AccessMode::Both, _) | (_, AccessMode::Both) => AccessMode::Both,
(AccessMode::Read, AccessMode::Write) | (AccessMode::Write, AccessMode::Read) => {
AccessMode::Both
}
_ => *self,
}
}
}
struct OpenState {
stateid: StateId,
access: AccessMode,
ref_count: u32,
}
pub(crate) struct StateManager {
open_files: RwLock<HashMap<Bytes, OpenState>>,
}
impl StateManager {
pub fn new() -> Self {
Self {
open_files: RwLock::new(HashMap::new()),
}
}
pub async fn get_stateid(&self, fh: &Bytes) -> StateId {
let files = self.open_files.read().await;
match files.get(fh) {
Some(state) => state.stateid.clone(),
None => StateId::anonymous(),
}
}
pub async fn register_open(&self, fh: &Bytes, stateid: StateId, access: AccessMode) {
let mut files = self.open_files.write().await;
match files.get_mut(fh) {
Some(existing) => {
existing.access = existing.access.upgrade(access);
existing.stateid = stateid;
existing.ref_count += 1;
debug!(
fh_len = fh.len(),
ref_count = existing.ref_count,
"state: reuse open"
);
}
None => {
files.insert(
fh.clone(),
OpenState {
stateid,
access,
ref_count: 1,
},
);
debug!(fh_len = fh.len(), "state: new open registered");
}
}
}
pub async fn release(&self, fh: &Bytes) -> Option<StateId> {
let mut files = self.open_files.write().await;
if let Some(state) = files.get_mut(fh) {
state.ref_count = state.ref_count.saturating_sub(1);
if state.ref_count == 0 {
let sid = state.stateid.clone();
files.remove(fh);
debug!(fh_len = fh.len(), "state: last ref released, should CLOSE");
return Some(sid);
}
debug!(
fh_len = fh.len(),
ref_count = state.ref_count,
"state: ref released"
);
}
None
}
pub async fn drain(&self) -> Vec<(Bytes, StateId)> {
let mut files = self.open_files.write().await;
let pairs: Vec<(Bytes, StateId)> = files
.iter()
.map(|(fh, state)| (fh.clone(), state.stateid.clone()))
.collect();
files.clear();
pairs
}
pub async fn clear(&self) {
let mut files = self.open_files.write().await;
files.clear();
}
pub async fn has_open(&self, fh: &Bytes, access: AccessMode) -> Option<StateId> {
let files = self.open_files.read().await;
files.get(fh).and_then(|state| {
if state.access.covers(access) {
Some(state.stateid.clone())
} else {
None
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn register_and_get() {
let mgr = StateManager::new();
let fh = Bytes::from_static(b"test_fh");
let sid = StateId::from_bytes(&[1u8; 16]);
assert_eq!(mgr.get_stateid(&fh).await, StateId::anonymous());
mgr.register_open(&fh, sid.clone(), AccessMode::Read).await;
assert_eq!(mgr.get_stateid(&fh).await, sid);
}
#[tokio::test]
async fn ref_count_tracking() {
let mgr = StateManager::new();
let fh = Bytes::from_static(b"test_fh");
let sid = StateId::from_bytes(&[2u8; 16]);
mgr.register_open(&fh, sid.clone(), AccessMode::Read).await;
mgr.register_open(&fh, sid.clone(), AccessMode::Read).await;
assert!(mgr.release(&fh).await.is_none());
assert_eq!(mgr.release(&fh).await, Some(sid));
assert_eq!(mgr.get_stateid(&fh).await, StateId::anonymous());
}
#[tokio::test]
async fn access_upgrade() {
let mgr = StateManager::new();
let fh = Bytes::from_static(b"test_fh");
let sid1 = StateId::from_bytes(&[3u8; 16]);
let sid2 = StateId::from_bytes(&[4u8; 16]);
mgr.register_open(&fh, sid1, AccessMode::Read).await;
assert!(mgr.has_open(&fh, AccessMode::Read).await.is_some());
assert!(mgr.has_open(&fh, AccessMode::Write).await.is_none());
mgr.register_open(&fh, sid2.clone(), AccessMode::Write)
.await;
assert!(mgr.has_open(&fh, AccessMode::Read).await.is_some());
assert!(mgr.has_open(&fh, AccessMode::Write).await.is_some());
assert!(mgr.has_open(&fh, AccessMode::Both).await.is_some());
}
#[tokio::test]
async fn drain_returns_all_and_clears() {
let mgr = StateManager::new();
let fh1 = Bytes::from_static(b"f1");
let fh2 = Bytes::from_static(b"f2");
let sid1 = StateId::from_bytes(&[10u8; 16]);
let sid2 = StateId::from_bytes(&[20u8; 16]);
mgr.register_open(&fh1, sid1.clone(), AccessMode::Read)
.await;
mgr.register_open(&fh2, sid2.clone(), AccessMode::Write)
.await;
let pairs = mgr.drain().await;
assert_eq!(pairs.len(), 2);
assert_eq!(mgr.get_stateid(&fh1).await, StateId::anonymous());
assert_eq!(mgr.get_stateid(&fh2).await, StateId::anonymous());
}
}