use crate::reader::RangeReader;
use std::collections::{BTreeSet, HashMap};
use std::sync::{Arc, Mutex};
pub const DEFAULT_BLOCK: u64 = 64 * 1024;
pub const DEFAULT_CACHE_CAP: u64 = 256 * 1024 * 1024;
pub fn auto_block(len: u64) -> u64 {
const MB: u64 = 1 << 20;
let mult: u64 = if len > 100 * MB {
2 } else {
1 };
mult * DEFAULT_BLOCK
}
struct CacheEntry {
data: Arc<[u8]>,
stamp: u64,
}
struct CacheState {
map: HashMap<u64, CacheEntry>,
used: u64,
tick: u64,
}
pub struct BlockCacheReader<R> {
inner: R,
block: u64,
len: u64,
cap: u64,
cache: Mutex<CacheState>,
}
impl<R: RangeReader> BlockCacheReader<R> {
pub fn new(inner: R, block: u64) -> Self {
let len = inner.len();
Self {
inner,
block: block.max(4096),
len,
cap: DEFAULT_CACHE_CAP,
cache: Mutex::new(CacheState {
map: HashMap::new(),
used: 0,
tick: 0,
}),
}
}
pub fn with_cache_cap(mut self, cap: u64) -> Self {
self.cap = cap;
self
}
pub fn cached_bytes(&self) -> u64 {
self.cache.lock().unwrap().used
}
fn bounds(&self, offset: u64, len: u64) -> std::io::Result<()> {
if offset.checked_add(len).is_none_or(|e| e > self.len) {
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"range out of bounds",
));
}
Ok(())
}
fn ensure(&self, want: &BTreeSet<u64>) -> std::io::Result<()> {
let missing: Vec<u64> = {
let mut st = self.cache.lock().unwrap();
st.tick += 1;
let tick = st.tick;
want.iter()
.copied()
.filter(|b| match st.map.get_mut(b) {
Some(e) => {
e.stamp = tick;
false
}
None => true,
})
.collect()
};
if missing.is_empty() {
return Ok(());
}
let mut spans: Vec<(u64, u64)> = Vec::new();
let mut runs: Vec<(u64, u64)> = Vec::new();
let mut i = 0;
while i < missing.len() {
let first = missing[i];
let mut last = first;
let mut j = i + 1;
while j < missing.len() && missing[j] == last + 1 {
last = missing[j];
j += 1;
}
let off = first * self.block;
let end = ((last + 1) * self.block).min(self.len);
spans.push((off, end - off));
runs.push((first, last));
i = j;
}
let blobs = self.inner.read_many(&spans)?;
if blobs.len() != spans.len() {
return Err(std::io::Error::other("block fetch returned wrong count"));
}
let mut st = self.cache.lock().unwrap();
st.tick += 1;
let tick = st.tick;
for (&(first, last), blob) in runs.iter().zip(blobs.into_iter()) {
let span_start = first * self.block;
for b in first..=last {
let lo = (b * self.block - span_start) as usize;
let hi = ((((b + 1) * self.block).min(self.len)) - span_start) as usize;
let hi = hi.min(blob.len());
let lo = lo.min(hi);
let data: Arc<[u8]> = Arc::from(&blob[lo..hi]);
st.used += data.len() as u64;
if let Some(old) = st.map.insert(b, CacheEntry { data, stamp: tick }) {
st.used -= old.data.len() as u64; }
}
}
Ok(())
}
fn trim(&self) {
let mut st = self.cache.lock().unwrap();
if st.used <= self.cap {
return;
}
let mut order: Vec<(u64, u64)> = st.map.iter().map(|(&b, e)| (e.stamp, b)).collect();
order.sort_unstable();
for (_, b) in order {
if st.used <= self.cap {
break;
}
if let Some(e) = st.map.remove(&b) {
st.used -= e.data.len() as u64;
}
}
}
fn assemble(&self, offset: u64, len: u64) -> std::io::Result<Vec<u8>> {
let first = offset / self.block;
let last = (offset + len - 1) / self.block;
let resident: Vec<Option<Arc<[u8]>>> = {
let st = self.cache.lock().unwrap();
(first..=last)
.map(|b| st.map.get(&b).map(|e| e.data.clone()))
.collect()
};
let mut out = Vec::with_capacity(len as usize);
let mut pos = offset;
let end = offset + len;
while pos < end {
let b = pos / self.block;
let block_start = b * self.block;
let within = (pos - block_start) as usize;
let fetched: Vec<u8>;
let blk: &[u8] = match &resident[(b - first) as usize] {
Some(data) => data,
None => {
let blen = ((b + 1) * self.block).min(self.len) - block_start;
fetched = self.inner.read_at(block_start, blen)?;
&fetched
}
};
let take = ((end - pos) as usize).min(blk.len().saturating_sub(within));
if take == 0 {
break;
}
out.extend_from_slice(&blk[within..within + take]);
pos += take as u64;
}
Ok(out)
}
}
impl<R: RangeReader> RangeReader for BlockCacheReader<R> {
fn len(&self) -> u64 {
self.len
}
fn concurrency(&self) -> usize {
self.inner.concurrency()
}
fn read_at(&self, offset: u64, len: u64) -> std::io::Result<Vec<u8>> {
if len == 0 {
return Ok(Vec::new());
}
self.bounds(offset, len)?;
let want: BTreeSet<u64> = (offset / self.block..=(offset + len - 1) / self.block).collect();
self.ensure(&want)?;
let out = self.assemble(offset, len)?;
self.trim();
Ok(out)
}
fn read_many(&self, ranges: &[(u64, u64)]) -> std::io::Result<Vec<Vec<u8>>> {
let mut want = BTreeSet::new();
for &(o, l) in ranges {
if l == 0 {
continue;
}
self.bounds(o, l)?;
for b in o / self.block..=(o + l - 1) / self.block {
want.insert(b);
}
}
self.ensure(&want)?;
let out = ranges
.iter()
.map(|&(o, l)| {
if l == 0 {
Ok(Vec::new())
} else {
self.assemble(o, l)
}
})
.collect::<std::io::Result<Vec<_>>>()?;
self.trim();
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::reader::{CountingReader, SliceReader};
#[test]
fn caches_blocks_and_serves_exact_bytes() {
let data: Vec<u8> = (0..100_000u32).map(|i| i as u8).collect();
let counting = Arc::new(CountingReader::new(SliceReader::new(&data)));
let r = BlockCacheReader::new(counting.clone(), 16 * 1024);
for off in [10u64, 200, 5000, 16_500, 17_000, 33_000, 33_100] {
assert_eq!(
r.read_at(off, 32).unwrap(),
data[off as usize..off as usize + 32]
);
}
assert!(
counting.requests() <= 3,
"physical fetches: {}",
counting.requests()
);
let before = counting.requests();
assert_eq!(r.read_at(200, 16).unwrap(), data[200..216]);
assert_eq!(counting.requests(), before);
}
#[test]
fn read_many_fetches_missing_blocks_once() {
let data: Vec<u8> = (0..200_000u32).map(|i| (i * 7) as u8).collect();
let counting = Arc::new(CountingReader::new(SliceReader::new(&data)));
let r = BlockCacheReader::new(counting.clone(), 32 * 1024);
let ranges = [(0u64, 8u64), (40_000, 16), (40_050, 16), (130_000, 64)];
let out = r.read_many(&ranges).unwrap();
for (&(o, l), got) in ranges.iter().zip(&out) {
assert_eq!(got, &data[o as usize..(o + l) as usize]);
}
assert!(
counting.requests() <= 3,
"physical: {}",
counting.requests()
);
}
#[test]
fn out_of_bounds_errors() {
let data = vec![0u8; 1000];
let r = BlockCacheReader::new(SliceReader::new(&data), 8192);
assert!(r.read_at(990, 20).is_err());
}
#[test]
fn eviction_caps_resident_bytes_and_keeps_reads_exact() {
let data: Vec<u8> = (0..1_000_000u32).map(|i| (i * 31) as u8).collect();
let counting = Arc::new(CountingReader::new(SliceReader::new(&data)));
let cap = 64 * 1024;
let r = BlockCacheReader::new(counting.clone(), 8192).with_cache_cap(cap);
for off in (0..1_000_000u64 - 64).step_by(37_777) {
assert_eq!(
r.read_at(off, 64).unwrap(),
data[off as usize..off as usize + 64]
);
assert!(
r.cached_bytes() <= cap,
"resident {} > cap {cap}",
r.cached_bytes()
);
}
}
#[test]
fn eviction_is_lru() {
let data: Vec<u8> = (0..64 * 1024u32).map(|i| i as u8).collect();
let counting = Arc::new(CountingReader::new(SliceReader::new(&data)));
let r = BlockCacheReader::new(counting.clone(), 8192).with_cache_cap(16 * 1024);
let block = |i: u64| i * 8192;
r.read_at(block(0), 16).unwrap(); r.read_at(block(1), 16).unwrap(); r.read_at(block(0), 16).unwrap();
let before = counting.requests();
r.read_at(block(2), 16).unwrap(); assert_eq!(counting.requests(), before + 1);
let before = counting.requests();
r.read_at(block(0), 16).unwrap(); assert_eq!(
counting.requests(),
before,
"recently-touched block evicted"
);
let before = counting.requests();
r.read_at(block(1), 16).unwrap(); assert_eq!(counting.requests(), before + 1);
}
#[test]
fn request_larger_than_cap_reads_exactly_then_trims() {
let data: Vec<u8> = (0..256 * 1024u32).map(|i| (i * 7) as u8).collect();
let cap = 16 * 1024;
let r = BlockCacheReader::new(SliceReader::new(&data), 8192).with_cache_cap(cap);
let (off, len) = (1000usize, 96 * 1024usize); let out = r.read_at(off as u64, len as u64).unwrap();
assert_eq!(out, &data[off..off + len]);
assert!(
r.cached_bytes() <= cap,
"resident {} > cap {cap} after the read",
r.cached_bytes()
);
}
}