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