use std::sync::Arc;
use http::Request;
use crate::{
Body,
client::layer::retry::{Action, Classifier, ClassifyFn, ReqRep, ScopeFn, Scoped},
};
#[derive(Clone)]
pub struct Policy {
pub(crate) budget: Option<f32>,
pub(crate) classifier: Classifier,
pub(crate) max_retries_per_request: u32,
pub(crate) scope: Scoped,
}
impl Policy {
#[inline]
pub fn never() -> Policy {
Self::scoped(|_| false).no_budget()
}
#[inline]
pub fn for_host<S>(host: S) -> Policy
where
S: for<'a> PartialEq<&'a str> + Send + Sync + 'static,
{
Self::scoped(move |req| {
req.uri()
.host()
.is_some_and(|request_host| host == request_host)
})
}
#[inline]
fn scoped<F>(func: F) -> Policy
where
F: Fn(&Request<Body>) -> bool + Send + Sync + 'static,
{
Self {
budget: Some(0.2),
classifier: Classifier::Never,
max_retries_per_request: 2,
scope: Scoped::Dyn(Arc::new(ScopeFn(func))),
}
}
#[inline]
pub fn no_budget(mut self) -> Self {
self.budget = None;
self
}
#[inline]
pub fn max_extra_load(mut self, extra_percent: f32) -> Self {
assert!(extra_percent >= 0.0);
assert!(extra_percent <= 1000.0);
self.budget = Some(extra_percent);
self
}
#[inline]
pub fn max_retries_per_request(mut self, max: u32) -> Self {
self.max_retries_per_request = max;
self
}
#[inline]
pub fn classify_fn<F>(mut self, func: F) -> Self
where
F: Fn(ReqRep<'_>) -> Action + Send + Sync + 'static,
{
self.classifier = Classifier::Dyn(Arc::new(ClassifyFn(func)));
self
}
}
impl Default for Policy {
fn default() -> Self {
Self {
budget: None,
classifier: Classifier::ProtocolNacks,
max_retries_per_request: 2,
scope: Scoped::Unscoped,
}
}
}
#[cfg(test)]
mod tests {
use super::Policy;
use crate::client::layer::retry::{Classifier, Scoped};
#[test]
fn default_policy_is_protocol_nacks_unscoped() {
let policy = Policy::default();
assert!(
matches!(policy.classifier, Classifier::ProtocolNacks),
"default classifier should be ProtocolNacks"
);
assert!(
matches!(policy.scope, Scoped::Unscoped),
"default scope should be Unscoped"
);
assert!(policy.budget.is_none(), "default budget should be None");
assert_eq!(policy.max_retries_per_request, 2);
}
#[test]
fn never_policy_has_never_classifier_no_budget() {
let policy = Policy::never();
assert!(
matches!(policy.classifier, Classifier::Never),
"never() classifier should be Never"
);
assert!(policy.budget.is_none(), "never() budget should be None");
}
#[test]
fn for_host_creates_scoped_policy() {
let policy = Policy::for_host("example.com");
assert!(
matches!(policy.scope, Scoped::Dyn(_)),
"for_host should create a scoped policy"
);
assert!(
matches!(policy.classifier, Classifier::Never),
"for_host classifier should be Never by default"
);
}
#[test]
fn classify_fn_sets_custom_classifier() {
let policy = Policy::never().classify_fn(|req_rep| {
if req_rep.method() == &http::Method::POST {
req_rep.retryable()
} else {
req_rep.success()
}
});
assert!(
matches!(policy.classifier, Classifier::Dyn(_)),
"classify_fn should set a Dyn classifier"
);
}
#[test]
fn max_retries_per_request_configuration() {
let policy = Policy::never().max_retries_per_request(5);
assert_eq!(policy.max_retries_per_request, 5);
}
#[test]
fn budget_configuration() {
let base = Policy::for_host("example.com");
let no_budget = base.clone().no_budget();
assert!(no_budget.budget.is_none());
let with_budget = base.max_extra_load(0.5);
assert_eq!(with_budget.budget, Some(0.5));
}
#[test]
fn combined_configuration() {
let policy = Policy::for_host("api.example.com")
.classify_fn(|req_rep| match req_rep.status() {
Some(s) if s.is_server_error() => req_rep.retryable(),
_ => req_rep.success(),
})
.max_retries_per_request(3)
.max_extra_load(0.3);
assert!(matches!(policy.classifier, Classifier::Dyn(_)));
assert!(matches!(policy.scope, Scoped::Dyn(_)));
assert_eq!(policy.max_retries_per_request, 3);
assert_eq!(policy.budget, Some(0.3));
}
mod proptests {
use proptest::prelude::*;
use super::Policy;
proptest! {
#[test]
fn max_extra_load_round_trip(budget in 0.0f32..=1000.0) {
let policy = Policy::for_host("example.com").max_extra_load(budget);
prop_assert_eq!(policy.budget, Some(budget));
}
#[test]
fn no_budget_clears_any_budget(budget in 0.0f32..=1000.0) {
let policy = Policy::for_host("example.com")
.max_extra_load(budget)
.no_budget();
prop_assert!(policy.budget.is_none());
}
}
#[test]
#[should_panic(expected = "assertion failed")]
fn max_extra_load_panics_on_negative() {
Policy::for_host("example.com").max_extra_load(-0.1);
}
#[test]
#[should_panic(expected = "assertion failed")]
fn max_extra_load_panics_on_over_limit() {
Policy::for_host("example.com").max_extra_load(1000.1);
}
#[test]
fn max_retries_never_negative() {
let policy = Policy::default();
assert!(policy.max_retries_per_request > 0);
let policy = Policy::never().max_retries_per_request(100);
assert_eq!(policy.max_retries_per_request, 100);
}
}
}