From c3566dd547352060abc6f8d67157febd84f78b10 Mon Sep 17 00:00:00 2001 From: John Tennant Date: Thu, 30 Jul 2026 11:23:02 -0400 Subject: [PATCH] feat(desktop): add local dictation pipeline Co-authored-by: Kenny Lopez Signed-off-by: John Tennant --- desktop/src-tauri/src/dictation.rs | 252 ++++++++++++++++++++++++++++ desktop/src-tauri/src/huddle/stt.rs | 104 ++++++++---- desktop/src-tauri/src/lib.rs | 1 + 3 files changed, 326 insertions(+), 31 deletions(-) create mode 100644 desktop/src-tauri/src/dictation.rs diff --git a/desktop/src-tauri/src/dictation.rs b/desktop/src-tauri/src/dictation.rs new file mode 100644 index 000000000..dc1aed227 --- /dev/null +++ b/desktop/src-tauri/src/dictation.rs @@ -0,0 +1,252 @@ +//! 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 shared huddle PTT gate; the composer controls its own session +//! - Slightly longer silence threshold for more coherent sentences +//! - Text goes to the frontend, not to the relay + +use std::sync::{Arc, LazyLock, Mutex}; + +use tauri::{Emitter, State}; + +use crate::app_state::AppState; +use crate::huddle::{models, stt::SttPipeline}; + +/// 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 app-wide 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>, + /// Monotonically increasing session counter. Included in all emitted events + /// so the frontend can ignore stale transcripts from a previous session's + /// forwarder that arrive after a new session has started. + session_id: u64, +} + +impl DictationState { + pub fn new() -> Self { + Self { + engine: None, + session_id: 0, + } + } +} + +static DICTATION_STATE: LazyLock> = + LazyLock::new(|| Mutex::new(DictationState::new())); + +/// `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 { + // 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(None); + + let (engine, text_rx) = SttPipeline::new_dictation(model_dir)?; + let engine = Arc::new(engine); + + // Store the engine in state and increment the session counter. + let session_id = { + let mut ds = DICTATION_STATE.lock().unwrap_or_else(|e| e.into_inner()); + ds.engine = Some(Arc::clone(&engine)); + ds.session_id += 1; + ds.session_id + }; + + // 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, + serde_json::json!({ "state": "started", "session": session_id }), + ); + spawn_dictation_forwarder(text_rx, handle, session_id); + } + + Ok(session_id) +} + +/// `stop_dictation` — stop the active dictation session. +/// +/// The final transcript (if any) is emitted asynchronously by the forwarder +/// task. The `dictation-state: stopped` event is emitted by the forwarder +/// after all pending transcripts have been forwarded, ensuring the frontend +/// receives the final text before the stopped signal. +/// +/// `session` scopes the stop to a specific session: the engine is only torn +/// down when the currently-stored `session_id` matches. This prevents a +/// delayed/fire-and-forget stop from an old session (e.g. one deferred behind +/// a final audio flush) from killing a *newer* session the user started in the +/// meantime. Pass `None` for an unconditional stop (used on cancel/unmount). +#[tauri::command] +pub fn stop_dictation(session: Option) -> Result<(), String> { + stop_dictation_inner(session); + // Note: `stopped` is emitted by the forwarder task after draining all + // pending transcripts — not here. This avoids a race where the frontend + // sees `stopped` before the final transcript arrives. + Ok(()) +} + +/// `push_dictation_audio` — feed raw PCM bytes into the dictation pipeline. +/// +/// Expects a raw binary body: an 8-byte little-endian `u64` session header +/// followed by f32 LE samples at 48 kHz mono. The header scopes the push to a +/// specific session — bytes are only fed to the engine when the header matches +/// the currently-stored `session_id`. This prevents late audio from a +/// just-stopped session (whose final `flushAudioBatch()` chunks are still +/// arriving) from being accepted by a *newer* session the user started in the +/// meantime and transcribed into the new draft. +/// +/// If no dictation session is active, or the session header doesn't match the +/// active session, the bytes are silently discarded. +#[tauri::command] +pub fn push_dictation_audio(request: tauri::ipc::Request<'_>) -> Result<(), String> { + /// Size of the leading little-endian `u64` session header. + const SESSION_HEADER_BYTES: usize = 8; + /// Maximum IPC audio batch size (audio payload only, excluding the header): 100 KB. + const MAX_AUDIO_BATCH_BYTES: usize = 100 * 1024; + + match request.body() { + tauri::ipc::InvokeBody::Raw(bytes) => { + if bytes.len() < SESSION_HEADER_BYTES { + return Err(format!( + "audio batch too small: {} bytes (need at least {} for session header)", + bytes.len(), + SESSION_HEADER_BYTES + )); + } + let (header, audio) = bytes.split_at(SESSION_HEADER_BYTES); + if audio.len() > MAX_AUDIO_BATCH_BYTES { + return Err(format!( + "audio batch too large: {} bytes (max {})", + audio.len(), + MAX_AUDIO_BATCH_BYTES + )); + } + // `split_at` guarantees `header` is exactly `SESSION_HEADER_BYTES` long. + let session = u64::from_le_bytes( + header + .try_into() + .map_err(|_| "invalid session header".to_string())?, + ); + let ds = DICTATION_STATE.lock().unwrap_or_else(|e| e.into_inner()); + // Only feed audio tagged with the currently-active session. Late + // chunks from an old session are silently dropped. + if session == ds.session_id { + if let Some(ref engine) = ds.engine { + engine.push_audio(audio.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() -> DictationStatus { + let model_ready = models::is_stt_ready(); + let is_active = 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(session: Option) { + let old_engine = { + let mut ds = DICTATION_STATE.lock().unwrap_or_else(|e| e.into_inner()); + // Session-scoped stop: only tear down when the requested session matches + // the one currently stored. A `None` session stops unconditionally. + match session { + Some(requested) if requested != ds.session_id => None, + _ => 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. +/// +/// Each event includes the `session` ID so the frontend can ignore stale +/// transcripts from a previous session's forwarder. When the channel closes +/// (engine stopped), the forwarder emits `dictation-state: stopped`. +fn spawn_dictation_forwarder( + mut text_rx: tokio::sync::mpsc::Receiver, + app_handle: tauri::AppHandle, + session_id: u64, +) { + tauri::async_runtime::spawn(async move { + while let Some(text) = text_rx.recv().await { + if text.is_empty() { + continue; + } + let payload = serde_json::json!({ "text": text, "session": session_id }); + if app_handle + .emit(DICTATION_TRANSCRIPT_EVENT, payload) + .is_err() + { + break; // App window closed. + } + } + // All transcripts forwarded — signal the frontend that dictation is done. + let _ = app_handle.emit( + DICTATION_STATE_EVENT, + serde_json::json!({ "state": "stopped", "session": session_id }), + ); + }); +} diff --git a/desktop/src-tauri/src/huddle/stt.rs b/desktop/src-tauri/src/huddle/stt.rs index 6f502ca72..67b639332 100644 --- a/desktop/src-tauri/src/huddle/stt.rs +++ b/desktop/src-tauri/src/huddle/stt.rs @@ -41,6 +41,23 @@ const AUDIO_QUEUE_DEPTH: usize = 50; /// Prevents OOM if VAD stays in speech mode (noisy environment). const MAX_SPEECH_SAMPLES: usize = 16_000 * 30; +/// Dictation waits slightly longer for a natural pause before flushing. +const DICTATION_SILENCE_FLUSH_FRAMES: usize = 25; + +/// Emit a partial dictation result after two seconds of uninterrupted speech. +const DICTATION_PARTIAL_FLUSH_SAMPLES: usize = 16_000 * 2; + +struct SttPipelineConfig { + model_dir: PathBuf, + silence_flush_frames: usize, + max_speech_samples: usize, + partial_flush_samples: Option, + tts_active: Arc, + tts_cancel: Option>, + ptt_active: Option>, + flush_on_shutdown: bool, +} + /// Handle to the running STT pipeline. /// /// Not Clone — wrap in `Arc` to share across threads. @@ -88,27 +105,46 @@ impl SttPipeline { tts_active: Arc, tts_cancel: Option>, ptt_active: Option>, + ) -> Result<(Self, tokio_mpsc::Receiver), String> { + Self::new_with_config(SttPipelineConfig { + model_dir, + silence_flush_frames: SILENCE_FLUSH_FRAMES, + max_speech_samples: MAX_SPEECH_SAMPLES, + partial_flush_samples: None, + tts_active, + tts_cancel, + ptt_active, + flush_on_shutdown: false, + }) + } + + /// Spawn an offline STT pipeline for message-composer dictation. + pub(crate) fn new_dictation( + model_dir: PathBuf, + ) -> Result<(Self, tokio_mpsc::Receiver), String> { + Self::new_with_config(SttPipelineConfig { + model_dir, + silence_flush_frames: DICTATION_SILENCE_FLUSH_FRAMES, + max_speech_samples: MAX_SPEECH_SAMPLES, + partial_flush_samples: Some(DICTATION_PARTIAL_FLUSH_SAMPLES), + tts_active: Arc::new(AtomicBool::new(false)), + tts_cancel: None, + ptt_active: None, + flush_on_shutdown: true, + }) + } + + fn new_with_config( + config: SttPipelineConfig, ) -> 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, - ) - }) + .spawn(move || stt_worker(config, audio_rx, text_tx, shutdown_worker)) .map_err(|e| format!("failed to spawn stt-worker thread: {e}"))?; let pipeline = Self { @@ -202,13 +238,10 @@ const TTS_COOLDOWN: Duration = Duration::from_millis(50); const STT_NUM_THREADS: i32 = 1; fn stt_worker( - model_dir: PathBuf, + config: SttPipelineConfig, 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}; @@ -235,12 +268,12 @@ fn stt_worker( // 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"); + 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", - model_dir.display() + config.model_dir.display() ); drain_until_shutdown(audio_rx, &shutdown); return; @@ -281,7 +314,8 @@ fn stt_worker( // ── 5. Main loop ────────────────────────────────────────────────────────── let mut tts_was_active = false; - let mut ptt_was_active = ptt_active + let mut ptt_was_active = config + .ptt_active .as_ref() .is_some_and(|p| p.load(Ordering::Acquire)); loop { @@ -291,7 +325,7 @@ fn stt_worker( } // Track TTS transitions to set the cooldown timer. - let tts_now = tts_active.load(Ordering::Acquire); + let tts_now = config.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()); @@ -302,7 +336,7 @@ fn stt_worker( // 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 { + if let Some(ref ptt) = config.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); @@ -345,18 +379,21 @@ fn stt_worker( &mut barge_in_frames, &recognizer, &text_tx, - &tts_active, - tts_cancel.as_deref(), + &config.tts_active, + config.tts_cancel.as_deref(), &mut tts_stopped_at, - ptt_active.as_ref(), + config.ptt_active.as_ref(), + config.silence_flush_frames, + config.max_speech_samples, + config.partial_flush_samples, ); } } } - // 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. + if config.flush_on_shutdown && !speech_buf.is_empty() { + flush_to_stt(&speech_buf, &recognizer, &text_tx); + } } /// Resample a mono 48 kHz chunk to 16 kHz using rubato. @@ -411,6 +448,9 @@ fn process_16k_samples( tts_cancel: Option<&AtomicBool>, tts_stopped_at: &mut Option, ptt_active: Option<&Arc>, + silence_flush_frames: usize, + max_speech_samples: usize, + partial_flush_samples: Option, ) { leftover.extend_from_slice(samples); @@ -494,7 +534,9 @@ fn process_16k_samples( 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 { + if speech_buf.len() >= max_speech_samples + || partial_flush_samples.is_some_and(|limit| speech_buf.len() >= limit) + { flush_to_stt(speech_buf, recognizer, text_tx); speech_buf.clear(); *silence_frames = 0; @@ -509,7 +551,7 @@ fn process_16k_samples( // 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 { + 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(); diff --git a/desktop/src-tauri/src/lib.rs b/desktop/src-tauri/src/lib.rs index 6814008f0..80b23b7b6 100644 --- a/desktop/src-tauri/src/lib.rs +++ b/desktop/src-tauri/src/lib.rs @@ -4,6 +4,7 @@ mod archive; mod builderlab; mod commands; mod deep_link; +mod dictation; mod egress_guard; mod event_sync; mod events;