use std::{
collections::BTreeMap,
num::NonZeroUsize,
ops::Range,
sync::{Arc, Weak},
};
use async_trait::async_trait;
use bytes::{Bytes, BytesMut};
use futures_util::{StreamExt, TryStreamExt, stream};
use parking_lot::Mutex;
use tokio::sync::{Mutex as AsyncMutex, Semaphore};
use crate::{RangeCache, RangeError, cache::ReadPlan};
#[async_trait]
pub trait RangeReader<K>: Send + Sync {
type Error: std::error::Error + Send + Sync + 'static;
async fn read_range(&self, key: &K, range: Range<usize>) -> Result<Bytes, Self::Error>;
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ReaderConfig {
max_fetch_concurrency: NonZeroUsize,
}
impl ReaderConfig {
#[must_use]
pub const fn new(max_fetch_concurrency: NonZeroUsize) -> Self {
Self {
max_fetch_concurrency,
}
}
#[must_use]
pub const fn max_fetch_concurrency(self) -> NonZeroUsize {
self.max_fetch_concurrency
}
}
#[derive(Debug, thiserror::Error)]
pub enum ReadError<E> {
#[error(transparent)]
Range(#[from] RangeError),
#[error("range source failed: {0}")]
Source(#[source] E),
#[error("source returned {actual} bytes for {range:?}; expected {expected}")]
ShortRead {
range: Range<usize>,
expected: usize,
actual: usize,
},
}
#[derive(Clone, Debug, Eq, Ord, PartialEq, PartialOrd)]
struct InFlightKey<K> {
key: K,
start: usize,
end: usize,
}
struct InFlightRegistry<K: Ord> {
entries: Mutex<BTreeMap<InFlightKey<K>, Weak<InFlightEntry>>>,
}
impl<K: Ord> Default for InFlightRegistry<K> {
fn default() -> Self {
Self {
entries: Mutex::new(BTreeMap::new()),
}
}
}
struct InFlightEntry {
response: AsyncMutex<Option<Bytes>>,
}
impl<K: Ord + Clone> InFlightRegistry<K> {
fn register(self: &Arc<Self>, request: InFlightKey<K>) -> InFlightRegistration<K> {
let entry = {
let mut entries = self.entries.lock();
let current = entries.get(&request).and_then(Weak::upgrade);
let entry = current.unwrap_or_else(|| {
Arc::new(InFlightEntry {
response: AsyncMutex::new(None),
})
});
entries.insert(request.clone(), Arc::downgrade(&entry));
entry
};
InFlightRegistration {
registry: Arc::clone(self),
request,
entry,
}
}
}
struct InFlightRegistration<K: Ord> {
registry: Arc<InFlightRegistry<K>>,
request: InFlightKey<K>,
entry: Arc<InFlightEntry>,
}
impl<K: Ord> Drop for InFlightRegistration<K> {
fn drop(&mut self) {
let mut entries = self.registry.entries.lock();
let registered_is_self = entries
.get(&self.request)
.is_some_and(|registered| registered.ptr_eq(&Arc::downgrade(&self.entry)));
if registered_is_self && Arc::strong_count(&self.entry) == 1 {
entries.remove(&self.request);
}
}
}
pub struct CachedReader<K: Ord, R: ?Sized> {
source: Arc<R>,
cache: RangeCache<K>,
config: ReaderConfig,
fetch_limit: Arc<Semaphore>,
in_flight: Arc<InFlightRegistry<K>>,
}
impl<K: Ord, R: ?Sized> Clone for CachedReader<K, R> {
fn clone(&self) -> Self {
Self {
source: Arc::clone(&self.source),
cache: self.cache.clone(),
config: self.config,
fetch_limit: Arc::clone(&self.fetch_limit),
in_flight: Arc::clone(&self.in_flight),
}
}
}
impl<K: Ord, R: ?Sized> CachedReader<K, R> {
#[must_use]
pub fn new(source: Arc<R>, cache: RangeCache<K>, config: ReaderConfig) -> Self {
Self {
source,
cache,
config,
fetch_limit: Arc::new(Semaphore::new(config.max_fetch_concurrency.get())),
in_flight: Arc::new(InFlightRegistry::default()),
}
}
#[must_use]
pub const fn cache(&self) -> &RangeCache<K> {
&self.cache
}
#[must_use]
pub const fn source(&self) -> &Arc<R> {
&self.source
}
#[must_use]
pub const fn config(&self) -> ReaderConfig {
self.config
}
}
impl<K, R> CachedReader<K, R>
where
K: Ord + Clone,
R: RangeReader<K> + ?Sized,
{
pub async fn read(&self, key: &K, range: Range<usize>) -> Result<Bytes, ReadError<R::Error>> {
let (mut cached, mut missing) = match self.cache.read_plan(key, range.clone())? {
ReadPlan::Complete(bytes) => return Ok(bytes),
ReadPlan::Fetch { cached, missing } => (cached, missing),
};
if missing.len() == 1 {
let gap = missing.pop().expect("one missing gap exists");
let fetched = self.fetch_gap(key, gap).await?;
if cached.is_empty() && fetched.0 == range {
return Ok(fetched.1);
}
let position =
cached.partition_point(|(chunk_range, _)| chunk_range.start < fetched.0.start);
cached.insert(position, fetched);
return Ok(reconstruct(range, cached));
}
let fetched = stream::iter(missing)
.map(|gap| self.fetch_gap(key, gap))
.buffer_unordered(self.config.max_fetch_concurrency.get())
.try_collect::<Vec<_>>()
.await?;
cached.extend(fetched);
cached.sort_unstable_by_key(|(chunk_range, _)| chunk_range.start);
Ok(reconstruct(range, cached))
}
async fn fetch_gap(
&self,
key: &K,
range: Range<usize>,
) -> Result<(Range<usize>, Bytes), ReadError<R::Error>> {
let registration = self.in_flight.register(InFlightKey {
key: key.clone(),
start: range.start,
end: range.end,
});
let mut response = registration.entry.response.lock().await;
if let Some(bytes) = response.clone() {
return Ok((range, bytes));
}
if let Some(bytes) = self.cache.get(key, range.clone())? {
return Ok((range, bytes));
}
let _fetch_permit = self
.fetch_limit
.acquire()
.await
.expect("private fetch semaphore remains open");
let bytes = self
.source
.read_range(key, range.clone())
.await
.map_err(ReadError::Source)?;
let expected = range.len();
if bytes.len() != expected {
return Err(ReadError::ShortRead {
range,
expected,
actual: bytes.len(),
});
}
let _ = self
.cache
.insert(key.clone(), range.clone(), bytes.clone())
.expect("validated source bytes match the requested range");
*response = Some(bytes.clone());
Ok((range, bytes))
}
}
fn reconstruct(range: Range<usize>, chunks: Vec<(Range<usize>, Bytes)>) -> Bytes {
let mut reconstructed = BytesMut::with_capacity(range.len());
let mut cursor = range.start;
for (chunk_range, bytes) in chunks {
assert_eq!(chunk_range.start, cursor, "read plan has no gaps");
assert_eq!(chunk_range.len(), bytes.len(), "chunk length is exact");
reconstructed.extend_from_slice(&bytes);
cursor = chunk_range.end;
}
assert_eq!(cursor, range.end, "read plan covers the request");
reconstructed.freeze()
}
#[async_trait]
impl<K, R> RangeReader<K> for CachedReader<K, R>
where
K: Ord + Clone + Send + Sync,
R: RangeReader<K> + ?Sized,
{
type Error = ReadError<R::Error>;
async fn read_range(&self, key: &K, range: Range<usize>) -> Result<Bytes, Self::Error> {
self.read(key, range).await
}
}
#[cfg(test)]
mod tests {
use std::{
num::NonZeroUsize,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
};
use async_trait::async_trait;
use bytes::Bytes;
use tokio::sync::Semaphore;
use super::{CachedReader, RangeReader, ReadError, ReaderConfig};
use crate::{CacheCapacity, RangeCache};
#[derive(Clone, Copy, Debug, thiserror::Error)]
#[error("controlled source failure")]
struct TestError;
struct ControlledSource {
started: Semaphore,
release: Semaphore,
calls: AtomicUsize,
fail: bool,
}
#[async_trait]
impl RangeReader<String> for ControlledSource {
type Error = TestError;
async fn read_range(
&self,
_key: &String,
range: std::ops::Range<usize>,
) -> Result<Bytes, Self::Error> {
self.calls.fetch_add(1, Ordering::SeqCst);
self.started.add_permits(1);
self.release
.acquire()
.await
.expect("source release semaphore remains open")
.forget();
if self.fail {
return Err(TestError);
}
Ok(Bytes::from(vec![0; range.len()]))
}
}
fn source(fail: bool) -> Arc<ControlledSource> {
Arc::new(ControlledSource {
started: Semaphore::new(0),
release: Semaphore::new(0),
calls: AtomicUsize::new(0),
fail,
})
}
async fn read_key(
reader: CachedReader<String, ControlledSource>,
) -> Result<Bytes, ReadError<TestError>> {
reader.read(&String::from("key"), 0..4).await
}
#[tokio::test]
async fn cancelled_leader_removes_its_in_flight_registration() {
let source = source(false);
let reader = CachedReader::new(
Arc::clone(&source),
RangeCache::new(CacheCapacity::Unbounded),
ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
);
let task_reader = reader.clone();
let task = tokio::spawn(read_key(task_reader));
source
.started
.acquire()
.await
.expect("source semaphore remains open")
.forget();
task.abort();
assert!(task.await.expect_err("task was cancelled").is_cancelled());
tokio::task::yield_now().await;
assert!(reader.in_flight.entries.lock().is_empty());
}
#[tokio::test]
async fn completed_requests_remove_their_in_flight_registration() {
let source = source(false);
let reader = CachedReader::new(
Arc::clone(&source),
RangeCache::new(CacheCapacity::Unbounded),
ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
);
let task_reader = reader.clone();
let task = tokio::spawn(read_key(task_reader));
source
.started
.acquire()
.await
.expect("source semaphore remains open")
.forget();
source.release.add_permits(1);
assert_eq!(
task.await
.expect("task completed")
.expect("source read succeeds"),
Bytes::from_static(&[0; 4])
);
assert!(reader.in_flight.entries.lock().is_empty());
}
#[tokio::test]
async fn failed_requests_remove_their_in_flight_registration() {
let source = source(true);
let reader = CachedReader::new(
Arc::clone(&source),
RangeCache::new(CacheCapacity::Unbounded),
ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
);
let task_reader = reader.clone();
let task = tokio::spawn(read_key(task_reader));
source
.started
.acquire()
.await
.expect("source semaphore remains open")
.forget();
source.release.add_permits(1);
assert_eq!(
task.await
.expect("task completed")
.expect_err("source read fails")
.to_string(),
"range source failed: controlled source failure"
);
assert!(reader.in_flight.entries.lock().is_empty());
}
#[tokio::test]
async fn fetch_gap_rechecks_the_cache_before_reading_the_source() {
let source = source(false);
let cache = RangeCache::new(CacheCapacity::Unbounded);
let key = String::from("key");
cache
.insert(key.clone(), 0..4, Bytes::from_static(b"data"))
.expect("valid insert");
let reader = CachedReader::new(
Arc::clone(&source),
cache,
ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
);
assert_eq!(
reader
.fetch_gap(&key, 0..4)
.await
.expect("cached gap succeeds"),
(0..4, Bytes::from_static(b"data"))
);
assert_eq!(source.calls.load(Ordering::SeqCst), 0);
assert!(reader.in_flight.entries.lock().is_empty());
}
#[tokio::test]
async fn fetch_gap_propagates_range_validation_errors() {
let source = source(false);
let reader = CachedReader::new(
Arc::clone(&source),
RangeCache::new(CacheCapacity::Unbounded),
ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
);
let reversed = std::ops::Range { start: 4, end: 3 };
assert_eq!(
reader
.fetch_gap(&String::from("key"), reversed)
.await
.expect_err("reversed range fails")
.to_string(),
"reversed byte range 4..3"
);
assert_eq!(source.calls.load(Ordering::SeqCst), 0);
assert!(reader.in_flight.entries.lock().is_empty());
}
}