use crate::block::{BlockDevice, BlockRead};
use crate::error::Result;
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
pub struct CachingDevice {
inner: Arc<dyn BlockRead>,
writable: Option<Arc<dyn BlockDevice>>,
block_size: u64,
state: Mutex<CacheState>,
}
struct CacheState {
entries: VecDeque<(u64, Arc<Vec<u8>>)>,
capacity: usize,
hits: u64,
misses: u64,
}
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,
state: Mutex::new(CacheState {
entries: VecDeque::with_capacity(capacity),
capacity,
hits: 0,
misses: 0,
}),
})
}
pub fn read_only(inner: Arc<dyn BlockRead>, block_size: u64, capacity: usize) -> Arc<Self> {
Arc::new(Self {
inner,
writable: None,
block_size,
state: Mutex::new(CacheState {
entries: VecDeque::with_capacity(capacity),
capacity,
hits: 0,
misses: 0,
}),
})
}
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();
}
fn invalidate_range(state: &mut CacheState, start: u64, end: u64, block_size: u64) {
state.entries.retain(|(off, _)| {
let block_end = off.saturating_add(block_size);
*off >= end || block_end <= start
});
}
}
impl CachingDevice {
fn block(&self, block_start: u64) -> Result<Arc<Vec<u8>>> {
{
let mut s = self.state.lock().unwrap();
if let Some(pos) = s.entries.iter().position(|(o, _)| *o == block_start) {
let entry = s.entries.remove(pos).expect("position just found it");
let data = entry.1.clone();
s.entries.push_front(entry);
s.hits += 1;
return Ok(data);
}
s.misses += 1;
}
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 s.entries.len() >= s.capacity {
s.entries.pop_back();
}
s.entries.push_front((block_start, data.clone()));
Ok(data)
}
}
impl BlockRead for CachingDevice {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
if buf.is_empty() {
return Ok(());
}
let bs = self.block_size;
let first = offset / bs;
let last = (offset + buf.len() as u64 - 1) / bs;
let spanned = (last - first + 1) as usize;
let sweeps_the_cache = {
let s = self.state.lock().unwrap();
spanned > 1 && spanned * 2 > s.capacity
};
if sweeps_the_cache {
return self.inner.read_at(offset, buf);
}
if offset.saturating_add(buf.len() as u64) > self.inner.size_bytes() {
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;
}
Ok(())
}
fn size_bytes(&self) -> u64 {
self.inner.size_bytes()
}
}
impl BlockDevice for CachingDevice {
fn write_at(&self, offset: u64, buf: &[u8]) -> Result<()> {
let end = offset.saturating_add(buf.len() as u64);
{
let mut s = self.state.lock().unwrap();
let bs = self.block_size;
Self::invalidate_range(&mut s, offset, end, bs);
}
let Some(writable) = self.writable.as_ref() else {
return Err(crate::error::Error::ReadOnly);
};
writable.write_at(offset, buf)
}
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())
}
}
#[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 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());
}
}