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
14pub 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 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#[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}