#![forbid(unsafe_code)]
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecordLock {
pub owner: u64,
pub start: u64,
pub end: u64,
pub typ: i32,
pub pid: u32,
}
#[derive(Debug, Default)]
pub struct LockTable {
locks: HashMap<u64, Vec<RecordLock>>,
}
pub const F_RDLCK: i32 = 0;
pub const F_WRLCK: i32 = 1;
pub const F_UNLCK: i32 = 2;
impl LockTable {
pub fn new() -> Self {
Self::default()
}
fn conflicts(typ: i32, existing: &RecordLock) -> bool {
typ == F_WRLCK || existing.typ == F_WRLCK
}
fn ranges_overlap(a: &RecordLock, start: u64, end: u64) -> bool {
a.start <= end && start <= a.end
}
pub fn setlk(
&mut self,
ino: u64,
owner: u64,
start: u64,
end: u64,
typ: i32,
pid: u32,
) -> Result<bool, String> {
if end < start {
return Err("lock end before start".into());
}
let locks = self.locks.entry(ino).or_default();
if typ == F_UNLCK {
locks.retain(|l| !(l.owner == owner && Self::ranges_overlap(l, start, end)));
return Ok(true);
}
if typ != F_RDLCK && typ != F_WRLCK {
return Err(format!("unknown lock type {typ}"));
}
let probe = RecordLock {
owner,
start,
end,
typ,
pid,
};
for existing in locks.iter() {
if existing.owner == owner {
continue;
}
if Self::ranges_overlap(existing, start, end) && Self::conflicts(typ, existing) {
return Ok(false);
}
}
locks.retain(|l| !(l.owner == owner && Self::ranges_overlap(l, start, end)));
locks.push(probe);
Ok(true)
}
pub fn getlk(
&self,
ino: u64,
owner: u64,
start: u64,
end: u64,
typ: i32,
) -> Option<RecordLock> {
let locks = self.locks.get(&ino)?;
for existing in locks.iter() {
if existing.owner == owner {
continue;
}
if Self::ranges_overlap(existing, start, end) && Self::conflicts(typ, existing) {
return Some(*existing);
}
}
None
}
pub fn release_owner(&mut self, owner: u64) {
for locks in self.locks.values_mut() {
locks.retain(|l| l.owner != owner);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn write_lock_blocks_read_and_write() {
let mut t = LockTable::new();
assert!(t.setlk(1, 10, 0, 100, F_WRLCK, 1).unwrap());
assert!(!t.setlk(1, 20, 50, 60, F_RDLCK, 2).unwrap());
assert!(!t.setlk(1, 20, 50, 60, F_WRLCK, 2).unwrap());
assert!(t.setlk(1, 20, 200, 300, F_RDLCK, 2).unwrap());
assert!(t.setlk(1, 10, 0, 199, F_WRLCK, 1).unwrap());
}
#[test]
fn read_locks_share() {
let mut t = LockTable::new();
assert!(t.setlk(1, 10, 0, 100, F_RDLCK, 1).unwrap());
assert!(t.setlk(1, 20, 0, 100, F_RDLCK, 2).unwrap());
assert!(!t.setlk(1, 30, 50, 60, F_WRLCK, 3).unwrap());
}
#[test]
fn unlock_and_getlk() {
let mut t = LockTable::new();
t.setlk(1, 10, 0, 100, F_WRLCK, 1).unwrap();
let got = t.getlk(1, 20, 50, 60, F_RDLCK);
assert!(got.is_some());
assert_eq!(got.unwrap().pid, 1);
t.setlk(1, 10, 0, 100, F_UNLCK, 1).unwrap();
assert!(t.getlk(1, 20, 50, 60, F_RDLCK).is_none());
}
#[test]
fn flush_releases_owner() {
let mut t = LockTable::new();
t.setlk(1, 10, 0, 100, F_WRLCK, 1).unwrap();
t.setlk(1, 20, 0, 100, F_RDLCK, 2).unwrap();
t.release_owner(10);
assert!(t.getlk(1, 30, 0, 100, F_RDLCK).is_none());
}
}