use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::{Arc, LazyLock};
use parking_lot::Mutex;
static REGISTRY: LazyLock<Mutex<HashMap<PathBuf, Arc<Mutex<()>>>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
struct Reclaim<'a> {
db_path: &'a Path,
lock: Arc<Mutex<()>>,
}
impl Drop for Reclaim<'_> {
fn drop(&mut self) {
let mut registry = REGISTRY.lock();
if Arc::strong_count(&self.lock) == 2 {
registry.remove(self.db_path);
}
}
}
pub fn with_data_frame_write<T>(db_path: &Path, work: impl FnOnce() -> T) -> T {
let reclaim = Reclaim {
db_path,
lock: REGISTRY
.lock()
.entry(db_path.to_path_buf())
.or_insert_with(|| Arc::new(Mutex::new(())))
.clone(),
};
let _guard = reclaim.lock.lock();
work()
}
#[cfg(test)]
fn registry_holds(db_path: &Path) -> bool {
REGISTRY.lock().contains_key(db_path)
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn test_writes_to_one_data_frame_do_not_overlap() {
let db_path = PathBuf::from("/data-frame-locks-test/one");
let inside = Arc::new(AtomicUsize::new(0));
let max_inside = Arc::new(AtomicUsize::new(0));
std::thread::scope(|scope| {
for _ in 0..8 {
let db_path = db_path.clone();
let inside = inside.clone();
let max_inside = max_inside.clone();
scope.spawn(move || {
for _ in 0..500 {
with_data_frame_write(&db_path, || {
let now = inside.fetch_add(1, Ordering::SeqCst) + 1;
max_inside.fetch_max(now, Ordering::SeqCst);
std::hint::spin_loop();
inside.fetch_sub(1, Ordering::SeqCst);
});
}
});
}
});
assert_eq!(
max_inside.load(Ordering::SeqCst),
1,
"two row writes on the same data frame overlapped"
);
}
#[test]
fn test_registry_reclaims_a_data_frame_once_its_writers_finish() {
let db_path = PathBuf::from("/data-frame-locks-test/reclaimed");
with_data_frame_write(&db_path, || {
assert!(
registry_holds(&db_path),
"a write in flight must hold its registry entry"
);
});
assert!(
!registry_holds(&db_path),
"the registry kept a lock for a data frame with no writers"
);
std::thread::scope(|scope| {
for _ in 0..8 {
let db_path = db_path.clone();
scope.spawn(move || {
for _ in 0..100 {
with_data_frame_write(&db_path, std::hint::spin_loop);
}
});
}
});
assert!(
!registry_holds(&db_path),
"the registry kept a lock after concurrent writers finished"
);
}
#[test]
fn test_a_panic_inside_the_guarded_work_still_reclaims() {
let db_path = PathBuf::from("/data-frame-locks-test/panicked");
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
with_data_frame_write(&db_path, || {
panic!("deliberate panic under the data frame lock")
});
}));
assert!(outcome.is_err(), "the panic should have propagated");
assert!(
!registry_holds(&db_path),
"a panic left its registry entry behind"
);
with_data_frame_write(&db_path, || {});
}
#[test]
fn test_writes_to_different_data_frames_take_different_locks() {
let b = PathBuf::from("/data-frame-locks-test/b");
let reached_inner =
with_data_frame_write(&PathBuf::from("/data-frame-locks-test/a"), || {
with_data_frame_write(&b, || true)
});
assert!(reached_inner);
}
}