feat(desktop): add local dictation pipeline

Co-authored-by: Kenny Lopez <klopez4212@gmail.com>
Signed-off-by: John Tennant <jtennant@squareup.com>
This commit is contained in:
John Tennant
2026-07-31 14:38:57 -04:00
co-authored by Kenny Lopez
parent 54c8ef30a9
commit c3566dd547
3 changed files with 326 additions and 31 deletions
+252
View File
@@ -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<Arc<SttPipeline>>,
/// 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<Mutex<DictationState>> =
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<u64, 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(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<u64>) -> 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<u64>) {
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<String>,
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 }),
);
});
}
+73 -31
View File
@@ -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<usize>,
tts_active: Arc<AtomicBool>,
tts_cancel: Option<Arc<AtomicBool>>,
ptt_active: Option<Arc<AtomicBool>>,
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<AtomicBool>,
tts_cancel: Option<Arc<AtomicBool>>,
ptt_active: Option<Arc<AtomicBool>>,
) -> Result<(Self, tokio_mpsc::Receiver<String>), 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>), 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>), 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,
)
})
.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<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};
@@ -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<std::time::Instant>,
ptt_active: Option<&Arc<AtomicBool>>,
silence_flush_frames: usize,
max_speech_samples: usize,
partial_flush_samples: Option<usize>,
) {
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();
+1
View File
@@ -4,6 +4,7 @@ mod archive;
mod builderlab;
mod commands;
mod deep_link;
mod dictation;
mod egress_guard;
mod event_sync;
mod events;