1use 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#[async_trait]
23pub trait RangeReader<K>: Send + Sync {
24 type Error: std::error::Error + Send + Sync + 'static;
26
27 async fn read_range(&self, key: &K, range: Range<usize>) -> Result<Bytes, Self::Error>;
29}
30
31#[derive(Clone, Copy, Debug, Eq, PartialEq)]
33pub struct ReaderConfig {
34 max_fetch_concurrency: NonZeroUsize,
35}
36
37impl ReaderConfig {
38 #[must_use]
40 pub const fn new(max_fetch_concurrency: NonZeroUsize) -> Self {
41 Self {
42 max_fetch_concurrency,
43 }
44 }
45
46 #[must_use]
48 pub const fn max_fetch_concurrency(self) -> NonZeroUsize {
49 self.max_fetch_concurrency
50 }
51}
52
53#[derive(Debug, thiserror::Error)]
55pub enum ReadError<E> {
56 #[error(transparent)]
58 Range(#[from] RangeError),
59 #[error("range source failed: {0}")]
61 Source(#[source] E),
62 #[error("source returned {actual} bytes for {range:?}; expected {expected}")]
64 ShortRead {
65 range: Range<usize>,
67 expected: usize,
69 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
136pub 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 #[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 #[must_use]
172 pub const fn cache(&self) -> &RangeCache<K> {
173 &self.cache
174 }
175
176 #[must_use]
178 pub const fn source(&self) -> &Arc<R> {
179 &self.source
180 }
181
182 #[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 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}