matrix_sdk_common/
failures_cache.rs1use std::{borrow::Borrow, collections::HashMap, hash::Hash, sync::Arc, time::Duration};
19
20use ruma::time::Instant;
21
22use super::locks::RwLock;
23
24const MAX_DELAY: u64 = 15 * 60;
25const MULTIPLIER: u64 = 15;
26
27#[derive(Clone, Debug)]
32pub struct FailuresCache<T: Eq + Hash> {
33 inner: Arc<InnerCache<T>>,
34}
35
36#[derive(Debug)]
37struct InnerCache<T: Eq + Hash> {
38 max_delay: Duration,
39 backoff_multiplier: u64,
40 items: RwLock<HashMap<T, FailuresItem>>,
41}
42
43impl<T: Eq + Hash> Default for InnerCache<T> {
44 fn default() -> Self {
45 Self {
46 max_delay: Duration::from_secs(MAX_DELAY),
47 backoff_multiplier: MULTIPLIER,
48 items: Default::default(),
49 }
50 }
51}
52
53#[derive(Debug, Clone, Copy)]
54struct FailuresItem {
55 insertion_time: Instant,
56 duration: Duration,
57
58 failure_count: u8,
61}
62
63impl FailuresItem {
64 fn expired(&self) -> bool {
66 self.insertion_time.elapsed() >= self.duration
67 }
68
69 fn expire(&mut self) {
74 self.duration = Duration::from_secs(0);
75 }
76}
77
78impl<T> FailuresCache<T>
79where
80 T: Eq + Hash,
81{
82 pub fn new() -> Self {
83 Self { inner: Default::default() }
84 }
85
86 pub fn with_settings(max_delay: Duration, multiplier: u8) -> Self {
87 Self {
88 inner: InnerCache {
89 max_delay,
90 backoff_multiplier: multiplier.into(),
91 items: Default::default(),
92 }
93 .into(),
94 }
95 }
96
97 pub fn contains<Q>(&self, key: &Q) -> bool
99 where
100 T: Borrow<Q>,
101 Q: Hash + Eq + ?Sized,
102 {
103 self.inner.items.read().get(key).is_some_and(|item| !item.expired())
104 }
105
106 pub fn failure_count<Q>(&self, key: &Q) -> Option<u8>
117 where
118 T: Borrow<Q>,
119 Q: Hash + Eq + ?Sized,
120 {
121 self.inner.items.read().get(key).map(|i| i.failure_count)
122 }
123
124 fn calculate_delay(&self, failure_count: u8) -> Duration {
131 let exponential_backoff = 2u64.saturating_pow(failure_count.into());
132 let delay = exponential_backoff.saturating_mul(self.inner.backoff_multiplier);
133
134 Duration::from_secs(delay).clamp(Duration::from_secs(1), self.inner.max_delay)
135 }
136
137 pub fn insert(&self, item: T) {
139 self.extend([item]);
140 }
141
142 pub fn extend(&self, iterator: impl IntoIterator<Item = T>) {
148 let mut lock = self.inner.items.write();
149
150 let now = Instant::now();
151
152 for key in iterator {
153 let failure_count = if let Some(value) = lock.get(&key) {
154 value.failure_count.saturating_add(1)
155 } else {
156 0
157 };
158
159 let delay = self.calculate_delay(failure_count);
160
161 let item = FailuresItem { insertion_time: now, duration: delay, failure_count };
162
163 lock.insert(key, item);
164 }
165 }
166
167 pub fn remove<'a, I, Q>(&'a self, iterator: I)
169 where
170 I: Iterator<Item = &'a Q>,
171 T: Borrow<Q>,
172 Q: Hash + Eq + 'a + ?Sized,
173 {
174 let mut lock = self.inner.items.write();
175
176 for item in iterator {
177 lock.remove(item);
178 }
179 }
180
181 #[doc(hidden)]
186 pub fn expire(&self, item: &T) {
187 self.inner.items.write().get_mut(item).map(FailuresItem::expire);
188 }
189}
190
191impl<T: Eq + Hash> Default for FailuresCache<T> {
192 fn default() -> Self {
193 Self::new()
194 }
195}
196
197#[cfg(test)]
198mod tests {
199 use std::time::Duration;
200
201 use proptest::prelude::*;
202
203 use super::FailuresCache;
204
205 #[test]
206 fn failures_cache() {
207 let cache = FailuresCache::new();
208
209 assert!(!cache.contains(&1));
210 cache.extend([1u8].iter());
211 assert!(cache.contains(&1));
212
213 cache.inner.items.write().get_mut(&1).unwrap().duration = Duration::from_secs(0);
214 assert!(!cache.contains(&1));
215
216 cache.remove([1u8].iter());
217 assert!(cache.inner.items.read().get(&1).is_none())
218 }
219
220 #[test]
221 fn failures_cache_timeout() {
222 let cache: FailuresCache<u8> = FailuresCache::new();
223
224 assert_eq!(cache.calculate_delay(0).as_secs(), 15);
225 assert_eq!(cache.calculate_delay(1).as_secs(), 30);
226 assert_eq!(cache.calculate_delay(2).as_secs(), 60);
227 assert_eq!(cache.calculate_delay(3).as_secs(), 120);
228 assert_eq!(cache.calculate_delay(4).as_secs(), 240);
229 assert_eq!(cache.calculate_delay(5).as_secs(), 480);
230 assert_eq!(cache.calculate_delay(6).as_secs(), 900);
231 assert_eq!(cache.calculate_delay(7).as_secs(), 900);
232 }
233
234 proptest! {
235 #[test]
236 fn failures_cache_proptest_timeout(count in 0..10u8) {
237 let cache: FailuresCache<u8> = FailuresCache::new();
238 let delay = cache.calculate_delay(count).as_secs();
239
240 assert!(delay <= 900);
241 assert!(delay >= 15);
242 }
243 }
244}