use crate::block_io::BlockDevice;
use crate::error::Result;
use std::collections::HashMap;
use std::sync::Mutex;
struct CacheState {
capacity: usize,
entries: HashMap<u64, (Vec<u8>, u64)>,
pinned: HashMap<u64, Vec<u8>>,
next_seq: u64,
hits: u64,
misses: u64,
}
impl CacheState {
fn new(capacity: usize) -> Self {
Self {
capacity,
entries: HashMap::with_capacity(capacity.min(1024)),
pinned: HashMap::new(),
next_seq: 0,
hits: 0,
misses: 0,
}
}
fn next_seq(&mut self) -> u64 {
let s = self.next_seq;
self.next_seq = s.wrapping_add(1);
s
}
fn get(&mut self, block: u64) -> Option<Vec<u8>> {
if let Some(bytes) = self.pinned.get(&block) {
self.hits += 1;
return Some(bytes.clone());
}
let seq = self.next_seq();
if let Some(slot) = self.entries.get_mut(&block) {
slot.1 = seq;
self.hits += 1;
Some(slot.0.clone())
} else {
self.misses += 1;
None
}
}
fn put(&mut self, block: u64, bytes: Vec<u8>) {
if let std::collections::hash_map::Entry::Occupied(mut e) = self.pinned.entry(block) {
e.insert(bytes);
return;
}
if self.entries.len() >= self.capacity {
if let Some((&victim, _)) = self.entries.iter().min_by_key(|(_, (_, seq))| *seq) {
self.entries.remove(&victim);
}
}
let seq = self.next_seq();
self.entries.insert(block, (bytes, seq));
}
fn pin(&mut self, block: u64, bytes: Vec<u8>) {
self.entries.remove(&block);
self.pinned.insert(block, bytes);
}
fn unpin_all(&mut self) {
let drained: Vec<(u64, Vec<u8>)> = self.pinned.drain().collect();
for (block, bytes) in drained {
if self.entries.len() >= self.capacity {
if let Some((&victim, _)) = self.entries.iter().min_by_key(|(_, (_, seq))| *seq) {
self.entries.remove(&victim);
}
}
let seq = self.next_seq();
self.entries.insert(block, (bytes, seq));
}
}
}
pub struct CachedDevice {
inner: std::sync::Arc<dyn BlockDevice>,
block_size: u32,
state: Mutex<CacheState>,
}
impl CachedDevice {
pub fn new(inner: std::sync::Arc<dyn BlockDevice>, block_size: u32, capacity: usize) -> Self {
Self {
inner,
block_size,
state: Mutex::new(CacheState::new(capacity.max(1))),
}
}
pub fn stats(&self) -> (u64, u64) {
let s = self.state.lock().expect("cache mutex poisoned");
(s.hits, s.misses)
}
}
impl BlockDevice for CachedDevice {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
let bs = self.block_size as u64;
let block = offset / bs;
let off_in_block = (offset % bs) as usize;
let len = buf.len();
if off_in_block + len <= bs as usize {
{
let mut state = self.state.lock().expect("cache mutex poisoned");
if let Some(blk) = state.get(block) {
buf.copy_from_slice(&blk[off_in_block..off_in_block + len]);
return Ok(());
}
}
let mut blk = vec![0u8; bs as usize];
self.inner.read_at(block * bs, &mut blk)?;
buf.copy_from_slice(&blk[off_in_block..off_in_block + len]);
let mut state = self.state.lock().expect("cache mutex poisoned");
state.put(block, blk);
return Ok(());
}
self.inner.read_at(offset, buf)
}
fn size_bytes(&self) -> u64 {
self.inner.size_bytes()
}
fn write_at(&self, offset: u64, buf: &[u8]) -> Result<()> {
let bs = self.block_size as u64;
self.inner.write_at(offset, buf)?;
let off_in_block = (offset % bs) as usize;
let len = buf.len();
if off_in_block == 0 && len == bs as usize {
let block = offset / bs;
let mut state = self.state.lock().expect("cache mutex poisoned");
state.put(block, buf.to_vec());
return Ok(());
}
let first_block = offset / bs;
let last_block = (offset + buf.len() as u64).saturating_sub(1) / bs;
let buf_end_byte = offset + buf.len() as u64;
let mut state = self.state.lock().expect("cache mutex poisoned");
for b in first_block..=last_block {
let block_start = b * bs;
let block_end = block_start + bs;
let write_start = offset.max(block_start);
let write_end = buf_end_byte.min(block_end);
let in_block_off = (write_start - block_start) as usize;
let in_block_end = (write_end - block_start) as usize;
let buf_start = (write_start - offset) as usize;
let buf_end = (write_end - offset) as usize;
if let Some(img) = state.pinned.get_mut(&b) {
img[in_block_off..in_block_end].copy_from_slice(&buf[buf_start..buf_end]);
}
state.entries.remove(&b);
}
Ok(())
}
fn flush(&self) -> Result<()> {
self.inner.flush()
}
fn is_writable(&self) -> bool {
self.inner.is_writable()
}
fn populate_cache(&self, block: u64, bytes: Vec<u8>) {
let mut state = self.state.lock().expect("cache mutex poisoned");
state.pin(block, bytes);
}
fn unpin_all(&self) {
let mut state = self.state.lock().expect("cache mutex poisoned");
state.unpin_all();
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
struct CountingDevice {
bytes: Mutex<Vec<u8>>,
reads: std::sync::atomic::AtomicU64,
writes: std::sync::atomic::AtomicU64,
writable: bool,
}
impl CountingDevice {
fn new(size: usize, writable: bool) -> Arc<Self> {
Arc::new(Self {
bytes: Mutex::new(vec![0u8; size]),
reads: std::sync::atomic::AtomicU64::new(0),
writes: std::sync::atomic::AtomicU64::new(0),
writable,
})
}
fn reads(&self) -> u64 {
self.reads.load(std::sync::atomic::Ordering::SeqCst)
}
fn writes(&self) -> u64 {
self.writes.load(std::sync::atomic::Ordering::SeqCst)
}
}
impl BlockDevice for CountingDevice {
fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<()> {
self.reads.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let b = self.bytes.lock().unwrap();
let off = offset as usize;
buf.copy_from_slice(&b[off..off + buf.len()]);
Ok(())
}
fn size_bytes(&self) -> u64 {
self.bytes.lock().unwrap().len() as u64
}
fn write_at(&self, offset: u64, buf: &[u8]) -> Result<()> {
self.writes
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let mut b = self.bytes.lock().unwrap();
let off = offset as usize;
b[off..off + buf.len()].copy_from_slice(buf);
Ok(())
}
fn is_writable(&self) -> bool {
self.writable
}
}
#[test]
fn second_read_of_same_block_is_a_cache_hit() {
let inner = CountingDevice::new(4096 * 16, false);
let cached = CachedDevice::new(inner.clone(), 4096, 8);
let mut buf = vec![0u8; 100];
cached.read_at(0, &mut buf).unwrap();
cached.read_at(0, &mut buf).unwrap();
cached.read_at(50, &mut buf).unwrap(); assert_eq!(inner.reads(), 1, "cache should serve all 3 from one read");
let (hits, misses) = cached.stats();
assert_eq!(hits, 2);
assert_eq!(misses, 1);
}
#[test]
fn unaligned_write_merges_into_pinned_block() {
let inner = CountingDevice::new(4096 * 16, true);
let cached = CachedDevice::new(inner.clone(), 4096, 8);
let mut journaled = vec![0u8; 4096];
journaled[0] = 0xAA;
journaled[3000] = 0xBB;
cached.populate_cache(5, journaled);
cached.write_at(5 * 4096 + 100, &[0xCCu8; 100]).unwrap();
let mut buf = vec![0u8; 4096];
for i in 0..4096 {
cached
.read_at(5 * 4096 + i as u64, &mut buf[i..i + 1])
.unwrap();
}
assert_eq!(buf[0], 0xAA, "pinned journaled byte 0 must survive");
assert_eq!(buf[100], 0xCC, "new sub-block bytes must be visible");
assert_eq!(buf[199], 0xCC);
assert_eq!(buf[200], 0x00, "untouched portion stays at journaled value");
assert_eq!(
buf[3000], 0xBB,
"pinned bytes outside the write window survive"
);
}
#[test]
fn unaligned_write_invalidates_cache_entry() {
let inner = CountingDevice::new(4096 * 16, true);
let cached = CachedDevice::new(inner.clone(), 4096, 8);
let mut buf = vec![0u8; 100];
cached.read_at(0, &mut buf).unwrap();
cached.write_at(0, &[42u8; 100]).unwrap(); cached.read_at(0, &mut buf).unwrap();
assert_eq!(inner.reads(), 2);
assert_eq!(buf[0], 42, "post-write read should see the new bytes");
assert_eq!(
inner.writes(),
1,
"an unaligned write must reach the device once, not be held in cache"
);
}
#[test]
fn aligned_write_updates_cache_entry() {
let inner = CountingDevice::new(4096 * 16, true);
let cached = CachedDevice::new(inner.clone(), 4096, 8);
let mut buf = vec![0u8; 100];
cached.read_at(0, &mut buf).unwrap();
assert_eq!(inner.reads(), 1);
let new_block = vec![0xCDu8; 4096];
cached.write_at(0, &new_block).unwrap();
let mut readback = vec![0u8; 100];
cached.read_at(0, &mut readback).unwrap();
assert_eq!(inner.reads(), 1, "aligned write-through skips disk read");
assert_eq!(readback[0], 0xCD);
let mut from_disk = vec![0u8; 100];
inner.read_at(0, &mut from_disk).unwrap();
assert_eq!(from_disk[0], 0xCD);
}
#[test]
fn populate_cache_pins_entry_against_lru() {
let inner = CountingDevice::new(4096 * 16, false);
let cached = CachedDevice::new(inner.clone(), 4096, 2); cached.populate_cache(7, vec![0xAAu8; 4096]);
let mut throwaway = vec![0u8; 8];
for blk in 0..5u64 {
cached.read_at(blk * 4096, &mut throwaway).unwrap();
}
let mut buf = vec![0u8; 8];
cached.read_at(7 * 4096, &mut buf).unwrap();
assert_eq!(buf, vec![0xAA; 8],
"pinned entry must survive LRU pressure — otherwise journaled writes vanish before checkpoint");
}
#[test]
fn unpin_all_lets_pinned_entries_evict_normally() {
let inner = CountingDevice::new(4096 * 16, false);
let cached = CachedDevice::new(inner.clone(), 4096, 1); cached.populate_cache(3, vec![0xBBu8; 4096]);
cached.unpin_all();
let mut throwaway = vec![0u8; 8];
cached.read_at(0, &mut throwaway).unwrap();
let mut buf3 = vec![0u8; 8];
cached.read_at(3 * 4096, &mut buf3).unwrap();
assert_eq!(
buf3,
vec![0; 8],
"after unpin + LRU eviction, the inner device's bytes win"
);
}
#[test]
fn multi_block_read_bypasses_cache() {
let inner = CountingDevice::new(4096 * 16, false);
let cached = CachedDevice::new(inner.clone(), 4096, 8);
let mut buf = vec![0u8; 8000]; cached.read_at(0, &mut buf).unwrap();
cached.read_at(0, &mut buf).unwrap();
assert_eq!(inner.reads(), 2);
let (hits, misses) = cached.stats();
assert_eq!(hits, 0);
assert_eq!(misses, 0, "multi-block reads bypass entirely");
}
#[test]
fn lru_evicts_oldest_when_capacity_exceeded() {
let inner = CountingDevice::new(4096 * 16, false);
let cached = CachedDevice::new(inner.clone(), 4096, 2); let mut buf = vec![0u8; 8];
for blk in 0..3u64 {
cached.read_at(blk * 4096, &mut buf).unwrap();
}
cached.read_at(0, &mut buf).unwrap();
cached.read_at(2 * 4096, &mut buf).unwrap();
let (hits, misses) = cached.stats();
assert_eq!(misses, 4, "blocks 0,1,2 + re-read of 0 (evicted)");
assert_eq!(hits, 1, "re-read of 2 still cached");
}
#[test]
fn lru_keeps_recently_touched_block_alive() {
let inner = CountingDevice::new(4096 * 16, false);
let cached = CachedDevice::new(inner.clone(), 4096, 2);
let mut buf = vec![0u8; 8];
cached.read_at(0, &mut buf).unwrap(); cached.read_at(4096, &mut buf).unwrap(); cached.read_at(0, &mut buf).unwrap(); cached.read_at(2 * 4096, &mut buf).unwrap(); cached.read_at(0, &mut buf).unwrap(); let (hits, misses) = cached.stats();
assert_eq!(hits, 2);
assert_eq!(misses, 3);
cached.read_at(4096, &mut buf).unwrap(); let (_, m_after) = cached.stats();
assert_eq!(m_after, 4, "block 1 was evicted, re-read is a miss");
}
}