use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use bytes::Bytes;
use tokio::sync::RwLock;
use tracing::debug;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub(crate) struct StateId {
pub raw: [u8; 16],
pub generation: u64,
}
impl StateId {
pub fn anonymous() -> Self {
Self {
raw: [0u8; 16],
generation: 0,
}
}
#[cfg(test)]
pub fn from_bytes(bytes: &[u8; 16]) -> Self {
Self {
raw: *bytes,
generation: 1,
}
}
pub fn from_bytes_at(bytes: &[u8; 16], generation: u64) -> Self {
Self {
raw: *bytes,
generation,
}
}
}
#[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,
generation: u64,
}
pub(crate) struct StateManager {
open_files: RwLock<HashMap<Bytes, OpenState>>,
active_generation: AtomicU64,
}
impl StateManager {
pub fn new() -> Self {
Self {
open_files: RwLock::new(HashMap::new()),
active_generation: AtomicU64::new(1),
}
}
pub async fn transition_to(&self, generation: u64) {
self.active_generation.store(generation, Ordering::Release);
let mut files = self.open_files.write().await;
files.clear();
}
pub async fn get_stateid(&self, fh: &Bytes) -> StateId {
let files = self.open_files.read().await;
match files.get(fh) {
Some(state) if state.generation == self.active_generation.load(Ordering::Acquire) => {
state.stateid.clone()
}
Some(_) | None => StateId::anonymous(),
}
}
pub async fn register_open(
&self,
fh: &Bytes,
stateid: StateId,
access: AccessMode,
) -> Result<(), crate::error::NfsError> {
let mut files = self.open_files.write().await;
let active = self.active_generation.load(Ordering::Acquire);
if stateid.generation != active {
return Err(crate::error::NfsError::Rpc(format!(
"cannot register OPEN from generation {}, active {active}",
stateid.generation
)));
}
let generation = stateid.generation;
match files.get_mut(fh) {
Some(existing) => {
existing.access = existing.access.upgrade(access);
existing.stateid = stateid;
existing.ref_count += 1;
existing.generation = generation;
debug!(
fh_len = fh.len(),
ref_count = existing.ref_count,
"state: reuse open"
);
}
None => {
files.insert(
fh.clone(),
OpenState {
stateid,
access,
ref_count: 1,
generation,
},
);
debug!(fh_len = fh.len(), "state: new open registered");
}
}
Ok(())
}
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.generation == self.active_generation.load(Ordering::Acquire)
&& 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
.unwrap();
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
.unwrap();
mgr.register_open(&fh, sid.clone(), AccessMode::Read)
.await
.unwrap();
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
.unwrap();
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
.unwrap();
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
.unwrap();
mgr.register_open(&fh2, sid2.clone(), AccessMode::Write)
.await
.unwrap();
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());
}
#[tokio::test]
async fn transition_rejects_late_open_from_old_generation() {
let mgr = StateManager::new();
let fh = Bytes::from_static(b"generation-fh");
let old = StateId::from_bytes_at(&[1; 16], 1);
mgr.register_open(&fh, old, AccessMode::Write)
.await
.unwrap();
mgr.transition_to(2).await;
assert_eq!(mgr.get_stateid(&fh).await, StateId::anonymous());
let late = StateId::from_bytes_at(&[2; 16], 1);
assert!(
mgr.register_open(&fh, late, AccessMode::Write)
.await
.is_err()
);
assert_eq!(mgr.get_stateid(&fh).await, StateId::anonymous());
let current = StateId::from_bytes_at(&[3; 16], 2);
mgr.register_open(&fh, current.clone(), AccessMode::Write)
.await
.unwrap();
assert_eq!(mgr.get_stateid(&fh).await, current);
}
}