use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use async_trait::async_trait;
use futures::channel::{mpsc, oneshot};
use futures::executor::block_on;
use futures::future::join;
use futures::sink::SinkExt;
use futures::stream::StreamExt;
use pipecrab_core::{
AudioChunk, AudioFormat, DataFrame, Decision, Direction, Processor, SystemFrame, Transcript,
};
use pipecrab_runtime::{Outbound, PipelineBuilder, Received, Stage, StageError};
struct BlockingStage {
block_rx: Mutex<Option<oneshot::Receiver<()>>>,
started: mpsc::Sender<()>,
interrupted: Arc<AtomicBool>,
}
impl Processor for BlockingStage {
type Effect = ();
fn decide_data(&mut self, _frame: &DataFrame) -> Decision<()> {
Decision::drop().emit(()) }
fn decide_system(&mut self, _dir: Direction, frame: &SystemFrame) -> Decision<()> {
if matches!(frame, SystemFrame::Interrupt) {
self.interrupted.store(true, Ordering::SeqCst);
}
Decision::drop()
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl Stage for BlockingStage {
async fn perform(&self, _effect: (), _out: &Outbound) -> Result<(), StageError> {
let _ = self.started.clone().send(()).await;
let rx = self
.block_rx
.lock()
.unwrap()
.take()
.expect("perform runs once");
let _ = rx.await; Ok(())
}
}
#[test]
fn interrupt_abandons_perform_and_runs_decide_system() {
block_on(async {
let interrupted = Arc::new(AtomicBool::new(false));
let (started_tx, mut started_rx) = mpsc::channel::<()>(1);
let (block_tx, block_rx) = oneshot::channel::<()>();
let stage = BlockingStage {
block_rx: Mutex::new(Some(block_rx)),
started: started_tx,
interrupted: interrupted.clone(),
};
let (ends, driver) = PipelineBuilder::new().stage(stage).build().start();
let input = ends.input; let _output = ends.output;
let feeder = async move {
input
.send_data(Transcript::user_final("go").into())
.await
.unwrap();
started_rx.next().await.expect("perform must start");
input
.send_system(Direction::Down, SystemFrame::Interrupt)
.await
.unwrap();
};
join(feeder, driver).await;
assert!(
interrupted.load(Ordering::SeqCst),
"decide_system(Interrupt) must have run"
);
assert!(
block_tx.is_canceled(),
"the in-flight perform must have been dropped (its receiver gone)",
);
});
}
struct CountingStage {
data_count: Arc<AtomicUsize>,
data_at_preempt: Arc<Mutex<Option<usize>>>,
}
impl Processor for CountingStage {
type Effect = ();
fn decide_data(&mut self, _frame: &DataFrame) -> Decision<()> {
self.data_count.fetch_add(1, Ordering::SeqCst);
Decision::drop() }
fn decide_system(&mut self, _dir: Direction, frame: &SystemFrame) -> Decision<()> {
if matches!(frame, SystemFrame::Start) {
*self.data_at_preempt.lock().unwrap() = Some(self.data_count.load(Ordering::SeqCst));
}
Decision::drop()
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl Stage for CountingStage {
async fn perform(&self, _effect: (), _out: &Outbound) -> Result<(), StageError> {
Ok(())
}
}
#[test]
fn sys_preempts_backed_up_data() {
block_on(async {
let data_count = Arc::new(AtomicUsize::new(0));
let data_at_preempt = Arc::new(Mutex::new(None));
let stage = CountingStage {
data_count: data_count.clone(),
data_at_preempt: data_at_preempt.clone(),
};
let (ends, driver) = PipelineBuilder::new().stage(stage).build().start();
let input = ends.input;
let _output = ends.output;
for i in 0..8 {
input
.send_data(Transcript::user_final(i.to_string()).into())
.await
.unwrap();
}
input
.send_system(Direction::Down, SystemFrame::Start)
.await
.unwrap();
drop(input);
driver.await;
assert_eq!(
data_at_preempt.lock().unwrap().clone(),
Some(0),
"the Start frame must jump the 8-frame data backlog",
);
assert_eq!(
data_count.load(Ordering::SeqCst),
8,
"all backed-up data is still processed afterward"
);
});
}
struct PassThrough;
impl Processor for PassThrough {
type Effect = ();
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl Stage for PassThrough {
async fn perform(&self, _effect: (), _out: &Outbound) -> Result<(), StageError> {
Ok(())
}
}
#[test]
fn pass_through_forwards_data() {
block_on(async {
let (ends, driver) = PipelineBuilder::new().stage(PassThrough).build().start();
let input = ends.input;
let mut output = ends.output;
let feeder = async move {
input
.send_data(Transcript::user_final("hi").into())
.await
.unwrap();
};
join(feeder, driver).await;
match output.recv().await {
Some(Received::Data(DataFrame::Transcript(s))) => assert_eq!(&*s.text, "hi"),
other => panic!("expected forwarded Transcript(hi), got {other:?}"),
}
});
}
fn input_audio(id: u8) -> DataFrame {
DataFrame::InputAudio {
bytes: Arc::from(&[id][..]),
sample_rate: 16_000,
num_channels: 1,
}
}
#[test]
fn interrupt_flushes_data_keeping_survivors_in_order() {
block_on(async {
let (ends, driver) = PipelineBuilder::new().stage(PassThrough).build().start();
let input = ends.input;
let mut output = ends.output;
input.send_data(input_audio(1)).await.unwrap();
input
.send_data(Transcript::user_final("drop me").into())
.await
.unwrap();
input.send_data(input_audio(2)).await.unwrap();
let audio = AudioChunk::new(Arc::from(&[0.0f32, 0.0][..]), AudioFormat::new(48_000, 1));
input.send_data(DataFrame::Audio(audio)).await.unwrap();
input
.send_system(Direction::Down, SystemFrame::Interrupt)
.await
.unwrap();
drop(input);
driver.await;
let mut ids = Vec::new();
while let Ok(frame) = output.data.try_recv() {
match frame {
DataFrame::InputAudio { bytes, .. } => ids.push(bytes[0]),
other => panic!("a non-survivor leaked past the flush: {other:?}"),
}
}
assert_eq!(
ids,
vec![1, 2],
"survivors kept in order; droppable frames flushed"
);
});
}
#[test]
fn nested_pipeline_forwards_through_both_levels() {
block_on(async {
let inner = PipelineBuilder::new().stage(PassThrough).build();
let (ends, driver) = PipelineBuilder::new()
.stage(inner)
.stage(PassThrough)
.build()
.start();
let input = ends.input;
let mut output = ends.output;
let feeder = async move {
input
.send_data(Transcript::user_final("deep").into())
.await
.unwrap();
};
join(feeder, driver).await;
match output.recv().await {
Some(Received::Data(DataFrame::Transcript(s))) => assert_eq!(&*s.text, "deep"),
other => panic!("expected forwarded Transcript(deep), got {other:?}"),
}
});
}
struct DistinctEffect;
struct DistinctEffectStage;
impl Processor for DistinctEffectStage {
type Effect = DistinctEffect;
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl Stage for DistinctEffectStage {
async fn perform(&self, _effect: DistinctEffect, _out: &Outbound) -> Result<(), StageError> {
Ok(())
}
}
#[test]
fn pipeline_composes_stages_with_distinct_effect_types() {
let _pipeline = PipelineBuilder::new()
.stage(PassThrough)
.stage(DistinctEffectStage)
.build();
}
#[cfg(not(target_arch = "wasm32"))]
#[test]
fn driver_is_send_on_native() {
fn assert_send<T: Send>(_: &T) {}
struct Noop;
impl Processor for Noop {
type Effect = ();
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl Stage for Noop {
async fn perform(&self, _e: (), _out: &Outbound) -> Result<(), StageError> {
Ok(())
}
}
let pipeline = PipelineBuilder::new().stage(Noop).build();
let (_ends, driver) = pipeline.start();
assert_send(&driver);
}