use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use pipecrab_core::AudioFormat;
use pipecrab_runtime::MaybeSendSync;
use crate::{SttError, Transcriber};
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
pub trait StreamingTranscriber: MaybeSendSync {
fn input_format(&self) -> AudioFormat;
async fn begin_utterance(&self) -> Result<(), SttError>;
async fn feed(&self, samples: Arc<[f32]>) -> Result<Vec<SttEvent>, SttError>;
async fn end_utterance(&self) -> Result<Vec<SttEvent>, SttError>;
fn cancel(&self);
}
#[derive(Clone, Debug, PartialEq)]
pub enum SttEvent {
Partial {
text: Arc<str>,
stable: usize,
},
Final(Arc<str>),
Endpoint,
}
pub struct Buffered<T: Transcriber> {
inner: T,
state: Mutex<BufferedState>,
}
struct BufferedState {
active: bool,
chunks: Vec<Arc<[f32]>>,
generation: u64,
}
impl<T: Transcriber> Buffered<T> {
pub fn new(transcriber: T) -> Self {
Self {
inner: transcriber,
state: Mutex::new(BufferedState {
active: false,
chunks: Vec::new(),
generation: 0,
}),
}
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl<T: Transcriber> StreamingTranscriber for Buffered<T> {
fn input_format(&self) -> AudioFormat {
self.inner.input_format()
}
async fn begin_utterance(&self) -> Result<(), SttError> {
let mut st = self.state.lock().expect("Buffered state mutex poisoned");
if st.active {
return Err(SttError::Engine(
"Buffered::begin_utterance called while an utterance is already active".into(),
));
}
st.active = true;
st.chunks.clear();
Ok(())
}
async fn feed(&self, samples: Arc<[f32]>) -> Result<Vec<SttEvent>, SttError> {
let mut st = self.state.lock().expect("Buffered state mutex poisoned");
if !st.active {
return Err(SttError::Engine(
"Buffered::feed called without an active utterance".into(),
));
}
st.chunks.push(samples);
Ok(Vec::new())
}
async fn end_utterance(&self) -> Result<Vec<SttEvent>, SttError> {
let (mut chunks, generation) = {
let mut st = self.state.lock().expect("Buffered state mutex poisoned");
if !st.active {
return Err(SttError::Engine(
"Buffered::end_utterance called without a begin_utterance".into(),
));
}
st.active = false;
(std::mem::take(&mut st.chunks), st.generation)
};
let samples = if chunks.len() == 1 {
chunks.pop().expect("one buffered chunk")
} else {
let sample_count = chunks.iter().map(|chunk| chunk.len()).sum();
let mut samples = Vec::with_capacity(sample_count);
for chunk in chunks {
samples.extend_from_slice(&chunk);
}
Arc::from(samples)
};
let text = self.inner.transcribe(samples).await?;
let stale = {
let st = self.state.lock().expect("Buffered state mutex poisoned");
st.generation != generation
};
if stale {
return Ok(Vec::new());
}
Ok(vec![SttEvent::Final(text.into())])
}
fn cancel(&self) {
let mut st = self.state.lock().expect("Buffered state mutex poisoned");
st.active = false;
st.chunks.clear();
st.generation = st.generation.wrapping_add(1);
}
}