use std::sync::Arc;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use crate::filter::TextFilter;
use crate::processor::{Correction, TextProcessor};
use crate::state::StateMachine;
use crate::traits::*;
use crate::types::*;
use crate::vad::{Segmenter, SegmenterConfig, SpeechSegment, VadBackend, VadStream};
pub struct PipelineBuilder {
asr: Option<Arc<dyn AsrAdapter>>,
filters: Vec<Arc<dyn TextFilter>>,
processors: Vec<Arc<dyn TextProcessor>>,
refiner: Option<Arc<dyn LlmRefiner>>,
context: Option<Arc<dyn ContextProvider>>,
emitter: Option<Arc<dyn OutputEmitter>>,
vad: Option<Arc<dyn VadBackend>>,
segmenter_config: SegmenterConfig,
final_pass: FinalPass,
audio_channel_size: usize,
asr_channel_size: usize,
}
impl PipelineBuilder {
pub fn new() -> Self {
Self {
asr: None,
filters: Vec::new(),
processors: Vec::new(),
refiner: None,
context: None,
emitter: None,
vad: None,
segmenter_config: SegmenterConfig::default(),
final_pass: FinalPass::SpeechOnly,
audio_channel_size: 32,
asr_channel_size: 8,
}
}
pub fn asr(mut self, asr: impl AsrAdapter + 'static) -> Self {
self.asr = Some(Arc::new(asr));
self
}
pub fn filter(mut self, filter: impl TextFilter + 'static) -> Self {
self.filters.push(Arc::new(filter));
self
}
pub fn refiner(mut self, refiner: impl LlmRefiner + 'static) -> Self {
self.refiner = Some(Arc::new(refiner));
self
}
pub fn processor(mut self, proc: impl TextProcessor + 'static) -> Self {
self.processors.push(Arc::new(proc));
self
}
pub fn context(mut self, ctx: impl ContextProvider + 'static) -> Self {
self.context = Some(Arc::new(ctx));
self
}
pub fn emitter(mut self, emitter: impl OutputEmitter + 'static) -> Self {
self.emitter = Some(Arc::new(emitter));
self
}
pub fn vad(mut self, vad: impl VadBackend + 'static) -> Self {
self.vad = Some(Arc::new(vad));
self
}
pub fn segmenter_config(mut self, config: SegmenterConfig) -> Self {
self.segmenter_config = config;
self
}
pub fn final_pass(mut self, policy: FinalPass) -> Self {
self.final_pass = policy;
self
}
pub fn audio_channel_size(mut self, size: usize) -> Self {
self.audio_channel_size = size;
self
}
pub fn asr_channel_size(mut self, size: usize) -> Self {
self.asr_channel_size = size;
self
}
pub fn build(self) -> Result<Pipeline, PipelineError> {
Ok(Pipeline {
asr: self.asr.ok_or(PipelineError::MissingComponent("asr"))?,
filters: self.filters,
processors: self.processors,
refiner: self.refiner,
context: self.context,
emitter: self.emitter,
vad: self.vad,
segmenter_config: self.segmenter_config,
final_pass: self.final_pass,
audio_channel_size: self.audio_channel_size,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum FinalPass {
#[default]
SpeechOnly,
WholeUtterance,
JoinSegments,
}
impl Default for PipelineBuilder {
fn default() -> Self {
Self::new()
}
}
pub struct Pipeline {
asr: Arc<dyn AsrAdapter>,
filters: Vec<Arc<dyn TextFilter>>,
processors: Vec<Arc<dyn TextProcessor>>,
refiner: Option<Arc<dyn LlmRefiner>>,
context: Option<Arc<dyn ContextProvider>>,
emitter: Option<Arc<dyn OutputEmitter>>,
vad: Option<Arc<dyn VadBackend>>,
segmenter_config: SegmenterConfig,
final_pass: FinalPass,
audio_channel_size: usize,
}
impl Pipeline {
pub fn builder() -> PipelineBuilder {
PipelineBuilder::new()
}
pub async fn transcribe(&self, audio: &[AudioChunk]) -> Result<SessionResult, PipelineError> {
let cancel = CancellationToken::new();
let Some(vad) = self.vad.as_deref() else {
return self.run(AsrSource::Audio(audio), Vec::new(), &cancel).await;
};
let samples = AudioChunk::concat(audio);
let sample_rate = AudioChunk::sample_rate_of(audio).unwrap_or(0);
let channels = audio.first().map(|c| c.channels).unwrap_or(1);
let segments =
match crate::vad::segment_buffer(vad, &samples, sample_rate, &self.segmenter_config) {
Ok(segments) => segments,
Err(e) => {
tracing::warn!(error = %e, "voice activity detection disabled");
return self
.run(AsrSource::Audio(audio), Vec::new(), &cancel)
.await
.map(|mut result| {
result.diagnostics.failures.push(StageFailure {
stage: Stage::Vad,
index: 0,
reason: e.to_string(),
});
result
});
}
};
let source = match self.final_pass {
FinalPass::WholeUtterance => AsrSource::Audio(audio),
FinalPass::SpeechOnly => {
AsrSource::Owned(speech_only(&samples, &segments, sample_rate, channels)?)
}
FinalPass::JoinSegments => {
let mut texts = Vec::new();
for segment in &segments {
let chunk = slice_chunk(&samples, *segment, sample_rate, channels);
match transcribe_segment(&self.asr, &chunk, &cancel).await? {
Some(text) => texts.push(text),
None => continue,
}
}
AsrSource::Text(texts.join(" "))
}
};
self.run(source, segments, &cancel).await
}
#[allow(clippy::needless_lifetimes)]
async fn run<'a>(
&self,
source: AsrSource<'a>,
segments: Vec<SpeechSegment>,
cancel: &CancellationToken,
) -> Result<SessionResult, PipelineError> {
run_session(
&self.asr,
&self.filters,
&self.processors,
self.refiner.as_deref(),
self.context.as_deref(),
self.emitter.as_deref(),
source,
segments,
cancel,
)
.await
}
pub fn session(&self) -> Session {
let (audio_tx, mut audio_rx) = mpsc::channel::<AudioChunk>(self.audio_channel_size);
let (partial_tx, partial_rx) = mpsc::channel::<Partial>(PARTIAL_CHANNEL_SIZE);
let (segment_tx, mut segment_rx) =
mpsc::channel::<(usize, SpeechSegment, AudioChunk)>(SEGMENT_CHANNEL_SIZE);
let cancel = CancellationToken::new();
let asr = Arc::clone(&self.asr);
let filters: Vec<Arc<dyn TextFilter>> = self.filters.iter().map(Arc::clone).collect();
let processors: Vec<Arc<dyn TextProcessor>> =
self.processors.iter().map(Arc::clone).collect();
let refiner = self.refiner.clone();
let context = self.context.clone();
let emitter = self.emitter.clone();
let vad = self.vad.clone();
let segmenter_config = self.segmenter_config.clone();
let final_pass = self.final_pass;
let cancel_inner = cancel.clone();
let partial_probe = partial_tx.clone();
let asr_partials = Arc::clone(&asr);
let cancel_partials = cancel.clone();
let asr_task = tokio::spawn(async move {
let mut texts: Vec<String> = Vec::new();
while let Some((index, segment, chunk)) = segment_rx.recv().await {
if cancel_partials.is_cancelled() {
break;
}
let rate = chunk.sample_rate;
match transcribe_segment(&asr_partials, &chunk, &cancel_partials).await {
Ok(Some(text)) => {
let start = segment.offset(rate);
let _ = partial_tx.try_send(Partial {
index,
text: text.clone(),
start,
end: start + segment.duration(rate),
});
texts.push(text);
}
Ok(None) => {}
Err(e) => {
tracing::warn!(error = %e, index, "utterance transcription failed");
}
}
}
texts
});
let handle = tokio::spawn(async move {
let mut chunks: Vec<AudioChunk> = Vec::new();
let mut samples: Vec<f32> = Vec::new();
let mut sample_rate = 0u32;
let mut channels = 1u16;
let mut live: Option<LiveSegmenter> = None;
let mut segments: Vec<SpeechSegment> = Vec::new();
let mut vad_active = vad.is_some();
let mut vad_failure: Option<String> = None;
let wanted = |probe: &mpsc::Sender<Partial>| {
final_pass == FinalPass::JoinSegments || !probe.is_closed()
};
let mut cancelled = false;
loop {
tokio::select! {
biased;
_ = cancel_inner.cancelled() => { cancelled = true; break; }
maybe = audio_rx.recv() => match maybe {
Some(chunk) => {
if sample_rate == 0 {
sample_rate = chunk.sample_rate;
channels = chunk.channels;
}
samples.extend_from_slice(&chunk.samples);
chunks.push(chunk);
if vad_active {
if live.is_none() {
let backend = vad.as_deref().expect("vad_active implies a backend");
match LiveSegmenter::new(backend, sample_rate, segmenter_config.clone()) {
Ok(segmenter) => live = Some(segmenter),
Err(e) => {
tracing::warn!(error = %e, "voice activity detection disabled");
vad_failure = Some(e.to_string());
vad_active = false;
}
}
}
if let Some(segmenter) = live.as_mut() {
for segment in segmenter.advance(&samples) {
let index = segments.len();
segments.push(segment);
if wanted(&partial_probe) {
let chunk = slice_chunk(&samples, segment, sample_rate, channels);
let _ = segment_tx.send((index, segment, chunk)).await;
}
}
}
}
}
None => break,
},
}
}
if !cancelled {
if let Some(segmenter) = live.as_mut() {
for segment in segmenter.flush(&samples) {
let index = segments.len();
segments.push(segment);
if wanted(&partial_probe) {
let chunk = slice_chunk(&samples, segment, sample_rate, channels);
let _ = segment_tx.send((index, segment, chunk)).await;
}
}
}
}
drop(segment_tx);
drop(partial_probe);
let segment_texts = asr_task.await.unwrap_or_default();
if cancelled {
return Err(PipelineError::Cancelled { during: "recording" });
}
let source = if !vad_active {
AsrSource::Owned(chunks)
} else {
match final_pass {
FinalPass::WholeUtterance => AsrSource::Owned(chunks),
FinalPass::SpeechOnly => {
AsrSource::Owned(speech_only(&samples, &segments, sample_rate, channels)?)
}
FinalPass::JoinSegments => AsrSource::Text(segment_texts.join(" ")),
}
};
run_session(
&asr,
&filters,
&processors,
refiner.as_deref(),
context.as_deref(),
emitter.as_deref(),
source,
segments,
&cancel_inner,
)
.await
.map(|mut result| {
if let Some(reason) = vad_failure {
result.diagnostics.failures.push(StageFailure {
stage: Stage::Vad,
index: 0,
reason,
});
}
result
})
});
Session {
audio: audio_tx,
partials: partial_rx,
cancel,
handle,
}
}
}
const PARTIAL_CHANNEL_SIZE: usize = 8;
const SEGMENT_CHANNEL_SIZE: usize = 4;
enum AsrSource<'a> {
Audio(&'a [AudioChunk]),
Owned(Vec<AudioChunk>),
Text(String),
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Partial {
pub index: usize,
pub text: String,
pub start: std::time::Duration,
pub end: std::time::Duration,
}
struct LiveSegmenter {
segmenter: Segmenter,
stream: Box<dyn VadStream>,
frame_size: usize,
consumed: usize,
}
impl LiveSegmenter {
fn new(
backend: &dyn VadBackend,
sample_rate: u32,
config: SegmenterConfig,
) -> Result<Self, crate::vad::VadError> {
Ok(Self {
segmenter: Segmenter::new(backend, sample_rate, config)?,
stream: backend.start(),
frame_size: backend.frame_size().max(1),
consumed: 0,
})
}
fn advance(&mut self, samples: &[f32]) -> Vec<SpeechSegment> {
let mut closed = Vec::new();
while samples.len() - self.consumed >= self.frame_size {
let frame = &samples[self.consumed..self.consumed + self.frame_size];
let probability = self.stream.speech_probability(frame);
self.consumed += self.frame_size;
if let Some(segment) = self.segmenter.push(probability) {
closed.push(segment);
}
}
closed
}
fn flush(&mut self, samples: &[f32]) -> Vec<SpeechSegment> {
let mut closed = self.advance(samples);
if samples.len() > self.consumed {
let mut frame = samples[self.consumed..].to_vec();
frame.resize(self.frame_size, 0.0);
let probability = self.stream.speech_probability(&frame);
self.consumed = samples.len();
if let Some(segment) = self.segmenter.push(probability) {
closed.push(segment);
}
}
closed.extend(self.segmenter.flush());
closed
}
}
fn slice_chunk(
samples: &[f32],
segment: SpeechSegment,
sample_rate: u32,
channels: u16,
) -> AudioChunk {
let start = segment.start.min(samples.len());
let end = segment.end.clamp(start, samples.len());
AudioChunk {
samples: samples[start..end].to_vec(),
sample_rate,
channels,
}
}
fn speech_only(
samples: &[f32],
segments: &[SpeechSegment],
sample_rate: u32,
channels: u16,
) -> Result<Vec<AudioChunk>, PipelineError> {
let mut speech: Vec<f32> = Vec::new();
let mut cursor = 0usize;
for segment in segments {
let start = segment.start.max(cursor).min(samples.len());
let end = segment.end.clamp(start, samples.len());
speech.extend_from_slice(&samples[start..end]);
cursor = end;
}
if speech.is_empty() {
return Err(PipelineError::NoSpeech);
}
Ok(vec![AudioChunk {
samples: speech,
sample_rate,
channels,
}])
}
async fn transcribe_segment(
asr: &Arc<dyn AsrAdapter>,
chunk: &AudioChunk,
cancel: &CancellationToken,
) -> Result<Option<String>, PipelineError> {
if chunk.samples.is_empty() {
return Ok(None);
}
let transcript = tokio::select! {
biased;
_ = cancel.cancelled() => return Err(PipelineError::Cancelled { during: "recording" }),
result = asr.transcribe(std::slice::from_ref(chunk)) => result?,
};
let text = transcript.text.trim().to_string();
Ok(if text.is_empty() { None } else { Some(text) })
}
pub struct Session {
pub audio: mpsc::Sender<AudioChunk>,
pub partials: mpsc::Receiver<Partial>,
pub cancel: CancellationToken,
handle: tokio::task::JoinHandle<Result<SessionResult, PipelineError>>,
}
impl Session {
pub async fn finish(self) -> Result<SessionResult, PipelineError> {
let Session { audio, handle, .. } = self;
drop(audio);
match handle.await {
Ok(result) => result,
Err(e) => Err(PipelineError::TaskFailed(e.to_string())),
}
}
pub async fn abort(self) {
self.cancel.cancel();
let _ = self.handle.await;
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct SessionResult {
pub raw_text: String,
pub output: RefinementOutput,
pub emit_result: Option<EmitResult>,
pub diagnostics: Diagnostics,
}
impl SessionResult {
pub fn text(&self) -> &str {
match &self.output {
RefinementOutput::TextInsertion { text, .. } => text,
RefinementOutput::StructuredInput { text, .. } => text.as_deref().unwrap_or_default(),
RefinementOutput::Command { .. } => "",
}
}
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct Diagnostics {
pub removed: Vec<String>,
pub corrections: Vec<Correction>,
pub failures: Vec<StageFailure>,
pub speech_segments: Vec<SpeechSegment>,
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct StageFailure {
pub stage: Stage,
pub index: usize,
pub reason: String,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum Stage {
Vad,
Filter,
Processor,
Refiner,
}
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum PipelineError {
#[error("pipeline is missing a required component: {0}")]
MissingComponent(&'static str),
#[error("no speech detected")]
NoSpeech,
#[error(transparent)]
Asr(#[from] AsrError),
#[error("cancelled during {during}")]
Cancelled { during: &'static str },
#[error("session task failed: {0}")]
TaskFailed(String),
#[error("invalid state transition: {0}")]
InvalidTransition(String),
}
#[allow(clippy::too_many_arguments)]
async fn run_session(
asr: &Arc<dyn AsrAdapter>,
filters: &[Arc<dyn TextFilter>],
processors: &[Arc<dyn TextProcessor>],
refiner: Option<&dyn LlmRefiner>,
context: Option<&dyn ContextProvider>,
emitter: Option<&dyn OutputEmitter>,
source: AsrSource<'_>,
speech_segments: Vec<SpeechSegment>,
cancel: &CancellationToken,
) -> Result<SessionResult, PipelineError> {
let mut sm = StateMachine::new();
let transition = |sm: &mut StateMachine, to| {
sm.transition(to)
.map(|_| ())
.map_err(|e| PipelineError::InvalidTransition(e.to_string()))
};
transition(&mut sm, PipelineState::Activating)?;
transition(&mut sm, PipelineState::Recording)?;
let raw_text = match &source {
AsrSource::Text(text) => text.trim().to_string(),
AsrSource::Audio(_) | AsrSource::Owned(_) => {
let audio: &[AudioChunk] = match &source {
AsrSource::Audio(audio) => audio,
AsrSource::Owned(audio) => audio,
AsrSource::Text(_) => unreachable!("matched above"),
};
let transcript = tokio::select! {
biased;
_ = cancel.cancelled() => {
sm.cancel().ok();
sm.reset();
return Err(PipelineError::Cancelled { during: "recording" });
}
result = asr.transcribe(audio) => result?,
};
transcript.text.trim().to_string()
}
};
if raw_text.is_empty() {
sm.reset();
return Err(PipelineError::NoSpeech);
}
transition(&mut sm, PipelineState::Processing)?;
let ctx = match context {
Some(provider) => tokio::select! {
biased;
_ = cancel.cancelled() => {
sm.cancel().ok();
sm.reset();
return Err(PipelineError::Cancelled { during: "context" });
}
snapshot = provider.get_context() => snapshot,
},
None => ContextSnapshot::default(),
};
let mut diagnostics = Diagnostics {
speech_segments,
..Diagnostics::default()
};
let mut text = raw_text.clone();
for (index, f) in filters.iter().enumerate() {
match f.filter(&text).await {
Ok(result) => {
tracing::debug!(before = %text, after = %result.text, removed = ?result.removed, "filter applied");
text = result.text;
diagnostics.removed.extend(result.removed);
}
Err(e) => {
tracing::warn!(error = %e, index, "filter failed, continuing with unfiltered text");
diagnostics.failures.push(StageFailure {
stage: Stage::Filter,
index,
reason: e.to_string(),
});
}
}
}
for (index, p) in processors.iter().enumerate() {
match p.process(&text, &ctx).await {
Ok(result) => {
tracing::debug!(before = %text, after = %result.text, corrections = ?result.corrections, "processor applied");
text = result.text;
diagnostics.corrections.extend(result.corrections);
}
Err(e) => {
tracing::warn!(error = %e, index, "processor failed, continuing with unprocessed text");
diagnostics.failures.push(StageFailure {
stage: Stage::Processor,
index,
reason: e.to_string(),
});
}
}
}
let output = match refiner {
None => RefinementOutput::TextInsertion {
text: text.clone(),
formatting: None,
},
Some(refiner) => {
let input = RefinementInput {
raw_text: text.clone(),
context: ctx,
mode: RefinementMode::Dictation,
};
tokio::select! {
biased;
_ = cancel.cancelled() => {
sm.cancel().ok();
sm.reset();
return Err(PipelineError::Cancelled { during: "refinement" });
}
result = refiner.refine(input) => match result {
Ok(output) => output,
Err(e) => {
tracing::warn!(error = %e, "refinement failed, falling back to processed text");
diagnostics.failures.push(StageFailure {
stage: Stage::Refiner,
index: 0,
reason: e.to_string(),
});
RefinementOutput::TextInsertion { text: text.clone(), formatting: None }
}
},
}
}
};
transition(&mut sm, PipelineState::Emitting)?;
let emit_result = match emitter {
Some(emitter) => Some(emitter.emit(output.clone()).await),
None => None,
};
transition(&mut sm, PipelineState::Idle)?;
Ok(SessionResult {
raw_text,
output,
emit_result,
diagnostics,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::mock::*;
#[tokio::test]
async fn full_pipeline_happy_path() {
let emitter = MockEmitter::new();
let outputs = emitter.outputs();
let pipeline = Pipeline::builder()
.asr(MockAsr::new("hello world"))
.refiner(MockRefiner::uppercase())
.context(MockContextProvider::new())
.emitter(emitter)
.build()
.unwrap();
let session = pipeline.session();
session
.audio
.send(AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
})
.await
.unwrap();
let result = session.finish().await.unwrap();
assert_eq!(result.raw_text, "hello world");
assert!(result.emit_result.as_ref().unwrap().success);
let buf = outputs.lock().await;
assert_eq!(buf.len(), 1);
match &buf[0] {
RefinementOutput::TextInsertion { text, .. } => {
assert_eq!(text, "HELLO WORLD");
}
_ => panic!("expected TextInsertion"),
}
}
#[tokio::test]
async fn graceful_degradation_on_llm_failure() {
let emitter = MockEmitter::new();
let outputs = emitter.outputs();
let pipeline = Pipeline::builder()
.asr(MockAsr::new("raw dictation text"))
.refiner(MockRefiner::failing("API timeout"))
.context(MockContextProvider::new())
.emitter(emitter)
.build()
.unwrap();
let session = pipeline.session();
session
.audio
.send(AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
})
.await
.unwrap();
let result = session.finish().await.unwrap();
let buf = outputs.lock().await;
match &buf[0] {
RefinementOutput::TextInsertion { text, .. } => {
assert_eq!(text, "raw dictation text");
}
_ => panic!("expected TextInsertion fallback"),
}
assert!(result.emit_result.as_ref().unwrap().success);
}
#[tokio::test]
async fn cancellation_during_recording() {
let pipeline = Pipeline::builder()
.asr(MockAsr::new("will be cancelled"))
.refiner(MockRefiner::passthrough())
.context(MockContextProvider::new())
.emitter(MockEmitter::new())
.build()
.unwrap();
let session = pipeline.session();
session
.audio
.send(AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
})
.await
.unwrap();
session.cancel.cancel();
let err = session.finish().await.unwrap_err();
assert!(
matches!(err, PipelineError::Cancelled { .. }),
"expected a cancellation, got: {err}"
);
}
#[tokio::test]
async fn build_without_asr_fails() {
let Err(err) = Pipeline::builder().build() else {
panic!("a pipeline with no ASR adapter must not build");
};
assert!(matches!(err, PipelineError::MissingComponent("asr")));
}
#[tokio::test]
async fn build_requires_only_asr() {
let pipeline = Pipeline::builder()
.asr(MockAsr::new("hello world"))
.build()
.expect("an ASR adapter alone must be enough");
let result = pipeline
.transcribe(&[AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
}])
.await
.expect("session must run without refiner, context or emitter");
assert_eq!(result.text(), "hello world");
assert!(
result.emit_result.is_none(),
"no emitter configured, so nothing should have been emitted"
);
}
#[tokio::test]
async fn spec_minimal_llm_free_pipeline_runs() {
use crate::filter::SimpleFillerFilter;
use crate::processor::{BasicPunctuationRestorer, SelfCorrectionDetector};
let pipeline = Pipeline::builder()
.asr(MockAsr::new("um so i think it works"))
.filter(SimpleFillerFilter::english())
.processor(SelfCorrectionDetector::new())
.processor(BasicPunctuationRestorer)
.build()
.expect("the spec's minimal pipeline must build");
let result = pipeline
.transcribe(&[AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
}])
.await
.expect("the spec's minimal pipeline must run");
assert!(
!result.text().is_empty(),
"expected text out of the Tier 1+2 path"
);
assert!(
result.diagnostics.removed.iter().any(|r| r.contains("um")),
"the filler filter should have reported what it removed, got {:?}",
result.diagnostics.removed
);
assert!(
result.diagnostics.failures.is_empty(),
"no stage should have failed: {:?}",
result.diagnostics.failures
);
}
#[tokio::test]
async fn pipeline_with_filler_filter() {
use crate::filter::SimpleFillerFilter;
let emitter = MockEmitter::new();
let outputs = emitter.outputs();
let pipeline = Pipeline::builder()
.asr(MockAsr::new("um I think uh we should deploy"))
.filter(SimpleFillerFilter::english())
.refiner(MockRefiner::passthrough())
.context(MockContextProvider::new())
.emitter(emitter)
.build()
.unwrap();
let session = pipeline.session();
session
.audio
.send(AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
})
.await
.unwrap();
let result = session.finish().await.unwrap();
assert_eq!(result.raw_text, "um I think uh we should deploy");
let buf = outputs.lock().await;
match &buf[0] {
RefinementOutput::TextInsertion { text, .. } => {
assert_eq!(text, "I think we should deploy");
}
_ => panic!("expected TextInsertion"),
}
}
#[tokio::test]
async fn pipeline_with_filter_and_processor() {
use crate::filter::SimpleFillerFilter;
use crate::processor::{BasicPunctuationRestorer, SelfCorrectionDetector};
let emitter = MockEmitter::new();
let outputs = emitter.outputs();
let pipeline = Pipeline::builder()
.asr(MockAsr::new("um I want to go to Boston no wait to Denver"))
.filter(SimpleFillerFilter::english())
.processor(SelfCorrectionDetector::new())
.processor(BasicPunctuationRestorer)
.refiner(MockRefiner::passthrough())
.context(MockContextProvider::new())
.emitter(emitter)
.build()
.unwrap();
let session = pipeline.session();
session
.audio
.send(AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
})
.await
.unwrap();
let _result = session.finish().await.unwrap();
let buf = outputs.lock().await;
match &buf[0] {
RefinementOutput::TextInsertion { text, .. } => {
assert!(!text.contains("um"), "filler should be removed: {text}");
assert!(
!text.contains("Boston"),
"reparandum should be removed: {text}"
);
assert!(text.contains("Denver"), "repair should be kept: {text}");
assert!(
text.starts_with(|c: char| c.is_uppercase()),
"should be capitalized: {text}"
);
assert!(text.ends_with('.'), "should have terminal period: {text}");
}
_ => panic!("expected TextInsertion"),
}
}
}