use std::{
fmt,
sync::atomic::{AtomicU64, Ordering},
};
use anyhow::Result;
use async_trait::async_trait;
use super::DataReaderTrait;
use crate::{Blob, ByteRange};
#[derive(Debug)]
pub(crate) struct SmallerRangeWontHelp;
impl fmt::Display for SmallerRangeWontHelp {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("retrying with smaller ranges would not help")
}
}
impl std::error::Error for SmallerRangeWontHelp {}
fn splitting_could_help(error: &anyhow::Error) -> bool {
error.downcast_ref::<SmallerRangeWontHelp>().is_none()
}
#[async_trait]
pub(crate) trait NetworkReader: DataReaderTrait {
async fn try_read_range(&self, range: &ByteRange) -> Result<Blob>;
fn max_request_bytes(&self) -> &AtomicU64;
async fn network_read_range(&self, range: &ByteRange) -> Result<Blob> {
log::trace!("network_read_range {range} ({} bytes)", range.length);
if range.length == 0 {
return Ok(Blob::default());
}
if range.length > self.max_request_bytes().load(Ordering::Relaxed) && range.length > 1 {
log::trace!(
"proactively splitting range {range} ({} bytes) based on previous failures",
range.length
);
return self.split_and_read(range).await;
}
match self.try_read_range(range).await {
Ok(blob) => Ok(blob),
Err(e) if range.length <= 1 || !splitting_could_help(&e) => Err(e),
Err(e) => {
self.max_request_bytes().fetch_min(range.length / 2, Ordering::Relaxed);
log::debug!(
"splitting failed range {range} ({} bytes) into two halves: {e}",
range.length
);
self.split_and_read(range).await
}
}
}
async fn split_and_read(&self, range: &ByteRange) -> Result<Blob> {
range.end()?;
let mid = range.offset + range.length / 2;
let left = ByteRange::new(range.offset, mid - range.offset);
let right = ByteRange::new(mid, range.offset + range.length - mid);
log::trace!("split_and_read {range} -> [{left}] + [{right}]");
let (blob_left, blob_right) = futures::future::try_join(self.read_range(&left), self.read_range(&right)).await?;
let mut data = blob_left.into_vec();
data.extend_from_slice(blob_right.as_slice());
Ok(Blob::from(data))
}
}
#[cfg(test)]
mod tests {
use std::{
fmt,
sync::{
Arc,
atomic::{AtomicU64, AtomicUsize, Ordering as AtomicOrdering},
},
time::Duration,
};
use super::*;
#[derive(Default)]
struct PeakState {
in_flight: AtomicUsize,
max_in_flight: AtomicUsize,
max_request: AtomicU64,
}
struct PeakNetReader {
state: Arc<PeakState>,
delay: Duration,
}
impl fmt::Debug for PeakNetReader {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PeakNetReader").finish()
}
}
#[async_trait]
impl DataReaderTrait for PeakNetReader {
async fn read_range(&self, range: &ByteRange) -> Result<Blob> {
self.network_read_range(range).await
}
async fn read_all(&self) -> Result<Blob> {
unreachable!("PeakNetReader only used for read_range")
}
fn name(&self) -> &str {
"peak-net"
}
}
#[async_trait]
impl NetworkReader for PeakNetReader {
async fn try_read_range(&self, range: &ByteRange) -> Result<Blob> {
let n = self.state.in_flight.fetch_add(1, AtomicOrdering::SeqCst) + 1;
self.state.max_in_flight.fetch_max(n, AtomicOrdering::SeqCst);
tokio::time::sleep(self.delay).await;
self.state.in_flight.fetch_sub(1, AtomicOrdering::SeqCst);
Ok(Blob::from(vec![0u8; usize::try_from(range.length).unwrap()]))
}
fn max_request_bytes(&self) -> &AtomicU64 {
&self.state.max_request
}
}
#[tokio::test]
async fn zero_length_read_short_circuits() {
let state = Arc::new(PeakState {
in_flight: AtomicUsize::new(0),
max_in_flight: AtomicUsize::new(0),
max_request: AtomicU64::new(u64::MAX),
});
let reader = PeakNetReader {
state: Arc::clone(&state),
delay: Duration::from_millis(10),
};
let blob = reader.network_read_range(&ByteRange::new(4117, 0)).await.unwrap();
assert_eq!(blob.len(), 0);
assert_eq!(state.max_in_flight.load(AtomicOrdering::SeqCst), 0);
}
struct FailingReader {
calls: AtomicUsize,
max_request: AtomicU64,
splittable: bool,
}
impl fmt::Debug for FailingReader {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("FailingReader").finish()
}
}
#[async_trait]
impl DataReaderTrait for FailingReader {
async fn read_range(&self, range: &ByteRange) -> Result<Blob> {
self.network_read_range(range).await
}
async fn read_all(&self) -> Result<Blob> {
unreachable!("FailingReader only used for read_range")
}
fn name(&self) -> &str {
"failing"
}
}
#[async_trait]
impl NetworkReader for FailingReader {
async fn try_read_range(&self, _range: &ByteRange) -> Result<Blob> {
self.calls.fetch_add(1, AtomicOrdering::SeqCst);
let error = anyhow::anyhow!("nope");
Err(if self.splittable {
error
} else {
error.context(SmallerRangeWontHelp)
})
}
fn max_request_bytes(&self) -> &AtomicU64 {
&self.max_request
}
}
#[tokio::test]
async fn an_unfixable_error_is_not_split() {
let reader = FailingReader {
calls: AtomicUsize::new(0),
max_request: AtomicU64::new(u64::MAX),
splittable: false,
};
let result = reader.network_read_range(&ByteRange::new(0, 1 << 20)).await;
assert!(result.is_err());
assert_eq!(
reader.calls.load(AtomicOrdering::SeqCst),
1,
"an unfixable error must not be retried in halves"
);
}
#[tokio::test]
async fn a_plain_error_still_splits() {
let reader = FailingReader {
calls: AtomicUsize::new(0),
max_request: AtomicU64::new(u64::MAX),
splittable: true,
};
let result = reader.network_read_range(&ByteRange::new(0, 64)).await;
assert!(result.is_err());
assert!(
reader.calls.load(AtomicOrdering::SeqCst) > 1,
"a size-related failure should still be retried in halves"
);
}
#[tokio::test]
async fn split_and_read_runs_halves_concurrently() {
let state = Arc::new(PeakState {
in_flight: AtomicUsize::new(0),
max_in_flight: AtomicUsize::new(0),
max_request: AtomicU64::new(10),
});
let reader = PeakNetReader {
state: Arc::clone(&state),
delay: Duration::from_millis(40),
};
let blob = reader.network_read_range(&ByteRange::new(0, 100)).await.unwrap();
assert_eq!(blob.len(), 100);
let peak = state.max_in_flight.load(AtomicOrdering::SeqCst);
assert!(peak >= 2, "expected concurrent split halves, saw peak {peak} in flight");
}
}