Skip to main content

synapse/
resilience.rs

1//! Retries + circuit breakers for synapse outbound HTTP clients.
2//!
3//! Mirrors talos-core `src/resilience.rs` (profiles, breakers, `is_retryable` rules).
4
5use std::sync::atomic::{AtomicU64, AtomicU8, Ordering};
6use std::time::{Duration, Instant};
7
8use backon::{ExponentialBuilder, Retryable};
9use parking_lot::Mutex;
10
11#[derive(Debug, Clone, Copy)]
12pub enum Profile {
13    Default,
14    Fast,
15    Aggressive,
16}
17
18#[derive(Debug, Clone)]
19pub struct ResiliencePolicy {
20    pub max_attempts: u32,
21    pub initial_delay: Duration,
22    pub max_delay: Duration,
23    pub multiplier: f32,
24    pub jitter: bool,
25    pub breaker_threshold: u32,
26    pub breaker_open_for: Duration,
27}
28
29impl Profile {
30    pub fn policy(self) -> ResiliencePolicy {
31        match self {
32            Profile::Default => ResiliencePolicy {
33                max_attempts: 3,
34                initial_delay: Duration::from_millis(200),
35                max_delay: Duration::from_secs(5),
36                multiplier: 2.0,
37                jitter: true,
38                breaker_threshold: 5,
39                breaker_open_for: Duration::from_secs(30),
40            },
41            Profile::Fast => ResiliencePolicy {
42                max_attempts: 1,
43                initial_delay: Duration::from_millis(0),
44                max_delay: Duration::from_millis(0),
45                multiplier: 1.0,
46                jitter: false,
47                breaker_threshold: 10,
48                breaker_open_for: Duration::from_secs(15),
49            },
50            Profile::Aggressive => ResiliencePolicy {
51                max_attempts: 5,
52                initial_delay: Duration::from_millis(100),
53                max_delay: Duration::from_secs(10),
54                multiplier: 2.0,
55                jitter: true,
56                breaker_threshold: 10,
57                breaker_open_for: Duration::from_secs(60),
58            },
59        }
60    }
61}
62
63#[derive(Debug, thiserror::Error)]
64pub enum ResilienceError<E = reqwest::Error> {
65    #[error("upstream call failed after retries: {0}")]
66    Exhausted(E),
67    #[error("circuit breaker {name} is open")]
68    CircuitOpen { name: String },
69}
70
71impl From<reqwest::Error> for ResilienceError {
72    fn from(e: reqwest::Error) -> Self {
73        ResilienceError::Exhausted(e)
74    }
75}
76
77pub async fn run<F, Fut, T>(
78    op: F,
79    profile: Profile,
80    breaker: &CircuitBreaker,
81    label: &'static str,
82) -> Result<T, ResilienceError>
83where
84    F: FnMut() -> Fut + Send,
85    Fut: std::future::Future<Output = Result<T, reqwest::Error>> + Send,
86{
87    run_with_classifier(op, profile, breaker, label, is_retryable_reqwest).await
88}
89
90pub async fn run_with_classifier<T, E, F, Fut, Cls>(
91    op: F,
92    profile: Profile,
93    breaker: &CircuitBreaker,
94    label: &'static str,
95    is_retryable: Cls,
96) -> Result<T, ResilienceError<E>>
97where
98    F: FnMut() -> Fut + Send,
99    Fut: std::future::Future<Output = Result<T, E>> + Send,
100    E: std::fmt::Debug + Send,
101    Cls: Fn(&E) -> bool + Send + Sync,
102{
103    let start = Instant::now();
104
105    if breaker.guard().is_err() {
106        record_call_metrics(label, "circuit_open", start.elapsed());
107        return Err(ResilienceError::CircuitOpen {
108            name: breaker.name.to_string(),
109        });
110    }
111
112    let policy = profile.policy();
113    let mut builder = ExponentialBuilder::default()
114        .with_max_times(policy.max_attempts.saturating_sub(1) as usize)
115        .with_min_delay(policy.initial_delay)
116        .with_max_delay(policy.max_delay)
117        .with_factor(policy.multiplier);
118    if policy.jitter {
119        builder = builder.with_jitter();
120    }
121
122    let mut attempt: u32 = 0;
123    let result = op
124        .retry(builder)
125        .when(&is_retryable)
126        .notify(|err, dur| {
127            attempt += 1;
128            metrics::counter!(
129                "synapse_resilience_retry_attempts_total",
130                "label" => label,
131            )
132            .increment(1);
133            tracing::warn!(
134                label,
135                attempt,
136                next_delay_ms = dur.as_millis() as u64,
137                error = ?err,
138                "retrying outbound call",
139            );
140        })
141        .await;
142
143    breaker.record(&result);
144
145    let outcome = if result.is_ok() {
146        "success"
147    } else {
148        "exhausted"
149    };
150    record_call_metrics(label, outcome, start.elapsed());
151
152    result.map_err(ResilienceError::Exhausted)
153}
154
155fn record_call_metrics(label: &'static str, outcome: &'static str, elapsed: Duration) {
156    metrics::counter!(
157        "synapse_resilience_calls_total",
158        "label" => label,
159        "outcome" => outcome,
160    )
161    .increment(1);
162    metrics::histogram!(
163        "synapse_resilience_call_duration_seconds",
164        "label" => label,
165        "outcome" => outcome,
166    )
167    .record(elapsed.as_secs_f64());
168}
169
170pub fn is_retryable_reqwest(e: &reqwest::Error) -> bool {
171    if e.is_timeout() || e.is_connect() {
172        return true;
173    }
174    if let Some(status) = e.status() {
175        return status.is_server_error()
176            || status == reqwest::StatusCode::REQUEST_TIMEOUT
177            || status == reqwest::StatusCode::TOO_MANY_REQUESTS;
178    }
179    false
180}
181
182pub struct CircuitBreaker {
183    pub(crate) name: &'static str,
184    pub(crate) threshold: u32,
185    pub(crate) open_for: Duration,
186    pub(crate) consecutive_failures: AtomicU64,
187    pub(crate) state: AtomicU8,
188    pub(crate) opened_at: Mutex<Option<Instant>>,
189}
190
191pub(crate) const STATE_CLOSED: u8 = 0;
192pub(crate) const STATE_OPEN: u8 = 1;
193pub(crate) const STATE_HALF_OPEN: u8 = 2;
194
195impl CircuitBreaker {
196    pub fn new(name: &'static str, profile: Profile) -> Self {
197        let p = profile.policy();
198        Self {
199            name,
200            threshold: p.breaker_threshold,
201            open_for: p.breaker_open_for,
202            consecutive_failures: AtomicU64::new(0),
203            state: AtomicU8::new(STATE_CLOSED),
204            opened_at: Mutex::new(None),
205        }
206    }
207
208    pub fn guard(&self) -> Result<(), ResilienceError> {
209        match self.state.load(Ordering::Acquire) {
210            STATE_CLOSED | STATE_HALF_OPEN => Ok(()),
211            STATE_OPEN => {
212                let opened_at = self.opened_at.lock();
213                let opened = opened_at.unwrap_or_else(Instant::now);
214                drop(opened_at);
215                if Instant::now().duration_since(opened) >= self.open_for {
216                    let prev = self.state.swap(STATE_HALF_OPEN, Ordering::AcqRel);
217                    tracing::info!(name = self.name, "circuit breaker half-open");
218                    record_breaker_transition(self.name, "half_open", prev, STATE_HALF_OPEN);
219                    Ok(())
220                } else {
221                    Err(ResilienceError::CircuitOpen {
222                        name: self.name.to_string(),
223                    })
224                }
225            }
226            _ => Ok(()),
227        }
228    }
229
230    pub fn record<T, E>(&self, result: &Result<T, E>) {
231        match result {
232            Ok(_) => {
233                self.consecutive_failures.store(0, Ordering::Release);
234                let prev = self.state.swap(STATE_CLOSED, Ordering::AcqRel);
235                if prev == STATE_HALF_OPEN {
236                    tracing::info!(name = self.name, "circuit breaker closed");
237                    record_breaker_transition(self.name, "closed", prev, STATE_CLOSED);
238                }
239            }
240            Err(_) => {
241                let n = self.consecutive_failures.fetch_add(1, Ordering::AcqRel) + 1;
242                if n >= self.threshold as u64 {
243                    let prev = self.state.swap(STATE_OPEN, Ordering::AcqRel);
244                    *self.opened_at.lock() = Some(Instant::now());
245                    if prev != STATE_OPEN {
246                        tracing::warn!(
247                            name = self.name,
248                            consecutive_failures = n,
249                            "circuit breaker opened",
250                        );
251                        record_breaker_transition(self.name, "open", prev, STATE_OPEN);
252                    }
253                }
254            }
255        }
256    }
257}
258
259fn record_breaker_transition(name: &'static str, transition: &'static str, prev: u8, new: u8) {
260    metrics::counter!(
261        "synapse_resilience_breaker_transitions_total",
262        "name" => name,
263        "transition" => transition,
264    )
265    .increment(1);
266    metrics::gauge!("synapse_resilience_breaker_state", "name" => name).set(new as f64);
267    let _ = prev;
268}
269
270#[cfg(test)]
271mod tests {
272    use super::*;
273    use std::sync::atomic::{AtomicU32, Ordering};
274
275    async fn fetch(url: &str) -> reqwest::Result<reqwest::Response> {
276        reqwest::Client::new()
277            .get(url)
278            .timeout(Duration::from_millis(50))
279            .send()
280            .await
281    }
282
283    #[tokio::test]
284    async fn is_retryable_timeout_is_true() {
285        let err = fetch("http://10.255.255.1:1/").await.unwrap_err();
286        assert!(
287            is_retryable_reqwest(&err),
288            "expected timeout/connect to be retryable, got {err:?}"
289        );
290    }
291
292    #[tokio::test]
293    async fn is_retryable_5xx_is_true() {
294        use wiremock::{matchers::method, Mock, MockServer, ResponseTemplate};
295        let server = MockServer::start().await;
296        Mock::given(method("GET"))
297            .respond_with(ResponseTemplate::new(503))
298            .mount(&server)
299            .await;
300        let resp = reqwest::Client::new()
301            .get(server.uri())
302            .send()
303            .await
304            .unwrap();
305        let err = resp.error_for_status().unwrap_err();
306        assert!(is_retryable_reqwest(&err));
307    }
308
309    #[tokio::test]
310    async fn run_retries_until_success() {
311        use wiremock::{matchers::method, Mock, MockServer, ResponseTemplate};
312
313        let server = MockServer::start().await;
314        Mock::given(method("GET"))
315            .respond_with(ResponseTemplate::new(503))
316            .up_to_n_times(2)
317            .mount(&server)
318            .await;
319        Mock::given(method("GET"))
320            .respond_with(ResponseTemplate::new(200))
321            .mount(&server)
322            .await;
323
324        let breaker = CircuitBreaker::new("test", Profile::Default);
325        let attempts = AtomicU32::new(0);
326        let result = run(
327            || async {
328                attempts.fetch_add(1, Ordering::SeqCst);
329                reqwest::Client::new()
330                    .get(server.uri())
331                    .send()
332                    .await?
333                    .error_for_status()
334            },
335            Profile::Default,
336            &breaker,
337            "test",
338        )
339        .await;
340        assert!(result.is_ok(), "got {result:?}");
341        assert_eq!(attempts.load(Ordering::SeqCst), 3);
342    }
343
344    #[test]
345    fn breaker_opens_after_threshold() {
346        let b = CircuitBreaker::new("t", Profile::Default);
347        for _ in 0..5 {
348            let err: Result<(), reqwest::Error> = Err(make_fake_reqwest_error());
349            b.record(&err);
350        }
351        assert!(matches!(
352            b.guard(),
353            Err(ResilienceError::CircuitOpen { .. })
354        ));
355    }
356
357    fn make_fake_reqwest_error() -> reqwest::Error {
358        reqwest::Client::new()
359            .get("http://invalid url")
360            .build()
361            .unwrap_err()
362    }
363
364    #[tokio::test]
365    async fn run_with_classifier_retries_on_classified_retryable() {
366        #[derive(Debug)]
367        struct MyErr(bool /* retryable */);
368
369        let breaker = CircuitBreaker::new("test-classifier", Profile::Fast);
370        let attempts = std::sync::atomic::AtomicU32::new(0);
371
372        let result: Result<u32, ResilienceError<MyErr>> = run_with_classifier(
373            || {
374                let n = attempts.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
375                async move {
376                    if n < 2 {
377                        Err(MyErr(true))
378                    } else {
379                        Ok(42u32)
380                    }
381                }
382            },
383            Profile::Default,
384            &breaker,
385            "test-classifier",
386            |e: &MyErr| e.0,
387        )
388        .await;
389
390        assert_eq!(result.unwrap(), 42);
391        assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 3);
392    }
393
394    #[tokio::test]
395    async fn run_with_classifier_fails_fast_on_non_retryable() {
396        #[derive(Debug)]
397        struct MyErr(bool);
398
399        let breaker = CircuitBreaker::new("test-fail-fast", Profile::Fast);
400        let attempts = std::sync::atomic::AtomicU32::new(0);
401
402        let result: Result<u32, ResilienceError<MyErr>> = run_with_classifier(
403            || {
404                attempts.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
405                async move { Err(MyErr(false)) }
406            },
407            Profile::Default,
408            &breaker,
409            "test-fail-fast",
410            |e: &MyErr| e.0,
411        )
412        .await;
413
414        assert!(matches!(
415            result,
416            Err(ResilienceError::Exhausted(MyErr(false)))
417        ));
418        assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 1);
419    }
420}