1use 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 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 );
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}