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::*;
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>>,
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,
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 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,
audio_channel_size: self.audio_channel_size,
})
}
}
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>>,
audio_channel_size: usize,
}
impl Pipeline {
pub fn builder() -> PipelineBuilder {
PipelineBuilder::new()
}
pub async fn transcribe(&self, audio: &[AudioChunk]) -> Result<SessionResult, PipelineError> {
run_session(
&self.asr,
&self.filters,
&self.processors,
self.refiner.as_deref(),
self.context.as_deref(),
self.emitter.as_deref(),
audio,
&CancellationToken::new(),
)
.await
}
pub fn session(&self) -> Session {
let (audio_tx, mut audio_rx) = mpsc::channel::<AudioChunk>(self.audio_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 cancel_inner = cancel.clone();
let handle = tokio::spawn(async move {
let mut chunks: Vec<AudioChunk> = Vec::new();
loop {
tokio::select! {
biased;
_ = cancel_inner.cancelled() => return Err(PipelineError::Cancelled { during: "recording" }),
maybe = audio_rx.recv() => match maybe {
Some(chunk) => chunks.push(chunk),
None => break,
},
}
}
run_session(
&asr,
&filters,
&processors,
refiner.as_deref(),
context.as_deref(),
emitter.as_deref(),
&chunks,
&cancel_inner,
)
.await
});
Session {
audio: audio_tx,
cancel,
handle,
}
}
}
pub struct Session {
pub audio: mpsc::Sender<AudioChunk>,
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>,
}
#[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 {
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>,
audio: &[AudioChunk],
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 transcript = tokio::select! {
biased;
_ = cancel.cancelled() => {
sm.cancel().ok();
sm.reset();
return Err(PipelineError::Cancelled { during: "recording" });
}
result = asr.transcribe(audio) => result?,
};
let raw_text = 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::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"),
}
}
}