1use std::sync::Arc;
2use std::sync::atomic::{AtomicBool, AtomicI32, AtomicUsize, Ordering};
3use std::time::Duration;
4
5use crate::ndarray::NDArray;
6
7pub struct QueuedArrayCounter {
10 count: AtomicUsize,
11 mutex: parking_lot::Mutex<()>,
12 condvar: parking_lot::Condvar,
13}
14
15impl QueuedArrayCounter {
16 pub fn new() -> Self {
18 Self {
19 count: AtomicUsize::new(0),
20 mutex: parking_lot::Mutex::new(()),
21 condvar: parking_lot::Condvar::new(),
22 }
23 }
24
25 pub fn increment(&self) {
27 self.count.fetch_add(1, Ordering::AcqRel);
28 }
29
30 pub fn decrement(&self) {
32 let prev = self.count.fetch_sub(1, Ordering::AcqRel);
33 if prev == 1 {
34 let _guard = self.mutex.lock();
35 self.condvar.notify_all();
36 }
37 }
38
39 pub fn get(&self) -> usize {
41 self.count.load(Ordering::Acquire)
42 }
43
44 pub fn wait_until_zero(&self, timeout: Duration) -> bool {
47 let mut guard = self.mutex.lock();
48 if self.count.load(Ordering::Acquire) == 0 {
49 return true;
50 }
51 !self
52 .condvar
53 .wait_while_for(
54 &mut guard,
55 |_| self.count.load(Ordering::Acquire) != 0,
56 timeout,
57 )
58 .timed_out()
59 }
60}
61
62impl Default for QueuedArrayCounter {
63 fn default() -> Self {
64 Self::new()
65 }
66}
67
68pub struct ArrayMessage {
72 pub array: Arc<NDArray>,
73 pub(crate) counter: Option<Arc<QueuedArrayCounter>>,
74 pub(crate) done_tx: Option<tokio::sync::oneshot::Sender<()>>,
77}
78
79impl Drop for ArrayMessage {
80 fn drop(&mut self) {
81 if let Some(tx) = self.done_tx.take() {
82 let _ = tx.send(());
83 }
84 if let Some(c) = self.counter.take() {
85 c.decrement();
86 }
87 }
88}
89
90#[derive(Debug, Clone, Copy, PartialEq, Eq)]
92pub enum PublishOutcome {
93 Delivered,
95 Disabled,
97 DroppedQueueFull,
100 ChannelClosed,
102}
103
104#[derive(Clone)]
120pub struct NDArraySender {
121 tx: tokio::sync::mpsc::Sender<ArrayMessage>,
122 port_name: String,
123 enabled: Arc<AtomicBool>,
124 blocking_mode: Arc<AtomicBool>,
125 queued_counter: Option<Arc<QueuedArrayCounter>>,
126 dropped_arrays: Arc<AtomicI32>,
132}
133
134impl NDArraySender {
135 pub async fn publish(&self, array: Arc<NDArray>) -> PublishOutcome {
143 self.publish_inner(array, true).await
144 }
145
146 pub async fn publish_scatter(&self, array: Arc<NDArray>, is_last: bool) -> PublishOutcome {
154 self.publish_inner(array, is_last).await
155 }
156
157 async fn publish_inner(&self, array: Arc<NDArray>, count_drop: bool) -> PublishOutcome {
162 if !self.enabled.load(Ordering::Acquire) {
163 return PublishOutcome::Disabled;
164 }
165
166 let blocking = self.blocking_mode.load(Ordering::Acquire);
167
168 if !blocking {
169 if let Some(ref c) = self.queued_counter {
172 c.increment();
173 }
174 let msg = ArrayMessage {
175 array,
176 counter: self.queued_counter.clone(),
177 done_tx: None,
178 };
179 return match self.tx.try_send(msg) {
180 Ok(()) => PublishOutcome::Delivered,
181 Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
183 if count_drop {
184 self.dropped_arrays.fetch_add(1, Ordering::AcqRel);
185 }
186 PublishOutcome::DroppedQueueFull
187 }
188 Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => {
189 PublishOutcome::ChannelClosed
190 }
191 };
192 }
193
194 if let Some(ref c) = self.queued_counter {
196 c.increment();
197 }
198 let (done_tx, done_rx) = tokio::sync::oneshot::channel();
199 let msg = ArrayMessage {
200 array,
201 counter: self.queued_counter.clone(),
202 done_tx: Some(done_tx),
203 };
204 if self.tx.send(msg).await.is_err() {
205 return PublishOutcome::ChannelClosed;
207 }
208 let _ = done_rx.await;
209 PublishOutcome::Delivered
210 }
211
212 pub fn is_enabled(&self) -> bool {
214 self.enabled.load(Ordering::Acquire)
215 }
216
217 pub fn is_blocking(&self) -> bool {
219 self.blocking_mode.load(Ordering::Acquire)
220 }
221
222 pub fn port_name(&self) -> &str {
223 &self.port_name
224 }
225
226 pub fn set_queued_counter(&mut self, counter: Arc<QueuedArrayCounter>) {
228 self.queued_counter = Some(counter);
229 }
230
231 pub fn set_dropped_arrays_counter(&mut self, counter: Arc<AtomicI32>) {
234 self.dropped_arrays = counter;
235 }
236
237 pub fn dropped_arrays_counter(&self) -> &Arc<AtomicI32> {
239 &self.dropped_arrays
240 }
241
242 pub fn capacity(&self) -> usize {
244 self.tx.capacity()
245 }
246
247 pub fn max_capacity(&self) -> usize {
249 self.tx.max_capacity()
250 }
251
252 pub(crate) fn set_mode_flags(
254 &mut self,
255 enabled: Arc<AtomicBool>,
256 blocking_mode: Arc<AtomicBool>,
257 ) {
258 self.enabled = enabled;
259 self.blocking_mode = blocking_mode;
260 }
261}
262
263pub struct NDArrayReceiver {
265 rx: tokio::sync::mpsc::Receiver<ArrayMessage>,
266}
267
268impl NDArrayReceiver {
269 pub fn pending(&self) -> usize {
271 self.rx.len()
272 }
273
274 pub fn max_capacity(&self) -> usize {
276 self.rx.max_capacity()
277 }
278
279 pub fn capacity(&self) -> usize {
281 self.rx.capacity()
282 }
283
284 pub fn blocking_recv(&mut self) -> Option<Arc<NDArray>> {
286 self.rx.blocking_recv().map(|msg| msg.array.clone())
287 }
288
289 pub async fn recv(&mut self) -> Option<Arc<NDArray>> {
291 self.rx.recv().await.map(|msg| msg.array.clone())
292 }
293
294 pub(crate) async fn recv_msg(&mut self) -> Option<ArrayMessage> {
297 self.rx.recv().await
298 }
299}
300
301pub fn ndarray_channel(port_name: &str, queue_size: usize) -> (NDArraySender, NDArrayReceiver) {
303 let (tx, rx) = tokio::sync::mpsc::channel(queue_size.max(1));
304 (
305 NDArraySender {
306 tx,
307 port_name: port_name.to_string(),
308 enabled: Arc::new(AtomicBool::new(true)),
309 blocking_mode: Arc::new(AtomicBool::new(false)),
310 queued_counter: None,
311 dropped_arrays: Arc::new(AtomicI32::new(0)),
312 },
313 NDArrayReceiver { rx },
314 )
315}
316
317pub struct NDArrayOutput {
319 senders: Vec<NDArraySender>,
320}
321
322impl NDArrayOutput {
323 pub fn new() -> Self {
324 Self {
325 senders: Vec::new(),
326 }
327 }
328
329 pub fn add(&mut self, sender: NDArraySender) {
330 self.senders.push(sender);
331 }
332
333 pub fn remove(&mut self, port_name: &str) {
334 self.senders.retain(|s| s.port_name != port_name);
335 }
336
337 pub fn take(&mut self, port_name: &str) -> Option<NDArraySender> {
339 let idx = self.senders.iter().position(|s| s.port_name == port_name)?;
340 Some(self.senders.swap_remove(idx))
341 }
342
343 pub async fn publish(&self, array: Arc<NDArray>) -> Vec<PublishOutcome> {
349 let futs = self.senders.iter().map(|s| s.publish(array.clone()));
350 futures_util::future::join_all(futs).await
351 }
352
353 pub async fn publish_to(&self, index: usize, array: Arc<NDArray>) -> Option<PublishOutcome> {
355 if let Some(sender) = self.senders.get(index % self.senders.len().max(1)) {
356 Some(sender.publish(array).await)
357 } else {
358 None
359 }
360 }
361
362 pub fn num_senders(&self) -> usize {
363 self.senders.len()
364 }
365
366 pub(crate) fn senders_clone(&self) -> Vec<NDArraySender> {
368 self.senders.clone()
369 }
370}
371
372#[derive(Clone)]
385pub struct ArrayPublisher {
386 output: Arc<parking_lot::Mutex<NDArrayOutput>>,
387}
388
389impl ArrayPublisher {
390 pub fn new(output: Arc<parking_lot::Mutex<NDArrayOutput>>) -> Self {
392 Self { output }
393 }
394
395 pub async fn publish(&self, array: Arc<NDArray>) -> Vec<PublishOutcome> {
402 let senders = self.output.lock().senders_clone();
403 let futs = senders.iter().map(|s| s.publish(array.clone()));
404 futures_util::future::join_all(futs).await
405 }
406}
407
408impl Default for NDArrayOutput {
409 fn default() -> Self {
410 Self::new()
411 }
412}
413
414#[cfg(test)]
415mod tests {
416 use super::*;
417 use crate::ndarray::{NDArray, NDDataType, NDDimension};
418
419 fn make_test_array(id: i32) -> Arc<NDArray> {
420 let mut arr = NDArray::new(vec![NDDimension::new(4)], NDDataType::UInt8);
421 arr.unique_id = id;
422 Arc::new(arr)
423 }
424
425 #[tokio::test]
426 async fn test_publish_receive_basic() {
427 let (sender, mut receiver) = ndarray_channel("TEST", 10);
428 sender.publish(make_test_array(1)).await;
429 sender.publish(make_test_array(2)).await;
430
431 let a1 = receiver.recv().await.unwrap();
432 assert_eq!(a1.unique_id, 1);
433 let a2 = receiver.recv().await.unwrap();
434 assert_eq!(a2.unique_id, 2);
435 }
436
437 #[tokio::test]
438 async fn test_publish_blocking_no_drop() {
439 let (sender, mut receiver) = ndarray_channel("TEST", 1);
442 sender.blocking_mode.store(true, Ordering::Release);
443
444 let s = sender.clone();
445 let pub_handle = tokio::spawn(async move {
446 s.publish(make_test_array(1)).await;
447 s.publish(make_test_array(2)).await;
448 s.publish(make_test_array(3)).await;
449 });
450
451 let a1 = receiver.recv().await.unwrap();
453 assert_eq!(a1.unique_id, 1);
454 let a2 = receiver.recv().await.unwrap();
455 assert_eq!(a2.unique_id, 2);
456 let a3 = receiver.recv().await.unwrap();
457 assert_eq!(a3.unique_id, 3);
458
459 pub_handle.await.unwrap();
460 }
461
462 #[tokio::test]
463 async fn test_publish_drops_on_full_queue() {
464 let (sender, _receiver) = ndarray_channel("TEST", 1);
467
468 assert_eq!(
470 sender.publish(make_test_array(1)).await,
471 PublishOutcome::Delivered
472 );
473 assert_eq!(
475 sender.publish(make_test_array(2)).await,
476 PublishOutcome::DroppedQueueFull
477 );
478 }
479
480 #[tokio::test]
481 async fn test_drop_on_full_does_not_leak_counter() {
482 let counter = Arc::new(QueuedArrayCounter::new());
484 let (mut sender, _receiver) = ndarray_channel("TEST", 1);
485 sender.set_queued_counter(counter.clone());
486
487 sender.publish(make_test_array(1)).await; assert_eq!(counter.get(), 1);
489 let outcome = sender.publish(make_test_array(2)).await; assert_eq!(outcome, PublishOutcome::DroppedQueueFull);
491 assert_eq!(counter.get(), 1);
493 }
494
495 #[tokio::test]
496 async fn test_blocking_callbacks_completion_wait() {
497 let (sender, mut receiver) = ndarray_channel("TEST", 10);
498 sender.blocking_mode.store(true, Ordering::Release);
499
500 let completed = Arc::new(AtomicBool::new(false));
501 let completed_clone = completed.clone();
502
503 let recv_handle = tokio::spawn(async move {
505 let msg = receiver.recv_msg().await.unwrap();
506 assert_eq!(msg.array.unique_id, 42);
507 tokio::time::sleep(Duration::from_millis(50)).await;
509 completed_clone.store(true, Ordering::Release);
510 });
512
513 sender.publish(make_test_array(42)).await;
515
516 assert!(completed.load(Ordering::Acquire));
518
519 recv_handle.await.unwrap();
520 }
521
522 #[tokio::test]
523 async fn test_fanout_three_receivers() {
524 let (s1, mut r1) = ndarray_channel("P1", 10);
525 let (s2, mut r2) = ndarray_channel("P2", 10);
526 let (s3, mut r3) = ndarray_channel("P3", 10);
527
528 let mut output = NDArrayOutput::new();
529 output.add(s1);
530 output.add(s2);
531 output.add(s3);
532
533 output.publish(make_test_array(42)).await;
534
535 assert_eq!(r1.recv().await.unwrap().unique_id, 42);
536 assert_eq!(r2.recv().await.unwrap().unique_id, 42);
537 assert_eq!(r3.recv().await.unwrap().unique_id, 42);
538 }
539
540 #[test]
541 fn test_blocking_recv() {
542 let rt = tokio::runtime::Builder::new_current_thread()
543 .enable_all()
544 .build()
545 .unwrap();
546 let (sender, mut receiver) = ndarray_channel("TEST", 10);
547
548 let handle = std::thread::spawn(move || {
549 let arr = receiver.blocking_recv().unwrap();
550 arr.unique_id
551 });
552
553 rt.block_on(sender.publish(make_test_array(99)));
554 let id = handle.join().unwrap();
555 assert_eq!(id, 99);
556 }
557
558 #[tokio::test]
559 async fn test_channel_closed_on_receiver_drop() {
560 let (sender, receiver) = ndarray_channel("TEST", 10);
561 drop(receiver);
562 sender.publish(make_test_array(1)).await;
564 }
565
566 #[test]
567 fn test_queued_counter_basic() {
568 let counter = QueuedArrayCounter::new();
569 assert_eq!(counter.get(), 0);
570 counter.increment();
571 assert_eq!(counter.get(), 1);
572 counter.increment();
573 assert_eq!(counter.get(), 2);
574 counter.decrement();
575 assert_eq!(counter.get(), 1);
576 counter.decrement();
577 assert_eq!(counter.get(), 0);
578 }
579
580 #[test]
581 fn test_queued_counter_wait_until_zero() {
582 let counter = Arc::new(QueuedArrayCounter::new());
583 counter.increment();
584 counter.increment();
585
586 let c = counter.clone();
587 let h = std::thread::spawn(move || {
588 std::thread::sleep(Duration::from_millis(10));
589 c.decrement();
590 std::thread::sleep(Duration::from_millis(10));
591 c.decrement();
592 });
593
594 assert!(counter.wait_until_zero(Duration::from_secs(5)));
595 h.join().unwrap();
596 }
597
598 #[test]
599 fn test_queued_counter_wait_timeout() {
600 let counter = Arc::new(QueuedArrayCounter::new());
601 counter.increment();
602 assert!(!counter.wait_until_zero(Duration::from_millis(10)));
603 }
604
605 #[tokio::test]
606 async fn test_publish_increments_counter() {
607 let counter = Arc::new(QueuedArrayCounter::new());
608 let (mut sender, mut _receiver) = ndarray_channel("TEST", 10);
609 sender.set_queued_counter(counter.clone());
610
611 sender.publish(make_test_array(1)).await;
612 assert_eq!(counter.get(), 1);
613 sender.publish(make_test_array(2)).await;
614 assert_eq!(counter.get(), 2);
615 }
616
617 #[tokio::test]
618 async fn test_message_drop_decrements() {
619 let counter = Arc::new(QueuedArrayCounter::new());
620 counter.increment();
621 let msg = ArrayMessage {
622 array: make_test_array(1),
623 counter: Some(counter.clone()),
624 done_tx: None,
625 };
626 assert_eq!(counter.get(), 1);
627 drop(msg);
628 assert_eq!(counter.get(), 0);
629 }
630}