use crate::block::{BlockDevice, BlockRead};
use crate::error::Result;
use std::collections::HashMap;
use std::sync::{Arc, Condvar, Mutex};
use std::thread::ThreadId;
pub const MAX_BLOCK_SIZE: u64 = 64 * 1024 * 1024;
pub struct CachingDevice {
inner: Arc<dyn BlockRead>,
writable: Option<Arc<dyn BlockDevice>>,
block_size: u64,
capacity: usize,
state: Mutex<CacheState>,
fetched: Condvar,
}
struct Node {
block_start: u64,
data: Arc<Vec<u8>>,
newer: Option<usize>,
older: Option<usize>,
}
struct Lru {
slots: Vec<Option<Node>>,
free: Vec<usize>,
index: HashMap<u64, usize>,
newest: Option<usize>,
oldest: Option<usize>,
}
impl Lru {
fn with_capacity(capacity: usize) -> Self {
Lru {
slots: Vec::with_capacity(capacity),
free: Vec::new(),
index: HashMap::with_capacity(capacity),
newest: None,
oldest: None,
}
}
fn len(&self) -> usize {
self.index.len()
}
fn unlink(&mut self, i: usize) {
let (newer, older) = {
let n = self.slots[i].as_ref().expect("unlink of a free slot");
(n.newer, n.older)
};
match newer {
Some(j) => self.slots[j].as_mut().expect("newer is live").older = older,
None => self.newest = older,
}
match older {
Some(j) => self.slots[j].as_mut().expect("older is live").newer = newer,
None => self.oldest = newer,
}
let n = self.slots[i].as_mut().expect("unlink of a free slot");
n.newer = None;
n.older = None;
}
fn link_newest(&mut self, i: usize) {
let old_head = self.newest;
{
let n = self.slots[i].as_mut().expect("link of a free slot");
n.newer = None;
n.older = old_head;
}
if let Some(j) = old_head {
self.slots[j].as_mut().expect("head is live").newer = Some(i);
} else {
self.oldest = Some(i);
}
self.newest = Some(i);
}
fn get(&mut self, block_start: u64) -> Option<Arc<Vec<u8>>> {
let i = *self.index.get(&block_start)?;
let data = self.slots[i]
.as_ref()
.expect("indexed slot is live")
.data
.clone();
if self.newest != Some(i) {
self.unlink(i);
self.link_newest(i);
}
Some(data)
}
fn remove(&mut self, block_start: u64) -> bool {
let Some(i) = self.index.remove(&block_start) else {
return false;
};
self.unlink(i);
self.slots[i] = None;
self.free.push(i);
true
}
fn insert(&mut self, block_start: u64, data: Arc<Vec<u8>>, capacity: usize) {
self.remove(block_start);
if self.len() >= capacity {
if let Some(oldest) = self.oldest {
let victim = self.slots[oldest]
.as_ref()
.expect("oldest is live")
.block_start;
self.remove(victim);
}
}
let node = Node {
block_start,
data,
newer: None,
older: None,
};
let i = match self.free.pop() {
Some(i) => {
self.slots[i] = Some(node);
i
}
None => {
self.slots.push(Some(node));
self.slots.len() - 1
}
};
self.index.insert(block_start, i);
self.link_newest(i);
}
fn clear(&mut self) {
self.slots.clear();
self.free.clear();
self.index.clear();
self.newest = None;
self.oldest = None;
}
fn retain_blocks(&mut self, keep: impl Fn(u64) -> bool) {
let doomed: Vec<u64> = self.index.keys().copied().filter(|b| !keep(*b)).collect();
for b in doomed {
self.remove(b);
}
}
#[cfg(test)]
fn recency_order(&self) -> Vec<u64> {
let mut out = Vec::with_capacity(self.len());
let mut cur = self.newest;
while let Some(i) = cur {
let n = self.slots[i].as_ref().expect("live");
out.push(n.block_start);
cur = n.older;
}
out
}
#[cfg(test)]
fn recency_order_reversed(&self) -> Vec<u64> {
let mut out = Vec::with_capacity(self.len());
let mut cur = self.oldest;
while let Some(i) = cur {
let n = self.slots[i].as_ref().expect("live");
out.push(n.block_start);
cur = n.newer;
}
out
}
}
struct CacheState {
entries: Lru,
hits: u64,
misses: u64,
generation: u64,
in_flight: Vec<(u64, ThreadId)>,
}
impl CachingDevice {
pub fn new(inner: Arc<dyn BlockDevice>, block_size: u64, capacity: usize) -> Arc<Self> {
Arc::new(Self {
inner: inner.clone(),
writable: Some(inner),
block_size,
capacity,
state: Mutex::new(CacheState {
entries: Lru::with_capacity(capacity),
hits: 0,
misses: 0,
generation: 0,
in_flight: Vec::new(),
}),
fetched: Condvar::new(),
})
}
pub fn read_only(inner: Arc<dyn BlockRead>, block_size: u64, capacity: usize) -> Arc<Self> {
Arc::new(Self {
inner,
writable: None,
block_size,
capacity,
state: Mutex::new(CacheState {
entries: Lru::with_capacity(capacity),
hits: 0,
misses: 0,
generation: 0,
in_flight: Vec::new(),
}),
fetched: Condvar::new(),
})
}
pub fn stats(&self) -> (u64, u64) {
let s = self.state.lock().unwrap();
(s.hits, s.misses)
}
pub fn invalidate_all(&self) {
let mut s = self.state.lock().unwrap();
s.entries.clear();
s.generation = s.generation.wrapping_add(1);
}
fn invalidate_range(state: &mut CacheState, start: u64, end: u64, block_size: u64) {
state.entries.retain_blocks(|off| {
let block_end = off.saturating_add(block_size);
off >= end || block_end <= start
});
state.generation = state.generation.wrapping_add(1);
}
fn invalidate_for_write(&self, start: u64, end: u64) {
let mut s = self.state.lock().unwrap();
let bs = self.block_size;
Self::invalidate_range(&mut s, start, end, bs);
}
fn check_block_size(&self) -> Result<()> {
if self.block_size == 0 {
return Err(crate::error::Error::Custom(
"cache block size is zero".to_string(),
));
}
if self.block_size > MAX_BLOCK_SIZE {
return Err(crate::error::Error::Custom(format!(
"cache block size {} exceeds the {MAX_BLOCK_SIZE}-byte ceiling",
self.block_size
)));
}
Ok(())
}
}
impl CachingDevice {
fn block(&self, block_start: u64) -> Result<Arc<Vec<u8>>> {
let generation_at_miss = {
let mut s = self.state.lock().unwrap();
loop {
if let Some(data) = s.entries.get(block_start) {
s.hits += 1;
return Ok(data);
}
let mine = std::thread::current().id();
let holding_a_fetch = s.in_flight.iter().any(|(_, owner)| *owner == mine);
let being_fetched = s.in_flight.iter().any(|(o, _)| *o == block_start);
if being_fetched && !holding_a_fetch {
s = self.fetched.wait(s).unwrap();
continue;
}
s.misses += 1;
s.in_flight.push((block_start, mine));
break s.generation;
}
};
let _fetch = FetchGuard {
device: self,
block_start,
owner: std::thread::current().id(),
};
let size = self.inner.size_bytes();
let end = block_start.saturating_add(self.block_size).min(size);
let len = end.saturating_sub(block_start) as usize;
let mut block = vec![0u8; len];
self.inner.read_at(block_start, &mut block)?;
let data = Arc::new(block);
{
let mut s = self.state.lock().unwrap();
if let Some(held) = s.entries.get(block_start) {
return Ok(held);
}
if s.generation == generation_at_miss {
s.entries.insert(block_start, data.clone(), self.capacity);
}
}
Ok(data)
}
}
impl BlockRead for CachingDevice {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
self.check_block_size()?;
if buf.is_empty() {
return Ok(());
}
let Some(end) = offset.checked_add(buf.len() as u64) else {
return Err(crate::error::Error::ShortRead {
offset,
want: buf.len(),
got: 0,
});
};
if end > self.inner.size_bytes() {
return self.inner.read_at(offset, buf);
}
let bs = self.block_size;
let first = offset / bs;
let last = (end - 1) / bs;
let spanned = (last - first + 1) as usize;
if spanned > 1 && spanned.saturating_mul(2) > self.capacity {
return self.inner.read_at(offset, buf);
}
let mut done = 0usize;
for index in first..=last {
let block_start = index * bs;
let block = self.block(block_start)?;
let from = (offset.max(block_start) - block_start) as usize;
let take = (block.len().saturating_sub(from)).min(buf.len() - done);
if take == 0 {
break;
}
buf[done..done + take].copy_from_slice(&block[from..from + take]);
done += take;
}
if done != buf.len() {
return Err(crate::error::Error::ShortRead {
offset,
want: buf.len(),
got: done,
});
}
Ok(())
}
fn size_bytes(&self) -> u64 {
self.inner.size_bytes()
}
}
impl BlockDevice for CachingDevice {
fn write_at(&self, offset: u64, buf: &[u8]) -> Result<()> {
let Some(writable) = self.writable.as_ref() else {
return Err(crate::error::Error::ReadOnly);
};
self.check_block_size()?;
let end = offset.saturating_add(buf.len() as u64);
self.invalidate_for_write(offset, end);
let result = writable.write_at(offset, buf);
self.invalidate_for_write(offset, end);
result
}
fn flush(&self) -> Result<()> {
match self.writable.as_ref() {
Some(writable) => writable.flush(),
None => Ok(()),
}
}
fn is_writable(&self) -> bool {
self.writable.as_ref().is_some_and(|w| w.is_writable())
}
fn set_len(&self, new_len: u64) -> Result<()> {
let Some(writable) = self.writable.as_ref() else {
return Err(crate::error::Error::ReadOnly);
};
self.check_block_size()?;
let from = self.inner.size_bytes().min(new_len);
self.invalidate_for_write(from, u64::MAX);
let result = writable.set_len(new_len);
self.invalidate_for_write(from, u64::MAX);
result
}
fn can_grow(&self) -> bool {
self.writable.as_ref().is_some_and(|w| w.can_grow())
}
}
struct FetchGuard<'a> {
device: &'a CachingDevice,
block_start: u64,
owner: ThreadId,
}
impl Drop for FetchGuard<'_> {
fn drop(&mut self) {
if let Ok(mut s) = self.device.state.lock() {
if let Some(pos) = s
.in_flight
.iter()
.position(|(o, owner)| *o == self.block_start && *owner == self.owner)
{
s.in_flight.remove(pos);
}
}
self.device.fetched.notify_all();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_device::Bytes;
const BS: u64 = 512;
fn backing() -> Arc<Bytes> {
Arc::new(Bytes::new((0..4096u32).map(|i| i as u8).collect()))
}
#[test]
fn a_read_only_device_can_be_cached() {
let inner = backing();
let cache = CachingDevice::read_only(inner, BS, 4);
let mut first = vec![0u8; BS as usize];
let mut again = vec![0u8; BS as usize];
cache.read_at(0, &mut first).expect("first read");
cache.read_at(0, &mut again).expect("second read");
assert_eq!(first, again, "the cache must serve what the device held");
assert_eq!(cache.stats(), (1, 1), "one hit after one miss");
}
#[test]
fn a_read_smaller_than_a_block_is_served_from_it() {
let cache = CachingDevice::read_only(backing(), BS, 8);
let mut whole = vec![0u8; BS as usize];
cache.read_at(0, &mut whole).expect("warm the block");
assert_eq!(cache.stats(), (0, 1), "one miss to fetch it");
for at in [0u64, 8, 100, 504] {
let mut small = [0u8; 8];
cache.read_at(at, &mut small).expect("sub-block read");
assert_eq!(
&small[..],
&whole[at as usize..at as usize + 8],
"the bytes must be the block's own, at the right offset"
);
}
assert_eq!(cache.stats(), (4, 1), "four hits, and no further misses");
}
#[test]
fn a_read_spanning_two_blocks_is_stitched() {
let inner = backing();
let mut direct = vec![0u8; 16];
inner.read_at(BS - 8, &mut direct).expect("read it plainly");
let cache = CachingDevice::read_only(backing(), BS, 8);
let mut across = vec![0u8; 16];
cache.read_at(BS - 8, &mut across).expect("spanning read");
assert_eq!(across, direct, "the same bytes the device would give");
assert_eq!(cache.stats(), (0, 2), "one miss per block touched");
cache.read_at(BS - 8, &mut across).expect("again");
assert_eq!(cache.stats(), (2, 2), "and both are held now");
}
#[test]
fn a_read_that_would_sweep_the_cache_passes_through() {
let cache = CachingDevice::read_only(backing(), BS, 4);
let mut big = vec![0u8; (BS * 4) as usize];
cache.read_at(0, &mut big).expect("a large read");
assert_eq!(
cache.stats(),
(0, 0),
"neither hit nor miss: it never consulted the cache"
);
}
#[test]
fn a_device_shorter_than_a_block_still_reads() {
let tiny: Arc<Bytes> = Arc::new(Bytes::new((0..100u32).map(|i| i as u8).collect()));
let cache = CachingDevice::read_only(tiny, BS, 4);
let mut buf = vec![0u8; 40];
cache
.read_at(10, &mut buf)
.expect("a read inside the device");
assert_eq!(buf[0], 10, "the wrong bytes came back");
assert_eq!(buf[39], 49);
cache.read_at(10, &mut buf).expect("again");
assert_eq!(cache.stats(), (1, 1));
}
#[test]
fn a_read_past_the_end_is_still_an_error() {
let tiny: Arc<Bytes> = Arc::new(Bytes::new(vec![0u8; 100]));
let cache = CachingDevice::read_only(tiny, BS, 4);
let mut buf = vec![0u8; 40];
assert!(
cache.read_at(80, &mut buf).is_err(),
"80 + 40 is past the end of a 100-byte device"
);
}
#[test]
fn a_bypassed_read_does_not_wait_for_the_cache_lock() {
use std::sync::mpsc;
use std::thread;
use std::time::Duration;
let cache = CachingDevice::read_only(backing(), BS, 4);
let held = cache.state.lock().expect("nothing else holds it yet");
let reader = Arc::clone(&cache);
let (tx, rx) = mpsc::channel();
thread::spawn(move || {
let mut big = vec![0u8; (BS * 4) as usize];
let outcome = reader.read_at(0, &mut big);
let _ = tx.send(outcome);
});
let outcome = rx.recv_timeout(Duration::from_secs(5)).expect(
"a read that bypasses the cache must not block on the cache lock; \
it timed out waiting for a mutex it has no reason to take",
);
outcome.expect("and the bypassed read itself must succeed");
drop(held);
assert_eq!(
cache.stats(),
(0, 0),
"it bypassed, so it neither hit nor missed"
);
}
#[test]
fn writing_through_a_read_only_cache_is_refused() {
let cache = CachingDevice::read_only(backing(), BS, 4);
assert!(!cache.is_writable());
assert!(matches!(
cache.write_at(0, &[1u8; 8]),
Err(crate::error::Error::ReadOnly)
));
assert!(cache.flush().is_ok());
}
}
#[cfg(test)]
mod lru_tests {
use super::*;
fn block(n: u8) -> Arc<Vec<u8>> {
Arc::new(vec![n; 4])
}
fn assert_consistent(lru: &Lru) {
let forward = lru.recency_order();
let mut backward = lru.recency_order_reversed();
backward.reverse();
assert_eq!(
forward, backward,
"the recency list disagrees with itself walked the other way"
);
assert_eq!(
forward.len(),
lru.len(),
"the list holds {} nodes and the index {} -- they have drifted",
forward.len(),
lru.len()
);
}
#[test]
fn a_hit_promotes_to_most_recently_used() {
let mut lru = Lru::with_capacity(4);
for i in 0..4u64 {
lru.insert(i * 100, block(i as u8), 4);
}
assert_eq!(lru.recency_order(), vec![300, 200, 100, 0]);
assert!(lru.get(100).is_some());
assert_eq!(lru.recency_order(), vec![100, 300, 200, 0]);
assert_consistent(&lru);
assert!(lru.get(100).is_some());
assert_eq!(lru.recency_order(), vec![100, 300, 200, 0]);
assert_consistent(&lru);
assert!(lru.get(0).is_some());
assert_eq!(lru.recency_order(), vec![0, 100, 300, 200]);
assert_consistent(&lru);
}
#[test]
fn the_least_recently_used_is_what_gets_evicted() {
let mut lru = Lru::with_capacity(3);
for i in 0..3u64 {
lru.insert(i * 100, block(i as u8), 3);
}
assert!(lru.get(0).is_some());
lru.insert(999, block(9), 3);
assert_eq!(lru.len(), 3);
assert_eq!(
lru.recency_order(),
vec![999, 0, 200],
"100 was least recently used and is the one that should be gone"
);
assert!(lru.get(100).is_none());
assert_consistent(&lru);
}
#[test]
fn evicted_slots_are_reused_rather_than_growing_the_slab() {
let mut lru = Lru::with_capacity(2);
for i in 0..20u64 {
lru.insert(i, block(i as u8), 2);
assert_consistent(&lru);
}
assert_eq!(lru.len(), 2);
assert!(
lru.slots.len() <= 3,
"twenty inserts at capacity 2 left {} slots: freed slots are not being \
reused, so the slab grows without bound",
lru.slots.len()
);
}
#[test]
fn a_re_insert_replaces_and_does_not_leave_the_old_node_linked() {
let mut lru = Lru::with_capacity(4);
lru.insert(10, block(1), 4);
lru.insert(20, block(2), 4);
lru.insert(10, block(3), 4);
assert_eq!(lru.len(), 1 + 1, "one entry per block, not one per insert");
assert_eq!(lru.recency_order(), vec![10, 20]);
assert_eq!(
lru.get(10).as_deref().map(|v| v[0]),
Some(3),
"the newer data wins"
);
assert_consistent(&lru);
}
#[test]
fn removing_and_retaining_keep_the_list_consistent() {
let mut lru = Lru::with_capacity(8);
for i in 0..8u64 {
lru.insert(i, block(i as u8), 8);
}
assert!(lru.remove(0), "the tail");
assert_consistent(&lru);
assert!(lru.remove(7), "the head");
assert_consistent(&lru);
assert!(lru.remove(4), "the middle");
assert_consistent(&lru);
assert!(!lru.remove(4), "already gone");
lru.retain_blocks(|b| b % 2 == 0);
assert_consistent(&lru);
assert_eq!(lru.recency_order(), vec![6, 2]);
lru.clear();
assert_eq!(lru.len(), 0);
assert!(lru.recency_order().is_empty());
assert_consistent(&lru);
}
#[test]
fn capacity_zero_behaves_as_it_did_before() {
let mut lru = Lru::with_capacity(0);
lru.insert(1, block(1), 0);
assert_eq!(lru.len(), 1);
lru.insert(2, block(2), 0);
assert_eq!(lru.len(), 1);
assert_eq!(lru.recency_order(), vec![2]);
assert_consistent(&lru);
}
}