diff --git a/src/main/speech/model-manager-progress-callback.test.ts b/src/main/speech/model-manager-progress-callback.test.ts index 8b566abad70..1c8633188e7 100644 --- a/src/main/speech/model-manager-progress-callback.test.ts +++ b/src/main/speech/model-manager-progress-callback.test.ts @@ -45,4 +45,30 @@ describe('ModelManager progress callbacks', () => { rmSync(dir, { recursive: true, force: true }) } }) + + it('coalesces per-chunk download progress to whole percent', () => { + const dir = mkdtempSync(join(tmpdir(), 'orca-model-manager-')) + try { + const manager = new ModelManager(dir) + const internals = manager as unknown as ModelManagerInternals + const listener = vi.fn() + manager.setProgressCallback(listener) + + // A 500MB model over a 64KB chunk stream reports this many times. + for (let chunk = 0; chunk < 8_000; chunk += 1) { + internals.updateState('model-a', 'downloading', chunk / 8_000) + } + + expect(listener.mock.calls.map(([, progress]) => progress)).toEqual( + Array.from({ length: 101 }, (_unused, percent) => percent / 100) + ) + const afterDownload = listener.mock.calls.length + internals.updateState('model-a', 'extracting') + internals.updateState('model-a', 'ready') + internals.updateState('model-a', 'ready') + expect(listener.mock.calls.length).toBe(afterDownload + 3) + } finally { + rmSync(dir, { recursive: true, force: true }) + } + }) }) diff --git a/src/main/speech/model-manager.ts b/src/main/speech/model-manager.ts index 9125ffa4962..f183805cddd 100644 --- a/src/main/speech/model-manager.ts +++ b/src/main/speech/model-manager.ts @@ -387,10 +387,24 @@ export class ModelManager { progress?: number, error?: string ): void { - const state: SpeechModelState = { id: modelId, status, progress, error } + const previous = this.modelStates.get(modelId) + // Whole-percent state matches the UI and prevents chunk-level IPC/poll churn. + const reportedProgress = + status === 'downloading' && progress !== undefined + ? Math.round(progress * 100) / 100 + : progress + if ( + status === 'downloading' && + previous?.status === 'downloading' && + previous.error === error && + previous.progress === reportedProgress + ) { + return + } + const state: SpeechModelState = { id: modelId, status, progress: reportedProgress, error } this.modelStates.set(modelId, state) - // Why: notify on every state change (not just progress) so extracting/ready/error transitions reach the UI. - const progressValue = progress ?? (status === 'extracting' ? 0.95 : -1) + // Repeated non-download states can be the requesting window's only resync signal. + const progressValue = reportedProgress ?? (status === 'extracting' ? 0.95 : -1) for (const callback of this.progressCallbacks) { callback(modelId, progressValue) } diff --git a/src/renderer/src/store/slices/dictation-model-state-stabilisation.test.ts b/src/renderer/src/store/slices/dictation-model-state-stabilisation.test.ts new file mode 100644 index 00000000000..3f8ad8eb59f --- /dev/null +++ b/src/renderer/src/store/slices/dictation-model-state-stabilisation.test.ts @@ -0,0 +1,62 @@ +/** @vitest-environment happy-dom */ +import { beforeEach, describe, expect, it, vi } from 'vitest' +import { create, type StateCreator } from 'zustand' +import type { SpeechModelState } from '../../../../shared/speech-types' +import type { AppState } from '../types' +import { createDictationSlice } from './dictation' + +type DictationTestStore = Pick +const dictationSlice = createDictationSlice as unknown as StateCreator + +let reply: SpeechModelState[] + +beforeEach(() => { + reply = [] + Object.assign(window, { + api: { + speech: { getModelStates: vi.fn(async () => reply.map((state) => ({ ...state }))) } + } + }) +}) + +describe('dictation model-state stabilisation', () => { + it('does not publish an unchanged reply', async () => { + reply = [{ id: 'whisper-tiny', status: 'downloading', progress: 0.42 }] + const store = create(dictationSlice) + await store.getState().refreshModelStates() + const previous = store.getState() + const subscriber = vi.fn() + store.subscribe(subscriber) + + await store.getState().refreshModelStates() + + expect(store.getState()).toBe(previous) + expect(subscriber).not.toHaveBeenCalled() + }) + + it.each([ + { changed: [{ id: 'whisper-tiny', status: 'downloading', progress: 0.43 }] }, + { changed: [{ id: 'whisper-tiny', status: 'ready' }] }, + { changed: [{ id: 'whisper-tiny', status: 'error', error: 'boom' }] }, + { + changed: [{ id: 'parakeet-tdt-0.6b-v3-int8', status: 'downloading', progress: 0.42 }] + }, + { + changed: [ + { id: 'whisper-tiny', status: 'downloading', progress: 0.42 }, + { id: 'parakeet-tdt-0.6b-v3-int8', status: 'not-downloaded' } + ] + } + ] as { changed: SpeechModelState[] }[])('publishes a changed reply', async ({ changed }) => { + const store = create(dictationSlice) + store.getState().setModelStates([{ id: 'whisper-tiny', status: 'downloading', progress: 0.42 }]) + const subscriber = vi.fn() + store.subscribe(subscriber) + + reply = changed + await store.getState().refreshModelStates() + + expect(store.getState().modelStates).toEqual(changed) + expect(subscriber).toHaveBeenCalledTimes(1) + }) +}) diff --git a/src/renderer/src/store/slices/dictation.ts b/src/renderer/src/store/slices/dictation.ts index 7b89cb69e35..e5cd28914c5 100644 --- a/src/renderer/src/store/slices/dictation.ts +++ b/src/renderer/src/store/slices/dictation.ts @@ -14,23 +14,49 @@ export type DictationSlice = { refreshModelStates: () => Promise } -export const createDictationSlice: StateCreator = (set) => ({ - dictationState: 'idle', - partialTranscript: '', - activeModelId: null, - modelStates: [], +function sameSpeechModelState(a: SpeechModelState, b: SpeechModelState): boolean { + return a.id === b.id && a.status === b.status && a.progress === b.progress && a.error === b.error +} - setDictationState: (state) => set({ dictationState: state }), - setPartialTranscript: (text) => set({ partialTranscript: text }), - setActiveModelId: (id) => set({ activeModelId: id }), - setModelStates: (states) => set({ modelStates: states }), +// Why: every getModelStates reply is a fresh array, so without this each no-op +// refresh re-renders every subscriber — including the open speech-model menu. +function resolveModelStates( + previous: SpeechModelState[], + next: SpeechModelState[] +): SpeechModelState[] { + if (previous.length !== next.length) { + return next + } + return previous.every((state, index) => sameSpeechModelState(state, next[index])) + ? previous + : next +} - refreshModelStates: async () => { - try { - const states = await window.api.speech.getModelStates() - set({ modelStates: states }) - } catch (err) { - console.error('Failed to fetch model states:', err) +export const createDictationSlice: StateCreator = (set) => { + const setModelStates = (states: SpeechModelState[]): void => { + set((prev) => { + const modelStates = resolveModelStates(prev.modelStates, states) + return modelStates === prev.modelStates ? prev : { modelStates } + }) + } + + return { + dictationState: 'idle', + partialTranscript: '', + activeModelId: null, + modelStates: [], + + setDictationState: (state) => set({ dictationState: state }), + setPartialTranscript: (text) => set({ partialTranscript: text }), + setActiveModelId: (id) => set({ activeModelId: id }), + setModelStates, + + refreshModelStates: async () => { + try { + setModelStates(await window.api.speech.getModelStates()) + } catch (err) { + console.error('Failed to fetch model states:', err) + } } } -}) +}