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