use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio_util::sync::CancellationToken;
use tracing::warn;
pub const BARRIER_WAIT_LIMIT: Duration = Duration::from_secs(15);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FrameOrder {
Barrier,
Follower,
Unordered,
}
#[derive(Default)]
struct Inner {
next_seq: u64,
tails: HashMap<String, (u64, CancellationToken)>,
}
#[derive(Default, Clone)]
pub struct SenderOrder {
inner: Arc<Mutex<Inner>>,
}
impl SenderOrder {
pub fn new() -> Self {
Self::default()
}
pub fn admit(&self, sender: Option<&str>, order: FrameOrder) -> Ticket {
let Some(sender) = sender else {
return Ticket::unordered();
};
let mut inner = self.inner.lock().unwrap_or_else(|p| p.into_inner());
match order {
FrameOrder::Unordered => Ticket::unordered(),
FrameOrder::Follower => Ticket {
wait: inner.tails.get(sender).map(|(_, t)| t.clone()),
_done: None,
},
FrameOrder::Barrier => {
let seq = inner.next_seq;
inner.next_seq = inner.next_seq.wrapping_add(1);
let token = CancellationToken::new();
let wait = inner
.tails
.insert(sender.to_string(), (seq, token.clone()))
.map(|(_, t)| t);
Ticket {
wait,
_done: Some(BarrierDone {
order: self.clone(),
sender: sender.to_string(),
seq,
token,
}),
}
}
}
}
pub fn tracked_senders(&self) -> usize {
self.inner
.lock()
.unwrap_or_else(|p| p.into_inner())
.tails
.len()
}
}
#[must_use = "a ticket orders nothing unless it is awaited and held while the frame is handled"]
pub struct Ticket {
wait: Option<CancellationToken>,
_done: Option<BarrierDone>,
}
impl Ticket {
fn unordered() -> Self {
Self {
wait: None,
_done: None,
}
}
pub async fn ready(&self) {
self.ready_within(BARRIER_WAIT_LIMIT).await;
}
async fn ready_within(&self, limit: Duration) {
let Some(wait) = &self.wait else { return };
if tokio::time::timeout(limit, wait.cancelled()).await.is_err() {
warn!(
waited = ?limit,
"an earlier relationship-control message from this sender is still being \
answered — handling this frame without waiting for it",
);
}
}
}
struct BarrierDone {
order: SenderOrder,
sender: String,
seq: u64,
token: CancellationToken,
}
impl Drop for BarrierDone {
fn drop(&mut self) {
self.token.cancel();
let mut inner = self.order.inner.lock().unwrap_or_else(|p| p.into_inner());
if inner
.tails
.get(&self.sender)
.is_some_and(|(seq, _)| *seq == self.seq)
{
inner.tails.remove(&self.sender);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::{mpsc, oneshot};
const PHONE: &str = "did:peer:2.phone";
const OTHER: &str = "did:peer:2.other";
fn handle(
ticket: Ticket,
label: &'static str,
gate: Option<oneshot::Receiver<()>>,
log: mpsc::UnboundedSender<&'static str>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
ticket.ready().await;
if let Some(gate) = gate {
let _ = gate.await;
}
let _ = log.send(label);
drop(ticket);
})
}
#[tokio::test]
async fn a_reply_waits_for_the_earlier_relationship_accept() {
let order = SenderOrder::new();
let (log_tx, mut log) = mpsc::unbounded_channel();
let (release_accept, accept_gate) = oneshot::channel();
let xrfi = order.admit(Some(PHONE), FrameOrder::Barrier);
let task = order.admit(Some(PHONE), FrameOrder::Follower);
let accept = handle(xrfi, "XRFA", Some(accept_gate), log_tx.clone());
let reply = handle(task, "reply", None, log_tx.clone());
for _ in 0..50 {
tokio::task::yield_now().await;
}
assert!(log.try_recv().is_err(), "the reply overtook the accept");
release_accept.send(()).unwrap();
accept.await.unwrap();
reply.await.unwrap();
assert_eq!(log.recv().await, Some("XRFA"));
assert_eq!(log.recv().await, Some("reply"));
assert_eq!(order.tracked_senders(), 0, "a finished barrier is evicted");
}
#[tokio::test]
async fn without_a_barrier_the_reply_overtakes() {
let order = SenderOrder::new();
let (log_tx, mut log) = mpsc::unbounded_channel();
let (release_accept, accept_gate) = oneshot::channel();
let xrfi = order.admit(Some(PHONE), FrameOrder::Unordered);
let task = order.admit(Some(PHONE), FrameOrder::Follower);
let accept = handle(xrfi, "XRFA", Some(accept_gate), log_tx.clone());
let reply = handle(task, "reply", None, log_tx.clone());
reply.await.unwrap();
assert_eq!(log.recv().await, Some("reply"));
release_accept.send(()).unwrap();
accept.await.unwrap();
}
#[tokio::test]
async fn other_senders_are_not_blocked() {
let order = SenderOrder::new();
let (log_tx, mut log) = mpsc::unbounded_channel();
let (release_accept, accept_gate) = oneshot::channel();
let xrfi = order.admit(Some(PHONE), FrameOrder::Barrier);
let other = order.admit(Some(OTHER), FrameOrder::Follower);
let accept = handle(xrfi, "XRFA", Some(accept_gate), log_tx.clone());
handle(other, "other", None, log_tx.clone()).await.unwrap();
assert_eq!(log.recv().await, Some("other"));
release_accept.send(()).unwrap();
accept.await.unwrap();
}
#[tokio::test]
async fn followers_run_concurrently() {
let order = SenderOrder::new();
let (log_tx, mut log) = mpsc::unbounded_channel();
let (release_first, first_gate) = oneshot::channel();
let first = order.admit(Some(PHONE), FrameOrder::Follower);
let second = order.admit(Some(PHONE), FrameOrder::Follower);
let first = handle(first, "first", Some(first_gate), log_tx.clone());
handle(second, "second", None, log_tx.clone())
.await
.unwrap();
assert_eq!(log.recv().await, Some("second"));
release_first.send(()).unwrap();
first.await.unwrap();
}
#[tokio::test]
async fn barriers_chain_in_arrival_order() {
let order = SenderOrder::new();
let (log_tx, mut log) = mpsc::unbounded_channel();
let (release_a, gate_a) = oneshot::channel();
let a = order.admit(Some(PHONE), FrameOrder::Barrier);
let b = order.admit(Some(PHONE), FrameOrder::Barrier);
let f = order.admit(Some(PHONE), FrameOrder::Follower);
let f = handle(f, "follower", None, log_tx.clone());
let b = handle(b, "b", None, log_tx.clone());
let a = handle(a, "a", Some(gate_a), log_tx.clone());
for _ in 0..50 {
tokio::task::yield_now().await;
}
assert!(log.try_recv().is_err());
release_a.send(()).unwrap();
for h in [a, b, f] {
h.await.unwrap();
}
assert_eq!(log.recv().await, Some("a"));
assert_eq!(log.recv().await, Some("b"));
assert_eq!(log.recv().await, Some("follower"));
assert_eq!(order.tracked_senders(), 0);
}
#[tokio::test]
async fn a_panicking_barrier_releases_its_followers() {
let order = SenderOrder::new();
let xrfi = order.admit(Some(PHONE), FrameOrder::Barrier);
let task = order.admit(Some(PHONE), FrameOrder::Follower);
let crashed = tokio::spawn(async move {
let _held = xrfi;
panic!("accept handler blew up");
});
assert!(crashed.await.is_err());
tokio::time::timeout(Duration::from_secs(1), task.ready())
.await
.expect("follower released");
assert_eq!(order.tracked_senders(), 0);
}
#[tokio::test]
async fn a_hung_barrier_does_not_strand_followers() {
let order = SenderOrder::new();
let _hung = order.admit(Some(PHONE), FrameOrder::Barrier);
let task = order.admit(Some(PHONE), FrameOrder::Follower);
tokio::time::timeout(
Duration::from_secs(5),
task.ready_within(Duration::from_millis(20)),
)
.await
.expect("the wait is bounded");
assert_eq!(order.tracked_senders(), 1, "still in flight, still tracked");
}
#[tokio::test]
async fn an_anonymous_frame_is_unordered() {
let order = SenderOrder::new();
let t = order.admit(None, FrameOrder::Barrier);
t.ready().await;
assert_eq!(order.tracked_senders(), 0);
}
}