Skip to main content

range_cache/
reader.rs

1//! Optional async read-through types.
2
3use std::{
4    collections::BTreeMap,
5    num::NonZeroUsize,
6    ops::Range,
7    sync::{Arc, Weak},
8};
9
10use async_trait::async_trait;
11use bytes::{Bytes, BytesMut};
12use futures_util::{StreamExt, TryStreamExt, stream};
13use parking_lot::Mutex;
14use tokio::sync::{Mutex as AsyncMutex, Semaphore};
15
16use crate::{RangeCache, RangeError, cache::ReadPlan};
17
18/// An immutable source capable of reading byte ranges.
19///
20/// Implementations must return exactly `range.len()` bytes on success. A
21/// [`CachedReader`] checks this contract before admitting a response.
22#[async_trait]
23pub trait RangeReader<K>: Send + Sync {
24    /// Source-specific error.
25    type Error: std::error::Error + Send + Sync + 'static;
26
27    /// Reads the requested byte range for `key`.
28    async fn read_range(&self, key: &K, range: Range<usize>) -> Result<Bytes, Self::Error>;
29}
30
31/// Configuration for [`CachedReader`].
32#[derive(Clone, Copy, Debug, Eq, PartialEq)]
33pub struct ReaderConfig {
34    max_fetch_concurrency: NonZeroUsize,
35}
36
37impl ReaderConfig {
38    /// Creates a configuration with an explicitly non-zero concurrency limit.
39    #[must_use]
40    pub const fn new(max_fetch_concurrency: NonZeroUsize) -> Self {
41        Self {
42            max_fetch_concurrency,
43        }
44    }
45
46    /// Returns the global maximum number of concurrent source fetches.
47    #[must_use]
48    pub const fn max_fetch_concurrency(self) -> NonZeroUsize {
49        self.max_fetch_concurrency
50    }
51}
52
53/// A read-through failure.
54#[derive(Debug, thiserror::Error)]
55pub enum ReadError<E> {
56    /// Invalid range or cache payload.
57    #[error(transparent)]
58    Range(#[from] RangeError),
59    /// Source-specific failure.
60    #[error("range source failed: {0}")]
61    Source(#[source] E),
62    /// The source returned fewer or more bytes than requested.
63    #[error("source returned {actual} bytes for {range:?}; expected {expected}")]
64    ShortRead {
65        /// Requested byte range.
66        range: Range<usize>,
67        /// Required byte length.
68        expected: usize,
69        /// Actual byte length.
70        actual: usize,
71    },
72}
73
74#[derive(Clone, Debug, Eq, Ord, PartialEq, PartialOrd)]
75struct InFlightKey<K> {
76    key: K,
77    start: usize,
78    end: usize,
79}
80
81struct InFlightRegistry<K: Ord> {
82    entries: Mutex<BTreeMap<InFlightKey<K>, Weak<InFlightEntry>>>,
83}
84
85impl<K: Ord> Default for InFlightRegistry<K> {
86    fn default() -> Self {
87        Self {
88            entries: Mutex::new(BTreeMap::new()),
89        }
90    }
91}
92
93struct InFlightEntry {
94    response: AsyncMutex<Option<Bytes>>,
95}
96
97impl<K: Ord + Clone> InFlightRegistry<K> {
98    fn register(self: &Arc<Self>, request: InFlightKey<K>) -> InFlightRegistration<K> {
99        let entry = {
100            let mut entries = self.entries.lock();
101            let current = entries.get(&request).and_then(Weak::upgrade);
102            let entry = current.unwrap_or_else(|| {
103                Arc::new(InFlightEntry {
104                    response: AsyncMutex::new(None),
105                })
106            });
107            entries.insert(request.clone(), Arc::downgrade(&entry));
108            entry
109        };
110        InFlightRegistration {
111            registry: Arc::clone(self),
112            request,
113            entry,
114        }
115    }
116}
117
118struct InFlightRegistration<K: Ord> {
119    registry: Arc<InFlightRegistry<K>>,
120    request: InFlightKey<K>,
121    entry: Arc<InFlightEntry>,
122}
123
124impl<K: Ord> Drop for InFlightRegistration<K> {
125    fn drop(&mut self) {
126        let mut entries = self.registry.entries.lock();
127        let registered_is_self = entries
128            .get(&self.request)
129            .is_some_and(|registered| registered.ptr_eq(&Arc::downgrade(&self.entry)));
130        if registered_is_self && Arc::strong_count(&self.entry) == 1 {
131            entries.remove(&self.request);
132        }
133    }
134}
135
136/// An async read-through adapter over a [`RangeReader`].
137pub struct CachedReader<K: Ord, R: ?Sized> {
138    source: Arc<R>,
139    cache: RangeCache<K>,
140    config: ReaderConfig,
141    fetch_limit: Arc<Semaphore>,
142    in_flight: Arc<InFlightRegistry<K>>,
143}
144
145impl<K: Ord, R: ?Sized> Clone for CachedReader<K, R> {
146    fn clone(&self) -> Self {
147        Self {
148            source: Arc::clone(&self.source),
149            cache: self.cache.clone(),
150            config: self.config,
151            fetch_limit: Arc::clone(&self.fetch_limit),
152            in_flight: Arc::clone(&self.in_flight),
153        }
154    }
155}
156
157impl<K: Ord, R: ?Sized> CachedReader<K, R> {
158    /// Wraps `source` with the provided cache and concurrency policy.
159    #[must_use]
160    pub fn new(source: Arc<R>, cache: RangeCache<K>, config: ReaderConfig) -> Self {
161        Self {
162            source,
163            cache,
164            config,
165            fetch_limit: Arc::new(Semaphore::new(config.max_fetch_concurrency.get())),
166            in_flight: Arc::new(InFlightRegistry::default()),
167        }
168    }
169
170    /// Returns the shared cache.
171    #[must_use]
172    pub const fn cache(&self) -> &RangeCache<K> {
173        &self.cache
174    }
175
176    /// Returns the wrapped source.
177    #[must_use]
178    pub const fn source(&self) -> &Arc<R> {
179        &self.source
180    }
181
182    /// Returns the reader configuration.
183    #[must_use]
184    pub const fn config(&self) -> ReaderConfig {
185        self.config
186    }
187}
188
189impl<K, R> CachedReader<K, R>
190where
191    K: Ord + Clone,
192    R: RangeReader<K> + ?Sized,
193{
194    /// Reads a range, fetching and caching only its missing gaps.
195    ///
196    /// Identical in-flight key-and-gap requests share one source fetch. Other
197    /// overlapping requests remain independent. Responses rejected by the
198    /// cache capacity policy are still returned to the caller.
199    ///
200    /// # Errors
201    ///
202    /// Returns a validation error, source error, or [`ReadError::ShortRead`].
203    /// Source failures and invalid response lengths are never cached.
204    ///
205    /// # Panics
206    ///
207    /// Panics only if the privately owned fetch semaphore is unexpectedly
208    /// closed or an internal read-plan coverage invariant is violated.
209    pub async fn read(&self, key: &K, range: Range<usize>) -> Result<Bytes, ReadError<R::Error>> {
210        let (mut cached, mut missing) = match self.cache.read_plan(key, range.clone())? {
211            ReadPlan::Complete(bytes) => return Ok(bytes),
212            ReadPlan::Fetch { cached, missing } => (cached, missing),
213        };
214
215        if missing.len() == 1 {
216            let gap = missing.pop().expect("one missing gap exists");
217            let fetched = self.fetch_gap(key, gap).await?;
218            if cached.is_empty() && fetched.0 == range {
219                return Ok(fetched.1);
220            }
221
222            let position =
223                cached.partition_point(|(chunk_range, _)| chunk_range.start < fetched.0.start);
224            cached.insert(position, fetched);
225            return Ok(reconstruct(range, cached));
226        }
227
228        let fetched = stream::iter(missing)
229            .map(|gap| self.fetch_gap(key, gap))
230            .buffer_unordered(self.config.max_fetch_concurrency.get())
231            .try_collect::<Vec<_>>()
232            .await?;
233        cached.extend(fetched);
234        cached.sort_unstable_by_key(|(chunk_range, _)| chunk_range.start);
235        Ok(reconstruct(range, cached))
236    }
237
238    async fn fetch_gap(
239        &self,
240        key: &K,
241        range: Range<usize>,
242    ) -> Result<(Range<usize>, Bytes), ReadError<R::Error>> {
243        let registration = self.in_flight.register(InFlightKey {
244            key: key.clone(),
245            start: range.start,
246            end: range.end,
247        });
248        let mut response = registration.entry.response.lock().await;
249
250        if let Some(bytes) = response.clone() {
251            return Ok((range, bytes));
252        }
253
254        if let Some(bytes) = self.cache.get(key, range.clone())? {
255            return Ok((range, bytes));
256        }
257
258        let _fetch_permit = self
259            .fetch_limit
260            .acquire()
261            .await
262            .expect("private fetch semaphore remains open");
263        let bytes = self
264            .source
265            .read_range(key, range.clone())
266            .await
267            .map_err(ReadError::Source)?;
268        let expected = range.len();
269        if bytes.len() != expected {
270            return Err(ReadError::ShortRead {
271                range,
272                expected,
273                actual: bytes.len(),
274            });
275        }
276
277        let _ = self
278            .cache
279            .insert(key.clone(), range.clone(), bytes.clone())
280            .expect("validated source bytes match the requested range");
281        *response = Some(bytes.clone());
282        Ok((range, bytes))
283    }
284}
285
286fn reconstruct(range: Range<usize>, chunks: Vec<(Range<usize>, Bytes)>) -> Bytes {
287    let mut reconstructed = BytesMut::with_capacity(range.len());
288    let mut cursor = range.start;
289    for (chunk_range, bytes) in chunks {
290        assert_eq!(chunk_range.start, cursor, "read plan has no gaps");
291        assert_eq!(chunk_range.len(), bytes.len(), "chunk length is exact");
292        reconstructed.extend_from_slice(&bytes);
293        cursor = chunk_range.end;
294    }
295    assert_eq!(cursor, range.end, "read plan covers the request");
296    reconstructed.freeze()
297}
298
299#[async_trait]
300impl<K, R> RangeReader<K> for CachedReader<K, R>
301where
302    K: Ord + Clone + Send + Sync,
303    R: RangeReader<K> + ?Sized,
304{
305    type Error = ReadError<R::Error>;
306
307    async fn read_range(&self, key: &K, range: Range<usize>) -> Result<Bytes, Self::Error> {
308        self.read(key, range).await
309    }
310}
311
312#[cfg(test)]
313mod tests {
314    use std::{
315        num::NonZeroUsize,
316        sync::{
317            Arc,
318            atomic::{AtomicUsize, Ordering},
319        },
320    };
321
322    use async_trait::async_trait;
323    use bytes::Bytes;
324    use tokio::sync::Semaphore;
325
326    use super::{CachedReader, RangeReader, ReadError, ReaderConfig};
327    use crate::{CacheCapacity, RangeCache};
328
329    #[derive(Clone, Copy, Debug, thiserror::Error)]
330    #[error("controlled source failure")]
331    struct TestError;
332
333    struct ControlledSource {
334        started: Semaphore,
335        release: Semaphore,
336        calls: AtomicUsize,
337        fail: bool,
338    }
339
340    #[async_trait]
341    impl RangeReader<String> for ControlledSource {
342        type Error = TestError;
343
344        async fn read_range(
345            &self,
346            _key: &String,
347            range: std::ops::Range<usize>,
348        ) -> Result<Bytes, Self::Error> {
349            self.calls.fetch_add(1, Ordering::SeqCst);
350            self.started.add_permits(1);
351            self.release
352                .acquire()
353                .await
354                .expect("source release semaphore remains open")
355                .forget();
356            if self.fail {
357                return Err(TestError);
358            }
359            Ok(Bytes::from(vec![0; range.len()]))
360        }
361    }
362
363    fn source(fail: bool) -> Arc<ControlledSource> {
364        Arc::new(ControlledSource {
365            started: Semaphore::new(0),
366            release: Semaphore::new(0),
367            calls: AtomicUsize::new(0),
368            fail,
369        })
370    }
371
372    async fn read_key(
373        reader: CachedReader<String, ControlledSource>,
374    ) -> Result<Bytes, ReadError<TestError>> {
375        reader.read(&String::from("key"), 0..4).await
376    }
377
378    #[tokio::test]
379    async fn cancelled_leader_removes_its_in_flight_registration() {
380        let source = source(false);
381        let reader = CachedReader::new(
382            Arc::clone(&source),
383            RangeCache::new(CacheCapacity::Unbounded),
384            ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
385        );
386        let task_reader = reader.clone();
387        let task = tokio::spawn(read_key(task_reader));
388        source
389            .started
390            .acquire()
391            .await
392            .expect("source semaphore remains open")
393            .forget();
394        task.abort();
395        assert!(task.await.expect_err("task was cancelled").is_cancelled());
396
397        tokio::task::yield_now().await;
398        assert!(reader.in_flight.entries.lock().is_empty());
399    }
400
401    #[tokio::test]
402    async fn completed_requests_remove_their_in_flight_registration() {
403        let source = source(false);
404        let reader = CachedReader::new(
405            Arc::clone(&source),
406            RangeCache::new(CacheCapacity::Unbounded),
407            ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
408        );
409        let task_reader = reader.clone();
410        let task = tokio::spawn(read_key(task_reader));
411        source
412            .started
413            .acquire()
414            .await
415            .expect("source semaphore remains open")
416            .forget();
417        source.release.add_permits(1);
418        assert_eq!(
419            task.await
420                .expect("task completed")
421                .expect("source read succeeds"),
422            Bytes::from_static(&[0; 4])
423        );
424        assert!(reader.in_flight.entries.lock().is_empty());
425    }
426
427    #[tokio::test]
428    async fn failed_requests_remove_their_in_flight_registration() {
429        let source = source(true);
430        let reader = CachedReader::new(
431            Arc::clone(&source),
432            RangeCache::new(CacheCapacity::Unbounded),
433            ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
434        );
435        let task_reader = reader.clone();
436        let task = tokio::spawn(read_key(task_reader));
437        source
438            .started
439            .acquire()
440            .await
441            .expect("source semaphore remains open")
442            .forget();
443        source.release.add_permits(1);
444        assert_eq!(
445            task.await
446                .expect("task completed")
447                .expect_err("source read fails")
448                .to_string(),
449            "range source failed: controlled source failure"
450        );
451        assert!(reader.in_flight.entries.lock().is_empty());
452    }
453
454    #[tokio::test]
455    async fn fetch_gap_rechecks_the_cache_before_reading_the_source() {
456        let source = source(false);
457        let cache = RangeCache::new(CacheCapacity::Unbounded);
458        let key = String::from("key");
459        cache
460            .insert(key.clone(), 0..4, Bytes::from_static(b"data"))
461            .expect("valid insert");
462        let reader = CachedReader::new(
463            Arc::clone(&source),
464            cache,
465            ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
466        );
467
468        assert_eq!(
469            reader
470                .fetch_gap(&key, 0..4)
471                .await
472                .expect("cached gap succeeds"),
473            (0..4, Bytes::from_static(b"data"))
474        );
475        assert_eq!(source.calls.load(Ordering::SeqCst), 0);
476        assert!(reader.in_flight.entries.lock().is_empty());
477    }
478
479    #[tokio::test]
480    async fn fetch_gap_propagates_range_validation_errors() {
481        let source = source(false);
482        let reader = CachedReader::new(
483            Arc::clone(&source),
484            RangeCache::new(CacheCapacity::Unbounded),
485            ReaderConfig::new(NonZeroUsize::new(1).expect("non-zero")),
486        );
487        let reversed = std::ops::Range { start: 4, end: 3 };
488
489        assert_eq!(
490            reader
491                .fetch_gap(&String::from("key"), reversed)
492                .await
493                .expect_err("reversed range fails")
494                .to_string(),
495            "reversed byte range 4..3"
496        );
497        assert_eq!(source.calls.load(Ordering::SeqCst), 0);
498        assert!(reader.in_flight.entries.lock().is_empty());
499    }
500}