diff --git a/desktop/src-tauri/src/dictation.rs b/desktop/src-tauri/src/dictation.rs index c3493ea22..a03840d02 100644 --- a/desktop/src-tauri/src/dictation.rs +++ b/desktop/src-tauri/src/dictation.rs @@ -35,11 +35,18 @@ const DICTATION_STATE_EVENT: &str = "dictation-state"; 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 } + Self { + engine: None, + session_id: 0, + } } } @@ -49,7 +56,7 @@ impl DictationState { /// `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> { +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. @@ -77,14 +84,16 @@ pub async fn start_dictation(state: State<'_, AppState>) -> Result<(), String> { let (engine, text_rx) = SttEngine::new(config)?; let engine = Arc::new(engine); - // Store the engine in state. - { + // Store the engine in state and increment the session counter. + let session_id = { let mut ds = state .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 @@ -94,11 +103,14 @@ pub async fn start_dictation(state: State<'_, AppState>) -> Result<(), String> { .clone(); if let Some(handle) = app_handle { - let _ = handle.emit(DICTATION_STATE_EVENT, "started"); - spawn_dictation_forwarder(text_rx, handle); + let _ = handle.emit( + DICTATION_STATE_EVENT, + serde_json::json!({ "state": "started", "session": session_id }), + ); + spawn_dictation_forwarder(text_rx, handle, session_id); } - Ok(()) + Ok(session_id) } /// `stop_dictation` — stop the active dictation session. @@ -195,23 +207,31 @@ fn stop_dictation_inner(state: &AppState) { /// Spawn an async task that reads transcribed text and emits Tauri events. /// -/// When the channel closes (engine stopped), the forwarder emits -/// `dictation-state: stopped` so the frontend knows all pending transcripts -/// have been delivered. +/// 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; } - if app_handle.emit(DICTATION_TRANSCRIPT_EVENT, &text).is_err() { + 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, "stopped"); + let _ = app_handle.emit( + DICTATION_STATE_EVENT, + serde_json::json!({ "state": "stopped", "session": session_id }), + ); }); } diff --git a/desktop/src/features/dictation/hooks/useLocalDictation.ts b/desktop/src/features/dictation/hooks/useLocalDictation.ts index 35f69c014..324c8ece0 100644 --- a/desktop/src/features/dictation/hooks/useLocalDictation.ts +++ b/desktop/src/features/dictation/hooks/useLocalDictation.ts @@ -70,9 +70,10 @@ export function useLocalDictation({ const unlistenStateRef = useRef(null); const onRecordingStartRef = useRef(onRecordingStart); const onTranscriptTextRef = useRef(onTranscriptText); - // Session ID — incremented on each start. Transcript events from a previous - // session's forwarder are ignored by comparing against the active session. - const sessionIdRef = useRef(0); + // Native session ID — set after `start_dictation` returns. Transcript and + // state events include this ID so we can definitively ignore stale events + // from a previous session's forwarder. + const nativeSessionRef = useRef(0); // Abort flag — set when stop/cancel is called while startRecording is still // awaiting async setup. The start resumes and bails before activating. const startAbortedRef = useRef(false); @@ -190,62 +191,16 @@ export function useLocalDictation({ const startRecording = useCallback(async () => { if (!isEnabled || isStarting || isRecording) return; - // Increment session ID and clear abort flag for this new start attempt. - const thisSession = ++sessionIdRef.current; + // Clear abort flag for this new start attempt. startAbortedRef.current = false; setIsStarting(true); onRecordingStartRef.current?.(); try { - // 1. Listen for transcript events from the native layer. - // Only accept events matching the current session to avoid stale - // transcripts from a previous session's forwarder leaking in. - const unlistenTranscript = await listen( - DICTATION_TRANSCRIPT_EVENT, - (event) => { - if (sessionIdRef.current !== thisSession) return; - if (event.payload) { - onTranscriptTextRef.current(event.payload); - } - }, - ); - // Bail if stop/cancel was called while we were awaiting. - if (startAbortedRef.current) { - unlistenTranscript(); - return; - } - unlistenTranscriptRef.current = unlistenTranscript; - - const unlistenState = await listen( - DICTATION_STATE_EVENT, - (event) => { - if (sessionIdRef.current !== thisSession) return; - if (event.payload === "stopped") { - setIsRecording(false); - setIsTranscribing(false); - // Clean up event listeners now that the session is fully done. - if (unlistenTranscriptRef.current) { - unlistenTranscriptRef.current(); - unlistenTranscriptRef.current = null; - } - if (unlistenStateRef.current) { - unlistenStateRef.current(); - unlistenStateRef.current = null; - } - } - }, - ); - // Bail if stop/cancel was called while we were awaiting. - if (startAbortedRef.current) { - unlistenTranscript(); - unlistenState(); - return; - } - unlistenStateRef.current = unlistenState; - - // 2. Start the native STT engine. - await invoke("start_dictation"); + // 1. Start the native STT engine — returns the session ID used to tag events. + const sessionId = await invoke("start_dictation"); + nativeSessionRef.current = sessionId; // Bail if aborted during engine start. if (startAbortedRef.current) { @@ -253,6 +208,56 @@ export function useLocalDictation({ return; } + // 2. Listen for transcript events from the native layer. + // Each event includes a `session` ID so we can definitively ignore stale + // transcripts from a previous session's forwarder. + const unlistenTranscript = await listen<{ + text: string; + session: number; + }>(DICTATION_TRANSCRIPT_EVENT, (event) => { + const { text, session } = event.payload; + if (session !== nativeSessionRef.current) return; + if (text) { + onTranscriptTextRef.current(text); + } + }); + // Bail if stop/cancel was called while we were awaiting. + if (startAbortedRef.current) { + unlistenTranscript(); + invoke("stop_dictation").catch(() => {}); + return; + } + unlistenTranscriptRef.current = unlistenTranscript; + + const unlistenState = await listen<{ + state: string; + session: number; + }>(DICTATION_STATE_EVENT, (event) => { + const { state, session } = event.payload; + if (session !== nativeSessionRef.current) return; + if (state === "stopped") { + setIsRecording(false); + setIsTranscribing(false); + // Clean up event listeners now that the session is fully done. + if (unlistenTranscriptRef.current) { + unlistenTranscriptRef.current(); + unlistenTranscriptRef.current = null; + } + if (unlistenStateRef.current) { + unlistenStateRef.current(); + unlistenStateRef.current = null; + } + } + }); + // Bail if stop/cancel was called while we were awaiting. + if (startAbortedRef.current) { + unlistenTranscript(); + unlistenState(); + invoke("stop_dictation").catch(() => {}); + return; + } + unlistenStateRef.current = unlistenState; + // 3. Capture mic audio. const stream = await navigator.mediaDevices.getUserMedia({ audio: {