Skip to main content

solti_model/domain/policy/
backoff.rs

1//! # Backoff policy
2//!
3//! [`BackoffPolicy`] configures retry delay growth and jitter.
4//! Direct deserialization validates fields.
5//! Struct literals must be checked with [`BackoffPolicy::validate`].
6
7use std::borrow::Cow;
8use std::hash::{Hash, Hasher};
9
10use serde::{Deserialize, Serialize};
11
12use crate::error::{ModelError, ModelResult};
13
14/// Exponential backoff configuration for task restart delays.
15///
16/// | Field      | Type           | Default   | Description                           |
17/// |------------|----------------|-----------|---------------------------------------|
18/// | `jitter`   | `JitterPolicy` | `Full`    | Randomness applied to each delay      |
19/// | `first_ms` | `u64`          | `1_000`   | Initial delay (ms)                    |
20/// | `max_ms`   | `u64`          | `30_000`  | Maximum delay cap (ms)                |
21/// | `factor`   | `f64`          | `2.0`     | Exponential growth multiplier         |
22///
23/// Before jitter, factor `2.0` grows `1s, 2s, 4s, 8s`.
24/// Growth stops at `max_ms`.
25///
26/// ## Example
27///
28/// ```
29/// use solti_model::{BackoffPolicy, JitterPolicy};
30///
31/// let backoff = BackoffPolicy {
32///     jitter: JitterPolicy::Equal,
33///     first_ms: 1_000,
34///     max_ms: 30_000,
35///     factor: 2.0,
36/// };
37///
38/// backoff.validate().unwrap();
39/// ```
40#[derive(Clone, Debug, Serialize, Deserialize)]
41#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
42#[cfg_attr(feature = "schema", schemars(!try_from, deny_unknown_fields))]
43#[serde(rename_all = "camelCase")]
44#[serde(try_from = "raw::BackoffPolicyRaw")]
45pub struct BackoffPolicy {
46    /// Jitter policy applied to each computed delay.
47    pub jitter: super::JitterPolicy,
48    /// Initial delay (ms) for exponential backoff.
49    #[cfg_attr(feature = "schema", schemars(range(min = 1)))]
50    pub first_ms: u64,
51    /// Maximum allowed delay (ms).
52    #[cfg_attr(feature = "schema", schemars(range(min = 1)))]
53    pub max_ms: u64,
54    /// Exponential growth multiplier.
55    #[cfg_attr(feature = "schema", schemars(range(min = 1.0)))]
56    pub factor: f64,
57}
58
59mod raw {
60    use super::*;
61
62    #[derive(Deserialize)]
63    #[serde(rename_all = "camelCase", deny_unknown_fields)]
64    pub(super) struct BackoffPolicyRaw {
65        pub jitter: super::super::JitterPolicy,
66        pub first_ms: u64,
67        pub max_ms: u64,
68        pub factor: f64,
69    }
70
71    impl TryFrom<BackoffPolicyRaw> for BackoffPolicy {
72        type Error = ModelError;
73
74        fn try_from(r: BackoffPolicyRaw) -> Result<Self, Self::Error> {
75            let p = BackoffPolicy {
76                jitter: r.jitter,
77                first_ms: r.first_ms,
78                max_ms: r.max_ms,
79                factor: r.factor,
80            };
81            p.validate()?;
82            Ok(p)
83        }
84    }
85}
86
87impl BackoffPolicy {
88    /// Validates backoff parameters.
89    ///
90    /// # Errors
91    ///
92    /// Returns [`ModelError::Invalid`] when `first_ms` is zero,
93    /// `max_ms` is below `first_ms`, or `factor` is not finite or below `1.0`.
94    ///
95    /// ## Example
96    ///
97    /// ```
98    /// use solti_model::BackoffPolicy;
99    ///
100    /// let mut backoff = BackoffPolicy::default();
101    /// backoff.first_ms = 0;
102    ///
103    /// assert!(backoff.validate().is_err());
104    /// ```
105    pub fn validate(&self) -> ModelResult<()> {
106        if self.first_ms == 0 {
107            return Err(ModelError::Invalid(Cow::Borrowed(
108                "backoff first_ms must be greater than zero",
109            )));
110        }
111        if self.max_ms < self.first_ms {
112            return Err(ModelError::Invalid(Cow::Borrowed(
113                "backoff max_ms must be >= first_ms",
114            )));
115        }
116        if !self.factor.is_finite() || self.factor < 1.0 {
117            return Err(ModelError::Invalid(Cow::Borrowed(
118                "backoff factor must be finite and >= 1.0",
119            )));
120        }
121        Ok(())
122    }
123}
124
125impl PartialEq for BackoffPolicy {
126    fn eq(&self, other: &Self) -> bool {
127        self.jitter == other.jitter
128            && self.factor.to_bits() == other.factor.to_bits()
129            && self.first_ms == other.first_ms
130            && self.max_ms == other.max_ms
131    }
132}
133
134impl Eq for BackoffPolicy {}
135
136impl Hash for BackoffPolicy {
137    fn hash<H: Hasher>(&self, state: &mut H) {
138        self.factor.to_bits().hash(state);
139        self.first_ms.hash(state);
140        self.jitter.hash(state);
141        self.max_ms.hash(state);
142    }
143}
144
145impl Default for BackoffPolicy {
146    /// Returns full jitter, a 1-second initial delay, a 30-second cap, and factor `2.0`.
147    fn default() -> Self {
148        Self {
149            jitter: super::JitterPolicy::Full,
150            first_ms: 1_000,
151            max_ms: 30_000,
152            factor: 2.0,
153        }
154    }
155}
156
157#[cfg(test)]
158mod tests {
159    use super::*;
160
161    #[test]
162    fn validation_accepts_defaults_and_rejects_invalid_fields() {
163        BackoffPolicy::default().validate().unwrap();
164
165        for invalid in [
166            BackoffPolicy {
167                first_ms: 0,
168                ..BackoffPolicy::default()
169            },
170            BackoffPolicy {
171                first_ms: 500,
172                max_ms: 100,
173                ..BackoffPolicy::default()
174            },
175            BackoffPolicy {
176                factor: 0.5,
177                ..BackoffPolicy::default()
178            },
179            BackoffPolicy {
180                factor: f64::NAN,
181                ..BackoffPolicy::default()
182            },
183        ] {
184            assert!(invalid.validate().is_err());
185        }
186    }
187
188    #[test]
189    fn serde_roundtrip_accepts_valid_policy() {
190        let policy = BackoffPolicy::default();
191        let json = serde_json::to_string(&policy).unwrap();
192        let back: BackoffPolicy = serde_json::from_str(&json).unwrap();
193        assert_eq!(back, policy);
194    }
195
196    #[test]
197    fn serde_rejects_every_invalid_field() {
198        for (json, field) in [
199            (
200                r#"{"jitter":"full","firstMs":0,"maxMs":30000,"factor":2.0}"#,
201                "first_ms",
202            ),
203            (
204                r#"{"jitter":"full","firstMs":1000,"maxMs":500,"factor":2.0}"#,
205                "max_ms",
206            ),
207            (
208                r#"{"jitter":"full","firstMs":1000,"maxMs":30000,"factor":0.5}"#,
209                "factor",
210            ),
211        ] {
212            let error = serde_json::from_str::<BackoffPolicy>(json).unwrap_err();
213            assert!(error.to_string().contains(field), "got: {error}");
214        }
215    }
216}