diff --git a/desktop/src-tauri/src/app_state.rs b/desktop/src-tauri/src/app_state.rs index 7421caa66..ce516d69f 100644 --- a/desktop/src-tauri/src/app_state.rs +++ b/desktop/src-tauri/src/app_state.rs @@ -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>, pub huddle_state: Mutex, - /// 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>, /// 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>, + /// Local dictation state (Parakeet STT for composer voice input). + pub dictation_state: Mutex, } /// 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()), } } diff --git a/desktop/src-tauri/src/dictation.rs b/desktop/src-tauri/src/dictation.rs new file mode 100644 index 000000000..fd4d7d562 --- /dev/null +++ b/desktop/src-tauri/src/dictation.rs @@ -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>, +} + +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, + 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. + } + } + }); +} diff --git a/desktop/src-tauri/src/huddle/mod.rs b/desktop/src-tauri/src/huddle/mod.rs index 5499ccab4..efcffca9e 100644 --- a/desktop/src-tauri/src/huddle/mod.rs +++ b/desktop/src-tauri/src/huddle/mod.rs @@ -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( - rx: std::sync::mpsc::Receiver, - 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 ──────────────────────────────────────────────────────────────── diff --git a/desktop/src-tauri/src/huddle/stt.rs b/desktop/src-tauri/src/huddle/stt.rs index 6f502ca72..5bbfe630f 100644 --- a/desktop/src-tauri/src/huddle/stt.rs +++ b/desktop/src-tauri/src/huddle/stt.rs @@ -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`) 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>, - /// Signals the worker thread to stop. - shutdown: Arc, - /// Worker thread handle — taken on drop to join cleanly. - thread: Option>, + 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` is returned separately so the - /// caller can move it directly into an async task. This avoids holding a - /// `Mutex` across await points (which would block a Tokio worker - /// thread on every `recv_timeout` call). pub fn new( model_dir: PathBuf, tts_active: Arc, tts_cancel: Option>, ptt_active: Option>, ) -> Result<(Self, tokio_mpsc::Receiver), String> { - let (audio_tx, audio_rx) = mpsc::sync_channel::>(AUDIO_QUEUE_DEPTH); - let (text_tx, text_rx) = tokio_mpsc::channel::(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) -> 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>, - text_tx: tokio_mpsc::Sender, - shutdown: Arc, - tts_active: Arc, - tts_cancel: Option>, - ptt_active: Option>, -) { - // ── 1. Initialise rubato resampler (48 kHz → 16 kHz, mono) ─────────────── - use rubato::{Fft, FixedSync, Resampler}; - - let mut resampler = match Fft::::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 = Vec::with_capacity(chunk_in * 2); - // Leftover 16 kHz samples that didn't fill a full VAD frame. - let mut leftover_16k: Vec = Vec::new(); - // Accumulated speech frames (16 kHz). - let mut speech_buf: Vec = 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 = 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 = 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, chunk_48k: &[f32]) -> Vec { - 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, - vad: &mut earshot::Detector, - speech_buf: &mut Vec, - silence_frames: &mut usize, - in_speech: &mut bool, - barge_in_frames: &mut usize, - recognizer: &sherpa_onnx::OfflineRecognizer, - text_tx: &tokio_mpsc::Sender, - tts_active: &Arc, - tts_cancel: Option<&AtomicBool>, - tts_stopped_at: &mut Option, - ptt_active: Option<&Arc>, -) { - leftover.extend_from_slice(samples); - - while leftover.len() >= VAD_FRAME_SAMPLES { - let frame: Vec = leftover.drain(..VAD_FRAME_SAMPLES).collect(); - let clamped: Vec = 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, -) { - 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 { - 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; diff --git a/desktop/src-tauri/src/lib.rs b/desktop/src-tauri/src/lib.rs index 2c110377c..6c31728a2 100644 --- a/desktop/src-tauri/src/lib.rs +++ b/desktop/src-tauri/src/lib.rs @@ -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"); diff --git a/desktop/src-tauri/src/stt_engine.rs b/desktop/src-tauri/src/stt_engine.rs new file mode 100644 index 000000000..7a501d273 --- /dev/null +++ b/desktop/src-tauri/src/stt_engine.rs @@ -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>, + /// Optional: TTS cancel flag — set by STT on barge-in detection. + pub tts_cancel: Option>, + /// Optional: push-to-talk flag — when `Some`, speech is only accumulated + /// while the flag is true. + pub ptt_active: Option>, +} + +// ── 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`) 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>, + /// Signals the worker thread to stop. + shutdown: Arc, + /// Worker thread handle — taken on drop to join cleanly. + thread: Option>, +} + +impl SttEngine { + /// Spawn the engine worker thread. + /// + /// Returns `(Self, Receiver)`. 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> { + let (audio_tx, audio_rx) = mpsc::sync_channel::>(AUDIO_QUEUE_DEPTH); + let (text_tx, text_rx) = tokio_mpsc::channel::(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) -> 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(rx: std::sync::mpsc::Receiver, 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>, + text_tx: tokio_mpsc::Sender, + shutdown: Arc, +) { + // ── 1. Initialise rubato resampler (48 kHz → 16 kHz, mono) ─────────────── + use rubato::{Fft, FixedSync, Resampler}; + + let mut resampler = match Fft::::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 = Vec::with_capacity(chunk_in * 2); + let mut leftover_16k: Vec = Vec::new(); + let mut speech_buf: Vec = 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 = 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 = 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, chunk_48k: &[f32]) -> Vec { + 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, + vad: &mut earshot::Detector, + speech_buf: &mut Vec, + silence_frames: &mut usize, + in_speech: &mut bool, + barge_in_frames: &mut usize, + recognizer: &sherpa_onnx::OfflineRecognizer, + text_tx: &tokio_mpsc::Sender, + has_tts: bool, + tts_active: &Arc, + tts_cancel: Option<&AtomicBool>, + tts_stopped_at: &mut Option, + ptt_active: Option<&Arc>, + silence_flush_threshold: usize, + max_speech_samples: usize, +) { + leftover.extend_from_slice(samples); + + while leftover.len() >= VAD_FRAME_SAMPLES { + let frame: Vec = leftover.drain(..VAD_FRAME_SAMPLES).collect(); + let clamped: Vec = 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, +) { + 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 { + bytes + .chunks_exact(4) + .map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]])) + .collect() +} diff --git a/desktop/src/features/dictation/hooks/useDictation.ts b/desktop/src/features/dictation/hooks/useDictation.ts index 5e21fed16..36050a086 100644 --- a/desktop/src/features/dictation/hooks/useDictation.ts +++ b/desktop/src/features/dictation/hooks/useDictation.ts @@ -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; diff --git a/desktop/src/features/dictation/hooks/useLocalDictation.ts b/desktop/src/features/dictation/hooks/useLocalDictation.ts new file mode 100644 index 000000000..266b70ff3 --- /dev/null +++ b/desktop/src/features/dictation/hooks/useLocalDictation.ts @@ -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 { + // 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(null); + const audioContextRef = useRef(null); + const workletRef = useRef(null); + const unlistenTranscriptRef = useRef(null); + const unlistenStateRef = useRef(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("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( + DICTATION_TRANSCRIPT_EVENT, + (event) => { + setIsTranscribing(false); + if (event.payload) { + onTranscriptTextRef.current(event.payload); + } + }, + ); + unlistenTranscriptRef.current = unlistenTranscript; + + const unlistenState = await listen( + 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) => { + 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, + }; +}