Skip to main content

rust_zero_core/
singleflight.rs

1use std::{
2    collections::HashMap,
3    error::Error,
4    fmt,
5    future::Future,
6    hash::Hash,
7    sync::{Arc, Mutex},
8};
9
10use tokio::sync::oneshot;
11
12type Waiter<V, E> = oneshot::Sender<Result<V, SingleFlightError<E>>>;
13
14/// Coalesces concurrent operations for the same key so only one operation executes.
15pub struct SingleFlight<K, V, E> {
16    flights: Mutex<HashMap<K, Vec<Waiter<V, E>>>>,
17}
18
19impl<K, V, E> Default for SingleFlight<K, V, E>
20where
21    K: Eq + Hash,
22{
23    fn default() -> Self {
24        Self {
25            flights: Mutex::new(HashMap::new()),
26        }
27    }
28}
29
30impl<K, V, E> SingleFlight<K, V, E>
31where
32    K: Clone + Eq + Hash,
33    V: Clone,
34{
35    pub fn new() -> Self {
36        Self::default()
37    }
38
39    /// Runs `operation` unless another call for `key` is already in progress.
40    ///
41    /// Concurrent callers receive the leader's result. If the leader is cancelled before it
42    /// finishes, waiting callers receive [`SingleFlightError::LeaderCancelled`].
43    pub async fn execute<F, Fut>(&self, key: K, operation: F) -> Result<V, SingleFlightError<E>>
44    where
45        F: FnOnce() -> Fut,
46        Fut: Future<Output = Result<V, E>>,
47    {
48        let receiver = {
49            let mut flights = self.flights.lock().expect("single-flight mutex poisoned");
50
51            if let Some(waiters) = flights.get_mut(&key) {
52                let (sender, receiver) = oneshot::channel();
53                waiters.push(sender);
54                Some(receiver)
55            } else {
56                flights.insert(key.clone(), Vec::new());
57                None
58            }
59        };
60
61        if let Some(receiver) = receiver {
62            return receiver
63                .await
64                .expect("single-flight leader must notify all waiting callers");
65        }
66
67        let mut guard = FlightGuard {
68            flights: &self.flights,
69            key: key.clone(),
70            armed: true,
71        };
72        let result = operation()
73            .await
74            .map_err(|error| SingleFlightError::Operation(Arc::new(error)));
75
76        let waiters = self
77            .flights
78            .lock()
79            .expect("single-flight mutex poisoned")
80            .remove(&key)
81            .expect("single-flight leader must have an active flight");
82        for waiter in waiters {
83            let _ = waiter.send(result.clone());
84        }
85        guard.armed = false;
86
87        result
88    }
89}
90
91struct FlightGuard<'a, K: Eq + Hash, V, E> {
92    flights: &'a Mutex<HashMap<K, Vec<Waiter<V, E>>>>,
93    key: K,
94    armed: bool,
95}
96
97impl<K, V, E> Drop for FlightGuard<'_, K, V, E>
98where
99    K: Eq + Hash,
100{
101    fn drop(&mut self) {
102        if !self.armed {
103            return;
104        }
105
106        if let Ok(mut flights) = self.flights.lock() {
107            if let Some(waiters) = flights.remove(&self.key) {
108                for waiter in waiters {
109                    let _ = waiter.send(Err(SingleFlightError::LeaderCancelled));
110                }
111            }
112        }
113    }
114}
115
116/// Errors returned by [`SingleFlight::execute`].
117#[derive(Debug, PartialEq, Eq)]
118pub enum SingleFlightError<E> {
119    Operation(Arc<E>),
120    LeaderCancelled,
121}
122
123impl<E> Clone for SingleFlightError<E> {
124    fn clone(&self) -> Self {
125        match self {
126            Self::Operation(error) => Self::Operation(Arc::clone(error)),
127            Self::LeaderCancelled => Self::LeaderCancelled,
128        }
129    }
130}
131
132impl<E> fmt::Display for SingleFlightError<E>
133where
134    E: fmt::Display,
135{
136    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
137        match self {
138            Self::Operation(error) => write!(formatter, "single-flight operation failed: {error}"),
139            Self::LeaderCancelled => formatter.write_str("single-flight leader was cancelled"),
140        }
141    }
142}
143
144impl<E> Error for SingleFlightError<E>
145where
146    E: Error + 'static,
147{
148    fn source(&self) -> Option<&(dyn Error + 'static)> {
149        match self {
150            Self::Operation(error) => Some(error.as_ref()),
151            Self::LeaderCancelled => None,
152        }
153    }
154}
155
156#[cfg(test)]
157mod tests {
158    use super::{SingleFlight, SingleFlightError};
159    use std::sync::{
160        atomic::{AtomicUsize, Ordering},
161        Arc,
162    };
163    use tokio::sync::Notify;
164
165    #[tokio::test]
166    async fn coalesces_concurrent_calls_for_the_same_key() {
167        let flights = Arc::new(SingleFlight::<String, usize, String>::new());
168        let calls = Arc::new(AtomicUsize::new(0));
169        let leader_started = Arc::new(Notify::new());
170        let release_leader = Arc::new(Notify::new());
171
172        let first = {
173            let flights = Arc::clone(&flights);
174            let calls = Arc::clone(&calls);
175            let leader_started = Arc::clone(&leader_started);
176            let release_leader = Arc::clone(&release_leader);
177            tokio::spawn(async move {
178                flights
179                    .execute("profile:42".to_owned(), || async move {
180                        calls.fetch_add(1, Ordering::SeqCst);
181                        leader_started.notify_one();
182                        release_leader.notified().await;
183                        Ok(42)
184                    })
185                    .await
186            })
187        };
188        leader_started.notified().await;
189
190        let second = {
191            let flights = Arc::clone(&flights);
192            let calls = Arc::clone(&calls);
193            tokio::spawn(async move {
194                flights
195                    .execute("profile:42".to_owned(), || async move {
196                        calls.fetch_add(1, Ordering::SeqCst);
197                        Ok(42)
198                    })
199                    .await
200            })
201        };
202
203        wait_for_waiter(&flights).await;
204        release_leader.notify_one();
205
206        assert_eq!(first.await.unwrap().unwrap(), 42);
207        assert_eq!(second.await.unwrap().unwrap(), 42);
208        assert_eq!(calls.load(Ordering::SeqCst), 1);
209    }
210
211    #[tokio::test]
212    async fn propagates_the_leader_error_to_waiting_callers() {
213        let flights = Arc::new(SingleFlight::<String, usize, String>::new());
214        let leader_started = Arc::new(Notify::new());
215        let release_leader = Arc::new(Notify::new());
216
217        let leader = {
218            let flights = Arc::clone(&flights);
219            let leader_started = Arc::clone(&leader_started);
220            let release_leader = Arc::clone(&release_leader);
221            tokio::spawn(async move {
222                flights
223                    .execute("profile:42".to_owned(), || async move {
224                        leader_started.notify_one();
225                        release_leader.notified().await;
226                        Err("database unavailable".to_owned())
227                    })
228                    .await
229            })
230        };
231        leader_started.notified().await;
232
233        let waiter = {
234            let flights = Arc::clone(&flights);
235            tokio::spawn(async move {
236                flights
237                    .execute("profile:42".to_owned(), || async { Ok(42) })
238                    .await
239            })
240        };
241
242        wait_for_waiter(&flights).await;
243        release_leader.notify_one();
244
245        for result in [leader.await.unwrap(), waiter.await.unwrap()] {
246            assert_eq!(
247                result,
248                Err(SingleFlightError::Operation(Arc::new(
249                    "database unavailable".to_owned()
250                )))
251            );
252        }
253    }
254
255    #[tokio::test]
256    async fn allows_a_new_call_after_a_completed_flight() {
257        let flights = SingleFlight::<String, usize, String>::new();
258        let calls = AtomicUsize::new(0);
259
260        for expected in [1, 2] {
261            let result = flights
262                .execute("profile:42".to_owned(), || async {
263                    Ok(calls.fetch_add(1, Ordering::SeqCst) + 1)
264                })
265                .await;
266            assert_eq!(result, Ok(expected));
267        }
268    }
269
270    async fn wait_for_waiter(flights: &SingleFlight<String, usize, String>) {
271        loop {
272            if flights
273                .flights
274                .lock()
275                .expect("single-flight mutex poisoned")
276                .get("profile:42")
277                .is_some_and(|waiters| waiters.len() == 1)
278            {
279                return;
280            }
281            tokio::task::yield_now().await;
282        }
283    }
284}