use fs_core::{BlockRead, CachingDevice, Result};
use std::sync::{Arc, Condvar, Mutex};
use std::time::{Duration, Instant};
mod common;
const BS: u64 = 8;
const TARGET: u64 = 24;
const RACERS: usize = 4;
const CAPACITY: usize = 4;
const PRIMED: [u64; 3] = [0, 8, 16];
const SETTLE: Duration = Duration::from_secs(2);
const JOIN_DEADLINE: Duration = Duration::from_secs(10);
struct GateState {
arrived: usize,
released: bool,
}
struct Gate {
state: Mutex<GateState>,
changed: Condvar,
want: usize,
}
impl Gate {
fn new(want: usize) -> Self {
Gate {
state: Mutex::new(GateState {
arrived: 0,
released: false,
}),
changed: Condvar::new(),
want,
}
}
fn park(&self) {
let mut g = self.state.lock().expect("gate lock");
g.arrived += 1;
self.changed.notify_all();
while g.arrived < self.want && !g.released {
g = self.changed.wait(g).expect("gate wait");
}
}
fn open_when_full_or_settled(&self) -> usize {
let mut g = self.state.lock().expect("gate lock");
let deadline = Instant::now() + SETTLE;
while g.arrived < self.want {
let left = deadline.saturating_duration_since(Instant::now());
if left.is_zero() {
break;
}
let (next, _) = self.changed.wait_timeout(g, left).expect("gate wait");
g = next;
}
let inside = g.arrived;
g.released = true;
self.changed.notify_all();
inside
}
}
struct GatedDevice {
bytes: Vec<u8>,
gate: Arc<Gate>,
target: u64,
reads: Mutex<Vec<u64>>,
}
impl GatedDevice {
fn new(bytes: Vec<u8>, gate: Arc<Gate>, target: u64) -> Self {
GatedDevice {
bytes,
gate,
target,
reads: Mutex::new(Vec::new()),
}
}
fn reads_at(&self, offset: u64) -> usize {
self.reads
.lock()
.expect("reads lock")
.iter()
.filter(|o| **o == offset)
.count()
}
fn total_reads(&self) -> usize {
self.reads.lock().expect("reads lock").len()
}
}
impl BlockRead for GatedDevice {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
self.reads.lock().expect("reads lock").push(offset);
if offset == self.target {
self.gate.park();
}
common::read_into(&self.bytes, offset, buf)
}
fn size_bytes(&self) -> u64 {
self.bytes.len() as u64
}
}
fn read_block(cache: &CachingDevice, offset: u64) -> Vec<u8> {
let mut buf = vec![0u8; BS as usize];
cache
.read_at(offset, &mut buf)
.unwrap_or_else(|e| panic!("read at {offset}: {e:?}"));
buf
}
fn expected_block(bytes: &[u8], offset: u64) -> Vec<u8> {
bytes[offset as usize..(offset + BS) as usize].to_vec()
}
#[test]
fn concurrent_misses_on_one_block_read_the_device_once_and_evict_nothing() {
let bytes: Vec<u8> = (0..64u8).collect();
let gate = Arc::new(Gate::new(RACERS));
let device = Arc::new(GatedDevice::new(bytes.clone(), Arc::clone(&gate), TARGET));
let cache = CachingDevice::read_only(Arc::clone(&device) as Arc<dyn BlockRead>, BS, CAPACITY);
for off in PRIMED {
assert_eq!(read_block(&cache, off), expected_block(&bytes, off));
}
assert_eq!(
device.total_reads(),
PRIMED.len(),
"priming should be one device read per block"
);
for off in PRIMED {
read_block(&cache, off);
}
assert_eq!(
device.total_reads(),
PRIMED.len(),
"the primed blocks must be served from the cache before the race, \
or the eviction assertion at the end proves nothing"
);
let (tx, rx) = std::sync::mpsc::channel();
for n in 0..RACERS {
let cache = Arc::clone(&cache);
let tx = tx.clone();
std::thread::spawn(move || {
let _ = tx.send((n, read_block(&cache, TARGET)));
});
}
drop(tx);
let inside = gate.open_when_full_or_settled();
for _ in 0..RACERS {
let (n, got) = rx.recv_timeout(JOIN_DEADLINE).expect(
"a racer never returned. A thread waiting for a fetch that is \
already in flight must be woken when it lands",
);
assert_eq!(
got,
expected_block(&bytes, TARGET),
"racer {n} got the wrong bytes"
);
}
assert_eq!(
inside, 1,
"{inside} of {RACERS} threads were inside the device read for block \
{TARGET} at once; a block already being fetched must be waited for, \
not fetched again"
);
assert_eq!(
device.reads_at(TARGET),
1,
"block {TARGET} was read from the device {} times for one miss",
device.reads_at(TARGET)
);
let before = device.total_reads();
for off in PRIMED {
assert_eq!(read_block(&cache, off), expected_block(&bytes, off));
}
assert_eq!(
device.total_reads(),
before,
"the {} primed entries did not survive one concurrent miss -- a \
capacity-{CAPACITY} cache evicted them to hold duplicates of block \
{TARGET}",
PRIMED.len()
);
let (hits, misses) = cache.stats();
let calls = (PRIMED.len() * 3 + RACERS) as u64;
assert_eq!(
hits + misses,
calls,
"every call is one hit or one miss: {hits} + {misses} != {calls}"
);
assert_eq!(
misses,
device.total_reads() as u64,
"one miss per device fetch: {misses} misses against {} reads",
device.total_reads()
);
}
#[test]
fn a_fetch_that_fails_releases_the_block_it_was_holding() {
struct FailsOnce {
bytes: Vec<u8>,
failed: Mutex<Vec<u64>>,
}
impl BlockRead for FailsOnce {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
let mut failed = self.failed.lock().expect("lock");
if !failed.contains(&offset) {
failed.push(offset);
return Err(fs_core::Error::ShortRead {
offset,
want: buf.len(),
got: 0,
});
}
drop(failed);
common::read_into(&self.bytes, offset, buf)
}
fn size_bytes(&self) -> u64 {
self.bytes.len() as u64
}
}
let device = Arc::new(FailsOnce {
bytes: (0..64u8).collect(),
failed: Mutex::new(Vec::new()),
});
let cache = CachingDevice::read_only(device as Arc<dyn BlockRead>, BS, CAPACITY);
let mut buf = vec![0u8; BS as usize];
assert!(
cache.read_at(TARGET, &mut buf).is_err(),
"the first read of this block fails at the device"
);
let (tx, rx) = std::sync::mpsc::channel();
let retry = Arc::clone(&cache);
std::thread::spawn(move || {
let mut buf = vec![0u8; BS as usize];
let outcome = retry.read_at(TARGET, &mut buf);
let _ = tx.send(outcome.map(|()| buf));
});
let got = rx
.recv_timeout(Duration::from_secs(10))
.expect(
"a read of a block whose previous fetch failed must not block; \
the failed fetch never released it",
)
.expect("and the retry itself succeeds");
assert_eq!(got, (24..32u8).collect::<Vec<u8>>());
}
#[test]
fn a_device_that_re_enters_the_cache_reads_again_instead_of_waiting_for_itself() {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::OnceLock;
struct Reentrant {
bytes: Vec<u8>,
cache: OnceLock<std::sync::Weak<CachingDevice>>,
reads: AtomicUsize,
saw_target: AtomicUsize,
nested: AtomicUsize,
}
impl BlockRead for Reentrant {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
self.reads.fetch_add(1, Ordering::SeqCst);
if offset == TARGET && self.saw_target.fetch_add(1, Ordering::SeqCst) == 0 {
self.nested.fetch_add(1, Ordering::SeqCst);
let cache = self
.cache
.get()
.expect("the test wires this before reading")
.upgrade()
.expect("the cache outlives the read");
let mut inner = vec![0u8; buf.len()];
cache
.read_at(offset, &mut inner)
.expect("the re-entrant read itself succeeds");
assert_eq!(
inner,
self.bytes[offset as usize..offset as usize + buf.len()],
"the re-entrant read returned the wrong bytes"
);
}
common::read_into(&self.bytes, offset, buf)
}
fn size_bytes(&self) -> u64 {
self.bytes.len() as u64
}
}
let device = Arc::new(Reentrant {
bytes: (0..64u8).collect(),
cache: OnceLock::new(),
reads: AtomicUsize::new(0),
saw_target: AtomicUsize::new(0),
nested: AtomicUsize::new(0),
});
let cache = CachingDevice::read_only(Arc::clone(&device) as Arc<dyn BlockRead>, BS, CAPACITY);
device.cache.set(Arc::downgrade(&cache)).expect("set once");
for off in PRIMED {
read_block(&cache, off);
}
let (tx, rx) = std::sync::mpsc::channel();
{
let cache = Arc::clone(&cache);
std::thread::spawn(move || {
let _ = tx.send(read_block(&cache, TARGET));
});
}
let got = rx.recv_timeout(JOIN_DEADLINE).expect(
"a device that read back through the cache wrapping it waited for its \
own fetch instead of reading again -- this is the deadlock the owner \
on each in-flight marker exists to prevent",
);
assert_eq!(got, expected_block(&device.bytes, TARGET));
assert_eq!(
device.nested.load(Ordering::SeqCst),
1,
"the device must actually have re-entered the cache once, or this \
test asserts nothing"
);
assert_eq!(
device.saw_target.load(Ordering::SeqCst),
2,
"the target must have been read twice: the outer fetch and the \
re-entrant one that could not wait for it"
);
assert_eq!(
device.reads.load(Ordering::SeqCst),
PRIMED.len() + 2,
"the primed blocks, plus the outer fetch and the re-entrant one: a \
thread that cannot wait for itself reads the device again"
);
let before = device.reads.load(Ordering::SeqCst);
for off in PRIMED {
read_block(&cache, off);
}
assert_eq!(
device.reads.load(Ordering::SeqCst),
before,
"a re-entrant fetch inserted block {TARGET} twice and evicted a primed \
entry to make room for the copy"
);
}
#[test]
fn two_threads_re_entering_into_each_others_fetches_do_not_deadlock() {
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Barrier, OnceLock};
const BLOCK_A: u64 = 24;
const BLOCK_B: u64 = 32;
const PRIMED_PAIR: [u64; 2] = [0, 8];
struct CrossReentrant {
bytes: Vec<u8>,
cache: OnceLock<std::sync::Weak<CachingDevice>>,
gate: Barrier,
recursed_a: AtomicBool,
recursed_b: AtomicBool,
nested: AtomicUsize,
reads: Mutex<Vec<u64>>,
}
impl CrossReentrant {
fn total_reads(&self) -> usize {
self.reads.lock().expect("reads lock").len()
}
}
impl BlockRead for CrossReentrant {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
self.reads.lock().expect("reads lock").push(offset);
let partner = match offset {
BLOCK_A if !self.recursed_a.swap(true, Ordering::SeqCst) => Some(BLOCK_B),
BLOCK_B if !self.recursed_b.swap(true, Ordering::SeqCst) => Some(BLOCK_A),
_ => None,
};
if let Some(partner) = partner {
self.gate.wait();
self.nested.fetch_add(1, Ordering::SeqCst);
let cache = self
.cache
.get()
.expect("the test wires this before reading")
.upgrade()
.expect("the cache outlives the read");
let mut inner = vec![0u8; buf.len()];
cache
.read_at(partner, &mut inner)
.expect("the re-entrant read itself succeeds");
assert_eq!(
inner,
self.bytes[partner as usize..partner as usize + buf.len()],
"the re-entrant read of block {partner} returned the wrong bytes"
);
}
common::read_into(&self.bytes, offset, buf)
}
fn size_bytes(&self) -> u64 {
self.bytes.len() as u64
}
}
let device = Arc::new(CrossReentrant {
bytes: (0..64u8).collect(),
cache: OnceLock::new(),
gate: Barrier::new(2),
recursed_a: AtomicBool::new(false),
recursed_b: AtomicBool::new(false),
nested: AtomicUsize::new(0),
reads: Mutex::new(Vec::new()),
});
let cache = CachingDevice::read_only(Arc::clone(&device) as Arc<dyn BlockRead>, BS, CAPACITY);
device.cache.set(Arc::downgrade(&cache)).expect("set once");
for off in PRIMED_PAIR {
read_block(&cache, off);
}
assert_eq!(
device.total_reads(),
PRIMED_PAIR.len(),
"priming should be one device read per block"
);
let (tx, rx) = std::sync::mpsc::channel();
for block in [BLOCK_A, BLOCK_B] {
let cache = Arc::clone(&cache);
let tx = tx.clone();
std::thread::spawn(move || {
let _ = tx.send((block, read_block(&cache, block)));
});
}
drop(tx);
let mut returned = Vec::new();
for _ in 0..2 {
let (block, got) = rx.recv_timeout(JOIN_DEADLINE).expect(
"one of the two threads never came back. Each re-entered the cache \
for the block the other was fetching, so each waited for a fetch \
the other was holding up and neither could ever be signalled. A \
thread that already owns a fetch must read again rather than wait, \
or a cycle of length two hangs both of them",
);
assert_eq!(
got,
expected_block(&device.bytes, block),
"block {block} came back wrong"
);
returned.push(block);
}
returned.sort_unstable();
assert_eq!(
returned,
vec![BLOCK_A, BLOCK_B],
"both threads must report, once each"
);
assert_eq!(
device.nested.load(Ordering::SeqCst),
2,
"both device reads must have re-entered the cache for the other's \
block, or the cycle this test exists for never formed"
);
let before = device.total_reads();
for off in PRIMED_PAIR.iter().copied().chain([BLOCK_A, BLOCK_B]) {
read_block(&cache, off);
}
assert_eq!(
device.total_reads(),
before,
"after the race all four blocks must be cached; a duplicate entry for \
block {BLOCK_A} or {BLOCK_B} evicted a primed one to make room"
);
}