use std::borrow::Cow;
use std::hash::{Hash, Hasher};
use serde::{Deserialize, Serialize};
use crate::error::{ModelError, ModelResult};
#[derive(Clone, Debug, Serialize, Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[cfg_attr(feature = "schema", schemars(!try_from, deny_unknown_fields))]
#[serde(rename_all = "camelCase")]
#[serde(try_from = "raw::BackoffPolicyRaw")]
pub struct BackoffPolicy {
pub jitter: super::JitterPolicy,
#[cfg_attr(feature = "schema", schemars(range(min = 1)))]
pub first_ms: u64,
#[cfg_attr(feature = "schema", schemars(range(min = 1)))]
pub max_ms: u64,
#[cfg_attr(feature = "schema", schemars(range(min = 1.0)))]
pub factor: f64,
}
mod raw {
use super::*;
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub(super) struct BackoffPolicyRaw {
pub jitter: super::super::JitterPolicy,
pub first_ms: u64,
pub max_ms: u64,
pub factor: f64,
}
impl TryFrom<BackoffPolicyRaw> for BackoffPolicy {
type Error = ModelError;
fn try_from(r: BackoffPolicyRaw) -> Result<Self, Self::Error> {
let p = BackoffPolicy {
jitter: r.jitter,
first_ms: r.first_ms,
max_ms: r.max_ms,
factor: r.factor,
};
p.validate()?;
Ok(p)
}
}
}
impl BackoffPolicy {
pub fn validate(&self) -> ModelResult<()> {
if self.first_ms == 0 {
return Err(ModelError::Invalid(Cow::Borrowed(
"backoff first_ms must be greater than zero",
)));
}
if self.max_ms < self.first_ms {
return Err(ModelError::Invalid(Cow::Borrowed(
"backoff max_ms must be >= first_ms",
)));
}
if !self.factor.is_finite() || self.factor < 1.0 {
return Err(ModelError::Invalid(Cow::Borrowed(
"backoff factor must be finite and >= 1.0",
)));
}
Ok(())
}
}
impl PartialEq for BackoffPolicy {
fn eq(&self, other: &Self) -> bool {
self.jitter == other.jitter
&& self.factor.to_bits() == other.factor.to_bits()
&& self.first_ms == other.first_ms
&& self.max_ms == other.max_ms
}
}
impl Eq for BackoffPolicy {}
impl Hash for BackoffPolicy {
fn hash<H: Hasher>(&self, state: &mut H) {
self.factor.to_bits().hash(state);
self.first_ms.hash(state);
self.jitter.hash(state);
self.max_ms.hash(state);
}
}
impl Default for BackoffPolicy {
fn default() -> Self {
Self {
jitter: super::JitterPolicy::Full,
first_ms: 1_000,
max_ms: 30_000,
factor: 2.0,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validation_accepts_defaults_and_rejects_invalid_fields() {
BackoffPolicy::default().validate().unwrap();
for invalid in [
BackoffPolicy {
first_ms: 0,
..BackoffPolicy::default()
},
BackoffPolicy {
first_ms: 500,
max_ms: 100,
..BackoffPolicy::default()
},
BackoffPolicy {
factor: 0.5,
..BackoffPolicy::default()
},
BackoffPolicy {
factor: f64::NAN,
..BackoffPolicy::default()
},
] {
assert!(invalid.validate().is_err());
}
}
#[test]
fn serde_roundtrip_accepts_valid_policy() {
let policy = BackoffPolicy::default();
let json = serde_json::to_string(&policy).unwrap();
let back: BackoffPolicy = serde_json::from_str(&json).unwrap();
assert_eq!(back, policy);
}
#[test]
fn serde_rejects_every_invalid_field() {
for (json, field) in [
(
r#"{"jitter":"full","firstMs":0,"maxMs":30000,"factor":2.0}"#,
"first_ms",
),
(
r#"{"jitter":"full","firstMs":1000,"maxMs":500,"factor":2.0}"#,
"max_ms",
),
(
r#"{"jitter":"full","firstMs":1000,"maxMs":30000,"factor":0.5}"#,
"factor",
),
] {
let error = serde_json::from_str::<BackoffPolicy>(json).unwrap_err();
assert!(error.to_string().contains(field), "got: {error}");
}
}
}