use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicI32, AtomicU64, AtomicUsize, Ordering};
use std::time::Duration;
use crate::ndarray::NDArray;
pub struct QueuedArrayCounter {
count: AtomicUsize,
mutex: parking_lot::Mutex<()>,
condvar: parking_lot::Condvar,
}
impl QueuedArrayCounter {
pub fn new() -> Self {
Self {
count: AtomicUsize::new(0),
mutex: parking_lot::Mutex::new(()),
condvar: parking_lot::Condvar::new(),
}
}
pub fn increment(&self) {
self.count.fetch_add(1, Ordering::AcqRel);
}
pub fn decrement(&self) {
let prev = self.count.fetch_sub(1, Ordering::AcqRel);
if prev == 1 {
let _guard = self.mutex.lock();
self.condvar.notify_all();
}
}
pub fn get(&self) -> usize {
self.count.load(Ordering::Acquire)
}
pub fn wait_until_zero(&self, timeout: Duration) -> bool {
let mut guard = self.mutex.lock();
if self.count.load(Ordering::Acquire) == 0 {
return true;
}
!self
.condvar
.wait_while_for(
&mut guard,
|_| self.count.load(Ordering::Acquire) != 0,
timeout,
)
.timed_out()
}
}
impl Default for QueuedArrayCounter {
fn default() -> Self {
Self::new()
}
}
pub struct ArrayMessage {
pub array: Arc<NDArray>,
pub(crate) counter: Option<Arc<QueuedArrayCounter>>,
pub(crate) done_tx: Option<tokio::sync::oneshot::Sender<()>>,
}
impl Drop for ArrayMessage {
fn drop(&mut self) {
if let Some(tx) = self.done_tx.take() {
let _ = tx.send(());
}
if let Some(c) = self.counter.take() {
c.decrement();
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PublishOutcome {
Delivered,
Disabled,
DroppedQueueFull,
DroppedCompressed,
Throttled,
ChannelClosed,
}
pub struct ArrayAdmission {
compression_aware: AtomicBool,
min_callback_time: AtomicU64,
last_process: parking_lot::Mutex<Option<std::time::Instant>>,
counted_drop: tokio::sync::Notify,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Admission {
Admit,
DropCompressed,
Throttled,
}
impl Default for ArrayAdmission {
fn default() -> Self {
Self {
compression_aware: AtomicBool::new(false),
min_callback_time: AtomicU64::new(0.0f64.to_bits()),
last_process: parking_lot::Mutex::new(None),
counted_drop: tokio::sync::Notify::new(),
}
}
}
impl ArrayAdmission {
pub fn set_compression_aware(&self, aware: bool) {
self.compression_aware.store(aware, Ordering::Release);
}
pub fn set_min_callback_time(&self, seconds: f64) {
self.min_callback_time
.store(seconds.to_bits(), Ordering::Release);
}
pub(crate) fn classify(&self, array: &NDArray) -> Admission {
if array.codec.is_some() && !self.compression_aware.load(Ordering::Acquire) {
return Admission::DropCompressed;
}
let min = f64::from_bits(self.min_callback_time.load(Ordering::Acquire));
let mut last = self.last_process.lock();
if min > 0.0
&& let Some(previous) = *last
&& previous.elapsed().as_secs_f64() < min
{
return Admission::Throttled;
}
*last = Some(std::time::Instant::now());
Admission::Admit
}
pub(crate) fn note_counted_drop(&self) {
self.counted_drop.notify_one();
}
pub(crate) async fn counted_drop(&self) {
self.counted_drop.notified().await;
}
}
#[derive(Debug, Default)]
pub(crate) struct OverflowEpisode(AtomicBool);
impl OverflowEpisode {
fn take(&self) -> bool {
self.0.swap(false, Ordering::AcqRel)
}
fn arm(&self) {
self.0.store(true, Ordering::Release);
}
fn disarm(&self) {
self.0.store(false, Ordering::Release);
}
}
fn try_send_arm(
tx: &parking_lot::RwLock<tokio::sync::mpsc::Sender<ArrayMessage>>,
queued_counter: &Option<Arc<QueuedArrayCounter>>,
dropped_arrays: &AtomicI32,
array: Arc<NDArray>,
episode: &OverflowEpisode,
admission: &ArrayAdmission,
) -> PublishOutcome {
let ignore_queue_full = episode.take();
if let Some(c) = queued_counter {
c.increment();
}
let msg = ArrayMessage {
array,
counter: queued_counter.clone(),
done_tx: None,
};
match tx.read().try_send(msg) {
Ok(()) => PublishOutcome::Delivered,
Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
episode.arm();
if !ignore_queue_full {
dropped_arrays.fetch_add(1, Ordering::AcqRel);
admission.note_counted_drop();
}
PublishOutcome::DroppedQueueFull
}
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => PublishOutcome::ChannelClosed,
}
}
#[derive(Clone)]
pub struct NDArraySender {
tx: Arc<parking_lot::RwLock<tokio::sync::mpsc::Sender<ArrayMessage>>>,
port_name: String,
enabled: Arc<AtomicBool>,
blocking_mode: Arc<AtomicBool>,
queued_counter: Option<Arc<QueuedArrayCounter>>,
dropped_arrays: Arc<AtomicI32>,
overflow: Arc<OverflowEpisode>,
admission: Arc<ArrayAdmission>,
}
impl NDArraySender {
pub async fn publish(&self, array: Arc<NDArray>) -> PublishOutcome {
self.publish_inner(array).await
}
pub async fn publish_scatter(&self, array: Arc<NDArray>, is_last: bool) -> PublishOutcome {
if is_last {
self.overflow.disarm();
} else {
self.overflow.arm();
}
self.publish_inner(array).await
}
async fn publish_inner(&self, array: Arc<NDArray>) -> PublishOutcome {
if !self.enabled.load(Ordering::Acquire) {
return PublishOutcome::Disabled;
}
match self.admission.classify(&array) {
Admission::Admit => {}
Admission::DropCompressed => {
self.dropped_arrays.fetch_add(1, Ordering::AcqRel);
self.admission.note_counted_drop();
return PublishOutcome::DroppedCompressed;
}
Admission::Throttled => return PublishOutcome::Throttled,
}
let blocking = self.blocking_mode.load(Ordering::Acquire);
if !blocking {
return try_send_arm(
&self.tx,
&self.queued_counter,
&self.dropped_arrays,
array,
&self.overflow,
&self.admission,
);
}
self.overflow.disarm();
if let Some(ref c) = self.queued_counter {
c.increment();
}
let (done_tx, done_rx) = tokio::sync::oneshot::channel();
let msg = ArrayMessage {
array,
counter: self.queued_counter.clone(),
done_tx: Some(done_tx),
};
let tx = self.tx.read().clone();
if tx.send(msg).await.is_err() {
return PublishOutcome::ChannelClosed;
}
let _ = done_rx.await;
PublishOutcome::Delivered
}
pub fn is_enabled(&self) -> bool {
self.enabled.load(Ordering::Acquire)
}
pub fn is_blocking(&self) -> bool {
self.blocking_mode.load(Ordering::Acquire)
}
pub fn port_name(&self) -> &str {
&self.port_name
}
pub fn set_queued_counter(&mut self, counter: Arc<QueuedArrayCounter>) {
self.queued_counter = Some(counter);
}
pub fn set_dropped_arrays_counter(&mut self, counter: Arc<AtomicI32>) {
self.dropped_arrays = counter;
}
pub fn dropped_arrays_counter(&self) -> &Arc<AtomicI32> {
&self.dropped_arrays
}
pub fn admission(&self) -> &Arc<ArrayAdmission> {
&self.admission
}
pub fn capacity(&self) -> usize {
self.tx.read().capacity()
}
pub fn max_capacity(&self) -> usize {
self.tx.read().max_capacity()
}
pub(crate) fn self_queue_handle(&self) -> SelfQueueHandle {
SelfQueueHandle {
tx: Arc::downgrade(&self.tx),
queued_counter: self.queued_counter.clone(),
dropped_arrays: self.dropped_arrays.clone(),
overflow: OverflowEpisode::default(),
admission: self.admission.clone(),
}
}
pub(crate) fn set_mode_flags(
&mut self,
enabled: Arc<AtomicBool>,
blocking_mode: Arc<AtomicBool>,
) {
self.enabled = enabled;
self.blocking_mode = blocking_mode;
}
}
pub struct NDArrayReceiver {
rx: tokio::sync::mpsc::Receiver<ArrayMessage>,
admission: Arc<ArrayAdmission>,
}
impl NDArrayReceiver {
pub fn admission(&self) -> &Arc<ArrayAdmission> {
&self.admission
}
pub fn pending(&self) -> usize {
self.rx.len()
}
pub fn max_capacity(&self) -> usize {
self.rx.max_capacity()
}
pub fn capacity(&self) -> usize {
self.rx.capacity()
}
pub fn blocking_recv(&mut self) -> Option<Arc<NDArray>> {
self.rx.blocking_recv().map(|msg| msg.array.clone())
}
pub async fn recv(&mut self) -> Option<Arc<NDArray>> {
self.rx.recv().await.map(|msg| msg.array.clone())
}
pub(crate) async fn recv_msg(&mut self) -> Option<ArrayMessage> {
self.rx.recv().await
}
pub(crate) fn try_recv_msg(&mut self) -> Option<ArrayMessage> {
self.rx.try_recv().ok()
}
}
pub(crate) struct SelfQueueHandle {
tx: std::sync::Weak<parking_lot::RwLock<tokio::sync::mpsc::Sender<ArrayMessage>>>,
queued_counter: Option<Arc<QueuedArrayCounter>>,
dropped_arrays: Arc<AtomicI32>,
overflow: OverflowEpisode,
admission: Arc<ArrayAdmission>,
}
impl SelfQueueHandle {
pub(crate) fn replace_queue(&self, capacity: usize) -> Option<NDArrayReceiver> {
let cell = self.tx.upgrade()?;
let (tx, rx) = tokio::sync::mpsc::channel(capacity.max(1));
*cell.write() = tx;
Some(NDArrayReceiver {
rx,
admission: Arc::clone(&self.admission),
})
}
pub(crate) fn try_enqueue(&self, array: Arc<NDArray>) -> Option<PublishOutcome> {
let cell = self.tx.upgrade()?;
match self.admission.classify(&array) {
Admission::Admit => {}
Admission::DropCompressed => {
self.dropped_arrays.fetch_add(1, Ordering::AcqRel);
self.admission.note_counted_drop();
return Some(PublishOutcome::DroppedCompressed);
}
Admission::Throttled => return Some(PublishOutcome::Throttled),
}
Some(try_send_arm(
&cell,
&self.queued_counter,
&self.dropped_arrays,
array,
&self.overflow,
&self.admission,
))
}
}
pub fn ndarray_channel(port_name: &str, queue_size: usize) -> (NDArraySender, NDArrayReceiver) {
let (tx, rx) = tokio::sync::mpsc::channel(queue_size.max(1));
let admission = Arc::new(ArrayAdmission::default());
(
NDArraySender {
tx: Arc::new(parking_lot::RwLock::new(tx)),
port_name: port_name.to_string(),
enabled: Arc::new(AtomicBool::new(true)),
blocking_mode: Arc::new(AtomicBool::new(false)),
queued_counter: None,
dropped_arrays: Arc::new(AtomicI32::new(0)),
overflow: Arc::new(OverflowEpisode::default()),
admission: Arc::clone(&admission),
},
NDArrayReceiver { rx, admission },
)
}
pub struct NDArrayOutput {
senders: Vec<NDArraySender>,
}
impl NDArrayOutput {
pub fn new() -> Self {
Self {
senders: Vec::new(),
}
}
pub fn add(&mut self, sender: NDArraySender) {
self.senders.push(sender);
}
pub fn remove(&mut self, port_name: &str) {
self.senders.retain(|s| s.port_name != port_name);
}
pub fn take(&mut self, port_name: &str) -> Option<NDArraySender> {
let idx = self.senders.iter().position(|s| s.port_name == port_name)?;
Some(self.senders.swap_remove(idx))
}
pub async fn publish(&self, array: Arc<NDArray>) -> Vec<PublishOutcome> {
let futs = self.senders.iter().map(|s| s.publish(array.clone()));
futures_util::future::join_all(futs).await
}
pub async fn publish_to(&self, index: usize, array: Arc<NDArray>) -> Option<PublishOutcome> {
if let Some(sender) = self.senders.get(index % self.senders.len().max(1)) {
Some(sender.publish(array).await)
} else {
None
}
}
pub fn num_senders(&self) -> usize {
self.senders.len()
}
pub(crate) fn senders_clone(&self) -> Vec<NDArraySender> {
self.senders.clone()
}
}
#[derive(Clone)]
pub struct ArrayPublisher {
output: Arc<parking_lot::Mutex<NDArrayOutput>>,
}
impl ArrayPublisher {
pub fn new(output: Arc<parking_lot::Mutex<NDArrayOutput>>) -> Self {
Self { output }
}
pub async fn publish(&self, array: Arc<NDArray>) -> Vec<PublishOutcome> {
let senders = self.output.lock().senders_clone();
let futs = senders.iter().map(|s| s.publish(array.clone()));
futures_util::future::join_all(futs).await
}
}
impl Default for NDArrayOutput {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ndarray::{NDArray, NDDataType, NDDimension};
fn make_test_array(id: i32) -> Arc<NDArray> {
let mut arr = NDArray::new(vec![NDDimension::new(4)], NDDataType::UInt8);
arr.unique_id = id;
Arc::new(arr)
}
#[tokio::test]
async fn test_publish_receive_basic() {
let (sender, mut receiver) = ndarray_channel("TEST", 10);
sender.publish(make_test_array(1)).await;
sender.publish(make_test_array(2)).await;
let a1 = receiver.recv().await.unwrap();
assert_eq!(a1.unique_id, 1);
let a2 = receiver.recv().await.unwrap();
assert_eq!(a2.unique_id, 2);
}
#[tokio::test]
async fn test_publish_blocking_no_drop() {
let (sender, mut receiver) = ndarray_channel("TEST", 1);
sender.blocking_mode.store(true, Ordering::Release);
let s = sender.clone();
let pub_handle = tokio::spawn(async move {
s.publish(make_test_array(1)).await;
s.publish(make_test_array(2)).await;
s.publish(make_test_array(3)).await;
});
let a1 = receiver.recv().await.unwrap();
assert_eq!(a1.unique_id, 1);
let a2 = receiver.recv().await.unwrap();
assert_eq!(a2.unique_id, 2);
let a3 = receiver.recv().await.unwrap();
assert_eq!(a3.unique_id, 3);
pub_handle.await.unwrap();
}
#[tokio::test]
async fn test_publish_drops_on_full_queue() {
let (sender, _receiver) = ndarray_channel("TEST", 1);
assert_eq!(
sender.publish(make_test_array(1)).await,
PublishOutcome::Delivered
);
assert_eq!(
sender.publish(make_test_array(2)).await,
PublishOutcome::DroppedQueueFull
);
}
#[tokio::test]
async fn test_drop_on_full_does_not_leak_counter() {
let counter = Arc::new(QueuedArrayCounter::new());
let (mut sender, _receiver) = ndarray_channel("TEST", 1);
sender.set_queued_counter(counter.clone());
sender.publish(make_test_array(1)).await; assert_eq!(counter.get(), 1);
let outcome = sender.publish(make_test_array(2)).await; assert_eq!(outcome, PublishOutcome::DroppedQueueFull);
assert_eq!(counter.get(), 1);
}
#[tokio::test]
async fn test_blocking_callbacks_completion_wait() {
let (sender, mut receiver) = ndarray_channel("TEST", 10);
sender.blocking_mode.store(true, Ordering::Release);
let completed = Arc::new(AtomicBool::new(false));
let completed_clone = completed.clone();
let recv_handle = tokio::spawn(async move {
let msg = receiver.recv_msg().await.unwrap();
assert_eq!(msg.array.unique_id, 42);
tokio::time::sleep(Duration::from_millis(50)).await;
completed_clone.store(true, Ordering::Release);
});
sender.publish(make_test_array(42)).await;
assert!(completed.load(Ordering::Acquire));
recv_handle.await.unwrap();
}
#[tokio::test]
async fn test_fanout_three_receivers() {
let (s1, mut r1) = ndarray_channel("P1", 10);
let (s2, mut r2) = ndarray_channel("P2", 10);
let (s3, mut r3) = ndarray_channel("P3", 10);
let mut output = NDArrayOutput::new();
output.add(s1);
output.add(s2);
output.add(s3);
output.publish(make_test_array(42)).await;
assert_eq!(r1.recv().await.unwrap().unique_id, 42);
assert_eq!(r2.recv().await.unwrap().unique_id, 42);
assert_eq!(r3.recv().await.unwrap().unique_id, 42);
}
#[test]
fn test_blocking_recv() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let (sender, mut receiver) = ndarray_channel("TEST", 10);
let handle = std::thread::spawn(move || {
let arr = receiver.blocking_recv().unwrap();
arr.unique_id
});
rt.block_on(sender.publish(make_test_array(99)));
let id = handle.join().unwrap();
assert_eq!(id, 99);
}
#[tokio::test]
async fn test_channel_closed_on_receiver_drop() {
let (sender, receiver) = ndarray_channel("TEST", 10);
drop(receiver);
sender.publish(make_test_array(1)).await;
}
#[test]
fn test_queued_counter_basic() {
let counter = QueuedArrayCounter::new();
assert_eq!(counter.get(), 0);
counter.increment();
assert_eq!(counter.get(), 1);
counter.increment();
assert_eq!(counter.get(), 2);
counter.decrement();
assert_eq!(counter.get(), 1);
counter.decrement();
assert_eq!(counter.get(), 0);
}
#[test]
fn test_queued_counter_wait_until_zero() {
let counter = Arc::new(QueuedArrayCounter::new());
counter.increment();
counter.increment();
let c = counter.clone();
let h = std::thread::spawn(move || {
std::thread::sleep(Duration::from_millis(10));
c.decrement();
std::thread::sleep(Duration::from_millis(10));
c.decrement();
});
assert!(counter.wait_until_zero(Duration::from_secs(5)));
h.join().unwrap();
}
#[test]
fn test_queued_counter_wait_timeout() {
let counter = Arc::new(QueuedArrayCounter::new());
counter.increment();
assert!(!counter.wait_until_zero(Duration::from_millis(10)));
}
#[tokio::test]
async fn test_publish_increments_counter() {
let counter = Arc::new(QueuedArrayCounter::new());
let (mut sender, mut _receiver) = ndarray_channel("TEST", 10);
sender.set_queued_counter(counter.clone());
sender.publish(make_test_array(1)).await;
assert_eq!(counter.get(), 1);
sender.publish(make_test_array(2)).await;
assert_eq!(counter.get(), 2);
}
#[tokio::test]
async fn test_message_drop_decrements() {
let counter = Arc::new(QueuedArrayCounter::new());
counter.increment();
let msg = ArrayMessage {
array: make_test_array(1),
counter: Some(counter.clone()),
done_tx: None,
};
assert_eq!(counter.get(), 1);
drop(msg);
assert_eq!(counter.get(), 0);
}
mod overflow_episode {
use super::*;
fn dropped(sender: &NDArraySender) -> i32 {
sender.dropped_arrays.load(Ordering::Acquire)
}
#[tokio::test]
async fn the_first_refusal_of_an_episode_counts() {
let (sender, _receiver) = ndarray_channel("TEST", 1);
sender.publish(make_test_array(1)).await; assert_eq!(
sender.publish(make_test_array(2)).await,
PublishOutcome::DroppedQueueFull
);
assert_eq!(dropped(&sender), 1);
}
#[tokio::test]
async fn consecutive_refusals_do_not_count_again() {
let (sender, _receiver) = ndarray_channel("TEST", 1);
sender.publish(make_test_array(1)).await;
for id in 2..=20 {
assert_eq!(
sender.publish(make_test_array(id)).await,
PublishOutcome::DroppedQueueFull
);
}
assert_eq!(dropped(&sender), 1, "19 dropped arrays, one episode");
}
#[tokio::test]
async fn a_successful_enqueue_ends_the_episode() {
let (sender, mut receiver) = ndarray_channel("TEST", 1);
sender.publish(make_test_array(1)).await;
sender.publish(make_test_array(2)).await; sender.publish(make_test_array(3)).await; assert_eq!(dropped(&sender), 1);
receiver.recv().await.unwrap(); sender.publish(make_test_array(4)).await; sender.publish(make_test_array(5)).await; assert_eq!(dropped(&sender), 2);
}
#[tokio::test]
async fn every_clone_of_a_sender_shares_one_episode() {
let (sender, _receiver) = ndarray_channel("TEST", 1);
let other = sender.clone();
sender.publish(make_test_array(1)).await;
sender.publish(make_test_array(2)).await;
other.publish(make_test_array(3)).await;
assert_eq!(dropped(&sender), 1);
}
#[tokio::test]
async fn the_reinjection_producer_carries_its_own_episode() {
let (sender, _receiver) = ndarray_channel("TEST", 1);
let handle = sender.self_queue_handle();
sender.publish(make_test_array(1)).await; sender.publish(make_test_array(2)).await; assert_eq!(dropped(&sender), 1);
assert_eq!(
handle.try_enqueue(make_test_array(3)),
Some(PublishOutcome::DroppedQueueFull)
);
assert_eq!(
dropped(&sender),
2,
"a separate producer, a separate episode"
);
handle.try_enqueue(make_test_array(4));
assert_eq!(dropped(&sender), 2, "…which then runs its own episode");
}
#[tokio::test]
async fn a_scatter_reroute_arms_the_episode_and_the_last_node_still_counts() {
let (sender, _receiver) = ndarray_channel("TEST", 1);
sender.publish(make_test_array(1)).await; assert_eq!(
sender.publish_scatter(make_test_array(2), false).await,
PublishOutcome::DroppedQueueFull
);
assert_eq!(dropped(&sender), 0, "rerouted past, not dropped");
assert_eq!(
sender.publish_scatter(make_test_array(3), true).await,
PublishOutcome::DroppedQueueFull
);
assert_eq!(dropped(&sender), 1, "the last node owns the drop");
}
#[tokio::test]
async fn the_blocking_arm_ends_an_open_episode() {
let (sender, mut receiver) = ndarray_channel("TEST", 1);
sender.publish(make_test_array(1)).await;
sender.publish(make_test_array(2)).await; assert_eq!(dropped(&sender), 1);
receiver.recv().await.unwrap();
sender.blocking_mode.store(true, Ordering::Release);
let s = sender.clone();
let pending = tokio::spawn(async move { s.publish(make_test_array(3)).await });
let msg = receiver.recv_msg().await.unwrap();
drop(msg);
pending.await.unwrap();
sender.blocking_mode.store(false, Ordering::Release);
sender.publish(make_test_array(4)).await; sender.publish(make_test_array(5)).await; assert_eq!(dropped(&sender), 2);
}
}
}