use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum BackpressurePolicy {
Block,
Shed,
#[default]
Fail,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct BackoffConfig {
pub initial_delay: Duration,
pub max_delay: Duration,
pub multiplier: f64,
pub max_retries: u32,
}
impl Default for BackoffConfig {
fn default() -> Self {
Self {
initial_delay: Duration::from_millis(200),
max_delay: Duration::from_secs(10),
multiplier: 2.0,
max_retries: 5,
}
}
}
impl BackoffConfig {
pub fn delay_for(&self, attempt: u32) -> Duration {
let factor = self.multiplier.powi(attempt.min(32) as i32);
let millis = (self.initial_delay.as_millis() as f64 * factor) as u64;
Duration::from_millis(millis).min(self.max_delay)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct BackpressureConfig {
pub spill_buffer_frames: usize,
pub policy: BackpressurePolicy,
pub backoff: BackoffConfig,
}
impl Default for BackpressureConfig {
fn default() -> Self {
Self {
spill_buffer_frames: 256,
policy: BackpressurePolicy::default(),
backoff: BackoffConfig::default(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct BackpressureReport {
pub policy: BackpressurePolicy,
pub high_water_frames: usize,
pub frames_shed: u64,
pub throttle_retries: u32,
}
pub(crate) fn is_throttling_error(err: &object_store::Error) -> bool {
let mut msg = err.to_string();
let mut cause = std::error::Error::source(err);
while let Some(c) = cause {
msg.push_str(" | ");
msg.push_str(&c.to_string());
cause = c.source();
}
const NEEDLES: [&str; 6] = [
"429",
"503",
"Too Many Requests",
"Service Unavailable",
"SlowDown",
"RequestThrottled",
];
NEEDLES.iter().any(|n| msg.contains(n))
}
pub(crate) async fn put_with_backoff(
store: &std::sync::Arc<dyn object_store::ObjectStore>,
key: &object_store::path::Path,
payload: object_store::PutPayload,
cfg: &BackpressureConfig,
retries_used: &mut u32,
) -> Result<object_store::PutResult, object_store::Error> {
use object_store::ObjectStoreExt;
let mut attempt: u32 = 0;
loop {
match store.put(key, payload.clone()).await {
Ok(r) => return Ok(r),
Err(e) if is_throttling_error(&e) => {
let unbounded = matches!(cfg.policy, BackpressurePolicy::Block);
if !unbounded && attempt >= cfg.backoff.max_retries {
return Err(e);
}
tokio::time::sleep(cfg.backoff.delay_for(attempt)).await;
attempt += 1;
*retries_used += 1;
}
Err(e) => return Err(e),
}
}
}
#[cfg(test)]
pub(crate) mod fault_injection {
use async_trait::async_trait;
use futures_util::stream::BoxStream;
use object_store::{
path::Path as ObjPath, CopyOptions, Error as OsError, GetOptions, GetResult, ListResult,
MultipartUpload, ObjectMeta, ObjectStore, PutMultipartOptions, PutOptions, PutPayload,
PutResult, Result as OsResult,
};
use std::collections::VecDeque;
use std::fmt;
use std::sync::{Arc, Mutex};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Fault {
TooManyRequests,
ServiceUnavailable,
Pass,
}
#[derive(Debug)]
struct SimulatedThrottle(&'static str);
impl fmt::Display for SimulatedThrottle {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.0)
}
}
impl std::error::Error for SimulatedThrottle {}
pub(crate) struct FaultyStore {
inner: Arc<dyn ObjectStore>,
faults: Mutex<VecDeque<Fault>>,
}
impl FaultyStore {
pub(crate) fn new(
inner: Arc<dyn ObjectStore>,
faults: impl IntoIterator<Item = Fault>,
) -> Self {
Self {
inner,
faults: Mutex::new(faults.into_iter().collect()),
}
}
}
impl fmt::Display for FaultyStore {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "FaultyStore({})", self.inner)
}
}
impl fmt::Debug for FaultyStore {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "FaultyStore({:?})", self.inner)
}
}
#[async_trait]
impl ObjectStore for FaultyStore {
async fn put_opts(
&self,
location: &ObjPath,
payload: PutPayload,
opts: PutOptions,
) -> OsResult<PutResult> {
let next = self.faults.lock().unwrap().pop_front();
if let Some(fault) = next.filter(|f| *f != Fault::Pass) {
let msg: &'static str = match fault {
Fault::TooManyRequests => {
"HTTP status client error (429 Too Many Requests) for url"
}
Fault::ServiceUnavailable => {
"HTTP status server error (503 Service Unavailable) for url"
}
Fault::Pass => unreachable!("filtered out above"),
};
return Err(OsError::Generic {
store: "faulty-test-store",
source: Box::new(SimulatedThrottle(msg)),
});
}
self.inner.put_opts(location, payload, opts).await
}
async fn put_multipart_opts(
&self,
location: &ObjPath,
opts: PutMultipartOptions,
) -> OsResult<Box<dyn MultipartUpload>> {
self.inner.put_multipart_opts(location, opts).await
}
async fn get_opts(&self, location: &ObjPath, options: GetOptions) -> OsResult<GetResult> {
self.inner.get_opts(location, options).await
}
fn delete_stream(
&self,
locations: BoxStream<'static, OsResult<ObjPath>>,
) -> BoxStream<'static, OsResult<ObjPath>> {
self.inner.delete_stream(locations)
}
fn list(&self, prefix: Option<&ObjPath>) -> BoxStream<'static, OsResult<ObjectMeta>> {
self.inner.list(prefix)
}
async fn list_with_delimiter(&self, prefix: Option<&ObjPath>) -> OsResult<ListResult> {
self.inner.list_with_delimiter(prefix).await
}
async fn copy_opts(
&self,
from: &ObjPath,
to: &ObjPath,
options: CopyOptions,
) -> OsResult<()> {
self.inner.copy_opts(from, to, options).await
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backoff_delay_grows_and_caps() {
let cfg = BackoffConfig {
initial_delay: Duration::from_millis(100),
max_delay: Duration::from_millis(500),
multiplier: 2.0,
max_retries: 10,
};
assert_eq!(cfg.delay_for(0), Duration::from_millis(100));
assert_eq!(cfg.delay_for(1), Duration::from_millis(200));
assert_eq!(cfg.delay_for(2), Duration::from_millis(400));
assert_eq!(cfg.delay_for(3), Duration::from_millis(500), "capped at max_delay");
assert_eq!(cfg.delay_for(10), Duration::from_millis(500));
}
#[test]
fn classifies_429_and_503_as_throttling_but_not_other_errors() {
let e429 = object_store::Error::Generic {
store: "t",
source: Box::new(std::io::Error::other("429 Too Many Requests")),
};
let e503 = object_store::Error::Generic {
store: "t",
source: Box::new(std::io::Error::other("503 Service Unavailable")),
};
let e_not_found = object_store::Error::NotFound {
path: "x".into(),
source: Box::new(std::io::Error::other("nope")),
};
assert!(is_throttling_error(&e429));
assert!(is_throttling_error(&e503));
assert!(!is_throttling_error(&e_not_found));
}
}