use std::sync::atomic::{AtomicU64, Ordering};
pub struct FileAllocator {
eof: AtomicU64,
alignment: u64,
}
impl FileAllocator {
pub fn new(initial_eof: u64) -> Self {
Self {
eof: AtomicU64::new(initial_eof),
alignment: 8,
}
}
pub fn allocate(&self, size: u64) -> u64 {
let mut cur = self.eof.load(Ordering::Acquire);
loop {
let aligned = (cur + self.alignment - 1) & !(self.alignment - 1);
let next = aligned + size;
match self
.eof
.compare_exchange_weak(cur, next, Ordering::AcqRel, Ordering::Acquire)
{
Ok(_) => return aligned,
Err(actual) => cur = actual,
}
}
}
pub fn eof(&self) -> u64 {
self.eof.load(Ordering::Acquire)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn basic_allocation() {
let alloc = FileAllocator::new(48);
let a = alloc.allocate(100);
assert_eq!(a, 48);
assert_eq!(alloc.eof(), 148);
}
#[test]
fn alignment() {
let alloc = FileAllocator::new(50); let a = alloc.allocate(10);
assert_eq!(a, 56); assert_eq!(alloc.eof(), 66);
}
#[test]
fn zero_size_allocation() {
let alloc = FileAllocator::new(48);
let a = alloc.allocate(0);
assert_eq!(a, 48);
assert_eq!(alloc.eof(), 48);
}
#[test]
fn successive_allocations() {
let alloc = FileAllocator::new(0);
let a1 = alloc.allocate(10);
let a2 = alloc.allocate(20);
let a3 = alloc.allocate(5);
assert_eq!(a1, 0);
assert_eq!(a2, 16); assert_eq!(a3, 40); }
#[test]
fn concurrent_allocations_are_disjoint() {
use std::sync::Arc;
use std::thread;
let alloc = Arc::new(FileAllocator::new(0));
let n_threads = 8;
let per_thread = 1000;
let size = 7u64;
let mut handles = Vec::new();
for _ in 0..n_threads {
let a = Arc::clone(&alloc);
handles.push(thread::spawn(move || {
let mut offs = Vec::with_capacity(per_thread);
for _ in 0..per_thread {
offs.push(a.allocate(size));
}
offs
}));
}
let mut all: Vec<u64> = handles
.into_iter()
.flat_map(|h| h.join().unwrap())
.collect();
all.sort_unstable();
for w in all.windows(2) {
assert_eq!(w[0] % 8, 0, "offset {} not 8-aligned", w[0]);
assert!(
w[1] >= w[0] + size,
"ranges overlap: {} + {} > {}",
w[0],
size,
w[1]
);
}
assert_eq!(all.len(), n_threads * per_thread);
let unique = all.iter().collect::<std::collections::HashSet<_>>().len();
assert_eq!(unique, all.len(), "duplicate offsets handed out");
}
}