use std::fmt::Debug;
use std::hash::Hash;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use dashmap::DashMap;
use dashmap::mapref::entry::Entry;
use futures::future::BoxFuture;
use tokio_util::task::TaskTracker;
use tracing::trace;
pub(crate) trait LaneObserver: Send + Sync {
fn lane_created(&self);
fn lane_closed(&self);
}
struct LaneState {
pending: AtomicUsize,
}
struct Lane<T> {
tx: flume::Sender<T>,
state: Arc<LaneState>,
}
pub(crate) struct LaneRouterConfig {
pub idle_ttl: Option<Duration>,
pub runtime: tokio::runtime::Handle,
pub tracker: TaskTracker,
pub observer: Option<Arc<dyn LaneObserver>>,
}
type Consumer<T> = Arc<dyn Fn(T) -> BoxFuture<'static, ()> + Send + Sync>;
pub(crate) struct LaneRouter<K, T>
where
K: Eq + Hash + Clone + Debug + Send + Sync + 'static,
T: Send + 'static,
{
lanes: Arc<DashMap<K, Lane<T>>>,
consumer: Consumer<T>,
config: LaneRouterConfig,
}
impl<K, T> LaneRouter<K, T>
where
K: Eq + Hash + Clone + Debug + Send + Sync + 'static,
T: Send + 'static,
{
pub(crate) fn new(consumer: Consumer<T>, config: LaneRouterConfig) -> Self {
Self {
lanes: Arc::new(DashMap::new()),
consumer,
config,
}
}
pub(crate) fn route(&self, key: K, item: T, capacity: Option<usize>) -> Result<usize, T> {
let lane = match self.lanes.entry(key.clone()) {
Entry::Occupied(occupied) => occupied.into_ref(),
Entry::Vacant(vacant) => {
let (tx, rx) = flume::unbounded();
let state = Arc::new(LaneState {
pending: AtomicUsize::new(0),
});
self.spawn_lane(key.clone(), rx, Arc::clone(&state));
vacant.insert(Lane { tx, state })
}
};
let depth = lane.state.pending.load(Ordering::Acquire);
if capacity.is_some_and(|capacity| depth >= capacity) {
return Err(item);
}
lane.state.pending.fetch_add(1, Ordering::AcqRel);
let _ = lane.tx.send(item);
Ok(depth)
}
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn lane_count(&self) -> usize {
self.lanes.len()
}
fn spawn_lane(&self, key: K, rx: flume::Receiver<T>, state: Arc<LaneState>) {
if let Some(observer) = self.config.observer.as_ref() {
observer.lane_created();
}
let task = LaneTask {
key,
rx,
state,
lanes: Arc::downgrade(&self.lanes),
consumer: Arc::clone(&self.consumer),
idle_ttl: self.config.idle_ttl,
observer: self.config.observer.clone(),
};
let tracked = self.config.tracker.track_future(task.run());
self.config.runtime.spawn(tracked);
}
}
struct LaneTask<K, T>
where
K: Eq + Hash + Clone + Debug + Send + Sync + 'static,
T: Send + 'static,
{
key: K,
rx: flume::Receiver<T>,
state: Arc<LaneState>,
lanes: std::sync::Weak<DashMap<K, Lane<T>>>,
consumer: Consumer<T>,
idle_ttl: Option<Duration>,
observer: Option<Arc<dyn LaneObserver>>,
}
impl<K, T> LaneTask<K, T>
where
K: Eq + Hash + Clone + Debug + Send + Sync + 'static,
T: Send + 'static,
{
async fn run(self) {
let _guard = LaneExitGuard {
key: self.key.clone(),
state: Arc::clone(&self.state),
lanes: self.lanes.clone(),
observer: self.observer.clone(),
};
loop {
let received = match self.idle_ttl {
Some(ttl) => match tokio::time::timeout(ttl, self.rx.recv_async()).await {
Ok(received) => received,
Err(_elapsed) => {
if self.try_reap() {
trace!(
target: "crate::messenger::lanes",
key = ?self.key,
"Reaping idle ordering lane"
);
break;
}
continue;
}
},
None => self.rx.recv_async().await,
};
let Ok(item) = received else {
break;
};
(self.consumer)(item).await;
self.state.pending.fetch_sub(1, Ordering::AcqRel);
}
}
fn try_reap(&self) -> bool {
let Some(lanes) = self.lanes.upgrade() else {
return true;
};
lanes
.remove_if(&self.key, |_, lane| {
Arc::ptr_eq(&lane.state, &self.state)
&& lane.state.pending.load(Ordering::Acquire) == 0
})
.is_some()
}
}
struct LaneExitGuard<K, T>
where
K: Eq + Hash + Clone + Debug + Send + Sync + 'static,
T: Send + 'static,
{
key: K,
state: Arc<LaneState>,
lanes: std::sync::Weak<DashMap<K, Lane<T>>>,
observer: Option<Arc<dyn LaneObserver>>,
}
impl<K, T> Drop for LaneExitGuard<K, T>
where
K: Eq + Hash + Clone + Debug + Send + Sync + 'static,
T: Send + 'static,
{
fn drop(&mut self) {
if let Some(lanes) = self.lanes.upgrade() {
lanes.remove_if(&self.key, |_, lane| Arc::ptr_eq(&lane.state, &self.state));
}
if let Some(observer) = self.observer.as_ref() {
observer.lane_closed();
}
}
}
#[cfg(test)]
#[derive(Default)]
struct CountingObserver {
created: AtomicUsize,
closed: AtomicUsize,
}
#[cfg(test)]
impl LaneObserver for CountingObserver {
fn lane_created(&self) {
self.created.fetch_add(1, Ordering::AcqRel);
}
fn lane_closed(&self) {
self.closed.fetch_add(1, Ordering::AcqRel);
}
}
#[cfg(test)]
fn group_by_key<K: Eq + Hash + Clone, V: Clone>(
log: &[(K, V)],
) -> std::collections::HashMap<K, Vec<V>> {
let mut grouped: std::collections::HashMap<K, Vec<V>> = std::collections::HashMap::new();
for (key, value) in log {
grouped.entry(key.clone()).or_default().push(value.clone());
}
grouped
}
#[cfg(test)]
mod tests {
use super::*;
use parking_lot::Mutex;
use std::sync::Arc;
use tokio::time::{Duration, timeout};
fn config(idle_ttl: Option<Duration>, observer: Arc<CountingObserver>) -> LaneRouterConfig {
LaneRouterConfig {
idle_ttl,
runtime: tokio::runtime::Handle::current(),
tracker: TaskTracker::new(),
observer: Some(observer),
}
}
async fn wait_for(label: &str, mut predicate: impl FnMut() -> bool) {
let deadline = Duration::from_secs(5);
timeout(deadline, async {
while !predicate() {
tokio::time::sleep(Duration::from_millis(1)).await;
}
})
.await
.unwrap_or_else(|_| panic!("timed out waiting for: {label}"));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn lane_preserves_per_key_order() {
let log = Arc::new(Mutex::new(Vec::new()));
let observer = Arc::new(CountingObserver::default());
let sink = Arc::clone(&log);
let router: LaneRouter<u64, (u64, u32)> = LaneRouter::new(
Arc::new(move |(key, seq)| {
let sink = Arc::clone(&sink);
Box::pin(async move {
tokio::task::yield_now().await;
sink.lock().push((key, seq));
})
}),
config(None, Arc::clone(&observer)),
);
const KEYS: u64 = 4;
const PER_KEY: u32 = 250;
for seq in 0..PER_KEY {
for key in 0..KEYS {
let _ = router.route(key, (key, seq), None);
}
}
wait_for("all items handled", || {
log.lock().len() == (KEYS as usize) * (PER_KEY as usize)
})
.await;
let grouped = group_by_key(&log.lock().clone());
for key in 0..KEYS {
let observed: Vec<u32> = grouped[&key].clone();
let expected: Vec<u32> = (0..PER_KEY).collect();
assert_eq!(observed, expected, "key {key} was reordered");
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn distinct_keys_run_concurrently() {
let barrier = Arc::new(tokio::sync::Barrier::new(2));
let observer = Arc::new(CountingObserver::default());
let done = Arc::new(AtomicUsize::new(0));
let consumer_barrier = Arc::clone(&barrier);
let consumer_done = Arc::clone(&done);
let router: LaneRouter<u64, ()> = LaneRouter::new(
Arc::new(move |()| {
let barrier = Arc::clone(&consumer_barrier);
let done = Arc::clone(&consumer_done);
Box::pin(async move {
barrier.wait().await;
done.fetch_add(1, Ordering::AcqRel);
})
}),
config(None, Arc::clone(&observer)),
);
let _ = router.route(1, (), None);
let _ = router.route(2, (), None);
wait_for("both lanes cleared the barrier", || {
done.load(Ordering::Acquire) == 2
})
.await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn idle_lane_is_reaped() {
let observer = Arc::new(CountingObserver::default());
let handled = Arc::new(AtomicUsize::new(0));
let consumer_handled = Arc::clone(&handled);
let router: LaneRouter<u64, ()> = LaneRouter::new(
Arc::new(move |()| {
let handled = Arc::clone(&consumer_handled);
Box::pin(async move {
handled.fetch_add(1, Ordering::AcqRel);
})
}),
config(Some(Duration::from_millis(50)), Arc::clone(&observer)),
);
let _ = router.route(7, (), None);
wait_for("item handled", || handled.load(Ordering::Acquire) == 1).await;
wait_for("lane reaped", || router.lane_count() == 0).await;
assert_eq!(observer.created.load(Ordering::Acquire), 1);
assert_eq!(observer.closed.load(Ordering::Acquire), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn no_reap_while_pending() {
let observer = Arc::new(CountingObserver::default());
let log = Arc::new(Mutex::new(Vec::new()));
let sink = Arc::clone(&log);
let router: LaneRouter<u64, u32> = LaneRouter::new(
Arc::new(move |seq| {
let sink = Arc::clone(&sink);
Box::pin(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
sink.lock().push(seq);
})
}),
config(Some(Duration::from_millis(1)), Arc::clone(&observer)),
);
for seq in 0..5 {
let _ = router.route(1, seq, None);
}
wait_for("all items handled", || log.lock().len() == 5).await;
assert_eq!(*log.lock(), vec![0, 1, 2, 3, 4]);
assert_eq!(
observer.created.load(Ordering::Acquire),
1,
"the lane must not be reaped and recreated while work is queued"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn no_reap_while_handler_in_flight() {
let observer = Arc::new(CountingObserver::default());
let finished = Arc::new(AtomicUsize::new(0));
let consumer_finished = Arc::clone(&finished);
let router: LaneRouter<u64, ()> = LaneRouter::new(
Arc::new(move |()| {
let finished = Arc::clone(&consumer_finished);
Box::pin(async move {
tokio::time::sleep(Duration::from_millis(150)).await;
finished.fetch_add(1, Ordering::AcqRel);
})
}),
config(Some(Duration::from_millis(5)), Arc::clone(&observer)),
);
let _ = router.route(1, (), None);
tokio::time::sleep(Duration::from_millis(60)).await;
assert_eq!(
router.lane_count(),
1,
"a lane with a handler in flight must not be reaped"
);
assert_eq!(finished.load(Ordering::Acquire), 0);
wait_for("handler finished", || finished.load(Ordering::Acquire) == 1).await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn reap_then_reuse_preserves_order() {
let observer = Arc::new(CountingObserver::default());
let log = Arc::new(Mutex::new(Vec::new()));
let sink = Arc::clone(&log);
let router: LaneRouter<u64, u32> = LaneRouter::new(
Arc::new(move |seq| {
let sink = Arc::clone(&sink);
Box::pin(async move {
sink.lock().push(seq);
})
}),
config(Some(Duration::from_millis(30)), Arc::clone(&observer)),
);
for seq in 0..5 {
let _ = router.route(1, seq, None);
}
wait_for("first burst handled", || log.lock().len() == 5).await;
wait_for("lane reaped", || router.lane_count() == 0).await;
for seq in 5..10 {
let _ = router.route(1, seq, None);
}
wait_for("second burst handled", || log.lock().len() == 10).await;
assert_eq!(*log.lock(), (0..10).collect::<Vec<_>>());
assert_eq!(
observer.created.load(Ordering::Acquire),
2,
"the reaped lane should have been rebuilt exactly once"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_producers_single_lane() {
let observer = Arc::new(CountingObserver::default());
let log = Arc::new(Mutex::new(Vec::new()));
let generation = Arc::new(AtomicUsize::new(0));
let sink = Arc::clone(&log);
let consumer_generation = Arc::clone(&generation);
let router: Arc<LaneRouter<u64, u32>> = Arc::new(LaneRouter::new(
Arc::new(move |item| {
let sink = Arc::clone(&sink);
let lane_gen = consumer_generation.load(Ordering::Acquire);
Box::pin(async move {
sink.lock().push((lane_gen, item));
})
}),
config(Some(Duration::from_millis(1)), Arc::clone(&observer)),
));
let generation_tracker = Arc::clone(&generation);
let observer_for_tracker = Arc::clone(&observer);
tokio::spawn(async move {
loop {
generation_tracker.store(
observer_for_tracker.created.load(Ordering::Acquire),
Ordering::Release,
);
tokio::time::sleep(Duration::from_millis(1)).await;
}
});
const PRODUCERS: u32 = 8;
const PER_PRODUCER: u32 = 200;
let mut handles = Vec::new();
for producer in 0..PRODUCERS {
let router = Arc::clone(&router);
handles.push(tokio::spawn(async move {
for i in 0..PER_PRODUCER {
let _ = router.route(1, producer * PER_PRODUCER + i, None);
if i % 32 == 0 {
tokio::task::yield_now().await;
}
}
}));
}
for handle in handles {
handle.await.expect("producer task");
}
wait_for("all items handled", || {
log.lock().len() == (PRODUCERS * PER_PRODUCER) as usize
})
.await;
let observed = log.lock().clone();
let generations: Vec<usize> = observed.iter().map(|(lane_gen, _)| *lane_gen).collect();
assert!(
generations.windows(2).all(|w| w[0] <= w[1]),
"two lanes for one key were live at the same time"
);
let mut items: Vec<u32> = observed.iter().map(|(_, item)| *item).collect();
items.sort_unstable();
assert_eq!(
items,
(0..PRODUCERS * PER_PRODUCER).collect::<Vec<_>>(),
"every item must be handled exactly once"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn dropping_router_stops_lanes() {
let observer = Arc::new(CountingObserver::default());
let handled = Arc::new(AtomicUsize::new(0));
let consumer_handled = Arc::clone(&handled);
let router: LaneRouter<u64, ()> = LaneRouter::new(
Arc::new(move |()| {
let handled = Arc::clone(&consumer_handled);
Box::pin(async move {
handled.fetch_add(1, Ordering::AcqRel);
})
}),
config(None, Arc::clone(&observer)),
);
for key in 0..3 {
let _ = router.route(key, (), None);
}
wait_for("items handled", || handled.load(Ordering::Acquire) == 3).await;
assert_eq!(observer.created.load(Ordering::Acquire), 3);
drop(router);
wait_for("all lanes exited", || {
observer.closed.load(Ordering::Acquire) == 3
})
.await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn capacity_is_enforced_per_key() {
let observer = Arc::new(CountingObserver::default());
let release = Arc::new(std::sync::atomic::AtomicBool::new(false));
let handled = Arc::new(Mutex::new(Vec::new()));
let consumer_release = Arc::clone(&release);
let sink = Arc::clone(&handled);
let router: LaneRouter<u64, u32> = LaneRouter::new(
Arc::new(move |item| {
let release = Arc::clone(&consumer_release);
let sink = Arc::clone(&sink);
Box::pin(async move {
while !release.load(Ordering::Acquire) {
tokio::time::sleep(Duration::from_millis(1)).await;
}
sink.lock().push(item);
})
}),
config(None, Arc::clone(&observer)),
);
assert_eq!(router.route(1, 10, Some(3)), Ok(0));
assert_eq!(router.route(1, 11, Some(3)), Ok(1));
assert_eq!(router.route(1, 12, Some(3)), Ok(2));
assert_eq!(
router.route(1, 13, Some(3)),
Err(13),
"a full lane must hand the item back"
);
assert_eq!(router.route(2, 20, Some(3)), Ok(0));
assert_eq!(router.route(1, 14, None), Ok(3));
release.store(true, Ordering::Release);
wait_for("all admitted items handled", || handled.lock().len() == 5).await;
let grouped = group_by_key(
&handled
.lock()
.iter()
.map(|item| (item / 10, *item))
.collect::<Vec<_>>(),
);
assert_eq!(grouped[&1], vec![10, 11, 12, 14], "key 1 lost or reordered");
assert_eq!(grouped[&2], vec![20]);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn capacity_frees_as_the_lane_drains() {
let observer = Arc::new(CountingObserver::default());
let handled = Arc::new(Mutex::new(Vec::new()));
let sink = Arc::clone(&handled);
let router: LaneRouter<u64, u32> = LaneRouter::new(
Arc::new(move |item| {
let sink = Arc::clone(&sink);
Box::pin(async move {
sink.lock().push(item);
})
}),
config(None, Arc::clone(&observer)),
);
assert!(router.route(1, 0, Some(1)).is_ok());
wait_for("first item handled", || handled.lock().len() == 1).await;
assert!(
router.route(1, 1, Some(1)).is_ok(),
"capacity must be reusable once the lane drains"
);
wait_for("second item handled", || handled.lock().len() == 2).await;
assert_eq!(*handled.lock(), vec![0, 1]);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn lane_survives_consumer_panic() {
let observer = Arc::new(CountingObserver::default());
let log = Arc::new(Mutex::new(Vec::new()));
let sink = Arc::clone(&log);
let router: LaneRouter<u64, u32> = LaneRouter::new(
Arc::new(move |seq| {
let sink = Arc::clone(&sink);
Box::pin(async move {
let result =
futures::FutureExt::catch_unwind(std::panic::AssertUnwindSafe(async {
assert_ne!(seq, 3, "deliberate panic");
seq
}))
.await;
if let Ok(seq) = result {
sink.lock().push(seq);
}
})
}),
config(None, Arc::clone(&observer)),
);
for seq in 0..8 {
let _ = router.route(1, seq, None);
}
wait_for("surviving items handled", || log.lock().len() == 7).await;
assert_eq!(*log.lock(), vec![0, 1, 2, 4, 5, 6, 7]);
assert_eq!(
observer.created.load(Ordering::Acquire),
1,
"the lane must survive a panicking item"
);
}
}