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: &[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,
buffer: Vec<f32>,
generation: u64,
}
impl<T: Transcriber> Buffered<T> {
pub fn new(transcriber: T) -> Self {
Self {
inner: transcriber,
state: Mutex::new(BufferedState { active: false, buffer: 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.buffer.clear();
Ok(())
}
async fn feed(&self, samples: &[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.buffer.extend_from_slice(samples);
Ok(Vec::new())
}
async fn end_utterance(&self) -> Result<Vec<SttEvent>, SttError> {
let (samples, 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.buffer), st.generation)
};
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.buffer.clear();
st.generation = st.generation.wrapping_add(1);
}
}