refactor(dictation): extract SttEngine, add local Parakeet dictation pipeline

Phase 1: Extract reusable SttEngine from huddle/stt.rs into stt_engine.rs.
The core STT logic (resample 48→16 kHz, earshot VAD, sherpa-onnx Parakeet
inference) is now a standalone component configurable via SttEngineConfig.
huddle/stt.rs becomes a thin wrapper that passes huddle-specific flags
(TTS barge-in, PTT gating). drain_until_shutdown moves to stt_engine and
is re-exported by huddle/mod.rs for backward compat.

Phase 2: Add Tauri dictation commands (start_dictation, stop_dictation,
push_dictation_audio, get_dictation_status) that create a standalone
SttEngine instance with dictation-tuned settings (longer silence threshold,
no TTS/PTT flags). Transcribed text is emitted to the frontend via
'dictation-transcript' Tauri events.

Phase 3: Add useLocalDictation hook that captures mic audio via AudioWorklet
and sends raw PCM to the native STT engine. useDictation now routes to
local STT when available (offline, no API key), falling back to cloud
(OpenAI Realtime via relay) when the model isn't downloaded.

Key wins:
- Works fully offline — no BUZZ_OPENAI_API_KEY needed
- Self-hosters get dictation for free
- Lower latency (no network round-trip)
- No relay billing concern
- Cloud fallback preserved for higher accuracy
This commit is contained in:
klopez4212
2026-07-11 16:18:40 +01:00
parent 1dcfbf25a2
commit d3fd07b10d
8 changed files with 1042 additions and 531 deletions
+5 -5
View File
@@ -12,6 +12,7 @@ use tauri::{AppHandle, Manager};
#[cfg(feature = "mesh-llm")]
use tokio::sync::Mutex as AsyncMutex;
use crate::dictation::DictationState;
use crate::huddle::HuddleState;
use crate::managed_agents::config_bridge::SessionConfigCache;
use crate::managed_agents::ManagedAgentProcess;
@@ -25,11 +26,7 @@ pub struct AppState {
pub channel_templates_store_lock: Mutex<()>,
pub managed_agent_processes: Mutex<HashMap<String, ManagedAgentProcess>>,
pub huddle_state: Mutex<HuddleState>,
/// Tauri app handle — stored after setup so huddle commands can emit
/// `huddle-state-changed` events without needing the handle threaded
/// through every call site.
///
/// Set once during `setup()` in `lib.rs`; never cleared.
/// Tauri app handle — set once during `setup()`, never cleared.
pub app_handle: Mutex<Option<AppHandle>>,
/// Selected audio output device name. `None` = system default.
/// Used by `connect_audio_relay` and TTS pipeline when opening sinks.
@@ -81,6 +78,8 @@ pub struct AppState {
/// listener is up before any restore/create can request a connection.
#[cfg(feature = "mesh-llm")]
pub mesh_coordinator: AsyncMutex<Option<crate::mesh_llm::MeshCoordinator>>,
/// Local dictation state (Parakeet STT for composer voice input).
pub dictation_state: Mutex<DictationState>,
}
/// Parse the `BUZZ_PRIVATE_KEY` env var into identity keys. `Some` means the
@@ -146,6 +145,7 @@ pub fn build_app_state() -> AppState {
mesh_llm_runtime: AsyncMutex::new(None),
#[cfg(feature = "mesh-llm")]
mesh_coordinator: AsyncMutex::new(None),
dictation_state: Mutex::new(DictationState::new()),
}
}
+212
View File
@@ -0,0 +1,212 @@
//! Local dictation pipeline — uses the Parakeet STT engine for offline
//! speech-to-text in the message composer.
//!
//! Unlike the huddle STT pipeline (which posts kind:9 events to the relay),
//! dictation emits transcribed text back to the frontend via Tauri events
//! so the composer can display it in real-time.
//!
//! Key differences from huddle STT:
//! - No TTS barge-in / echo gating (no agent voice in composer context)
//! - No PTT (dictation uses a toggle button, not push-to-talk)
//! - Slightly longer silence threshold for more coherent sentences
//! - Text goes to the frontend, not to the relay
use std::sync::Arc;
use tauri::{Emitter, State};
use crate::app_state::AppState;
use crate::huddle::models;
use crate::stt_engine::{
SttEngine, SttEngineConfig, DEFAULT_MAX_SPEECH_SAMPLES, DICTATION_SILENCE_FLUSH_FRAMES,
};
/// Tauri event name emitted when a dictation transcript segment is ready.
const DICTATION_TRANSCRIPT_EVENT: &str = "dictation-transcript";
/// Tauri event name emitted when dictation state changes (started/stopped).
const DICTATION_STATE_EVENT: &str = "dictation-state";
/// State for the active dictation session.
///
/// Stored in `AppState` behind a `Mutex`. Only one dictation session can be
/// active at a time (starting a new one stops the previous).
pub(crate) struct DictationState {
/// The running STT engine, if dictation is active.
engine: Option<Arc<SttEngine>>,
}
impl DictationState {
pub fn new() -> Self {
Self { engine: None }
}
}
/// `start_dictation` — begin local STT dictation.
///
/// Starts the Parakeet STT engine and spawns a task that emits
/// `dictation-transcript` events to the frontend as text is recognized.
/// Returns an error if models are not downloaded yet.
#[tauri::command]
pub async fn start_dictation(state: State<'_, AppState>) -> Result<(), String> {
// Check if models are ready.
if !models::is_stt_ready() {
// Kick off download if not already in progress.
if let Some(mgr) = models::global_model_manager() {
mgr.start_stt_download(state.http_client.clone());
}
return Err("STT model not ready — download in progress".to_string());
}
let model_dir = models::stt_model_dir().ok_or("STT model directory not found")?;
// Stop any existing dictation session first.
stop_dictation_inner(&state);
let config = SttEngineConfig {
model_dir,
silence_flush_frames: DICTATION_SILENCE_FLUSH_FRAMES,
max_speech_samples: DEFAULT_MAX_SPEECH_SAMPLES,
tts_active: None,
tts_cancel: None,
ptt_active: None,
};
let (engine, text_rx) = SttEngine::new(config)?;
let engine = Arc::new(engine);
// Store the engine in state.
{
let mut ds = state
.dictation_state
.lock()
.unwrap_or_else(|e| e.into_inner());
ds.engine = Some(Arc::clone(&engine));
}
// Spawn a task that forwards transcribed text to the frontend.
let app_handle = state
.app_handle
.lock()
.unwrap_or_else(|e| e.into_inner())
.clone();
if let Some(handle) = app_handle {
let _ = handle.emit(DICTATION_STATE_EVENT, "started");
spawn_dictation_forwarder(text_rx, handle);
}
Ok(())
}
/// `stop_dictation` — stop the active dictation session.
#[tauri::command]
pub fn stop_dictation(state: State<'_, AppState>) -> Result<(), String> {
stop_dictation_inner(&state);
// Emit state change to frontend.
if let Some(handle) = state
.app_handle
.lock()
.unwrap_or_else(|e| e.into_inner())
.as_ref()
{
let _ = handle.emit(DICTATION_STATE_EVENT, "stopped");
}
Ok(())
}
/// `push_dictation_audio` — feed raw PCM bytes into the dictation pipeline.
///
/// Expects a raw binary body of f32 LE samples at 48 kHz mono.
/// If no dictation session is active, the bytes are silently discarded.
#[tauri::command]
pub fn push_dictation_audio(
request: tauri::ipc::Request<'_>,
state: State<'_, AppState>,
) -> Result<(), String> {
/// Maximum IPC audio batch size: 100 KB.
const MAX_AUDIO_BATCH_BYTES: usize = 100 * 1024;
match request.body() {
tauri::ipc::InvokeBody::Raw(bytes) => {
if bytes.len() > MAX_AUDIO_BATCH_BYTES {
return Err(format!(
"audio batch too large: {} bytes (max {})",
bytes.len(),
MAX_AUDIO_BATCH_BYTES
));
}
let ds = state
.dictation_state
.lock()
.unwrap_or_else(|e| e.into_inner());
if let Some(ref engine) = ds.engine {
engine.push_audio(bytes.to_vec())?;
}
Ok(())
}
_ => Err("expected raw binary body".to_string()),
}
}
/// `get_dictation_status` — check if local dictation is available and/or active.
#[tauri::command]
pub fn get_dictation_status(state: State<'_, AppState>) -> DictationStatus {
let model_ready = models::is_stt_ready();
let is_active = state
.dictation_state
.lock()
.unwrap_or_else(|e| e.into_inner())
.engine
.is_some();
DictationStatus {
available: model_ready,
active: is_active,
}
}
/// Response for `get_dictation_status`.
#[derive(serde::Serialize, Clone)]
pub struct DictationStatus {
/// Whether the local STT model is downloaded and ready.
pub available: bool,
/// Whether a dictation session is currently active.
pub active: bool,
}
// ── Internal helpers ──────────────────────────────────────────────────────────
fn stop_dictation_inner(state: &AppState) {
let old_engine = {
let mut ds = state
.dictation_state
.lock()
.unwrap_or_else(|e| e.into_inner());
ds.engine.take()
};
if let Some(engine) = old_engine {
engine.shutdown();
// Drop outside the lock — thread join may block briefly.
drop(engine);
}
}
/// Spawn an async task that reads transcribed text and emits Tauri events.
fn spawn_dictation_forwarder(
mut text_rx: tokio::sync::mpsc::Receiver<String>,
app_handle: tauri::AppHandle,
) {
tauri::async_runtime::spawn(async move {
while let Some(text) = text_rx.recv().await {
if text.is_empty() {
continue;
}
if app_handle.emit(DICTATION_TRANSCRIPT_EVENT, &text).is_err() {
break; // App window closed.
}
}
});
}
+2 -17
View File
@@ -40,23 +40,8 @@ pub mod wire;
// ── Shared utilities ──────────────────────────────────────────────────────────
/// Drain and discard all pending messages until shutdown or disconnect.
/// Shared by both the STT and TTS worker threads for graceful degradation
/// when model files are missing or initialization fails.
pub(super) fn drain_until_shutdown<T>(
rx: std::sync::mpsc::Receiver<T>,
shutdown: &std::sync::atomic::AtomicBool,
) {
loop {
if shutdown.load(std::sync::atomic::Ordering::Acquire) {
break;
}
match rx.recv_timeout(std::time::Duration::from_millis(100)) {
Ok(_) => continue,
Err(_) => break,
}
}
}
/// Re-export from `stt_engine` for backward compatibility with `tts.rs`.
pub(super) use crate::stt_engine::drain_until_shutdown;
// ── Re-exports ────────────────────────────────────────────────────────────────
+31 -503
View File
@@ -1,12 +1,14 @@
//! Speech-to-Text pipeline for huddle voice transcription.
//!
//! Mental model:
//! This is a thin wrapper around `crate::stt_engine::SttEngine` configured
//! with huddle-specific settings (TTS barge-in, PTT gating, huddle silence
//! threshold).
//!
//! ```text
//! AudioWorklet (48 kHz f32 PCM)
//! → push_audio_pcm (Tauri cmd)
//! → SttPipeline::push_audio [bounded sync_channel]
//! → stt_worker thread
//! → SttPipeline::push_audio
//! → SttEngine worker thread
//! rubato: 48 kHz → 16 kHz mono
//! earshot VAD: accumulate speech frames
//! sherpa-onnx Parakeet TDT-CTC 110M: transcribe on silence
@@ -14,57 +16,36 @@
//! → tokio task (start_stt_pipeline)
//! builds kind:9 event → relay
//! ```
//!
//! The worker runs on a dedicated `std::thread` (not async) because
//! sherpa-onnx is CPU-bound and not Send-safe across await points.
use std::{
path::PathBuf,
sync::{
atomic::{AtomicBool, Ordering},
mpsc::{self, Receiver, SyncSender},
Arc,
},
thread,
time::Duration,
sync::{atomic::AtomicBool, Arc},
};
use tokio::sync::mpsc as tokio_mpsc;
use crate::stt_engine::{
SttEngine, SttEngineConfig, DEFAULT_MAX_SPEECH_SAMPLES, DEFAULT_SILENCE_FLUSH_FRAMES,
};
// ── Public pipeline handle ────────────────────────────────────────────────────
/// Bounded audio queue capacity.
/// 100 ms batches at 48 kHz ≈ 19 KB each → 50 slots ≈ 5 s / ~1 MB max backlog.
const AUDIO_QUEUE_DEPTH: usize = 50;
/// Maximum speech buffer size: 30 seconds at 16 kHz.
/// Prevents OOM if VAD stays in speech mode (noisy environment).
const MAX_SPEECH_SAMPLES: usize = 16_000 * 30;
/// Handle to the running STT pipeline.
/// Handle to the running huddle STT pipeline.
///
/// Wraps `SttEngine` with huddle-specific construction (TTS flags, PTT).
/// Not Clone — wrap in `Arc` to share across threads.
///
/// The text receiver (`tokio::sync::mpsc::Receiver<String>`) is returned
/// separately from `new()` so the caller can move it directly into an async
/// task without holding a Mutex across await points.
#[derive(Debug)]
pub struct SttPipeline {
/// Send raw PCM bytes (f32 LE, 48 kHz mono) into the pipeline.
audio_tx: SyncSender<Vec<u8>>,
/// Signals the worker thread to stop.
shutdown: Arc<AtomicBool>,
/// Worker thread handle — taken on drop to join cleanly.
thread: Option<thread::JoinHandle<()>>,
engine: SttEngine,
}
impl SttPipeline {
/// Spawn the pipeline thread.
/// Spawn the huddle STT pipeline.
///
/// `tts_active` is a shared flag set by the TTS pipeline while audio is
/// playing. The STT worker uses it to:
/// - discard accumulated speech (echo prevention / barge-in gating)
/// - apply a 200 ms cooldown after TTS stops before re-enabling STT
/// - apply a cooldown after TTS stops before re-enabling STT
/// - detect barge-in: speech onset during TTS → set `tts_cancel`
///
/// `tts_cancel` (optional) is the TTS pipeline's cancel flag. When the STT
@@ -76,492 +57,39 @@ impl SttPipeline {
/// When `None`, the pipeline runs in continuous VAD mode.
///
/// Returns `Err` only if the thread cannot be spawned (OS error).
/// If model files are missing, the worker logs and exits cleanly —
/// the pipeline handle is still returned but will never produce text.
///
/// The `tokio::sync::mpsc::Receiver<String>` is returned separately so the
/// caller can move it directly into an async task. This avoids holding a
/// `Mutex<Receiver>` across await points (which would block a Tokio worker
/// thread on every `recv_timeout` call).
pub fn new(
model_dir: PathBuf,
tts_active: Arc<AtomicBool>,
tts_cancel: Option<Arc<AtomicBool>>,
ptt_active: Option<Arc<AtomicBool>>,
) -> Result<(Self, tokio_mpsc::Receiver<String>), String> {
let (audio_tx, audio_rx) = mpsc::sync_channel::<Vec<u8>>(AUDIO_QUEUE_DEPTH);
let (text_tx, text_rx) = tokio_mpsc::channel::<String>(64);
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_worker = Arc::clone(&shutdown);
let tts_cancel_worker = tts_cancel.as_ref().map(Arc::clone);
let ptt_active_worker = ptt_active.as_ref().map(Arc::clone);
let handle = thread::Builder::new()
.name("stt-worker".into())
.spawn(move || {
stt_worker(
model_dir,
audio_rx,
text_tx,
shutdown_worker,
tts_active,
tts_cancel_worker,
ptt_active_worker,
)
})
.map_err(|e| format!("failed to spawn stt-worker thread: {e}"))?;
let pipeline = Self {
audio_tx,
shutdown,
thread: Some(handle),
let config = SttEngineConfig {
model_dir,
silence_flush_frames: DEFAULT_SILENCE_FLUSH_FRAMES,
max_speech_samples: DEFAULT_MAX_SPEECH_SAMPLES,
tts_active: Some(tts_active),
tts_cancel,
ptt_active,
};
Ok((pipeline, text_rx))
let (engine, text_rx) = SttEngine::new(config)?;
Ok((Self { engine }, text_rx))
}
/// Signal the worker thread to stop.
pub fn shutdown(&self) {
self.shutdown.store(true, Ordering::Release);
self.engine.shutdown();
}
/// Returns `true` if the worker thread has exited (init failure, crash, or normal exit).
/// Used by hot-start to detect dead pipelines and clear them for retry.
/// Returns `true` if the worker thread has exited.
pub fn is_finished(&self) -> bool {
self.thread.as_ref().is_none_or(|h| h.is_finished())
self.engine.is_finished()
}
/// Feed raw PCM bytes into the pipeline.
/// Feed raw PCM bytes (f32 LE, 48 kHz mono) into the pipeline.
///
/// Non-blocking. Drops audio silently if the pipeline can't keep up —
/// better to lose frames than to stall the UI thread.
/// Non-blocking. Drops audio silently if the pipeline can't keep up.
pub fn push_audio(&self, pcm_bytes: Vec<u8>) -> Result<(), String> {
// Reject non-4-byte-aligned input — would silently truncate in bytes_to_f32.
if !pcm_bytes.len().is_multiple_of(4) {
return Err(format!(
"audio input not 4-byte aligned ({} bytes) — expected f32 LE samples",
pcm_bytes.len()
));
}
// Drop audio if the pipeline can't keep up — better than blocking the UI.
let _ = self.audio_tx.try_send(pcm_bytes);
Ok(())
self.engine.push_audio(pcm_bytes)
}
}
impl Drop for SttPipeline {
fn drop(&mut self) {
// Signal the worker to stop.
self.shutdown.store(true, Ordering::Release);
// Dropping `audio_tx` (implicitly when self is dropped after this fn)
// unblocks the worker's recv_timeout loop. Join to ensure clean exit.
if let Some(thread) = self.thread.take() {
let _ = thread.join();
}
}
}
// ── Worker thread ─────────────────────────────────────────────────────────────
/// How many 16 kHz samples of silence before we flush to STT.
/// 300 ms × 16 000 Hz / 256 samples-per-frame ≈ 19 frames.
/// Previous value (28 frames / 450 ms) felt sluggish in conversation.
const SILENCE_FLUSH_FRAMES: usize = 19;
/// Consecutive VAD speech frames required before triggering barge-in during TTS.
/// 20 frames × 256 samples / 16 kHz ≈ 320 ms — must be long enough to filter
/// speaker-to-mic feedback (TTS audio bleeding through the mic) while still
/// catching real human interruptions. 80 ms (previous: 5 frames) was too
/// aggressive — laptop speakers without headphones triggered false barge-in
/// within the first word of TTS playback.
const BARGE_IN_DEBOUNCE_FRAMES: usize = 20;
/// earshot requires exactly 256 samples per frame at 16 kHz.
const VAD_FRAME_SAMPLES: usize = 256;
/// VAD probability threshold — above this is considered speech.
const VAD_THRESHOLD: f32 = 0.5;
/// How long the worker waits on the audio channel before checking the shutdown flag.
const RECV_TIMEOUT: Duration = Duration::from_millis(50);
/// 50 ms cooldown after TTS stops before STT re-enables.
/// Prevents the tail of TTS audio from being transcribed as speech.
/// Previous value (200 ms) was eating the first word when the user spoke
/// immediately after the agent finished.
const TTS_COOLDOWN: Duration = Duration::from_millis(50);
/// Number of ONNX Runtime intra-op threads used by the offline recognizer.
///
/// Held at 1 (conservative) until we have a local A/B on real huddle audio.
/// Sherpa-onnx's Parakeet example uses 2 and most published RTF numbers are
/// at 2 threads on x86_64 server class hardware, but the encoder runs only
/// on VAD chunk boundaries on a dedicated thread, so the threading knob
/// trades worker latency against potential oversubscription with the audio
/// worklet on small Macs (4-core Intel especially). Bump to 2 once the A/B
/// shows it's safe on the minimum-spec target.
const STT_NUM_THREADS: i32 = 1;
fn stt_worker(
model_dir: PathBuf,
audio_rx: Receiver<Vec<u8>>,
text_tx: tokio_mpsc::Sender<String>,
shutdown: Arc<AtomicBool>,
tts_active: Arc<AtomicBool>,
tts_cancel: Option<Arc<AtomicBool>>,
ptt_active: Option<Arc<AtomicBool>>,
) {
// ── 1. Initialise rubato resampler (48 kHz → 16 kHz, mono) ───────────────
use rubato::{Fft, FixedSync, Resampler};
let mut resampler = match Fft::<f32>::new(48_000, 16_000, 1024, 2, 1, FixedSync::Input) {
Ok(r) => r,
Err(e) => {
eprintln!("buzz-desktop: STT resampler init failed: {e}");
return;
}
};
let chunk_in = resampler.input_frames_next();
// ── 2. Initialise earshot VAD ─────────────────────────────────────────────
use earshot::{DefaultPredictor, Detector};
let mut vad = Detector::new(DefaultPredictor::new());
// ── 3. Initialise sherpa-onnx recognizer ─────────────────────────────────
//
// Parakeet TDT-CTC 110M ships as a single `model.int8.onnx` (CTC head) plus
// `tokens.txt`. sherpa-onnx infers the model family from which inner config
// has a `model` path set, so we don't need to set `model_type` explicitly.
// (See rust-api-examples/parakeet_tdt_ctc_simulate_streaming_microphone.rs
// in k2-fsa/sherpa-onnx.)
use sherpa_onnx::{OfflineRecognizer, OfflineRecognizerConfig};
let tokens_path = model_dir.join("tokens.txt");
let model_path = model_dir.join("model.int8.onnx");
if !tokens_path.exists() || !model_path.exists() {
eprintln!(
"buzz-desktop: STT model not found at {} — STT disabled",
model_dir.display()
);
drain_until_shutdown(audio_rx, &shutdown);
return;
}
let mut cfg = OfflineRecognizerConfig::default();
cfg.model_config.nemo_ctc.model = Some(model_path.to_string_lossy().into_owned());
cfg.model_config.tokens = Some(tokens_path.to_string_lossy().into_owned());
cfg.model_config.num_threads = STT_NUM_THREADS;
// Explicit — defaults are not part of the API contract, and noisy debug
// logging in release builds would be expensive on every VAD chunk.
cfg.model_config.debug = false;
let recognizer = match OfflineRecognizer::create(&cfg) {
Some(r) => r,
None => {
eprintln!("buzz-desktop: OfflineRecognizer::create returned None — STT disabled");
drain_until_shutdown(audio_rx, &shutdown);
return;
}
};
// ── 4. Processing state ───────────────────────────────────────────────────
// Leftover 48 kHz samples that didn't fill a full resampler chunk.
let mut input_buf_48k: Vec<f32> = Vec::with_capacity(chunk_in * 2);
// Leftover 16 kHz samples that didn't fill a full VAD frame.
let mut leftover_16k: Vec<f32> = Vec::new();
// Accumulated speech frames (16 kHz).
let mut speech_buf: Vec<f32> = Vec::new();
// Consecutive silence frame count.
let mut silence_frames: usize = 0;
// Whether we're currently in a speech segment.
let mut in_speech = false;
// Consecutive speech frames seen during TTS — used for barge-in debounce.
let mut barge_in_frames: usize = 0;
// Timestamp when TTS last stopped — used for the 200 ms cooldown.
let mut tts_stopped_at: Option<std::time::Instant> = None;
// ── 5. Main loop ──────────────────────────────────────────────────────────
let mut tts_was_active = false;
let mut ptt_was_active = ptt_active
.as_ref()
.is_some_and(|p| p.load(Ordering::Acquire));
loop {
// Check shutdown flag before blocking.
if shutdown.load(Ordering::Acquire) {
break;
}
// Track TTS transitions to set the cooldown timer.
let tts_now = tts_active.load(Ordering::Acquire);
if tts_was_active && !tts_now {
// TTS just stopped — record the timestamp for the cooldown window.
tts_stopped_at = Some(std::time::Instant::now());
}
tts_was_active = tts_now;
// Track PTT transitions — flush accumulated speech when key is released.
// The worklet stops sending frames when PTT is inactive, so the normal
// silence-accumulation flush path never runs. We must flush here on the
// active→inactive edge to avoid buffering speech across PTT presses.
if let Some(ref ptt) = ptt_active {
let ptt_now = ptt.load(Ordering::Acquire);
if ptt_was_active && !ptt_now && in_speech && !speech_buf.is_empty() {
flush_to_stt(&speech_buf, &recognizer, &text_tx);
speech_buf.clear();
silence_frames = 0;
in_speech = false;
}
ptt_was_active = ptt_now;
}
// Use recv_timeout so we can periodically check the shutdown flag.
let bytes = match audio_rx.recv_timeout(RECV_TIMEOUT) {
Ok(b) => b,
Err(mpsc::RecvTimeoutError::Timeout) => continue,
Err(mpsc::RecvTimeoutError::Disconnected) => break, // Sender dropped.
};
// Drain any additional pending messages to batch-process.
let mut batch = vec![bytes];
while let Ok(b) = audio_rx.try_recv() {
batch.push(b);
}
for bytes in batch {
// Convert raw bytes to f32 samples (little-endian).
let samples_48k = bytes_to_f32(&bytes);
input_buf_48k.extend_from_slice(&samples_48k);
// Resample in chunk_in-sized blocks.
while input_buf_48k.len() >= chunk_in {
let chunk: Vec<f32> = input_buf_48k.drain(..chunk_in).collect();
let resampled = resample_chunk(&mut resampler, &chunk);
process_16k_samples(
&resampled,
&mut leftover_16k,
&mut vad,
&mut speech_buf,
&mut silence_frames,
&mut in_speech,
&mut barge_in_frames,
&recognizer,
&text_tx,
&tts_active,
tts_cancel.as_deref(),
&mut tts_stopped_at,
ptt_active.as_ref(),
);
}
}
}
// No final flush — leave_huddle/end_huddle emit lifecycle events before
// the STT worker exits, so a final flush would post a kind:9 message AFTER
// the user has "left." Losing the last partial utterance is acceptable.
}
/// Resample a mono 48 kHz chunk to 16 kHz using rubato.
/// Returns the resampled samples (may be empty on error).
fn resample_chunk(resampler: &mut rubato::Fft<f32>, chunk_48k: &[f32]) -> Vec<f32> {
use audioadapter_buffers::direct::InterleavedSlice;
use rubato::Resampler;
// rubato expects interleaved layout even for mono.
let input = match InterleavedSlice::new(chunk_48k, 1, chunk_48k.len()) {
Ok(a) => a,
Err(e) => {
eprintln!("buzz-desktop: STT resample input error: {e}");
return Vec::new();
}
};
match resampler.process(&input, 0, None) {
Ok(out) => out.take_data(),
Err(e) => {
eprintln!("buzz-desktop: STT resample error: {e}");
Vec::new()
}
}
}
/// Feed 16 kHz samples through the VAD and accumulate speech.
/// Flushes to STT when silence exceeds threshold.
///
/// When `tts_active` is set:
/// - In PTT mode: skip accumulation (PTT press handles TTS cancellation).
/// - In VAD mode: speech onset triggers barge-in via `tts_cancel`.
/// - After TTS stops, a cooldown prevents tail audio from being transcribed.
///
/// When `ptt_active` is `Some`:
/// - VAD `is_speech` is ANDed with the PTT flag — when the key is released,
/// `is_speech` becomes false, silence_frames accumulates, and the existing
/// flush logic kicks in naturally. The 200 ms release delay + ~300 ms
/// silence flush gives a natural utterance tail.
#[allow(clippy::too_many_arguments)]
fn process_16k_samples(
samples: &[f32],
leftover: &mut Vec<f32>,
vad: &mut earshot::Detector<earshot::DefaultPredictor>,
speech_buf: &mut Vec<f32>,
silence_frames: &mut usize,
in_speech: &mut bool,
barge_in_frames: &mut usize,
recognizer: &sherpa_onnx::OfflineRecognizer,
text_tx: &tokio_mpsc::Sender<String>,
tts_active: &Arc<AtomicBool>,
tts_cancel: Option<&AtomicBool>,
tts_stopped_at: &mut Option<std::time::Instant>,
ptt_active: Option<&Arc<AtomicBool>>,
) {
leftover.extend_from_slice(samples);
while leftover.len() >= VAD_FRAME_SAMPLES {
let frame: Vec<f32> = leftover.drain(..VAD_FRAME_SAMPLES).collect();
let clamped: Vec<f32> = frame.iter().map(|&s| s.clamp(-1.0, 1.0)).collect();
let prob = vad.predict_f32(&clamped);
let is_speech = prob > VAD_THRESHOLD;
// PTT gating: when PTT key is not held, treat as silence.
// This causes natural flush when the key is released — silence_frames
// accumulates and the existing flush logic kicks in after
// SILENCE_FLUSH_FRAMES. The 200 ms release delay + ~300 ms silence
// flush gives a natural utterance tail.
let is_speech = if let Some(ptt) = ptt_active {
is_speech && ptt.load(Ordering::Acquire)
} else {
is_speech
};
let tts_playing = tts_active.load(Ordering::Acquire);
// While TTS is playing: skip accumulation (echo prevention).
if tts_playing {
if ptt_active.is_some() {
// PTT mode — PTT press handles TTS cancellation directly
// (via the global shortcut handler). Just skip accumulation.
*in_speech = false;
*barge_in_frames = 0;
speech_buf.clear();
*silence_frames = 0;
continue;
}
// VAD mode — barge-in detection.
// Without acoustic echo cancellation, this requires a longer
// debounce (BARGE_IN_DEBOUNCE_FRAMES ≈ 320 ms) to filter
// speaker-to-mic feedback.
if is_speech {
*barge_in_frames += 1;
if *barge_in_frames >= BARGE_IN_DEBOUNCE_FRAMES {
// Real speech detected during TTS — trigger barge-in.
if let Some(cancel) = tts_cancel {
cancel.store(true, Ordering::Release);
}
*barge_in_frames = 0;
}
} else {
*barge_in_frames = 0;
}
// Don't accumulate speech during TTS (echo prevention).
*in_speech = false;
speech_buf.clear();
*silence_frames = 0;
continue;
}
// TTS not playing — check cooldown window.
if let Some(stopped) = *tts_stopped_at {
if stopped.elapsed() < TTS_COOLDOWN {
// Still in cooldown — discard but keep tracking speech state.
if !is_speech {
*in_speech = false;
}
speech_buf.clear();
*silence_frames = 0;
*barge_in_frames = 0;
continue;
} else {
// Cooldown expired — clear the timer and reset all segment state.
*tts_stopped_at = None;
*in_speech = false;
*silence_frames = 0;
*barge_in_frames = 0;
}
}
if is_speech {
*silence_frames = 0;
*in_speech = true;
speech_buf.extend_from_slice(&frame);
// OOM guard: flush and reset if the buffer exceeds 30 s of audio.
if speech_buf.len() >= MAX_SPEECH_SAMPLES {
flush_to_stt(speech_buf, recognizer, text_tx);
speech_buf.clear();
*silence_frames = 0;
*in_speech = false;
}
} else if *in_speech {
// Still accumulate during brief silence gaps.
speech_buf.extend_from_slice(&frame);
*silence_frames += 1;
// In PTT mode, don't flush on silence — accumulate the entire
// key-hold as one utterance. The PTT release edge in the main
// loop handles the flush. In VAD mode, flush after the silence
// threshold so each natural pause becomes a separate message.
if ptt_active.is_none() && *silence_frames >= SILENCE_FLUSH_FRAMES {
// End of utterance — transcribe.
flush_to_stt(speech_buf, recognizer, text_tx);
speech_buf.clear();
*silence_frames = 0;
*in_speech = false;
}
}
// If not in speech and not accumulating, just discard the frame.
}
}
/// Run sherpa-onnx on the accumulated speech buffer and send the text.
///
/// Uses `blocking_send` because this runs on a `std::thread` (not async).
/// The tokio channel's `blocking_send` is safe to call from sync contexts.
fn flush_to_stt(
speech_buf: &[f32],
recognizer: &sherpa_onnx::OfflineRecognizer,
text_tx: &tokio_mpsc::Sender<String>,
) {
if speech_buf.is_empty() {
return;
}
let stream = recognizer.create_stream();
stream.accept_waveform(16_000, speech_buf);
recognizer.decode(&stream);
let text = stream
.get_result()
.map(|r| r.text.trim().to_string())
.unwrap_or_default();
if !text.is_empty() {
if let Err(e) = text_tx.blocking_send(text) {
eprintln!("buzz-desktop: STT text channel closed: {e}");
}
}
}
/// Convert raw bytes (f32 LE) to f32 samples.
/// Caller should ensure `bytes.len() % 4 == 0`; extra bytes are silently truncated.
///
/// Assumes little-endian — matches all current Tauri targets (macOS ARM64,
/// Windows/Linux x86). The JS AudioWorklet's Float32Array uses platform-native
/// byte order, which is LE on all supported platforms.
fn bytes_to_f32(bytes: &[u8]) -> Vec<f32> {
bytes
.chunks_exact(4)
.map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.collect()
}
// drain_until_shutdown lives in super (huddle/mod.rs) — shared with tts.rs.
use super::drain_until_shutdown;
+6
View File
@@ -2,6 +2,7 @@ mod app_state;
mod archive;
mod commands;
mod deep_link;
mod dictation;
mod event_sync;
mod events;
mod huddle;
@@ -20,6 +21,7 @@ mod ptt_shortcut;
mod relay;
mod secret_store;
mod shutdown;
mod stt_engine;
mod templates;
mod util;
@@ -757,6 +759,10 @@ pub fn run() {
archive::read_unindexed_observer_rows,
is_auto_update_supported,
set_window_vibrancy,
dictation::start_dictation,
dictation::stop_dictation,
dictation::push_dictation_audio,
dictation::get_dictation_status,
])
.build(tauri::generate_context!())
.expect("error while building tauri application");
+506
View File
@@ -0,0 +1,506 @@
//! Reusable Speech-to-Text engine backed by sherpa-onnx (Parakeet TDT-CTC 110M).
//!
//! This module extracts the core STT logic — resample, VAD, inference — into a
//! standalone component that can be instantiated by both the huddle pipeline and
//! the composer dictation feature with different configurations.
//!
//! ```text
//! Caller (48 kHz f32 PCM)
//! → SttEngine::push_audio [bounded sync_channel]
//! → stt_worker thread
//! rubato: 48 kHz → 16 kHz mono
//! earshot VAD: accumulate speech frames
//! sherpa-onnx Parakeet TDT-CTC 110M: transcribe on silence
//! → text_rx [tokio mpsc channel]
//! → caller's async task
//! ```
//!
//! The worker runs on a dedicated `std::thread` (not async) because
//! sherpa-onnx is CPU-bound and not Send-safe across await points.
use std::{
path::PathBuf,
sync::{
atomic::{AtomicBool, Ordering},
mpsc::{self, Receiver, SyncSender},
Arc,
},
thread,
time::Duration,
};
use tokio::sync::mpsc as tokio_mpsc;
// ── Configuration ─────────────────────────────────────────────────────────────
/// Default silence frames before flush (~300ms at 16 kHz / 256 samples per frame).
/// Used by the huddle pipeline.
pub const DEFAULT_SILENCE_FLUSH_FRAMES: usize = 19;
/// Silence frames for dictation (~400ms) — slightly longer for more coherent sentences.
pub const DICTATION_SILENCE_FLUSH_FRAMES: usize = 25;
/// Default maximum speech buffer: 30 seconds at 16 kHz.
pub const DEFAULT_MAX_SPEECH_SAMPLES: usize = 16_000 * 30;
/// Configuration for the STT engine.
///
/// Allows callers to tune VAD behavior and optionally wire in TTS/PTT flags
/// for huddle-specific features (barge-in, echo gating, push-to-talk).
#[derive(Clone)]
pub struct SttEngineConfig {
/// Path to the directory containing `model.int8.onnx` and `tokens.txt`.
pub model_dir: PathBuf,
/// Number of consecutive silence frames before flushing to STT.
/// ~300ms = 19 frames, ~400ms = 25 frames at 16 kHz / 256 samples per frame.
pub silence_flush_frames: usize,
/// Maximum speech buffer size in samples (OOM guard).
pub max_speech_samples: usize,
/// Optional: shared flag set by TTS while audio is playing (echo prevention).
pub tts_active: Option<Arc<AtomicBool>>,
/// Optional: TTS cancel flag — set by STT on barge-in detection.
pub tts_cancel: Option<Arc<AtomicBool>>,
/// Optional: push-to-talk flag — when `Some`, speech is only accumulated
/// while the flag is true.
pub ptt_active: Option<Arc<AtomicBool>>,
}
// ── Public engine handle ──────────────────────────────────────────────────────
/// Bounded audio queue capacity.
/// 100 ms batches at 48 kHz ≈ 19 KB each → 50 slots ≈ 5 s / ~1 MB max backlog.
const AUDIO_QUEUE_DEPTH: usize = 50;
/// Handle to the running STT engine.
///
/// Not Clone — wrap in `Arc` to share across threads.
///
/// The text receiver (`tokio::sync::mpsc::Receiver<String>`) is returned
/// separately from `new()` so the caller can move it directly into an async
/// task without holding a Mutex across await points.
#[derive(Debug)]
pub(crate) struct SttEngine {
/// Send raw PCM bytes (f32 LE, 48 kHz mono) into the engine.
audio_tx: SyncSender<Vec<u8>>,
/// Signals the worker thread to stop.
shutdown: Arc<AtomicBool>,
/// Worker thread handle — taken on drop to join cleanly.
thread: Option<thread::JoinHandle<()>>,
}
impl SttEngine {
/// Spawn the engine worker thread.
///
/// Returns `(Self, Receiver<String>)`. The receiver yields transcribed text
/// segments. It is returned separately so the caller can move it into an
/// async task without holding a Mutex.
///
/// Returns `Err` only if the thread cannot be spawned (OS error).
/// If model files are missing, the worker logs and exits cleanly —
/// the engine handle is still returned but will never produce text.
pub fn new(config: SttEngineConfig) -> Result<(Self, tokio_mpsc::Receiver<String>), String> {
let (audio_tx, audio_rx) = mpsc::sync_channel::<Vec<u8>>(AUDIO_QUEUE_DEPTH);
let (text_tx, text_rx) = tokio_mpsc::channel::<String>(64);
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_worker = Arc::clone(&shutdown);
let handle = thread::Builder::new()
.name("stt-engine".into())
.spawn(move || {
stt_worker(config, audio_rx, text_tx, shutdown_worker);
})
.map_err(|e| format!("failed to spawn stt-engine thread: {e}"))?;
let engine = Self {
audio_tx,
shutdown,
thread: Some(handle),
};
Ok((engine, text_rx))
}
/// Signal the worker thread to stop.
pub fn shutdown(&self) {
self.shutdown.store(true, Ordering::Release);
}
/// Returns `true` if the worker thread has exited (init failure, crash, or normal exit).
pub fn is_finished(&self) -> bool {
self.thread.as_ref().is_none_or(|h| h.is_finished())
}
/// Feed raw PCM bytes (f32 LE, 48 kHz mono) into the engine.
///
/// Non-blocking. Drops audio silently if the engine can't keep up —
/// better to lose frames than to stall the caller.
pub fn push_audio(&self, pcm_bytes: Vec<u8>) -> Result<(), String> {
if !pcm_bytes.len().is_multiple_of(4) {
return Err(format!(
"audio input not 4-byte aligned ({} bytes) — expected f32 LE samples",
pcm_bytes.len()
));
}
let _ = self.audio_tx.try_send(pcm_bytes);
Ok(())
}
}
impl Drop for SttEngine {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::Release);
if let Some(thread) = self.thread.take() {
let _ = thread.join();
}
}
}
// ── Shared utility ────────────────────────────────────────────────────────────
/// Drain and discard all pending messages until shutdown or disconnect.
///
/// Used by both the STT and TTS worker threads for graceful degradation
/// when model files are missing or initialization fails.
pub(crate) fn drain_until_shutdown<T>(rx: std::sync::mpsc::Receiver<T>, shutdown: &AtomicBool) {
loop {
if shutdown.load(Ordering::Acquire) {
break;
}
match rx.recv_timeout(Duration::from_millis(100)) {
Ok(_) => continue,
Err(_) => break,
}
}
}
// ── Worker thread ─────────────────────────────────────────────────────────────
/// Consecutive VAD speech frames required before triggering barge-in during TTS.
/// 20 frames × 256 samples / 16 kHz ≈ 320 ms — filters speaker-to-mic feedback.
const BARGE_IN_DEBOUNCE_FRAMES: usize = 20;
/// earshot requires exactly 256 samples per frame at 16 kHz.
const VAD_FRAME_SAMPLES: usize = 256;
/// VAD probability threshold — above this is considered speech.
const VAD_THRESHOLD: f32 = 0.5;
/// How long the worker waits on the audio channel before checking the shutdown flag.
const RECV_TIMEOUT: Duration = Duration::from_millis(50);
/// Cooldown after TTS stops before STT re-enables.
/// Prevents the tail of TTS audio from being transcribed as speech.
const TTS_COOLDOWN: Duration = Duration::from_millis(50);
/// Number of ONNX Runtime intra-op threads used by the offline recognizer.
const STT_NUM_THREADS: i32 = 1;
fn stt_worker(
config: SttEngineConfig,
audio_rx: Receiver<Vec<u8>>,
text_tx: tokio_mpsc::Sender<String>,
shutdown: Arc<AtomicBool>,
) {
// ── 1. Initialise rubato resampler (48 kHz → 16 kHz, mono) ───────────────
use rubato::{Fft, FixedSync, Resampler};
let mut resampler = match Fft::<f32>::new(48_000, 16_000, 1024, 2, 1, FixedSync::Input) {
Ok(r) => r,
Err(e) => {
eprintln!("buzz-desktop: STT resampler init failed: {e}");
return;
}
};
let chunk_in = resampler.input_frames_next();
// ── 2. Initialise earshot VAD ─────────────────────────────────────────────
use earshot::{DefaultPredictor, Detector};
let mut vad = Detector::new(DefaultPredictor::new());
// ── 3. Initialise sherpa-onnx recognizer ─────────────────────────────────
use sherpa_onnx::{OfflineRecognizer, OfflineRecognizerConfig};
let tokens_path = config.model_dir.join("tokens.txt");
let model_path = config.model_dir.join("model.int8.onnx");
if !tokens_path.exists() || !model_path.exists() {
eprintln!(
"buzz-desktop: STT model not found at {} — STT disabled",
config.model_dir.display()
);
drain_until_shutdown(audio_rx, &shutdown);
return;
}
let mut cfg = OfflineRecognizerConfig::default();
cfg.model_config.nemo_ctc.model = Some(model_path.to_string_lossy().into_owned());
cfg.model_config.tokens = Some(tokens_path.to_string_lossy().into_owned());
cfg.model_config.num_threads = STT_NUM_THREADS;
cfg.model_config.debug = false;
let recognizer = match OfflineRecognizer::create(&cfg) {
Some(r) => r,
None => {
eprintln!("buzz-desktop: OfflineRecognizer::create returned None — STT disabled");
drain_until_shutdown(audio_rx, &shutdown);
return;
}
};
// ── 4. Processing state ───────────────────────────────────────────────────
let mut input_buf_48k: Vec<f32> = Vec::with_capacity(chunk_in * 2);
let mut leftover_16k: Vec<f32> = Vec::new();
let mut speech_buf: Vec<f32> = Vec::new();
let mut silence_frames: usize = 0;
let mut in_speech = false;
let mut barge_in_frames: usize = 0;
let mut tts_stopped_at: Option<std::time::Instant> = None;
// ── 5. Main loop ──────────────────────────────────────────────────────────
let has_tts = config.tts_active.is_some();
let tts_active_flag = config
.tts_active
.unwrap_or_else(|| Arc::new(AtomicBool::new(false)));
let tts_cancel_flag = config.tts_cancel;
let ptt_active_flag = config.ptt_active;
let mut tts_was_active = false;
let mut ptt_was_active = ptt_active_flag
.as_ref()
.is_some_and(|p| p.load(Ordering::Acquire));
loop {
if shutdown.load(Ordering::Acquire) {
break;
}
// Track TTS transitions (only relevant when TTS flags are wired).
if has_tts {
let tts_now = tts_active_flag.load(Ordering::Acquire);
if tts_was_active && !tts_now {
tts_stopped_at = Some(std::time::Instant::now());
}
tts_was_active = tts_now;
}
// Track PTT transitions — flush on key release.
if let Some(ref ptt) = ptt_active_flag {
let ptt_now = ptt.load(Ordering::Acquire);
if ptt_was_active && !ptt_now && in_speech && !speech_buf.is_empty() {
flush_to_stt(&speech_buf, &recognizer, &text_tx);
speech_buf.clear();
silence_frames = 0;
in_speech = false;
}
ptt_was_active = ptt_now;
}
let bytes = match audio_rx.recv_timeout(RECV_TIMEOUT) {
Ok(b) => b,
Err(mpsc::RecvTimeoutError::Timeout) => continue,
Err(mpsc::RecvTimeoutError::Disconnected) => break,
};
// Drain any additional pending messages to batch-process.
let mut batch = vec![bytes];
while let Ok(b) = audio_rx.try_recv() {
batch.push(b);
}
for bytes in batch {
let samples_48k = bytes_to_f32(&bytes);
input_buf_48k.extend_from_slice(&samples_48k);
while input_buf_48k.len() >= chunk_in {
let chunk: Vec<f32> = input_buf_48k.drain(..chunk_in).collect();
let resampled = resample_chunk(&mut resampler, &chunk);
process_16k_samples(
&resampled,
&mut leftover_16k,
&mut vad,
&mut speech_buf,
&mut silence_frames,
&mut in_speech,
&mut barge_in_frames,
&recognizer,
&text_tx,
has_tts,
&tts_active_flag,
tts_cancel_flag.as_deref(),
&mut tts_stopped_at,
ptt_active_flag.as_ref(),
config.silence_flush_frames,
config.max_speech_samples,
);
}
}
}
}
/// Resample a mono 48 kHz chunk to 16 kHz using rubato.
fn resample_chunk(resampler: &mut rubato::Fft<f32>, chunk_48k: &[f32]) -> Vec<f32> {
use audioadapter_buffers::direct::InterleavedSlice;
use rubato::Resampler;
let input = match InterleavedSlice::new(chunk_48k, 1, chunk_48k.len()) {
Ok(a) => a,
Err(e) => {
eprintln!("buzz-desktop: STT resample input error: {e}");
return Vec::new();
}
};
match resampler.process(&input, 0, None) {
Ok(out) => out.take_data(),
Err(e) => {
eprintln!("buzz-desktop: STT resample error: {e}");
Vec::new()
}
}
}
/// Feed 16 kHz samples through the VAD and accumulate speech.
/// Flushes to STT when silence exceeds the configured threshold.
#[allow(clippy::too_many_arguments)]
fn process_16k_samples(
samples: &[f32],
leftover: &mut Vec<f32>,
vad: &mut earshot::Detector<earshot::DefaultPredictor>,
speech_buf: &mut Vec<f32>,
silence_frames: &mut usize,
in_speech: &mut bool,
barge_in_frames: &mut usize,
recognizer: &sherpa_onnx::OfflineRecognizer,
text_tx: &tokio_mpsc::Sender<String>,
has_tts: bool,
tts_active: &Arc<AtomicBool>,
tts_cancel: Option<&AtomicBool>,
tts_stopped_at: &mut Option<std::time::Instant>,
ptt_active: Option<&Arc<AtomicBool>>,
silence_flush_threshold: usize,
max_speech_samples: usize,
) {
leftover.extend_from_slice(samples);
while leftover.len() >= VAD_FRAME_SAMPLES {
let frame: Vec<f32> = leftover.drain(..VAD_FRAME_SAMPLES).collect();
let clamped: Vec<f32> = frame.iter().map(|&s| s.clamp(-1.0, 1.0)).collect();
let prob = vad.predict_f32(&clamped);
let is_speech = prob > VAD_THRESHOLD;
// PTT gating: when PTT key is not held, treat as silence.
let is_speech = if let Some(ptt) = ptt_active {
is_speech && ptt.load(Ordering::Acquire)
} else {
is_speech
};
// TTS echo prevention (only when TTS flags are wired).
if has_tts {
let tts_playing = tts_active.load(Ordering::Acquire);
if tts_playing {
if ptt_active.is_some() {
// PTT mode — skip accumulation.
*in_speech = false;
*barge_in_frames = 0;
speech_buf.clear();
*silence_frames = 0;
continue;
}
// VAD mode — barge-in detection.
if is_speech {
*barge_in_frames += 1;
if *barge_in_frames >= BARGE_IN_DEBOUNCE_FRAMES {
if let Some(cancel) = tts_cancel {
cancel.store(true, Ordering::Release);
}
*barge_in_frames = 0;
}
} else {
*barge_in_frames = 0;
}
*in_speech = false;
speech_buf.clear();
*silence_frames = 0;
continue;
}
// TTS cooldown window.
if let Some(stopped) = *tts_stopped_at {
if stopped.elapsed() < TTS_COOLDOWN {
if !is_speech {
*in_speech = false;
}
speech_buf.clear();
*silence_frames = 0;
*barge_in_frames = 0;
continue;
} else {
*tts_stopped_at = None;
*in_speech = false;
*silence_frames = 0;
*barge_in_frames = 0;
}
}
}
if is_speech {
*silence_frames = 0;
*in_speech = true;
speech_buf.extend_from_slice(&frame);
// OOM guard.
if speech_buf.len() >= max_speech_samples {
flush_to_stt(speech_buf, recognizer, text_tx);
speech_buf.clear();
*silence_frames = 0;
*in_speech = false;
}
} else if *in_speech {
speech_buf.extend_from_slice(&frame);
*silence_frames += 1;
// In PTT mode, don't flush on silence — the PTT release edge handles it.
if ptt_active.is_none() && *silence_frames >= silence_flush_threshold {
flush_to_stt(speech_buf, recognizer, text_tx);
speech_buf.clear();
*silence_frames = 0;
*in_speech = false;
}
}
}
}
/// Run sherpa-onnx on the accumulated speech buffer and send the text.
fn flush_to_stt(
speech_buf: &[f32],
recognizer: &sherpa_onnx::OfflineRecognizer,
text_tx: &tokio_mpsc::Sender<String>,
) {
if speech_buf.is_empty() {
return;
}
let stream = recognizer.create_stream();
stream.accept_waveform(16_000, speech_buf);
recognizer.decode(&stream);
let text = stream
.get_result()
.map(|r| r.text.trim().to_string())
.unwrap_or_default();
if !text.is_empty() {
if let Err(e) = text_tx.blocking_send(text) {
eprintln!("buzz-desktop: STT text channel closed: {e}");
}
}
}
/// Convert raw bytes (f32 LE) to f32 samples.
fn bytes_to_f32(bytes: &[u8]) -> Vec<f32> {
bytes
.chunks_exact(4)
.map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.collect()
}
@@ -6,6 +6,7 @@ import {
parseAutoSubmitPhrases,
replaceTrailingTranscribedText,
} from "../lib/voiceInput";
import { useLocalDictation } from "./useLocalDictation";
import { useRealtimeDictation } from "./useRealtimeDictation";
interface UseDictationOptions {
@@ -63,11 +64,6 @@ export function useDictation({
return;
}
// Set the text first so the composer shows the final dictated content,
// then trigger send. We intentionally do NOT clear the composer here —
// the send flow in MessageComposer handles clearing on successful send.
// If a mention dialog opens (non-member mention), the text stays in the
// composer so the user doesn't lose their dictated message.
setText(textWithoutPhrase.trim());
onSend(textWithoutPhrase.trim());
lastTranscriptRef.current = "";
@@ -75,12 +71,25 @@ export function useDictation({
[autoSubmitPhrases, getText, onSend, isSendBlockedRef, setText],
);
const dictation = useRealtimeDictation({
// Try local STT first (offline, no API key needed).
const localDictation = useLocalDictation({
onRecordingStart: () => {
lastTranscriptRef.current = "";
},
onTranscriptText: handleTranscript,
});
// Fall back to cloud (OpenAI Realtime) if local is unavailable.
const cloudDictation = useRealtimeDictation({
disabled: localDictation.isEnabled, // Disable cloud when local is available
onRecordingStart: () => {
lastTranscriptRef.current = "";
},
onTranscriptText: handleTranscript,
});
// Use whichever is enabled — local takes priority.
const dictation = localDictation.isEnabled ? localDictation : cloudDictation;
stopRecordingRef.current = dictation.stopRecording;
return dictation;
@@ -0,0 +1,265 @@
import { invoke } from "@tauri-apps/api/core";
import { listen, type UnlistenFn } from "@tauri-apps/api/event";
import { useCallback, useEffect, useRef, useState } from "react";
import { toast } from "sonner";
/**
* Raw binary invoke — uses Tauri's internal IPC for zero-copy ArrayBuffer transfer.
* Same pattern as huddle's audioWorklet.ts.
*/
function invokeRawBinary(cmd: string, payload: Uint8Array): Promise<unknown> {
// biome-ignore lint/suspicious/noExplicitAny: Tauri internals have no public type definition
const internals = (window as any).__TAURI_INTERNALS__;
if (!internals?.invoke) {
return Promise.reject(new Error("Tauri internals not available"));
}
return internals.invoke(cmd, payload);
}
interface UseLocalDictationOptions {
disabled?: boolean;
onRecordingStart?: () => void;
onTranscriptText: (text: string) => void;
}
interface DictationStatus {
available: boolean;
active: boolean;
}
const DICTATION_TRANSCRIPT_EVENT = "dictation-transcript";
const DICTATION_STATE_EVENT = "dictation-state";
/**
* Local STT dictation hook using the Parakeet model via Tauri native commands.
*
* Works fully offline — no relay or OpenAI API key needed. Uses the same
* sherpa-onnx Parakeet TDT-CTC 110M model as huddle transcription.
*
* Audio capture uses the Web Audio API (AudioWorklet) on the frontend side,
* then sends raw PCM bytes to the native STT engine via `push_dictation_audio`.
*/
export function useLocalDictation({
disabled = false,
onRecordingStart,
onTranscriptText,
}: UseLocalDictationOptions) {
const [isRecording, setIsRecording] = useState(false);
const [isStarting, setIsStarting] = useState(false);
const [isTranscribing, setIsTranscribing] = useState(false);
const [isAvailable, setIsAvailable] = useState(false);
const streamRef = useRef<MediaStream | null>(null);
const audioContextRef = useRef<AudioContext | null>(null);
const workletRef = useRef<AudioWorkletNode | null>(null);
const unlistenTranscriptRef = useRef<UnlistenFn | null>(null);
const unlistenStateRef = useRef<UnlistenFn | null>(null);
const onRecordingStartRef = useRef(onRecordingStart);
const onTranscriptTextRef = useRef(onTranscriptText);
onRecordingStartRef.current = onRecordingStart;
onTranscriptTextRef.current = onTranscriptText;
const isEnabled = !disabled && isAvailable;
// Check availability on mount.
useEffect(() => {
let cancelled = false;
invoke<DictationStatus>("get_dictation_status")
.then((status) => {
if (!cancelled) setIsAvailable(status.available);
})
.catch(() => {
if (!cancelled) setIsAvailable(false);
});
return () => {
cancelled = true;
};
}, []);
const cleanup = useCallback(() => {
// Stop mic.
if (streamRef.current) {
for (const track of streamRef.current.getTracks()) {
track.stop();
}
streamRef.current = null;
}
// Disconnect audio worklet.
if (workletRef.current) {
workletRef.current.disconnect();
workletRef.current = null;
}
// Close audio context.
if (audioContextRef.current) {
void audioContextRef.current.close();
audioContextRef.current = null;
}
// Unlisten events.
if (unlistenTranscriptRef.current) {
unlistenTranscriptRef.current();
unlistenTranscriptRef.current = null;
}
if (unlistenStateRef.current) {
unlistenStateRef.current();
unlistenStateRef.current = null;
}
}, []);
// Cleanup on unmount.
useEffect(() => cleanup, [cleanup]);
const startRecording = useCallback(async () => {
if (!isEnabled || isStarting || isRecording) return;
setIsStarting(true);
onRecordingStartRef.current?.();
try {
// 1. Listen for transcript events from the native layer.
const unlistenTranscript = await listen<string>(
DICTATION_TRANSCRIPT_EVENT,
(event) => {
setIsTranscribing(false);
if (event.payload) {
onTranscriptTextRef.current(event.payload);
}
},
);
unlistenTranscriptRef.current = unlistenTranscript;
const unlistenState = await listen<string>(
DICTATION_STATE_EVENT,
(event) => {
if (event.payload === "stopped") {
setIsRecording(false);
setIsTranscribing(false);
}
},
);
unlistenStateRef.current = unlistenState;
// 2. Start the native STT engine.
await invoke("start_dictation");
// 3. Capture mic audio.
const stream = await navigator.mediaDevices.getUserMedia({
audio: {
autoGainControl: true,
echoCancellation: true,
noiseSuppression: true,
},
});
streamRef.current = stream;
// 4. Set up AudioWorklet to send PCM to native layer.
const audioContext = new AudioContext({ sampleRate: 48000 });
audioContextRef.current = audioContext;
// Create a simple processor that sends raw f32 PCM to the native side.
const processorCode = `
class DictationProcessor extends AudioWorkletProcessor {
process(inputs) {
const input = inputs[0];
if (input && input[0] && input[0].length > 0) {
// Send f32 samples as raw bytes.
this.port.postMessage(input[0].buffer);
}
return true;
}
}
registerProcessor('dictation-processor', DictationProcessor);
`;
const blob = new Blob([processorCode], {
type: "application/javascript",
});
const blobUrl = URL.createObjectURL(blob);
try {
await audioContext.audioWorklet.addModule(blobUrl);
} finally {
URL.revokeObjectURL(blobUrl);
}
const source = audioContext.createMediaStreamSource(stream);
const worklet = new AudioWorkletNode(audioContext, "dictation-processor");
workletRef.current = worklet;
// Forward PCM bytes to the native STT engine via raw binary IPC.
worklet.port.onmessage = (event: MessageEvent<ArrayBuffer>) => {
const float32 = new Float32Array(event.data);
const bytes = new Uint8Array(
float32.buffer,
float32.byteOffset,
float32.byteLength,
);
invokeRawBinary("push_dictation_audio", bytes).catch(() => {});
};
source.connect(worklet);
worklet.connect(audioContext.destination);
setIsRecording(true);
setIsTranscribing(true);
} catch (error) {
cleanup();
setIsRecording(false);
setIsTranscribing(false);
const message =
error instanceof Error ? error.message : "Local dictation failed";
if (/not allowed|denied|permission/i.test(message)) {
toast.error("Microphone access denied", {
description:
"Allow microphone access in System Settings to use dictation.",
});
} else if (/not found|no audio/i.test(message)) {
toast.error("No microphone found", {
description: "Connect a microphone and try again.",
});
} else if (/model not ready/i.test(message)) {
toast.error("Voice model downloading", {
description:
"The speech model is still downloading. Try again shortly.",
});
} else {
toast.error("Dictation failed", { description: message });
}
} finally {
setIsStarting(false);
}
}, [cleanup, isEnabled, isRecording, isStarting]);
const stopRecording = useCallback(() => {
cleanup();
invoke("stop_dictation").catch(() => {});
setIsRecording(false);
// Keep isTranscribing briefly — final segment may still arrive.
setTimeout(() => setIsTranscribing(false), 500);
}, [cleanup]);
const cancelRecording = useCallback(() => {
cleanup();
invoke("stop_dictation").catch(() => {});
setIsRecording(false);
setIsTranscribing(false);
}, [cleanup]);
const toggleRecording = useCallback(() => {
if (isRecording || isStarting) {
stopRecording();
return;
}
void startRecording();
}, [isRecording, isStarting, startRecording, stopRecording]);
return {
isEnabled,
isRecording,
isStarting,
isTranscribing,
startRecording,
stopRecording,
cancelRecording,
toggleRecording,
};
}