use std::collections::HashMap;
use std::hash::Hash;
use std::sync::atomic::{AtomicU8, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, MutexGuard};
use tokio::sync::watch;
use tokio_util::sync::{CancellationToken, WaitForCancellationFuture};
const RUNNING: u8 = 0;
const CANCELLED: u8 = 1;
const PUBLISHING: u8 = 2;
pub struct SingleFlight<K, V> {
inner: Arc<Table<K, V>>,
}
struct Table<K, V> {
inflight: Mutex<HashMap<K, Arc<Slot<V>>>>,
}
impl<K, V> Clone for SingleFlight<K, V> {
fn clone(&self) -> SingleFlight<K, V> {
SingleFlight {
inner: Arc::clone(&self.inner),
}
}
}
impl<K, V> Default for SingleFlight<K, V> {
fn default() -> SingleFlight<K, V> {
SingleFlight {
inner: Arc::new(Table {
inflight: Mutex::new(HashMap::new()),
}),
}
}
}
pub struct Slot<V> {
state: AtomicU8,
waiters: AtomicUsize,
cancel: CancellationToken,
outcome: watch::Sender<Option<V>>,
}
impl<V> Slot<V> {
pub(crate) fn new() -> Slot<V> {
Slot {
state: AtomicU8::new(RUNNING),
waiters: AtomicUsize::new(0),
cancel: CancellationToken::new(),
outcome: watch::Sender::new(None),
}
}
pub fn begin_publishing(&self) -> bool {
self.state
.compare_exchange(RUNNING, PUBLISHING, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
}
pub fn begin_cancel(&self) -> bool {
self.state
.compare_exchange(RUNNING, CANCELLED, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
}
pub fn cancelled(&self) -> WaitForCancellationFuture<'_> {
self.cancel.cancelled()
}
}
impl<V: Clone> Slot<V> {
pub async fn wait(&self) -> Option<V> {
let mut outcome = self.outcome.subscribe();
loop {
let done = outcome.borrow_and_update().clone();
if done.is_some() {
return done;
}
if outcome.changed().await.is_err() {
return None;
}
}
}
}
impl<K: Eq + Hash + Clone, V> SingleFlight<K, V> {
pub fn new() -> SingleFlight<K, V> {
SingleFlight::default()
}
fn table(&self) -> MutexGuard<'_, HashMap<K, Arc<Slot<V>>>> {
self.inner
.inflight
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub fn waiting_on(&self, key: &K) -> usize {
self.table()
.get(key)
.map(|slot| slot.waiters.load(Ordering::Acquire))
.unwrap_or(0)
}
pub fn publishing(&self, key: &K) -> bool {
self.table()
.get(key)
.is_some_and(|slot| slot.state.load(Ordering::Acquire) == PUBLISHING)
}
pub fn join(&self, key: K) -> (Arc<Slot<V>>, bool, Waiter<K, V>) {
let mut inflight = self.table();
let (slot, leader) = match inflight.get(&key) {
Some(slot) => (Arc::clone(slot), false),
None => {
let slot = Arc::new(Slot::new());
inflight.insert(key.clone(), Arc::clone(&slot));
(slot, true)
}
};
slot.waiters.fetch_add(1, Ordering::AcqRel);
(
Arc::clone(&slot),
leader,
Waiter {
flight: self.clone(),
key,
slot,
},
)
}
fn finish(&self, key: &K, slot: &Arc<Slot<V>>) {
let mut inflight = self.table();
if inflight
.get(key)
.is_some_and(|held| Arc::ptr_eq(held, slot))
{
inflight.remove(key);
}
}
}
pub struct Waiter<K: Eq + Hash + Clone, V> {
flight: SingleFlight<K, V>,
key: K,
slot: Arc<Slot<V>>,
}
impl<K: Eq + Hash + Clone, V> Drop for Waiter<K, V> {
fn drop(&mut self) {
let mut inflight = self.flight.table();
if self.slot.waiters.fetch_sub(1, Ordering::AcqRel) != 1 {
return;
}
if self.slot.begin_cancel() {
self.slot.cancel.cancel();
if inflight
.get(&self.key)
.is_some_and(|held| Arc::ptr_eq(held, &self.slot))
{
inflight.remove(&self.key);
}
}
}
}
pub struct Resolution<K: Eq + Hash + Clone, V: Clone> {
flight: SingleFlight<K, V>,
key: K,
slot: Arc<Slot<V>>,
abandoned: Option<V>,
subject: String,
}
impl<K: Eq + Hash + Clone, V: Clone> Resolution<K, V> {
pub fn new(
flight: SingleFlight<K, V>,
key: K,
slot: Arc<Slot<V>>,
abandoned: V,
subject: String,
) -> Resolution<K, V> {
Resolution {
flight,
key,
slot,
abandoned: Some(abandoned),
subject,
}
}
pub fn answer(&mut self, outcome: V) {
self.slot.outcome.send_replace(Some(outcome));
self.abandoned = None;
}
}
impl<K: Eq + Hash + Clone, V: Clone> Drop for Resolution<K, V> {
fn drop(&mut self) {
if let Some(abandoned) = self.abandoned.take() {
tracing::error!(
subject = %self.subject,
"coalesced upstream work ended without an outcome; failing its waiters \
and retiring the slot so the key stays servable"
);
self.slot.outcome.send_replace(Some(abandoned));
}
self.flight.finish(&self.key, &self.slot);
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
const BOUND: Duration = Duration::from_secs(5);
const WHY: &str = "the waiter never resolved: `Resolution` must publish with `send_replace`, not \
`send`, which is a no-op when no receiver has subscribed yet";
#[tokio::test]
async fn an_answer_reaches_a_waiter_that_arrived_before_it_and_one_that_arrived_after() {
let flight: SingleFlight<u8, &'static str> = SingleFlight::new();
let (slot, leader, _waiter) = flight.join(1);
assert!(leader);
let early = {
let slot = Arc::clone(&slot);
tokio::spawn(async move { tokio::time::timeout(BOUND, slot.wait()).await })
};
let mut resolution = Resolution::new(
flight.clone(),
1,
Arc::clone(&slot),
"abandoned",
"1".to_owned(),
);
resolution.answer("done");
assert_eq!(
early.await.expect("the waiter task finishes").expect(WHY),
Some("done")
);
assert_eq!(
tokio::time::timeout(BOUND, slot.wait()).await.expect(WHY),
Some("done"),
"and a waiter that subscribes after the answer reads the answer that is there"
);
}
#[tokio::test]
async fn an_abandoned_resolution_answers_its_waiters_and_retires_the_slot() {
let flight: SingleFlight<u8, &'static str> = SingleFlight::new();
let (slot, _leader, waiter) = flight.join(1);
drop(Resolution::new(
flight.clone(),
1,
Arc::clone(&slot),
"abandoned",
"1".to_owned(),
));
assert_eq!(
tokio::time::timeout(BOUND, slot.wait()).await.expect(WHY),
Some("abandoned")
);
drop(waiter);
let (_slot, leader, _waiter) = flight.join(1);
assert!(leader, "and the next request starts work of its own");
}
}