From e36129c538edcdcc5e3204414a4c55dfd5203016 Mon Sep 17 00:00:00 2001 From: Eduardo Jose Costa <59846713+EduCosta85@users.noreply.github.com> Date: Sat, 12 Sep 2026 07:45:01 -0300 Subject: [PATCH 1/2] feat(stt): add local speech dictation with Parakeet TDT --- bun.lock | 22 +- package.json | 8 +- .../desktop/src/common/adapter/ipcBridge.ts | 22 ++ .../src/common/types/provider/speech.ts | 24 +- packages/desktop/src/process/bridge/index.ts | 3 + .../src/process/bridge/speechBridge.ts | 37 +++ .../src/process/services/speech/index.ts | 11 + .../services/speech/localSpeechService.ts | 167 +++++++++++ .../process/services/speech/modelCatalog.ts | 66 +++++ .../process/services/speech/modelManager.ts | 269 ++++++++++++++++++ .../process/services/speech/sherpaLoader.ts | 57 ++++ .../src/process/services/speech/types.ts | 32 +++ .../VoiceInputSection/SpeechTestPanel.tsx | 3 + .../VoiceInputSection/index.tsx | 204 ++++++++++--- .../VoiceInputSection/speechSettingsUtils.ts | 31 +- .../renderer/services/SpeechToTextService.ts | 44 +++ .../src/renderer/services/i18n/i18n-keys.d.ts | 4 + .../services/i18n/locales/de-DE/settings.json | 6 +- .../services/i18n/locales/en-US/settings.json | 6 +- .../services/i18n/locales/es-ES/settings.json | 6 +- .../services/i18n/locales/fa-IR/settings.json | 6 +- .../services/i18n/locales/fr-FR/settings.json | 6 +- .../services/i18n/locales/ja-JP/settings.json | 6 +- .../services/i18n/locales/ko-KR/settings.json | 6 +- .../services/i18n/locales/pt-BR/settings.json | 6 +- .../services/i18n/locales/ru-RU/settings.json | 6 +- .../services/i18n/locales/tr-TR/settings.json | 6 +- .../services/i18n/locales/uk-UA/settings.json | 6 +- .../services/i18n/locales/zh-CN/settings.json | 6 +- .../services/i18n/locales/zh-TW/settings.json | 6 +- .../renderer/services/speech/speechModels.ts | 2 + .../services/speech/speechStreamPolicy.ts | 7 + .../unit/renderer/speechSettingsUtils.test.ts | 17 ++ tests/unit/speech/localSpeechService.test.ts | 69 +++++ tests/unit/speech/modelCatalog.test.ts | 48 ++++ tests/unit/speech/modelManager.test.ts | 39 +++ tests/unit/speech/sherpaLoader.test.ts | 34 +++ 37 files changed, 1235 insertions(+), 63 deletions(-) create mode 100644 packages/desktop/src/process/bridge/speechBridge.ts create mode 100644 packages/desktop/src/process/services/speech/index.ts create mode 100644 packages/desktop/src/process/services/speech/localSpeechService.ts create mode 100644 packages/desktop/src/process/services/speech/modelCatalog.ts create mode 100644 packages/desktop/src/process/services/speech/modelManager.ts create mode 100644 packages/desktop/src/process/services/speech/sherpaLoader.ts create mode 100644 packages/desktop/src/process/services/speech/types.ts create mode 100644 tests/unit/speech/localSpeechService.test.ts create mode 100644 tests/unit/speech/modelCatalog.test.ts create mode 100644 tests/unit/speech/modelManager.test.ts create mode 100644 tests/unit/speech/sherpaLoader.test.ts diff --git a/bun.lock b/bun.lock index e442869ce9e..3aa98dc322d 100644 --- a/bun.lock +++ b/bun.lock @@ -167,6 +167,12 @@ "optionalDependencies": { "@rollup/rollup-win32-x64-msvc": "^4.46.2", "electron-winstaller": "^5.4.0", + "sherpa-onnx": "1.13.8", + "sherpa-onnx-darwin-arm64": "1.13.8", + "sherpa-onnx-darwin-x64": "1.13.8", + "sherpa-onnx-linux-arm64": "1.13.8", + "sherpa-onnx-linux-x64": "1.13.8", + "sherpa-onnx-win-x64": "1.13.8", }, }, "packages/desktop": { @@ -1987,8 +1993,6 @@ "encodeurl": ["encodeurl@2.0.0", "", {}, "sha512-Q0n9HRi4m6JuGIV1eFlmvJB7ZEVxu93IrMyiMsGC0lrMJMWzRgx6WGquyfQgZVb31vhGgXnfmPNNXmxnOkRBrg=="], - "encoding": ["encoding@0.1.13", "", { "dependencies": { "iconv-lite": "^0.6.2" } }, "sha512-ETBauow1T35Y/WZMkio9jiM0Z5xjHHmJ4XmjZOq1l/dXz3lr2sRn87nJy20RupqSh1F2m3HHPSp8ShIPQJrJ3A=="], - "end-of-stream": ["end-of-stream@1.4.5", "", { "dependencies": { "once": "^1.4.0" } }, "sha512-ooEGc6HP26xXq/N+GCGOT0JKCLDGrq2bQUZrQ7gyrJiZANJ/8YDTxTpQBXGMn+WbIQXNVpyWymm7KYVICQnyOg=="], "entities": ["entities@6.0.1", "", {}, "sha512-aN97NXWF6AWBTahfVOIrB/NShkzi5H7F9r1s9mD3cDj4Ko5f2qhhVoYMibXF7GlLveb/D2ioWay8lxI97Ven3g=="], @@ -3073,6 +3077,18 @@ "shell-path": ["shell-path@3.1.0", "", { "dependencies": { "shell-env": "^4.0.1" } }, "sha512-s/9q9PEtcRmDTz69+cJ3yYBAe9yGrL7e46gm2bU4pQ9N48ecPK9QrGFnLwYgb4smOHskx4PL7wCNMktW2AoD+g=="], + "sherpa-onnx": ["sherpa-onnx@1.13.8", "", {}, "sha512-bsZ7K55hNakFbCVIrCsuwEFLLiWa/1/KI6dXEK3CwtI6N/j/Zc1LMheo1syMJhG0e4gu7MXPwbtKfbm44BQYIg=="], + + "sherpa-onnx-darwin-arm64": ["sherpa-onnx-darwin-arm64@1.13.8", "", { "os": "darwin", "cpu": "arm64" }, "sha512-FPNgJMgnWVl/KhRTIhG3KL3A4Om63Rn4YKXc9/uHY7SzLcvqLJLc/h7UBWJwduXvv7K18t5NpxHR6XgXn4sjWw=="], + + "sherpa-onnx-darwin-x64": ["sherpa-onnx-darwin-x64@1.13.8", "", { "os": "darwin", "cpu": "x64" }, "sha512-7BLRpjM6w4f9W46/nmkmq8lEKUayhebvcpslCVQ+6QN2uReYlZEMDZlSpXMjme+hUFrPfRz8P3UNq8ep/4d19g=="], + + "sherpa-onnx-linux-arm64": ["sherpa-onnx-linux-arm64@1.13.8", "", { "os": "linux", "cpu": "arm64" }, "sha512-Tlg7a70b/Wge3OF8IgTHF9jhSVCsLyKQKhwc4BsJ5A+dL/SrFtGBjzuHp4XeLhiiOT7afCxX5PdSn/D4c8Lnuw=="], + + "sherpa-onnx-linux-x64": ["sherpa-onnx-linux-x64@1.13.8", "", { "os": "linux", "cpu": "x64" }, "sha512-6plnhjagsSeTntCgnlag86hWbs/uZE9Crms1LgOb68/1nKsIQjMd+WG519m+aPwT6TrsBOiEMzrx41t8sL5L5g=="], + + "sherpa-onnx-win-x64": ["sherpa-onnx-win-x64@1.13.8", "", { "os": "win32", "cpu": "x64" }, "sha512-oZF1c9VPOKtMwn83Bboc5XSWL+76BRoyB3eUuVnCknBKxwSULZU2Foia9VHWzU+n4I12rPsP6z6H9Rp1hD9o8g=="], + "shiki": ["shiki@3.23.0", "", { "dependencies": { "@shikijs/core": "3.23.0", "@shikijs/engine-javascript": "3.23.0", "@shikijs/engine-oniguruma": "3.23.0", "@shikijs/langs": "3.23.0", "@shikijs/themes": "3.23.0", "@shikijs/types": "3.23.0", "@shikijs/vscode-textmate": "^10.0.2", "@types/hast": "^3.0.4" } }, "sha512-55Dj73uq9ZXL5zyeRPzHQsK7Nbyt6Y10k5s7OjuFZGMhpp4r/rsLBH0o/0fstIzX1Lep9VxefWljK/SKCzygIA=="], "side-channel": ["side-channel@1.1.0", "", { "dependencies": { "es-errors": "^1.3.0", "object-inspect": "^1.13.3", "side-channel-list": "^1.0.0", "side-channel-map": "^1.0.1", "side-channel-weakmap": "^1.0.2" } }, "sha512-ZX99e6tRweoUXqR+VBrslhda51Nh5MTQwou5tnUDgbtyM0dBgmhEDtWGP/xbKn6hqfPRHujUNwz5fy/wbbhnpw=="], @@ -3643,8 +3659,6 @@ "electron-winstaller/fs-extra": ["fs-extra@7.0.1", "", { "dependencies": { "graceful-fs": "^4.1.2", "jsonfile": "^4.0.0", "universalify": "^0.1.0" } }, "sha512-YJDaCJZEnBmcbw13fvdAM9AwNOJwOzrE4pqMqBq5nFiEqXUqHwlK4B+3pUw6JNvfSPtX05xFHtYy/1ni01eGCw=="], - "encoding/iconv-lite": ["iconv-lite@0.6.3", "", { "dependencies": { "safer-buffer": ">= 2.1.2 < 3.0.0" } }, "sha512-4fCk79wshMdzMp2rH06qWrJE4iolqLhCUH+OiuIgU++RB0+94NlDL81atO7GX55uUKueo0txHNtvEyI6D7WdMw=="], - "execa/get-stream": ["get-stream@6.0.1", "", {}, "sha512-ts6Wi+2j3jQjqi70w5AlN8DFnkSwC+MqmxEzdEALB2qXZYV3X/b1CTfgPLGJNMeAWxdPfU8FO1ms3NUfaHCPYg=="], "express/cookie": ["cookie@0.7.2", "", {}, "sha512-yki5XnKuf750l50uGTllt6kKILY4nQ1eNIQatoXEByZ5dWgnKqbnqmTrBE5B4N7lrMJKQ2ytWMiTO2o0v6Ew/w=="], diff --git a/package.json b/package.json index e0f58b5fa9c..d9ff7930ce7 100644 --- a/package.json +++ b/package.json @@ -236,7 +236,13 @@ }, "optionalDependencies": { "@rollup/rollup-win32-x64-msvc": "^4.46.2", - "electron-winstaller": "^5.4.0" + "electron-winstaller": "^5.4.0", + "sherpa-onnx": "1.13.8", + "sherpa-onnx-darwin-arm64": "1.13.8", + "sherpa-onnx-darwin-x64": "1.13.8", + "sherpa-onnx-linux-arm64": "1.13.8", + "sherpa-onnx-linux-x64": "1.13.8", + "sherpa-onnx-win-x64": "1.13.8" }, "resolutions": { "@codemirror/language": "6.12.3", diff --git a/packages/desktop/src/common/adapter/ipcBridge.ts b/packages/desktop/src/common/adapter/ipcBridge.ts index f0d1e9f066f..20ec1cd0c90 100644 --- a/packages/desktop/src/common/adapter/ipcBridge.ts +++ b/packages/desktop/src/common/adapter/ipcBridge.ts @@ -49,6 +49,7 @@ import type { ProviderHealthCheckResponse, UpdateProviderRequest, } from '../types/provider/providerApi'; +import type { SpeechModelDownloadStatus, SpeechToTextResult } from '../types/provider/speech'; import type { ITeamAgentRemovedEvent, ITeamAgentRenamedEvent, @@ -2504,3 +2505,24 @@ export const sidebar = { (p) => `/api/sidebar/archived/project/${encodeURIComponent(p.project_id)}` ), }; + +// --------------------------------------------------------------------------- +// Local Speech-to-Text (Parakeet TDT) +// --------------------------------------------------------------------------- + +export const speech = { + checkModel: bridge.buildProvider<{ isReady: boolean; status: SpeechModelDownloadStatus }, { modelId: string }>( + 'speech:checkModel' + ), + downloadModel: bridge.buildProvider('speech:downloadModel'), + cancelDownload: bridge.buildProvider('speech:cancelDownload'), + transcribe: bridge.buildProvider( + 'speech:transcribe' + ), + onDownloadProgress: bridge.buildEmitter<{ + downloadedBytes: number; + modelId: string; + percent: number; + totalBytes: number; + }>('speech:downloadProgress'), +}; diff --git a/packages/desktop/src/common/types/provider/speech.ts b/packages/desktop/src/common/types/provider/speech.ts index ddbdc5b349a..a476b1a7e66 100644 --- a/packages/desktop/src/common/types/provider/speech.ts +++ b/packages/desktop/src/common/types/provider/speech.ts @@ -4,7 +4,24 @@ * SPDX-License-Identifier: Apache-2.0 */ -export type SpeechToTextProvider = 'openai' | 'deepgram'; +export type SpeechToTextProvider = 'openai' | 'deepgram' | 'local'; + +export type LocalSpeechModelId = 'parakeet-tdt-0.6b-v3-int8' | 'parakeet-tdt-0.6b-v2-int8'; + +export type LocalSpeechToTextConfig = { + hotwords?: string; + language?: string; + model: LocalSpeechModelId | string; +}; + +export type SpeechModelDownloadStatus = { + downloadedBytes: number; + error?: string; + modelId: string; + progress: number; + status: 'idle' | 'downloading' | 'ready' | 'error'; + totalBytes: number; +}; export type OpenAISpeechToTextConfig = { api_key: string; @@ -27,10 +44,11 @@ export type DeepgramSpeechToTextConfig = { export type SpeechToTextConfig = { autoSend?: boolean; - enabled: boolean; - provider: SpeechToTextProvider; deepgram?: DeepgramSpeechToTextConfig; + enabled: boolean; + local?: LocalSpeechToTextConfig; openai?: OpenAISpeechToTextConfig; + provider: SpeechToTextProvider; }; export type SpeechToTextAudioBuffer = Uint8Array | number[] | Record; diff --git a/packages/desktop/src/process/bridge/index.ts b/packages/desktop/src/process/bridge/index.ts index f2a571aa7d3..0696cd165cb 100644 --- a/packages/desktop/src/process/bridge/index.ts +++ b/packages/desktop/src/process/bridge/index.ts @@ -12,6 +12,7 @@ import { initWindowControlsBridge } from './windowControlsBridge'; import { initNotificationBridge } from './notificationBridge'; import { initWebuiBridge } from './webuiBridge'; import { initThemeBridge } from './themeBridge'; +import { initSpeechBridge } from './speechBridge'; export type BridgeDependencies = Record; @@ -24,12 +25,14 @@ export function initAllBridges(_deps: BridgeDependencies = {}): void { initNotificationBridge(); initWebuiBridge(); initThemeBridge(); + initSpeechBridge(); } export { initApplicationBridge, initDialogBridge, initNotificationBridge, + initSpeechBridge, initSystemSettingsBridge, initThemeBridge, initUpdateBridge, diff --git a/packages/desktop/src/process/bridge/speechBridge.ts b/packages/desktop/src/process/bridge/speechBridge.ts new file mode 100644 index 00000000000..7f232d2dc8b --- /dev/null +++ b/packages/desktop/src/process/bridge/speechBridge.ts @@ -0,0 +1,37 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { ipcBridge } from '@/common'; +import { + cancelModelDownload, + downloadModel, + getModelStatus, + isModelReady, + transcribeLocalAudio, +} from '../services/speech'; + +export function initSpeechBridge(): void { + ipcBridge.speech.checkModel.provider(async ({ modelId }) => { + const isReady = await isModelReady(modelId); + const status = getModelStatus(modelId); + return { isReady, status }; + }); + + ipcBridge.speech.downloadModel.provider(async ({ modelId }) => { + return downloadModel(modelId, (progress) => { + ipcBridge.speech.onDownloadProgress.emit(progress); + }); + }); + + ipcBridge.speech.cancelDownload.provider(({ modelId }) => { + return cancelModelDownload(modelId); + }); + + ipcBridge.speech.transcribe.provider(async ({ audioBuffer, modelId }) => { + const bytes = audioBuffer instanceof Uint8Array ? audioBuffer : new Uint8Array(audioBuffer); + return transcribeLocalAudio(bytes, modelId); + }); +} diff --git a/packages/desktop/src/process/services/speech/index.ts b/packages/desktop/src/process/services/speech/index.ts new file mode 100644 index 00000000000..e78149b080f --- /dev/null +++ b/packages/desktop/src/process/services/speech/index.ts @@ -0,0 +1,11 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +export * from './types'; +export * from './modelCatalog'; +export * from './sherpaLoader'; +export * from './modelManager'; +export * from './localSpeechService'; diff --git a/packages/desktop/src/process/services/speech/localSpeechService.ts b/packages/desktop/src/process/services/speech/localSpeechService.ts new file mode 100644 index 00000000000..172228abf6b --- /dev/null +++ b/packages/desktop/src/process/services/speech/localSpeechService.ts @@ -0,0 +1,167 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import path from 'node:path'; +import type { SpeechToTextResult } from '@/common/types/provider/speech'; +import { getLocalSpeechModelManifest } from './modelCatalog'; +import { getModelStorageDir, isModelReady } from './modelManager'; +import { loadSherpaAddon } from './sherpaLoader'; + +type CachedRecognizer = { + modelId: string; + recognizer: any; +}; + +let activeRecognizerCache: CachedRecognizer | null = null; + +function getOrInitRecognizer(modelId: string, sherpa: any): any { + if (activeRecognizerCache && activeRecognizerCache.modelId === modelId) { + return activeRecognizerCache.recognizer; + } + + const manifest = getLocalSpeechModelManifest(modelId); + if (!manifest) { + throw new Error(`STT_UNKNOWN_MODEL: ${modelId}`); + } + + const modelDir = getModelStorageDir(modelId); + const encoderPath = path.join(modelDir, 'encoder.int8.onnx'); + const decoderPath = path.join(modelDir, 'decoder.int8.onnx'); + const joinerPath = path.join(modelDir, 'joiner.int8.onnx'); + const tokensPath = path.join(modelDir, 'tokens.txt'); + + const config = { + decodingMethod: 'greedy_search', + featConfig: { + featureDim: 80, + sampleRate: manifest.sampleRate, + }, + modelConfig: { + debug: 0, + numThreads: 4, + provider: 'cpu', + tokens: tokensPath, + transducer: { + decoder: decoderPath, + encoder: encoderPath, + joiner: joinerPath, + }, + }, + }; + + const recognizer = sherpa.createOfflineRecognizer(config); + activeRecognizerCache = { + modelId, + recognizer, + }; + + return recognizer; +} + +function toBuffer(audioBytes: Uint8Array | Buffer | number[] | Record): Buffer { + if (Buffer.isBuffer(audioBytes)) { + return audioBytes; + } + if (audioBytes instanceof Uint8Array) { + return Buffer.from(audioBytes.buffer, audioBytes.byteOffset, audioBytes.byteLength); + } + if (Array.isArray(audioBytes)) { + return Buffer.from(audioBytes); + } + if (typeof audioBytes === 'object' && audioBytes !== null) { + const keys = Object.keys(audioBytes); + const buf = Buffer.alloc(keys.length); + for (let i = 0; i < keys.length; i++) { + buf[i] = (audioBytes as Record)[i]; + } + return buf; + } + return Buffer.alloc(0); +} + +export async function transcribeLocalAudio( + audioBytes: Uint8Array | Buffer | number[] | Record, + modelId = 'parakeet-tdt-0.6b-v3-int8' +): Promise { + const ready = await isModelReady(modelId); + if (!ready) { + throw new Error('STT_LOCAL_MODEL_NOT_DOWNLOADED'); + } + + const sherpa = loadSherpaAddon(); + const recognizer = getOrInitRecognizer(modelId, sherpa); + const manifest = getLocalSpeechModelManifest(modelId); + const targetSampleRate = manifest?.sampleRate ?? 16000; + + const nodeBuffer = toBuffer(audioBytes); + + let samples: Float32Array; + let sampleRate: number; + + try { + if (typeof sherpa.readWaveFromBinary === 'function') { + const wave = sherpa.readWaveFromBinary(nodeBuffer); + samples = wave.samples; + sampleRate = wave.sampleRate; + } else if (typeof sherpa.readWaveFromBinaryData === 'function') { + const wave = sherpa.readWaveFromBinaryData(nodeBuffer); + samples = wave.samples; + sampleRate = wave.sampleRate; + } else { + throw new Error('No wave parser available'); + } + } catch { + // If not a valid WAV header or parsing failed, treat as raw 16-bit PCM (little endian) + const int16 = new Int16Array(nodeBuffer.buffer, nodeBuffer.byteOffset, nodeBuffer.byteLength / 2); + samples = new Float32Array(int16.length); + for (let i = 0; i < int16.length; i++) { + samples[i] = int16[i] / 32768.0; + } + sampleRate = targetSampleRate; + } + + let text = ''; + if (typeof recognizer.createStream === 'function') { + // OOP JS / WASM interface + const stream = recognizer.createStream(); + try { + stream.acceptWaveform(sampleRate, samples); + recognizer.decode(stream); + const result = recognizer.getResult(stream); + const parsed = typeof result === 'string' ? JSON.parse(result) : result; + text = (parsed?.text ?? '').trim(); + } finally { + try { + stream.free?.(); + } catch { + // ignore + } + } + } else { + // Native C++ NAPI addon interface + const stream = sherpa.createOfflineStream(recognizer); + try { + sherpa.acceptWaveformOffline(stream, { sampleRate, samples }); + sherpa.decodeOfflineStream(recognizer, stream); + const resultJson = sherpa.getOfflineStreamResultAsJson(stream); + const parsed = typeof resultJson === 'string' ? JSON.parse(resultJson) : resultJson; + text = (parsed?.text ?? '').trim(); + } finally { + // Native addon manages memory or GC + } + } + + return { + language: manifest?.language, + model: modelId, + provider: 'local', + text, + }; +} + +export function clearLocalRecognizerCache(): void { + activeRecognizerCache = null; +} diff --git a/packages/desktop/src/process/services/speech/modelCatalog.ts b/packages/desktop/src/process/services/speech/modelCatalog.ts new file mode 100644 index 00000000000..8c47f901120 --- /dev/null +++ b/packages/desktop/src/process/services/speech/modelCatalog.ts @@ -0,0 +1,66 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import type { LocalSpeechModelManifest, SpeechModelFileSpec } from './types'; + +const hfFiles = ( + repo: string, + revision: string, + specs: Array<[name: string, sizeBytes: number, sha256: string]> +): SpeechModelFileSpec[] => + specs.map(([name, sizeBytes, sha256]) => ({ + name, + sizeBytes, + sha256, + url: `https://huggingface.co/${repo}/resolve/${revision}/${encodeURIComponent(name)}?download=true`, + })); + +export const LOCAL_SPEECH_MODELS: LocalSpeechModelManifest[] = [ + { + id: 'parakeet-tdt-0.6b-v3-int8', + label: 'Parakeet TDT v3', + description: 'High accuracy for 25 languages (EN, PT, ES, FR, DE, etc.). Automatic punctuation and capitalization.', + type: 'transducer', + language: 'multilingual', + sampleRate: 16000, + modelingUnit: 'bpe', + recommended: true, + totalSizeBytes: 652184281 + 11845275 + 6355277 + 93939, + files: hfFiles( + 'csukuangfj/sherpa-onnx-nemo-parakeet-tdt-0.6b-v3-int8', + '2bda32ec70b097a55adaa07d9a7173915b43cc78', + [ + ['encoder.int8.onnx', 652184281, 'acfc2b4456377e15d04f0243af540b7fe7c992f8d898d751cf134c3a55fd2247'], + ['decoder.int8.onnx', 11845275, '179e50c43d1a9de79c8a24149a2f9bac6eb5981823f2a2ed88d655b24248db4e'], + ['joiner.int8.onnx', 6355277, '3164c13fc2821009440d20fcb5fdc78bff28b4db2f8d0f0b329101719c0948b3'], + ['tokens.txt', 93939, 'd58544679ea4bc6ac563d1f545eb7d474bd6cfa467f0a6e2c1dc1c7d37e3c35d'], + ] + ), + }, + { + id: 'parakeet-tdt-0.6b-v2-int8', + label: 'Parakeet TDT v2', + description: 'English-only FastConformer TDT 0.6B INT8. Faster and lighter English transcription.', + type: 'transducer', + language: 'en', + sampleRate: 16000, + modelingUnit: 'bpe', + totalSizeBytes: 652184296 + 7257753 + 1739080 + 93939, + files: hfFiles( + 'csukuangfj/sherpa-onnx-nemo-parakeet-tdt-0.6b-v2-int8', + '1ab9323565ddb038682214b292f588070a538ce2', + [ + ['encoder.int8.onnx', 652184296, 'a32b12d17bbbc309d0686fbbcc2987b5e9b8333a7da83fa6b089f0a2acd651ab'], + ['decoder.int8.onnx', 7257753, 'b6bb64963457237b900e496ee9994b59294526439fbcc1fecf705b31a15c6b4e'], + ['joiner.int8.onnx', 1739080, '7946164367946e7f9f29a122407c3252b680dbae9a51343eb2488d057c3c43d2'], + ['tokens.txt', 93939, 'd58544679ea4bc6ac563d1f545eb7d474bd6cfa467f0a6e2c1dc1c7d37e3c35d'], + ] + ), + }, +]; + +export const getLocalSpeechModelManifest = (modelId: string): LocalSpeechModelManifest | undefined => + LOCAL_SPEECH_MODELS.find((m) => m.id === modelId) ?? LOCAL_SPEECH_MODELS[0]; diff --git a/packages/desktop/src/process/services/speech/modelManager.ts b/packages/desktop/src/process/services/speech/modelManager.ts new file mode 100644 index 00000000000..e8204acff44 --- /dev/null +++ b/packages/desktop/src/process/services/speech/modelManager.ts @@ -0,0 +1,269 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { createHash } from 'node:crypto'; +import { createReadStream, createWriteStream, existsSync, promises as fs } from 'node:fs'; +import os from 'node:os'; +import path from 'node:path'; +import { Readable } from 'node:stream'; +import { pipeline } from 'node:stream/promises'; +import type { SpeechModelDownloadStatus } from '@/common/types/provider/speech'; +import { getLocalSpeechModelManifest } from './modelCatalog'; +import type { ModelDownloadProgressCallback, SpeechModelFileSpec } from './types'; + +function getStorageBaseDir(): string { + try { + // eslint-disable-next-line @typescript-eslint/no-require-imports + const electron = require('electron'); + const electronApp = electron?.app || electron; + if (typeof electronApp?.getPath === 'function') { + return electronApp.getPath('userData'); + } + } catch { + // Electron app may not be initialized in tests or worker threads + } + return path.join(os.homedir(), '.aionui'); +} + +export function getModelStorageDir(modelId: string): string { + return path.join(getStorageBaseDir(), 'models', 'stt', modelId); +} + +async function verifyFileSha256(filePath: string, expectedSha256: string): Promise { + if (!existsSync(filePath)) return false; + return new Promise((resolve) => { + const hash = createHash('sha256'); + const stream = createReadStream(filePath); + stream.on('data', (chunk) => hash.update(chunk)); + stream.on('end', () => { + const digest = hash.digest('hex'); + resolve(digest.toLowerCase() === expectedSha256.toLowerCase()); + }); + stream.on('error', () => resolve(false)); + }); +} + +export async function isModelReady(modelId: string): Promise { + const manifest = getLocalSpeechModelManifest(modelId); + if (!manifest) return false; + + const modelDir = getModelStorageDir(modelId); + if (!existsSync(modelDir)) return false; + + for (const file of manifest.files) { + const filePath = path.join(modelDir, file.name); + if (!existsSync(filePath)) return false; + try { + const stat = await fs.stat(filePath); + if (stat.size !== file.sizeBytes) return false; + } catch { + return false; + } + } + + return true; +} + +const activeDownloads = new Map(); + +export function getModelStatus(modelId: string): SpeechModelDownloadStatus { + const active = activeDownloads.get(modelId); + if (active) { + return active.progress; + } + const manifest = getLocalSpeechModelManifest(modelId); + const totalBytes = manifest?.totalSizeBytes ?? 0; + return { + modelId, + status: 'idle', + progress: 0, + downloadedBytes: 0, + totalBytes, + }; +} + +const DOWNLOAD_USER_AGENT = + 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36'; + +async function downloadSingleFile( + spec: SpeechModelFileSpec, + destPath: string, + onChunk: (bytesRead: number) => void, + signal: AbortSignal +): Promise { + const tmpPath = `${destPath}.tmp`; + const primaryUrl = spec.url; + const mirrorUrl = spec.url.replace('https://huggingface.co/', 'https://hf-mirror.com/'); + + let response: Response | null = null; + try { + response = await fetch(primaryUrl, { signal, headers: { 'User-Agent': DOWNLOAD_USER_AGENT } }); + if (!response.ok) { + throw new Error(`HTTP ${response.status} ${response.statusText}`); + } + } catch (primaryError) { + // Try fallback mirror + try { + response = await fetch(mirrorUrl, { signal, headers: { 'User-Agent': DOWNLOAD_USER_AGENT } }); + if (!response.ok) { + throw new Error(`HTTP ${response.status} ${response.statusText}`, { cause: primaryError }); + } + } catch { + throw new Error(`Failed to download ${spec.name} from primary and mirror: ${String(primaryError)}`); + } + } + + if (!response.body) { + throw new Error(`Empty response body for ${spec.name}`); + } + + const fileStream = createWriteStream(tmpPath); + const nodeReadable = Readable.fromWeb(response.body as any); + + nodeReadable.on('data', (chunk: Buffer) => { + onChunk(chunk.length); + }); + + try { + await pipeline(nodeReadable, fileStream); + } catch (err) { + try { + await fs.unlink(tmpPath); + } catch { + // ignore + } + throw err; + } + + // Verify sha256 + const isValid = await verifyFileSha256(tmpPath, spec.sha256); + if (!isValid) { + try { + await fs.unlink(tmpPath); + } catch { + // ignore + } + throw new Error(`Checksum mismatch for file ${spec.name}`); + } + + await fs.rename(tmpPath, destPath); +} + +export async function downloadModel( + modelId: string, + onProgress?: ModelDownloadProgressCallback +): Promise { + const manifest = getLocalSpeechModelManifest(modelId); + if (!manifest) { + throw new Error(`Unknown model ID: ${modelId}`); + } + + if (await isModelReady(modelId)) { + const readyStatus: SpeechModelDownloadStatus = { + modelId, + status: 'ready', + progress: 100, + downloadedBytes: manifest.totalSizeBytes, + totalBytes: manifest.totalSizeBytes, + }; + onProgress?.({ + modelId, + percent: 100, + downloadedBytes: manifest.totalSizeBytes, + totalBytes: manifest.totalSizeBytes, + }); + return readyStatus; + } + + if (activeDownloads.has(modelId)) { + return activeDownloads.get(modelId)!.progress; + } + + const abortController = new AbortController(); + const progressStatus: SpeechModelDownloadStatus = { + modelId, + status: 'downloading', + progress: 0, + downloadedBytes: 0, + totalBytes: manifest.totalSizeBytes, + }; + + activeDownloads.set(modelId, { abortController, progress: progressStatus }); + + const modelDir = getModelStorageDir(modelId); + await fs.mkdir(modelDir, { recursive: true }); + + let cumulativeBytes = 0; + + try { + for (const file of manifest.files) { + const destPath = path.join(modelDir, file.name); + + // Check if this specific file already exists and is valid + if (existsSync(destPath)) { + const stat = await fs.stat(destPath); + if (stat.size === file.sizeBytes) { + cumulativeBytes += file.sizeBytes; + progressStatus.downloadedBytes = cumulativeBytes; + progressStatus.progress = Math.round((cumulativeBytes / manifest.totalSizeBytes) * 100); + onProgress?.({ + modelId, + percent: progressStatus.progress, + downloadedBytes: cumulativeBytes, + totalBytes: manifest.totalSizeBytes, + }); + continue; + } + } + + await downloadSingleFile( + file, + destPath, + (bytesRead) => { + cumulativeBytes += bytesRead; + progressStatus.downloadedBytes = cumulativeBytes; + progressStatus.progress = Math.min(99, Math.round((cumulativeBytes / manifest.totalSizeBytes) * 100)); + onProgress?.({ + modelId, + percent: progressStatus.progress, + downloadedBytes: cumulativeBytes, + totalBytes: manifest.totalSizeBytes, + }); + }, + abortController.signal + ); + } + + progressStatus.status = 'ready'; + progressStatus.progress = 100; + progressStatus.downloadedBytes = manifest.totalSizeBytes; + + onProgress?.({ + modelId, + percent: 100, + downloadedBytes: manifest.totalSizeBytes, + totalBytes: manifest.totalSizeBytes, + }); + + return progressStatus; + } catch (error) { + progressStatus.status = 'error'; + progressStatus.error = error instanceof Error ? error.message : String(error); + throw error; + } finally { + activeDownloads.delete(modelId); + } +} + +export function cancelModelDownload(modelId: string): boolean { + const active = activeDownloads.get(modelId); + if (active) { + active.abortController.abort(); + activeDownloads.delete(modelId); + return true; + } + return false; +} diff --git a/packages/desktop/src/process/services/speech/sherpaLoader.ts b/packages/desktop/src/process/services/speech/sherpaLoader.ts new file mode 100644 index 00000000000..525ae15e70b --- /dev/null +++ b/packages/desktop/src/process/services/speech/sherpaLoader.ts @@ -0,0 +1,57 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { createRequire } from 'node:module'; + +const req = createRequire(import.meta.url); + +export function getSherpaPackageName(): string | null { + const { platform, arch } = process; + if (platform === 'darwin') { + return arch === 'arm64' ? 'sherpa-onnx-darwin-arm64' : 'sherpa-onnx-darwin-x64'; + } + if (platform === 'linux') { + return arch === 'arm64' ? 'sherpa-onnx-linux-arm64' : 'sherpa-onnx-linux-x64'; + } + if (platform === 'win32') { + return arch === 'x64' ? 'sherpa-onnx-win-x64' : null; + } + return null; +} + +let cachedSherpaAddon: unknown = null; + +export function loadSherpaAddon(): any { + if (cachedSherpaAddon) { + return cachedSherpaAddon; + } + + const pkgName = getSherpaPackageName(); + if (!pkgName) { + throw new Error(`Unsupported platform/architecture for local speech: ${process.platform}-${process.arch}`); + } + + try { + cachedSherpaAddon = req(pkgName); + return cachedSherpaAddon; + } catch (err1) { + try { + cachedSherpaAddon = req('sherpa-onnx'); + return cachedSherpaAddon; + } catch { + throw new Error(`Failed to load native speech module ${pkgName}: ${String(err1)}`); + } + } +} + +export function isSherpaSupported(): boolean { + try { + loadSherpaAddon(); + return true; + } catch { + return false; + } +} diff --git a/packages/desktop/src/process/services/speech/types.ts b/packages/desktop/src/process/services/speech/types.ts new file mode 100644 index 00000000000..a0806fad61f --- /dev/null +++ b/packages/desktop/src/process/services/speech/types.ts @@ -0,0 +1,32 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +export type SpeechModelFileSpec = { + name: string; + sha256: string; + sizeBytes: number; + url: string; +}; + +export type LocalSpeechModelManifest = { + description: string; + files: SpeechModelFileSpec[]; + id: string; + label: string; + language: string; + modelingUnit: 'bpe' | 'cjkchar' | 'cjkchar+bpe'; + recommended?: boolean; + sampleRate: number; + totalSizeBytes: number; + type: 'transducer'; +}; + +export type ModelDownloadProgressCallback = (progress: { + downloadedBytes: number; + modelId: string; + percent: number; + totalBytes: number; +}) => void; diff --git a/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/SpeechTestPanel.tsx b/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/SpeechTestPanel.tsx index 8f9dd525966..b31412be1d9 100644 --- a/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/SpeechTestPanel.tsx +++ b/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/SpeechTestPanel.tsx @@ -62,6 +62,9 @@ const SpeechTestPanel: React.FC = ({ config, source }) => }, [isRecording, recordingDurationMs, stopRecording]); const validate = useCallback((): string | null => { + if (source === 'local') { + return null; + } if (source === 'custom') { if (!isValidHttpUrl(config.openai?.base_url ?? '')) { return t('settings.speechToTextBaseUrlInvalid'); diff --git a/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/index.tsx b/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/index.tsx index 50a9e29884e..b351a9e45d4 100644 --- a/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/index.tsx +++ b/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/index.tsx @@ -4,18 +4,20 @@ * SPDX-License-Identifier: Apache-2.0 */ -import type { SpeechToTextConfig } from '@/common/types/provider/speech'; +import { ipcBridge } from '@/common'; +import type { SpeechModelDownloadStatus, SpeechToTextConfig } from '@/common/types/provider/speech'; import AionSelect from '@/renderer/components/base/AionSelect'; import { SPEECH_TO_TEXT_CONFIG_CHANGED_EVENT } from '@/renderer/services/SpeechToTextService'; import { getClientBusinessSetting, setClientBusinessSetting } from '@/renderer/services/clientBusinessSettings'; import { getModelStreamCapability } from '@/renderer/services/speech/speechStreamPolicy'; -import { Divider, Form, Input, Switch } from '@arco-design/web-react'; +import { Button, Divider, Form, Input, Message, Progress, Switch } from '@arco-design/web-react'; import React, { useCallback, useEffect, useRef, useState } from 'react'; import { useTranslation } from 'react-i18next'; import SpeechTestPanel from './SpeechTestPanel'; import { DEEPGRAM_SPEECH_MODEL_PRESETS, DEFAULT_SPEECH_TO_TEXT_CONFIG, + LOCAL_SPEECH_MODEL_PRESETS, OPENAI_SPEECH_MODEL_PRESETS, SPEECH_LANGUAGE_OPTIONS, applySpeechSource, @@ -134,28 +136,122 @@ const VoiceInputSection: React.FC = () => { [updateConfig] ); + type LocalField = keyof NonNullable; + + const handleLocalChange = useCallback( + (field: LocalField, value: string) => { + updateConfig( + (current) => + ({ + ...current, + local: { ...DEFAULT_SPEECH_TO_TEXT_CONFIG.local, ...current.local, [field]: value }, + }) as SpeechToTextConfig + ); + }, + [updateConfig] + ); + + const isLocal = source === 'local'; const isDeepgram = source === 'deepgram'; const isCustom = source === 'custom'; - const activeLanguage = (isDeepgram ? config.deepgram?.language : config.openai?.language) ?? ''; - const activeModel = (isDeepgram ? config.deepgram?.model : config.openai?.model) ?? ''; + const activeLanguage = + (isLocal ? config.local?.language : isDeepgram ? config.deepgram?.language : config.openai?.language) ?? ''; + const activeModel = + (isLocal ? config.local?.model : isDeepgram ? config.deepgram?.model : config.openai?.model) ?? + (isLocal ? 'parakeet-tdt-0.6b-v3-int8' : ''); const activeApiKey = (isDeepgram ? config.deepgram?.api_key : config.openai?.api_key) ?? ''; - const modelPresets = isDeepgram ? DEEPGRAM_SPEECH_MODEL_PRESETS : OPENAI_SPEECH_MODEL_PRESETS; + const modelPresets = isLocal + ? LOCAL_SPEECH_MODEL_PRESETS + : isDeepgram + ? DEEPGRAM_SPEECH_MODEL_PRESETS + : OPENAI_SPEECH_MODEL_PRESETS; const customBaseUrl = config.openai?.base_url ?? ''; const isBaseUrlInvalid = isCustom && customBaseUrl.trim() !== '' && !isValidHttpUrl(customBaseUrl); + const [localModelStatus, setLocalModelStatus] = useState(null); + const [isDownloading, setIsDownloading] = useState(false); + + useEffect(() => { + if (!isLocal || !activeModel || !ipcBridge.speech) { + return; + } + let cancelled = false; + + const check = async () => { + try { + const res = await ipcBridge.speech?.checkModel?.invoke({ modelId: activeModel }); + if (!cancelled && res?.status) { + setLocalModelStatus(res.status); + } + } catch { + // ignore + } + }; + + void check(); + + const off = ipcBridge.speech?.onDownloadProgress?.on?.((progress) => { + if (progress.modelId === activeModel) { + setLocalModelStatus({ + modelId: progress.modelId, + status: progress.percent >= 100 ? 'ready' : 'downloading', + progress: progress.percent, + downloadedBytes: progress.downloadedBytes, + totalBytes: progress.totalBytes, + }); + if (progress.percent >= 100) { + setIsDownloading(false); + } + } + }); + + return () => { + cancelled = true; + off?.(); + }; + }, [activeModel, isLocal]); + + const handleDownloadModel = useCallback(async () => { + if (!activeModel || isDownloading || !ipcBridge.speech) { + return; + } + try { + setIsDownloading(true); + await ipcBridge.speech.downloadModel.invoke({ modelId: activeModel }); + setLocalModelStatus({ + modelId: activeModel, + status: 'ready', + progress: 100, + downloadedBytes: 0, + totalBytes: 0, + }); + Message.success(t('settings.speechToTextLocalModelReady')); + } catch (err) { + Message.error(err instanceof Error ? err.message : String(err)); + } finally { + setIsDownloading(false); + } + }, [activeModel, isDownloading, t]); + const handleModelChange = useCallback( (value: string) => { - if (isDeepgram) { + if (isLocal) { + handleLocalChange('model', value); + } else if (isDeepgram) { handleDeepgramChange('model', value); } else { handleOpenAIChange('model', value); } }, - [handleDeepgramChange, handleOpenAIChange, isDeepgram] + [handleDeepgramChange, handleLocalChange, handleOpenAIChange, isDeepgram, isLocal] ); const handleLanguageChange = useCallback( (value: string) => { + if (isLocal) { + handleLocalChange('language', value); + return; + } if (isDeepgram) { handleDeepgramChange('language', value); return; @@ -175,7 +271,7 @@ const VoiceInputSection: React.FC = () => { }) as SpeechToTextConfig ); }, - [handleDeepgramChange, isDeepgram, updateConfig] + [handleDeepgramChange, handleLocalChange, isDeepgram, isLocal, updateConfig] ); const handleApiKeyChange = useCallback( @@ -209,13 +305,14 @@ const VoiceInputSection: React.FC = () => {
+ {t('settings.speechToTextSourceLocal')} {t('settings.speechToTextSourceOpenAI')} {t('settings.speechToTextSourceDeepgram')} {t('settings.speechToTextSourceCustom')} - {isCustom && ( + {!isLocal && isCustom && ( } validateStatus={isBaseUrlInvalid ? 'error' : undefined} @@ -229,38 +326,67 @@ const VoiceInputSection: React.FC = () => { )} - - } - > - - + {!isLocal && ( + + } + > + + + )} - - {buildModelOptions(modelPresets, activeModel).map((model) => { - const capability = getModelStreamCapability(source, model); - const badgeText = - capability === 'supported' - ? t('settings.speechToTextStreamingBadge') - : capability === 'unsupported' - ? t('settings.speechToTextWholeBadge') - : null; - return ( - - {model} - {badgeText !== null && {badgeText}} - - ); - })} - +
+ + {buildModelOptions(modelPresets, activeModel).map((model) => { + const capability = isLocal ? 'unsupported' : getModelStreamCapability(source, model); + const badgeText = + capability === 'supported' + ? t('settings.speechToTextStreamingBadge') + : capability === 'unsupported' + ? t('settings.speechToTextWholeBadge') + : null; + return ( + + {model} + {badgeText !== null && {badgeText}} + + ); + })} + + + {isLocal && ( +
+ {localModelStatus?.status === 'ready' ? ( + + ✓ {t('settings.speechToTextLocalModelReady')} + + ) : isDownloading || localModelStatus?.status === 'downloading' ? ( +
+
+ {t('settings.speechToTextDownloadingModel')} + {localModelStatus?.progress ?? 0}% +
+ +
+ ) : ( +
+ Parakeet TDT (~670 MB) + +
+ )} +
+ )} +
diff --git a/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/speechSettingsUtils.ts b/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/speechSettingsUtils.ts index e3fed78ed65..989730c3663 100644 --- a/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/speechSettingsUtils.ts +++ b/packages/desktop/src/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection/speechSettingsUtils.ts @@ -5,10 +5,14 @@ */ import type { SpeechToTextConfig } from '@/common/types/provider/speech'; -export { DEEPGRAM_SPEECH_MODEL_PRESETS, OPENAI_SPEECH_MODEL_PRESETS } from '@renderer/services/speech/speechModels'; +export { + DEEPGRAM_SPEECH_MODEL_PRESETS, + LOCAL_SPEECH_MODEL_PRESETS, + OPENAI_SPEECH_MODEL_PRESETS, +} from '@renderer/services/speech/speechModels'; /** UI-level service source. 'custom' is stored as provider:'openai' + non-empty base_url. */ -export type SpeechSource = 'openai' | 'deepgram' | 'custom'; +export type SpeechSource = 'openai' | 'deepgram' | 'custom' | 'local'; /** Language autonyms are intentionally not translated. Empty value = auto detect. */ export const SPEECH_LANGUAGE_OPTIONS: Array<{ value: string; label?: string }> = [ @@ -59,7 +63,11 @@ export const migrateSpeechLanguage = (config: SpeechToTextConfig): SpeechToTextC export const DEFAULT_SPEECH_TO_TEXT_CONFIG: SpeechToTextConfig = { enabled: false, - provider: 'openai', + provider: 'local', + local: { + model: 'parakeet-tdt-0.6b-v3-int8', + language: '', + }, openai: { api_key: '', base_url: '', @@ -80,6 +88,10 @@ export const DEFAULT_SPEECH_TO_TEXT_CONFIG: SpeechToTextConfig = { export const normalizeSpeechToTextConfig = (config?: Partial): SpeechToTextConfig => ({ ...DEFAULT_SPEECH_TO_TEXT_CONFIG, ...config, + local: { + ...DEFAULT_SPEECH_TO_TEXT_CONFIG.local, + ...config?.local, + }, openai: { ...DEFAULT_SPEECH_TO_TEXT_CONFIG.openai, ...config?.openai, @@ -91,6 +103,9 @@ export const normalizeSpeechToTextConfig = (config?: Partial }); export const deriveSpeechSource = (config: SpeechToTextConfig): SpeechSource => { + if (config.provider === 'local') { + return 'local'; + } if (config.provider === 'deepgram') { return 'deepgram'; } @@ -107,6 +122,16 @@ export const applySpeechSource = ( source: SpeechSource, rememberedCustomBaseUrl = '' ): SpeechToTextConfig => { + if (source === 'local') { + return { + ...config, + provider: 'local', + local: { + ...DEFAULT_SPEECH_TO_TEXT_CONFIG.local, + ...config.local, + }, + }; + } if (source === 'deepgram') { return { ...config, provider: 'deepgram' }; } diff --git a/packages/desktop/src/renderer/services/SpeechToTextService.ts b/packages/desktop/src/renderer/services/SpeechToTextService.ts index 3ab48424303..a0f14d04741 100644 --- a/packages/desktop/src/renderer/services/SpeechToTextService.ts +++ b/packages/desktop/src/renderer/services/SpeechToTextService.ts @@ -4,8 +4,11 @@ * SPDX-License-Identifier: Apache-2.0 */ +import { ipcBridge } from '@/common'; import { getBaseUrl } from '@/common/adapter/httpBridge'; import type { SpeechToTextResult } from '@/common/types/provider/speech'; +import { getClientBusinessSetting } from './clientBusinessSettings'; +import { encodeWavPcm16, floatTo16BitPcm, resampleLinear } from './speech/pcmRecorder'; /** Dispatched on window whenever the speech-to-text config is saved. */ export const SPEECH_TO_TEXT_CONFIG_CHANGED_EVENT = 'aionui:speech-to-text-config-changed'; @@ -76,9 +79,50 @@ const parseErrorResponse = (response: XMLHttpRequest): Error => { return new Error(`STT_REQUEST_FAILED:${response.status} ${response.statusText}`); }; +async function convertBlobToWav(blob: Blob, targetSampleRate = 16000): Promise { + try { + const arrayBuffer = await blob.arrayBuffer(); + const AudioContextClass = + window.AudioContext || (window as unknown as { webkitAudioContext: typeof AudioContext }).webkitAudioContext; + if (!AudioContextClass) { + return blob; + } + const audioContext = new AudioContextClass(); + try { + const decoded = await audioContext.decodeAudioData(arrayBuffer.slice(0)); + const channelData = decoded.getChannelData(0); + const resampled = resampleLinear(channelData, decoded.sampleRate, targetSampleRate); + const pcm16 = floatTo16BitPcm(resampled); + return encodeWavPcm16(new Uint8Array(pcm16.buffer), targetSampleRate, 1); + } finally { + void audioContext.close(); + } + } catch { + return blob; + } +} + export async function transcribeAudioBlob(blob: Blob, languageHint?: string): Promise { ensureAudioSize(blob); + try { + const config = await getClientBusinessSetting('tools.speechToText'); + if (config?.provider === 'local') { + const wavBlob = await convertBlobToWav(blob, 16000); + const arrayBuffer = await wavBlob.arrayBuffer(); + const modelId = config.local?.model || 'parakeet-tdt-0.6b-v3-int8'; + return await ipcBridge.speech.transcribe.invoke({ + audioBuffer: Array.from(new Uint8Array(arrayBuffer)), + modelId, + }); + } + } catch (error) { + if (error instanceof Error && error.message.includes('STT_LOCAL_MODEL_NOT_DOWNLOADED')) { + throw error; + } + // If not local or failed reading settings, continue to cloud fallback + } + const mimeType = blob.type || 'audio/webm'; const file_name = createAudioFileName(mimeType); diff --git a/packages/desktop/src/renderer/services/i18n/i18n-keys.d.ts b/packages/desktop/src/renderer/services/i18n/i18n-keys.d.ts index 677ae71324f..a4a4fde898b 100644 --- a/packages/desktop/src/renderer/services/i18n/i18n-keys.d.ts +++ b/packages/desktop/src/renderer/services/i18n/i18n-keys.d.ts @@ -2543,8 +2543,11 @@ export type I18nKey = | 'settings.speechToTextBaseUrlInvalid' | 'settings.speechToTextBaseUrlPlaceholder' | 'settings.speechToTextDescription' + | 'settings.speechToTextDownloadModel' + | 'settings.speechToTextDownloadingModel' | 'settings.speechToTextLanguage' | 'settings.speechToTextLanguageAuto' + | 'settings.speechToTextLocalModelReady' | 'settings.speechToTextModel' | 'settings.speechToTextModelPlaceholder' | 'settings.speechToTextOptional' @@ -2552,6 +2555,7 @@ export type I18nKey = | 'settings.speechToTextSource' | 'settings.speechToTextSourceCustom' | 'settings.speechToTextSourceDeepgram' + | 'settings.speechToTextSourceLocal' | 'settings.speechToTextSourceOpenAI' | 'settings.speechToTextStreamingBadge' | 'settings.speechToTextTest' diff --git a/packages/desktop/src/renderer/services/i18n/locales/de-DE/settings.json b/packages/desktop/src/renderer/services/i18n/locales/de-DE/settings.json index ff225d976aa..76c859f6b24 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/de-DE/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/de-DE/settings.json @@ -1414,5 +1414,9 @@ "crossSessionMessageDesc": "Erlaubt Agenten, Nachrichten an Ihre anderen Unterhaltungen zu senden. Beim Ausschalten wird auch die @@-Erwähnung im Eingabefeld deaktiviert.", "crossSessionMessageDisabledBanner": "Nachrichten zwischen Unterhaltungen sind aus, die @@-Erwähnung ist nicht verfügbar.", "crossSessionMessageResume": "Wieder einschalten", - "crossSessionMessageUpdateFailed": "Einstellung konnte nicht aktualisiert werden" + "crossSessionMessageUpdateFailed": "Einstellung konnte nicht aktualisiert werden", + "speechToTextSourceLocal": "Lokal (Offline - Parakeet TDT)", + "speechToTextDownloadModel": "Modell herunterladen (~670 MB)", + "speechToTextDownloadingModel": "Modell wird heruntergeladen...", + "speechToTextLocalModelReady": "Modell bereit für Offline-Diktat" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/en-US/settings.json b/packages/desktop/src/renderer/services/i18n/locales/en-US/settings.json index 63fe6cb5a45..9cd150e21b8 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/en-US/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/en-US/settings.json @@ -1414,5 +1414,9 @@ "crossSessionMessageDesc": "Let agents send messages to your other conversations. Turning this off also disables the @@ conversation mention in the input box.", "crossSessionMessageDisabledBanner": "Cross-conversation messages are off, so the @@ conversation mention is unavailable.", "crossSessionMessageResume": "Turn back on", - "crossSessionMessageUpdateFailed": "Failed to update the cross-conversation message setting" + "crossSessionMessageUpdateFailed": "Failed to update the cross-conversation message setting", + "speechToTextSourceLocal": "Local (Offline - Parakeet TDT)", + "speechToTextDownloadModel": "Download Model (~670 MB)", + "speechToTextDownloadingModel": "Downloading model...", + "speechToTextLocalModelReady": "Model ready for offline dictation" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/es-ES/settings.json b/packages/desktop/src/renderer/services/i18n/locales/es-ES/settings.json index a0390e0e500..85d6f65f68b 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/es-ES/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/es-ES/settings.json @@ -1414,5 +1414,9 @@ "crossSessionMessageDesc": "Permite que los agentes envíen mensajes a tus otras conversaciones. Al desactivarlo también se desactiva la mención @@ en el cuadro de texto.", "crossSessionMessageDisabledBanner": "Los mensajes entre conversaciones están desactivados, la mención @@ no está disponible.", "crossSessionMessageResume": "Volver a activar", - "crossSessionMessageUpdateFailed": "No se pudo actualizar el ajuste de mensajes entre conversaciones" + "crossSessionMessageUpdateFailed": "No se pudo actualizar el ajuste de mensajes entre conversaciones", + "speechToTextSourceLocal": "Local (Sin conexión - Parakeet TDT)", + "speechToTextDownloadModel": "Descargar modelo (~670 MB)", + "speechToTextDownloadingModel": "Descargando modelo...", + "speechToTextLocalModelReady": "Modelo listo para dictado sin conexión" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/fa-IR/settings.json b/packages/desktop/src/renderer/services/i18n/locales/fa-IR/settings.json index 4299495a08e..e2a03b46fc0 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/fa-IR/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/fa-IR/settings.json @@ -1414,5 +1414,9 @@ "crossSessionMessageDesc": "به عامل‌ها اجازه می‌دهد به گفتگوهای دیگر شما پیام بفرستند. با خاموش کردن، منشن @@ در کادر ورودی نیز غیرفعال می‌شود.", "crossSessionMessageDisabledBanner": "پیام‌های بین گفتگوها خاموش است، منشن @@ در دسترس نیست.", "crossSessionMessageResume": "روشن کردن دوباره", - "crossSessionMessageUpdateFailed": "به‌روزرسانی تنظیم پیام‌های بین گفتگوها ناموفق بود" + "crossSessionMessageUpdateFailed": "به‌روزرسانی تنظیم پیام‌های بین گفتگوها ناموفق بود", + "speechToTextSourceLocal": "محلی (آفلاین - Parakeet TDT)", + "speechToTextDownloadModel": "دانلود مدل (~670 مگابایت)", + "speechToTextDownloadingModel": "در حال دانلود مدل...", + "speechToTextLocalModelReady": "مدل برای دیکته آفلاین آماده است" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/fr-FR/settings.json b/packages/desktop/src/renderer/services/i18n/locales/fr-FR/settings.json index 1ef4101f40b..c22b759bbaa 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/fr-FR/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/fr-FR/settings.json @@ -1414,5 +1414,9 @@ "crossSessionMessageDesc": "Permet aux agents d'envoyer des messages à vos autres conversations. La désactivation coupe aussi la mention @@ dans la zone de saisie.", "crossSessionMessageDisabledBanner": "Les messages entre conversations sont désactivés, la mention @@ est indisponible.", "crossSessionMessageResume": "Réactiver", - "crossSessionMessageUpdateFailed": "Échec de la mise à jour du réglage des messages entre conversations" + "crossSessionMessageUpdateFailed": "Échec de la mise à jour du réglage des messages entre conversations", + "speechToTextSourceLocal": "Local (Hors ligne - Parakeet TDT)", + "speechToTextDownloadModel": "Télécharger le modèle (~670 Mo)", + "speechToTextDownloadingModel": "Téléchargement du modèle...", + "speechToTextLocalModelReady": "Modèle prêt pour la dictée hors ligne" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/ja-JP/settings.json b/packages/desktop/src/renderer/services/i18n/locales/ja-JP/settings.json index cba69252b22..bdf1375088c 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/ja-JP/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/ja-JP/settings.json @@ -1416,5 +1416,9 @@ "crossSessionMessageDesc": "エージェントが他の会話にメッセージを送れるようにします。オフにすると入力欄の @@ 会話メンションも無効になります。", "crossSessionMessageDisabledBanner": "会話間メッセージはオフです。入力欄の @@ 会話メンションは使用できません。", "crossSessionMessageResume": "再度オンにする", - "crossSessionMessageUpdateFailed": "会話間メッセージ設定の更新に失敗しました" + "crossSessionMessageUpdateFailed": "会話間メッセージ設定の更新に失敗しました", + "speechToTextSourceLocal": "ローカル(オフライン - Parakeet TDT)", + "speechToTextDownloadModel": "モデルをダウンロード (~670 MB)", + "speechToTextDownloadingModel": "モデルをダウンロード中...", + "speechToTextLocalModelReady": "オフライン音声認識の準備が完了しました" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/ko-KR/settings.json b/packages/desktop/src/renderer/services/i18n/locales/ko-KR/settings.json index a1545a23ec0..dc69ada5b78 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/ko-KR/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/ko-KR/settings.json @@ -1416,5 +1416,9 @@ "crossSessionMessageDesc": "에이전트가 다른 대화로 메시지를 보낼 수 있게 합니다. 끄면 입력창의 @@ 대화 멘션도 비활성화됩니다.", "crossSessionMessageDisabledBanner": "대화 간 메시지가 꺼져 있어 입력창의 @@ 대화 멘션을 사용할 수 없습니다.", "crossSessionMessageResume": "다시 켜기", - "crossSessionMessageUpdateFailed": "대화 간 메시지 설정 업데이트에 실패했습니다" + "crossSessionMessageUpdateFailed": "대화 간 메시지 설정 업데이트에 실패했습니다", + "speechToTextSourceLocal": "로컬 (오프라인 - Parakeet TDT)", + "speechToTextDownloadModel": "모델 다운로드 (~670 MB)", + "speechToTextDownloadingModel": "모델 다운로드 중...", + "speechToTextLocalModelReady": "오프라인 음성 인식이 준비되었습니다" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/pt-BR/settings.json b/packages/desktop/src/renderer/services/i18n/locales/pt-BR/settings.json index 6579095343b..72dbc60e8e0 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/pt-BR/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/pt-BR/settings.json @@ -1417,5 +1417,9 @@ "crossSessionMessageDesc": "Permite que agentes enviem mensagens para suas outras conversas. Desativar também desliga a menção @@ na caixa de entrada.", "crossSessionMessageDisabledBanner": "As mensagens entre conversas estão desativadas, a menção @@ não está disponível.", "crossSessionMessageResume": "Ativar novamente", - "crossSessionMessageUpdateFailed": "Falha ao atualizar a configuração de mensagens entre conversas" + "crossSessionMessageUpdateFailed": "Falha ao atualizar a configuração de mensagens entre conversas", + "speechToTextSourceLocal": "Local (Offline - Parakeet TDT)", + "speechToTextDownloadModel": "Baixar Modelo (~670 MB)", + "speechToTextDownloadingModel": "Baixando modelo...", + "speechToTextLocalModelReady": "Modelo pronto para uso offline" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/ru-RU/settings.json b/packages/desktop/src/renderer/services/i18n/locales/ru-RU/settings.json index 9b2ab753978..b302cd16c2f 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/ru-RU/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/ru-RU/settings.json @@ -1428,5 +1428,9 @@ "crossSessionMessageDesc": "Разрешить агентам отправлять сообщения в ваши другие беседы. При отключении также отключается упоминание @@ в поле ввода.", "crossSessionMessageDisabledBanner": "Сообщения между беседами отключены, упоминание @@ недоступно.", "crossSessionMessageResume": "Включить снова", - "crossSessionMessageUpdateFailed": "Не удалось обновить настройку сообщений между беседами" + "crossSessionMessageUpdateFailed": "Не удалось обновить настройку сообщений между беседами", + "speechToTextSourceLocal": "Локально (Офлайн - Parakeet TDT)", + "speechToTextDownloadModel": "Скачать модель (~670 МБ)", + "speechToTextDownloadingModel": "Скачивание модели...", + "speechToTextLocalModelReady": "Модель готова к офлайн-диктованию" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/tr-TR/settings.json b/packages/desktop/src/renderer/services/i18n/locales/tr-TR/settings.json index 6287aea490e..1bcb736af48 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/tr-TR/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/tr-TR/settings.json @@ -1414,5 +1414,9 @@ "crossSessionMessageDesc": "Aracıların diğer sohbetlerinize mesaj göndermesine izin verir. Kapatmak giriş alanındaki @@ sohbet etiketini de devre dışı bırakır.", "crossSessionMessageDisabledBanner": "Sohbetler arası mesajlar kapalı, @@ sohbet etiketi kullanılamıyor.", "crossSessionMessageResume": "Yeniden aç", - "crossSessionMessageUpdateFailed": "Sohbetler arası mesaj ayarı güncellenemedi" + "crossSessionMessageUpdateFailed": "Sohbetler arası mesaj ayarı güncellenemedi", + "speechToTextSourceLocal": "Yerel (Çevrimdışı - Parakeet TDT)", + "speechToTextDownloadModel": "Modeli İndir (~670 MB)", + "speechToTextDownloadingModel": "Model indiriliyor...", + "speechToTextLocalModelReady": "Model çevrimdışı dikte için hazır" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/uk-UA/settings.json b/packages/desktop/src/renderer/services/i18n/locales/uk-UA/settings.json index c49e9fddb15..a74a38c1f70 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/uk-UA/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/uk-UA/settings.json @@ -1429,5 +1429,9 @@ "crossSessionMessageDesc": "Дозволяє агентам надсилати повідомлення до ваших інших розмов. Вимкнення також відключає згадку @@ у полі введення.", "crossSessionMessageDisabledBanner": "Повідомлення між розмовами вимкнено, згадка @@ недоступна.", "crossSessionMessageResume": "Увімкнути знову", - "crossSessionMessageUpdateFailed": "Не вдалося оновити налаштування повідомлень між розмовами" + "crossSessionMessageUpdateFailed": "Не вдалося оновити налаштування повідомлень між розмовами", + "speechToTextSourceLocal": "Локально (Офлайн - Parakeet TDT)", + "speechToTextDownloadModel": "Завантажити модель (~670 МБ)", + "speechToTextDownloadingModel": "Завантаження моделі...", + "speechToTextLocalModelReady": "Модель готова до офлайн-диктування" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/zh-CN/settings.json b/packages/desktop/src/renderer/services/i18n/locales/zh-CN/settings.json index f2ffc1b94d5..bdfabf37d15 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/zh-CN/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/zh-CN/settings.json @@ -1415,5 +1415,9 @@ "crossSessionMessageDesc": "允许 Agent 向你的其他会话发送消息。关闭后,输入框的 @@ 会话提及与投递能力同时停用。", "crossSessionMessageDisabledBanner": "跨会话消息已关闭,输入框的 @@ 会话提及不可用。", "crossSessionMessageResume": "重新开启", - "crossSessionMessageUpdateFailed": "跨会话消息设置更新失败" + "crossSessionMessageUpdateFailed": "跨会话消息设置更新失败", + "speechToTextSourceLocal": "本地(离线 - Parakeet TDT)", + "speechToTextDownloadModel": "下载模型 (~670 MB)", + "speechToTextDownloadingModel": "正在下载模型...", + "speechToTextLocalModelReady": "模型已准备就绪,支持离线语音识别" } diff --git a/packages/desktop/src/renderer/services/i18n/locales/zh-TW/settings.json b/packages/desktop/src/renderer/services/i18n/locales/zh-TW/settings.json index aebb9bd2052..c1a763dfe91 100644 --- a/packages/desktop/src/renderer/services/i18n/locales/zh-TW/settings.json +++ b/packages/desktop/src/renderer/services/i18n/locales/zh-TW/settings.json @@ -1416,5 +1416,9 @@ "crossSessionMessageDesc": "允許 Agent 向你的其他對話傳送訊息。關閉後,輸入框的 @@ 對話提及與投遞功能會同時停用。", "crossSessionMessageDisabledBanner": "跨對話訊息已關閉,輸入框的 @@ 對話提及不可用。", "crossSessionMessageResume": "重新開啟", - "crossSessionMessageUpdateFailed": "跨對話訊息設定更新失敗" + "crossSessionMessageUpdateFailed": "跨對話訊息設定更新失敗", + "speechToTextSourceLocal": "本地(離線 - Parakeet TDT)", + "speechToTextDownloadModel": "下載模型 (~670 MB)", + "speechToTextDownloadingModel": "正在下載模型...", + "speechToTextLocalModelReady": "模型已準備就緒,支援離線語音識別" } diff --git a/packages/desktop/src/renderer/services/speech/speechModels.ts b/packages/desktop/src/renderer/services/speech/speechModels.ts index 46cc5c78ee2..2461cf87b2d 100644 --- a/packages/desktop/src/renderer/services/speech/speechModels.ts +++ b/packages/desktop/src/renderer/services/speech/speechModels.ts @@ -8,3 +8,5 @@ export const OPENAI_SPEECH_MODEL_PRESETS = ['gpt-4o-transcribe', 'gpt-4o-mini-transcribe', 'whisper-1']; export const DEEPGRAM_SPEECH_MODEL_PRESETS = ['nova-3', 'nova-2']; + +export const LOCAL_SPEECH_MODEL_PRESETS = ['parakeet-tdt-0.6b-v3-int8', 'parakeet-tdt-0.6b-v2-int8']; diff --git a/packages/desktop/src/renderer/services/speech/speechStreamPolicy.ts b/packages/desktop/src/renderer/services/speech/speechStreamPolicy.ts index 67a8656344f..4967a04a3a8 100644 --- a/packages/desktop/src/renderer/services/speech/speechStreamPolicy.ts +++ b/packages/desktop/src/renderer/services/speech/speechStreamPolicy.ts @@ -67,6 +67,10 @@ export const getModelStreamCapability = (source: 'openai' | 'deepgram' | 'custom * - openai with a custom base_url → unknown (custom endpoint behaviour varies) */ export const getStreamCapability = (config: SpeechToTextConfig): StreamCapability => { + if (config.provider === 'local') { + return 'unsupported'; + } + if (config.provider === 'deepgram') { return getModelStreamCapability('deepgram', config.deepgram?.model ?? ''); } @@ -83,6 +87,9 @@ export const getStreamCapability = (config: SpeechToTextConfig): StreamCapabilit /** Derive a stable string key for the active provider sub-config. */ const streamMemoryEntry = (config: SpeechToTextConfig): string => { + if (config.provider === 'local') { + return `local||${config.local?.model ?? ''}`; + } if (config.provider === 'deepgram') { return `deepgram||${config.deepgram?.model ?? ''}`; } diff --git a/tests/unit/renderer/speechSettingsUtils.test.ts b/tests/unit/renderer/speechSettingsUtils.test.ts index 1b2e48fd632..f92f7e9546d 100644 --- a/tests/unit/renderer/speechSettingsUtils.test.ts +++ b/tests/unit/renderer/speechSettingsUtils.test.ts @@ -8,6 +8,7 @@ import { describe, expect, it } from 'vitest'; import { DEEPGRAM_SPEECH_MODEL_PRESETS, DEFAULT_SPEECH_TO_TEXT_CONFIG, + LOCAL_SPEECH_MODEL_PRESETS, OPENAI_SPEECH_MODEL_PRESETS, SPEECH_LANGUAGE_OPTIONS, applySpeechSource, @@ -47,6 +48,11 @@ describe('deriveSpeechSource', () => { }); expect(deriveSpeechSource(config)).toBe('openai'); }); + + it('returns local when provider is local', () => { + const config = normalizeSpeechToTextConfig({ enabled: true, provider: 'local' }); + expect(deriveSpeechSource(config)).toBe('local'); + }); }); describe('applySpeechSource', () => { @@ -79,9 +85,20 @@ describe('applySpeechSource', () => { const next = applySpeechSource(customConfig, 'custom', 'https://other/v1'); expect(next.openai?.base_url).toBe('https://my-host/v1'); }); + + it('switching to local sets provider to local and keeps local model', () => { + const next = applySpeechSource(customConfig, 'local'); + expect(next.provider).toBe('local'); + expect(next.local?.model).toBe('parakeet-tdt-0.6b-v3-int8'); + }); }); describe('model presets', () => { + it('local presets include Parakeet TDT v3 and v2', () => { + expect(LOCAL_SPEECH_MODEL_PRESETS).toContain('parakeet-tdt-0.6b-v3-int8'); + expect(LOCAL_SPEECH_MODEL_PRESETS).toContain('parakeet-tdt-0.6b-v2-int8'); + }); + it('openai presets exclude realtime-only models in phase 1', () => { expect(OPENAI_SPEECH_MODEL_PRESETS).toContain('gpt-4o-transcribe'); expect(OPENAI_SPEECH_MODEL_PRESETS).toContain('whisper-1'); diff --git a/tests/unit/speech/localSpeechService.test.ts b/tests/unit/speech/localSpeechService.test.ts new file mode 100644 index 00000000000..511dd710ef1 --- /dev/null +++ b/tests/unit/speech/localSpeechService.test.ts @@ -0,0 +1,69 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { describe, expect, it, vi } from 'vitest'; + +const { isModelReadyMock, mockStream, mockRecognizer, mockSherpa } = vi.hoisted(() => { + const stream = { + acceptWaveform: vi.fn(), + free: vi.fn(), + }; + const recognizer = { + createStream: vi.fn(() => stream), + decode: vi.fn(), + getResult: vi.fn(() => ({ text: 'Ola mundo' })), + }; + const sherpa = { + readWaveFromBinaryData: vi.fn(() => ({ + samples: new Float32Array(16000), + sampleRate: 16000, + })), + createOfflineRecognizer: vi.fn(() => recognizer), + }; + return { + isModelReadyMock: vi.fn((_id: string) => Promise.resolve(false)), + mockStream: stream, + mockRecognizer: recognizer, + mockSherpa: sherpa, + }; +}); + +vi.mock('@/process/services/speech/modelManager', () => ({ + isModelReady: isModelReadyMock, + getModelStorageDir: vi.fn(() => '/mock/model/dir'), +})); + +vi.mock('@/process/services/speech/sherpaLoader', () => ({ + loadSherpaAddon: vi.fn(() => mockSherpa), +})); + +import { clearLocalRecognizerCache, transcribeLocalAudio } from '@/process/services/speech/localSpeechService'; + +describe('localSpeechService', () => { + it('throws STT_LOCAL_MODEL_NOT_DOWNLOADED if model files are missing', async () => { + isModelReadyMock.mockResolvedValueOnce(false); + await expect(transcribeLocalAudio(new Uint8Array([1, 2, 3]), 'parakeet-tdt-0.6b-v3-int8')).rejects.toThrow( + 'STT_LOCAL_MODEL_NOT_DOWNLOADED' + ); + }); + + it('allows clearing local recognizer cache without throwing', () => { + expect(() => clearLocalRecognizerCache()).not.toThrow(); + }); + + it('successfully transcribes audio when model is ready and sherpa is loaded', async () => { + isModelReadyMock.mockResolvedValueOnce(true); + clearLocalRecognizerCache(); + const result = await transcribeLocalAudio(new Uint8Array([1, 2, 3, 4]), 'parakeet-tdt-0.6b-v3-int8'); + + expect(result.text).toBe('Ola mundo'); + expect(result.provider).toBe('local'); + expect(mockRecognizer.createStream).toHaveBeenCalled(); + expect(mockStream.acceptWaveform).toHaveBeenCalled(); + expect(mockRecognizer.decode).toHaveBeenCalledWith(mockStream); + expect(mockStream.free).toHaveBeenCalled(); + }); +}); diff --git a/tests/unit/speech/modelCatalog.test.ts b/tests/unit/speech/modelCatalog.test.ts new file mode 100644 index 00000000000..a9d6cbbc77f --- /dev/null +++ b/tests/unit/speech/modelCatalog.test.ts @@ -0,0 +1,48 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { describe, expect, it } from 'vitest'; +import { getLocalSpeechModelManifest, LOCAL_SPEECH_MODELS } from '@/process/services/speech/modelCatalog'; + +describe('modelCatalog', () => { + it('contains Parakeet TDT v3 as recommended model', () => { + const v3 = getLocalSpeechModelManifest('parakeet-tdt-0.6b-v3-int8'); + expect(v3).toBeDefined(); + expect(v3?.id).toBe('parakeet-tdt-0.6b-v3-int8'); + expect(v3?.label).toBe('Parakeet TDT v3'); + expect(v3?.recommended).toBe(true); + expect(v3?.language).toBe('multilingual'); + expect(v3?.sampleRate).toBe(16000); + }); + + it('contains Parakeet TDT v2 model', () => { + const v2 = getLocalSpeechModelManifest('parakeet-tdt-0.6b-v2-int8'); + expect(v2).toBeDefined(); + expect(v2?.id).toBe('parakeet-tdt-0.6b-v2-int8'); + expect(v2?.label).toBe('Parakeet TDT v2'); + expect(v2?.language).toBe('en'); + }); + + it('has valid download file specifications for Parakeet TDT v3', () => { + const v3 = getLocalSpeechModelManifest('parakeet-tdt-0.6b-v3-int8')!; + const fileNames = v3.files.map((f) => f.name); + expect(fileNames).toContain('encoder.int8.onnx'); + expect(fileNames).toContain('decoder.int8.onnx'); + expect(fileNames).toContain('joiner.int8.onnx'); + expect(fileNames).toContain('tokens.txt'); + + for (const file of v3.files) { + expect(file.url).toContain('huggingface.co'); + expect(file.sha256).toHaveLength(64); + expect(file.sizeBytes).toBeGreaterThan(0); + } + }); + + it('falls back to recommended model when unknown id is passed', () => { + const unknown = getLocalSpeechModelManifest('non-existent-model'); + expect(unknown).toBe(LOCAL_SPEECH_MODELS[0]); + }); +}); diff --git a/tests/unit/speech/modelManager.test.ts b/tests/unit/speech/modelManager.test.ts new file mode 100644 index 00000000000..38f5c7c7091 --- /dev/null +++ b/tests/unit/speech/modelManager.test.ts @@ -0,0 +1,39 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { describe, expect, it } from 'vitest'; +import { + cancelModelDownload, + getModelStatus, + getModelStorageDir, + isModelReady, +} from '@/process/services/speech/modelManager'; + +describe('modelManager', () => { + it('computes expected storage dir under models/stt', () => { + const dir = getModelStorageDir('parakeet-tdt-0.6b-v3-int8'); + expect(dir).toContain('models'); + expect(dir).toContain('stt'); + expect(dir).toContain('parakeet-tdt-0.6b-v3-int8'); + }); + + it('returns false for isModelReady when model files are not yet downloaded', async () => { + const ready = await isModelReady('non-existent-test-model'); + expect(ready).toBe(false); + }); + + it('returns idle status when model is not downloading', () => { + const status = getModelStatus('parakeet-tdt-0.6b-v3-int8'); + expect(status.status).toBe('idle'); + expect(status.modelId).toBe('parakeet-tdt-0.6b-v3-int8'); + expect(status.totalBytes).toBeGreaterThan(0); + }); + + it('cancelModelDownload returns false when no download is in progress', () => { + const cancelled = cancelModelDownload('parakeet-tdt-0.6b-v3-int8'); + expect(cancelled).toBe(false); + }); +}); diff --git a/tests/unit/speech/sherpaLoader.test.ts b/tests/unit/speech/sherpaLoader.test.ts new file mode 100644 index 00000000000..fdcb4790026 --- /dev/null +++ b/tests/unit/speech/sherpaLoader.test.ts @@ -0,0 +1,34 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { describe, expect, it } from 'vitest'; +import { getSherpaPackageName, isSherpaSupported, loadSherpaAddon } from '@/process/services/speech/sherpaLoader'; + +describe('sherpaLoader', () => { + it('returns expected package name based on platform and architecture', () => { + const pkg = getSherpaPackageName(); + if (process.platform === 'darwin' && process.arch === 'arm64') { + expect(pkg).toBe('sherpa-onnx-darwin-arm64'); + } else if (process.platform === 'darwin' && process.arch === 'x64') { + expect(pkg).toBe('sherpa-onnx-darwin-x64'); + } else if (process.platform === 'linux' && process.arch === 'x64') { + expect(pkg).toBe('sherpa-onnx-linux-x64'); + } else if (process.platform === 'win32' && process.arch === 'x64') { + expect(pkg).toBe('sherpa-onnx-win-x64'); + } + }); + + it('loads sherpa addon on supported platform without error', () => { + if (process.platform === 'darwin' && process.arch === 'arm64') { + expect(isSherpaSupported()).toBe(true); + const addon = loadSherpaAddon(); + expect(addon).toBeDefined(); + expect(typeof addon.createOfflineRecognizer).toBe('function'); + expect(typeof addon.createOfflineStream).toBe('function'); + expect(typeof addon.readWaveFromBinary).toBe('function'); + } + }); +}); From 6c32620216e24c27f8395b850564967ef6b0b36d Mon Sep 17 00:00:00 2001 From: Eduardo Jose Costa <59846713+EduCosta85@users.noreply.github.com> Date: Sat, 12 Sep 2026 09:13:04 -0300 Subject: [PATCH 2/2] test(stt): increase unit and dom test coverage for local speech dictation --- .../renderer/speechStreamPolicy.dom.test.ts | 22 ++ .../renderer/speechTestPanel.dom.test.tsx | 13 + .../unit/renderer/speechToTextService.test.ts | 130 +++++++- .../renderer/voiceInputSection.dom.test.tsx | 159 ++++++++++ tests/unit/speech/localSpeechService.test.ts | 89 +++++- tests/unit/speech/modelManager.test.ts | 296 +++++++++++++++++- tests/unit/speech/sherpaLoader.test.ts | 96 ++++-- tests/unit/speech/speechBridge.test.ts | 126 ++++++++ tests/unit/speech/speechIndex.test.ts | 25 ++ 9 files changed, 901 insertions(+), 55 deletions(-) create mode 100644 tests/unit/speech/speechBridge.test.ts create mode 100644 tests/unit/speech/speechIndex.test.ts diff --git a/tests/unit/renderer/speechStreamPolicy.dom.test.ts b/tests/unit/renderer/speechStreamPolicy.dom.test.ts index c33fb52f6ef..1732dbf8e4e 100644 --- a/tests/unit/renderer/speechStreamPolicy.dom.test.ts +++ b/tests/unit/renderer/speechStreamPolicy.dom.test.ts @@ -139,6 +139,17 @@ describe('getStreamCapability', () => { expect(getStreamCapability(config)).toBe('unsupported'); }); }); + + describe('local provider', () => { + it('returns unsupported for local speech provider', () => { + const config: SpeechToTextConfig = { + enabled: true, + provider: 'local', + local: { model: 'parakeet-tdt-0.6b-v3-int8', language: '' }, + }; + expect(getStreamCapability(config)).toBe('unsupported'); + }); + }); }); // --------------------------------------------------------------------------- @@ -185,6 +196,17 @@ describe('shouldTryStreaming', () => { it('unknown capability with no memory → true (optimistic)', () => { expect(shouldTryStreaming(openaiCustom('gpt-4o-transcribe'))).toBe(true); }); + + it('local provider always returns false for shouldTryStreaming', () => { + const localConfig: SpeechToTextConfig = { + enabled: true, + provider: 'local', + local: { model: 'parakeet-tdt-0.6b-v3-int8', language: '' }, + }; + expect(shouldTryStreaming(localConfig)).toBe(false); + rememberStreamUnsupported(localConfig); + expect(shouldTryStreaming(localConfig)).toBe(false); + }); }); // --------------------------------------------------------------------------- diff --git a/tests/unit/renderer/speechTestPanel.dom.test.tsx b/tests/unit/renderer/speechTestPanel.dom.test.tsx index 0787e82a324..35b6789e22f 100644 --- a/tests/unit/renderer/speechTestPanel.dom.test.tsx +++ b/tests/unit/renderer/speechTestPanel.dom.test.tsx @@ -67,4 +67,17 @@ describe('SpeechTestPanel', () => { expect(speechSettingsMocks.setClientBusinessSetting).toHaveBeenCalledWith('tools.speechToText', config) ); }); + + it('validates and proceeds without API key when source is local', async () => { + const localConfig: SpeechToTextConfig = { + enabled: true, + provider: 'local', + local: { model: 'parakeet-tdt-0.6b-v3-int8', language: '' }, + }; + render(); + fireEvent.click(screen.getByText('settings.speechToTextTest')); + await waitFor(() => + expect(speechSettingsMocks.setClientBusinessSetting).toHaveBeenCalledWith('tools.speechToText', localConfig) + ); + }); }); diff --git a/tests/unit/renderer/speechToTextService.test.ts b/tests/unit/renderer/speechToTextService.test.ts index fb11fd64f03..29ec243d472 100644 --- a/tests/unit/renderer/speechToTextService.test.ts +++ b/tests/unit/renderer/speechToTextService.test.ts @@ -4,15 +4,34 @@ * SPDX-License-Identifier: Apache-2.0 * * Unit tests for renderer/services/SpeechToTextService.ts. - * Regression tests for voice input failing with 400: the backend /api/stt - * endpoint only accepts multipart with fields `file`, `fileName`, `mimeType`, - * `languageHint` — the previous code sent JSON (Electron) or a wrong - * `audio` multipart field (WebUI). * * @vitest-environment node */ -import { describe, it, expect, beforeEach, afterEach, vi } from 'vitest'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; + +const settingsMocks = vi.hoisted(() => ({ + getClientBusinessSetting: vi.fn(), +})); + +const ipcMocks = vi.hoisted(() => ({ + transcribeInvoke: vi.fn(), +})); + +vi.mock('@/renderer/services/clientBusinessSettings', () => ({ + getClientBusinessSetting: settingsMocks.getClientBusinessSetting, +})); + +vi.mock('@/common', () => ({ + ipcBridge: { + speech: { + transcribe: { + invoke: ipcMocks.transcribeInvoke, + }, + }, + }, +})); + import { transcribeAudioBlob } from '@/renderer/services/SpeechToTextService'; type XhrListener = () => void; @@ -50,6 +69,14 @@ class FakeXMLHttpRequest { this.responseText = responseText; this.listeners.load?.(); } + + triggerError() { + this.listeners.error?.(); + } + + triggerAbort() { + this.listeners.abort?.(); + } } const waitForRequest = async (): Promise => { @@ -65,10 +92,61 @@ describe('SpeechToTextService.transcribeAudioBlob', () => { beforeEach(() => { FakeXMLHttpRequest.instances = []; vi.stubGlobal('XMLHttpRequest', FakeXMLHttpRequest); + settingsMocks.getClientBusinessSetting.mockResolvedValue(undefined); + ipcMocks.transcribeInvoke.mockReset(); }); afterEach(() => { vi.unstubAllGlobals(); + vi.clearAllMocks(); + }); + + it('rejects with STT_FILE_TOO_LARGE when audio blob exceeds 30MB', async () => { + const hugeBlob = { + size: 31 * 1024 * 1024, + type: 'audio/webm', + } as unknown as Blob; + + await expect(transcribeAudioBlob(hugeBlob)).rejects.toThrow('STT_FILE_TOO_LARGE'); + }); + + it('transcribes via local IPC bridge when provider is local', async () => { + settingsMocks.getClientBusinessSetting.mockResolvedValue({ + provider: 'local', + local: { + model: 'parakeet-tdt-0.6b-v3-int8', + }, + }); + + ipcMocks.transcribeInvoke.mockResolvedValueOnce({ + model: 'parakeet-tdt-0.6b-v3-int8', + provider: 'local', + text: 'Local dictation result', + }); + + const blob = new Blob(['pcm-audio-data'], { type: 'audio/wav' }); + const result = await transcribeAudioBlob(blob); + + expect(result.text).toBe('Local dictation result'); + expect(ipcMocks.transcribeInvoke).toHaveBeenCalledWith( + expect.objectContaining({ + modelId: 'parakeet-tdt-0.6b-v3-int8', + }) + ); + }); + + it('re-throws STT_LOCAL_MODEL_NOT_DOWNLOADED when local model is missing', async () => { + settingsMocks.getClientBusinessSetting.mockResolvedValue({ + provider: 'local', + local: { + model: 'parakeet-tdt-0.6b-v3-int8', + }, + }); + + ipcMocks.transcribeInvoke.mockRejectedValueOnce(new Error('STT_LOCAL_MODEL_NOT_DOWNLOADED')); + + const blob = new Blob(['pcm-audio-data'], { type: 'audio/wav' }); + await expect(transcribeAudioBlob(blob)).rejects.toThrow('STT_LOCAL_MODEL_NOT_DOWNLOADED'); }); it('sends multipart fields matching the backend contract (file/fileName/mimeType/languageHint)', async () => { @@ -78,8 +156,6 @@ describe('SpeechToTextService.transcribeAudioBlob', () => { const xhr = await waitForRequest(); expect(xhr.method).toBe('POST'); expect(xhr.url).toContain('/api/stt'); - // Credentialed cross-origin requests are rejected by the browser because - // the desktop backend responds with Access-Control-Allow-Origin: * expect(xhr.withCredentials).toBe(false); expect(xhr.sentBody).toBeInstanceOf(FormData); @@ -98,6 +174,46 @@ describe('SpeechToTextService.transcribeAudioBlob', () => { await expect(pending).resolves.toEqual({ model: 'whisper-1', provider: 'openai', text: 'hello' }); }); + it('correctly maps audio extensions for different mime types', async () => { + const typesAndExtensions = [ + { mime: 'audio/mp4', ext: 'm4a' }, + { mime: 'audio/mpeg', ext: 'mp3' }, + { mime: 'audio/ogg', ext: 'ogg' }, + { mime: 'audio/wav', ext: 'wav' }, + ]; + + for (const item of typesAndExtensions) { + FakeXMLHttpRequest.instances = []; + const blob = new Blob(['audio'], { type: item.mime }); + const pending = transcribeAudioBlob(blob); + const xhr = await waitForRequest(); + const formData = xhr.sentBody as FormData; + expect(formData.get('fileName')).toBe(`speech-input.${item.ext}`); + xhr.respond(200, JSON.stringify({ success: true, data: { text: 'ok' } })); + await pending; + } + }); + + it('handles network error', async () => { + const blob = new Blob(['fake-audio'], { type: 'audio/webm' }); + const pending = transcribeAudioBlob(blob); + + const xhr = await waitForRequest(); + xhr.triggerError(); + + await expect(pending).rejects.toThrow('STT_NETWORK_ERROR'); + }); + + it('handles abort error', async () => { + const blob = new Blob(['fake-audio'], { type: 'audio/webm' }); + const pending = transcribeAudioBlob(blob); + + const xhr = await waitForRequest(); + xhr.triggerAbort(); + + await expect(pending).rejects.toThrow('STT_ABORTED'); + }); + it('rejects with the backend error code so STT errors map correctly', async () => { const blob = new Blob(['fake-audio'], { type: 'audio/webm' }); const pending = transcribeAudioBlob(blob); diff --git a/tests/unit/renderer/voiceInputSection.dom.test.tsx b/tests/unit/renderer/voiceInputSection.dom.test.tsx index e16a3d720fb..d7c8349c0da 100644 --- a/tests/unit/renderer/voiceInputSection.dom.test.tsx +++ b/tests/unit/renderer/voiceInputSection.dom.test.tsx @@ -7,6 +7,7 @@ import { fireEvent, render, screen, waitFor } from '@testing-library/react'; import React from 'react'; import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { ipcBridge } from '@/common'; import type { SpeechToTextConfig } from '@/common/types/provider/speech'; const configStore: { value?: SpeechToTextConfig } = {}; @@ -25,11 +26,14 @@ vi.mock('react-i18next', () => ({ useTranslation: () => ({ t: (key: string) => key, i18n: { language: 'en-US' } }), })); +import { Message } from '@arco-design/web-react'; import VoiceInputSection from '@/renderer/components/settings/SettingsModal/contents/SystemModalContent/VoiceInputSection'; describe('VoiceInputSection', () => { beforeEach(() => { configStore.value = undefined; + vi.spyOn(Message, 'success').mockImplementation(() => undefined as never); + vi.spyOn(Message, 'error').mockImplementation(() => undefined as never); speechSettingsMocks.getClientBusinessSetting.mockResolvedValue(undefined); speechSettingsMocks.setClientBusinessSetting.mockResolvedValue(undefined); // jsdom does not implement matchMedia; arco-design's responsive Grid needs it @@ -160,4 +164,159 @@ describe('VoiceInputSection', () => { expect(screen.queryByText('settings.speechToTextPunctuate')).toBeNull(); expect(screen.queryByText('settings.speechToTextSmartFormat')).toBeNull(); }); + + it('local provider renders download button and handles download model click', async () => { + configStore.value = { + enabled: true, + provider: 'local', + local: { + model: 'parakeet-tdt-0.6b-v3-int8', + language: '', + }, + }; + speechSettingsMocks.getClientBusinessSetting.mockResolvedValue(configStore.value); + + const checkModelSpy = vi.spyOn(ipcBridge.speech.checkModel, 'invoke').mockResolvedValue({ + isReady: false, + status: { + modelId: 'parakeet-tdt-0.6b-v3-int8', + status: 'idle', + progress: 0, + downloadedBytes: 0, + totalBytes: 670000000, + }, + }); + + const downloadModelSpy = vi.spyOn(ipcBridge.speech.downloadModel, 'invoke').mockResolvedValue({ + modelId: 'parakeet-tdt-0.6b-v3-int8', + status: 'ready', + progress: 100, + downloadedBytes: 670000000, + totalBytes: 670000000, + }); + + render(); + + await waitFor(() => expect(screen.getByText('settings.speechToTextDownloadModel')).toBeTruthy()); + expect(screen.queryByText('settings.speechToTextApiKey')).toBeNull(); + expect(screen.queryByText('settings.speechToTextBaseUrl')).toBeNull(); + expect(checkModelSpy).toHaveBeenCalledWith({ modelId: 'parakeet-tdt-0.6b-v3-int8' }); + + // Click download button + const downloadBtn = screen.getByText('settings.speechToTextDownloadModel'); + fireEvent.click(downloadBtn); + + await waitFor(() => { + expect(downloadModelSpy).toHaveBeenCalledWith({ modelId: 'parakeet-tdt-0.6b-v3-int8' }); + }); + }); + + it('local provider renders ready checkmark when model is already downloaded', async () => { + configStore.value = { + enabled: true, + provider: 'local', + local: { + model: 'parakeet-tdt-0.6b-v3-int8', + language: '', + }, + }; + speechSettingsMocks.getClientBusinessSetting.mockResolvedValue(configStore.value); + + vi.spyOn(ipcBridge.speech.checkModel, 'invoke').mockResolvedValue({ + isReady: true, + status: { + modelId: 'parakeet-tdt-0.6b-v3-int8', + status: 'ready', + progress: 100, + downloadedBytes: 670000000, + totalBytes: 670000000, + }, + }); + + render(); + + await waitFor(() => expect(screen.getByText(/settings\.speechToTextLocalModelReady/)).toBeTruthy()); + expect(screen.queryByText('settings.speechToTextDownloadModel')).toBeNull(); + }); + + it('shows error message when model download fails', async () => { + configStore.value = { + enabled: true, + provider: 'local', + local: { + model: 'parakeet-tdt-0.6b-v3-int8', + language: '', + }, + }; + speechSettingsMocks.getClientBusinessSetting.mockResolvedValue(configStore.value); + + vi.spyOn(ipcBridge.speech.checkModel, 'invoke').mockResolvedValue({ + isReady: false, + status: { + modelId: 'parakeet-tdt-0.6b-v3-int8', + status: 'idle', + progress: 0, + downloadedBytes: 0, + totalBytes: 670000000, + }, + }); + + vi.spyOn(ipcBridge.speech.downloadModel, 'invoke').mockRejectedValue(new Error('Network error')); + const messageErrorSpy = vi.spyOn(Message, 'error'); + + render(); + + await waitFor(() => expect(screen.getByText('settings.speechToTextDownloadModel')).toBeTruthy()); + fireEvent.click(screen.getByText('settings.speechToTextDownloadModel')); + + await waitFor(() => { + expect(messageErrorSpy).toHaveBeenCalledWith('Network error'); + }); + }); + + it('updates local model download status when onDownloadProgress emits', async () => { + configStore.value = { + enabled: true, + provider: 'local', + local: { + model: 'parakeet-tdt-0.6b-v3-int8', + language: '', + }, + }; + speechSettingsMocks.getClientBusinessSetting.mockResolvedValue(configStore.value); + + let progressListener: Function | null = null; + vi.spyOn(ipcBridge.speech.onDownloadProgress, 'on').mockImplementation((cb: any) => { + progressListener = cb; + return () => {}; + }); + + vi.spyOn(ipcBridge.speech.checkModel, 'invoke').mockResolvedValue({ + isReady: false, + status: { + modelId: 'parakeet-tdt-0.6b-v3-int8', + status: 'idle', + progress: 0, + downloadedBytes: 0, + totalBytes: 670000000, + }, + }); + + render(); + + await waitFor(() => expect(progressListener).toBeTruthy()); + + // Emit progress event + await waitFor(() => { + progressListener!({ + modelId: 'parakeet-tdt-0.6b-v3-int8', + percent: 45, + downloadedBytes: 300000000, + totalBytes: 670000000, + }); + }); + + await waitFor(() => expect(screen.getByText('45%')).toBeTruthy()); + expect(screen.getByText('settings.speechToTextDownloadingModel')).toBeTruthy(); + }); }); diff --git a/tests/unit/speech/localSpeechService.test.ts b/tests/unit/speech/localSpeechService.test.ts index 511dd710ef1..09a188084af 100644 --- a/tests/unit/speech/localSpeechService.test.ts +++ b/tests/unit/speech/localSpeechService.test.ts @@ -4,7 +4,7 @@ * SPDX-License-Identifier: Apache-2.0 */ -import { describe, expect, it, vi } from 'vitest'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; const { isModelReadyMock, mockStream, mockRecognizer, mockSherpa } = vi.hoisted(() => { const stream = { @@ -17,11 +17,19 @@ const { isModelReadyMock, mockStream, mockRecognizer, mockSherpa } = vi.hoisted( getResult: vi.fn(() => ({ text: 'Ola mundo' })), }; const sherpa = { - readWaveFromBinaryData: vi.fn(() => ({ + readWaveFromBinary: vi.fn((buf: Buffer) => ({ + samples: new Float32Array(16000), + sampleRate: 16000, + })), + readWaveFromBinaryData: vi.fn((buf: Buffer) => ({ samples: new Float32Array(16000), sampleRate: 16000, })), createOfflineRecognizer: vi.fn(() => recognizer), + createOfflineStream: vi.fn(() => ({})), + acceptWaveformOffline: vi.fn(), + decodeOfflineStream: vi.fn(), + getOfflineStreamResultAsJson: vi.fn(() => JSON.stringify({ text: 'Native result' })), }; return { isModelReadyMock: vi.fn((_id: string) => Promise.resolve(false)), @@ -43,6 +51,17 @@ vi.mock('@/process/services/speech/sherpaLoader', () => ({ import { clearLocalRecognizerCache, transcribeLocalAudio } from '@/process/services/speech/localSpeechService'; describe('localSpeechService', () => { + beforeEach(() => { + vi.clearAllMocks(); + clearLocalRecognizerCache(); + mockRecognizer.createStream = vi.fn(() => mockStream); + mockRecognizer.getResult = vi.fn(() => ({ text: 'Ola mundo' })); + mockSherpa.readWaveFromBinary = vi.fn(() => ({ + samples: new Float32Array(16000), + sampleRate: 16000, + })); + }); + it('throws STT_LOCAL_MODEL_NOT_DOWNLOADED if model files are missing', async () => { isModelReadyMock.mockResolvedValueOnce(false); await expect(transcribeLocalAudio(new Uint8Array([1, 2, 3]), 'parakeet-tdt-0.6b-v3-int8')).rejects.toThrow( @@ -54,16 +73,66 @@ describe('localSpeechService', () => { expect(() => clearLocalRecognizerCache()).not.toThrow(); }); - it('successfully transcribes audio when model is ready and sherpa is loaded', async () => { - isModelReadyMock.mockResolvedValueOnce(true); - clearLocalRecognizerCache(); + it('successfully transcribes audio with Uint8Array and reuses cached recognizer on second call', async () => { + isModelReadyMock.mockResolvedValue(true); + + const result1 = await transcribeLocalAudio(new Uint8Array([1, 2, 3, 4]), 'parakeet-tdt-0.6b-v3-int8'); + expect(result1.text).toBe('Ola mundo'); + expect(result1.provider).toBe('local'); + expect(mockSherpa.createOfflineRecognizer).toHaveBeenCalledTimes(1); + + // Second call should reuse activeRecognizerCache + const result2 = await transcribeLocalAudio(Buffer.from([1, 2, 3, 4]), 'parakeet-tdt-0.6b-v3-int8'); + expect(result2.text).toBe('Ola mundo'); + expect(mockSherpa.createOfflineRecognizer).toHaveBeenCalledTimes(1); + }); + + it('handles number array and object record inputs to toBuffer', async () => { + isModelReadyMock.mockResolvedValue(true); + + const arrayInput = [0, 1, 2, 3]; + const resArray = await transcribeLocalAudio(arrayInput, 'parakeet-tdt-0.6b-v3-int8'); + expect(resArray.text).toBe('Ola mundo'); + + const recordInput = { 0: 10, 1: 20, 2: 30 }; + const resRecord = await transcribeLocalAudio(recordInput as any, 'parakeet-tdt-0.6b-v3-int8'); + expect(resRecord.text).toBe('Ola mundo'); + }); + + it('handles string JSON result in recognizer.getResult', async () => { + isModelReadyMock.mockResolvedValue(true); + mockRecognizer.getResult = vi.fn(() => JSON.stringify({ text: 'Parsed string result' })); + const result = await transcribeLocalAudio(new Uint8Array([1, 2, 3, 4]), 'parakeet-tdt-0.6b-v3-int8'); + expect(result.text).toBe('Parsed string result'); + }); + it('falls back to raw 16-bit PCM conversion when wave parsing throws', async () => { + isModelReadyMock.mockResolvedValue(true); + mockSherpa.readWaveFromBinary = vi.fn(() => { + throw new Error('Not a WAV header'); + }); + + // Create 100 bytes of 16-bit PCM + const pcm = new Int16Array(50); + for (let i = 0; i < 50; i++) pcm[i] = i * 100; + const rawBytes = new Uint8Array(pcm.buffer); + + const result = await transcribeLocalAudio(rawBytes, 'parakeet-tdt-0.6b-v3-int8'); expect(result.text).toBe('Ola mundo'); - expect(result.provider).toBe('local'); - expect(mockRecognizer.createStream).toHaveBeenCalled(); - expect(mockStream.acceptWaveform).toHaveBeenCalled(); - expect(mockRecognizer.decode).toHaveBeenCalledWith(mockStream); - expect(mockStream.free).toHaveBeenCalled(); + expect(mockStream.acceptWaveform).toHaveBeenCalledWith(16000, expect.any(Float32Array)); + }); + + it('supports native C++ NAPI addon interface when recognizer.createStream is absent', async () => { + isModelReadyMock.mockResolvedValue(true); + // Remove createStream to exercise native addon path + delete (mockRecognizer as any).createStream; + + const result = await transcribeLocalAudio(new Uint8Array([1, 2, 3, 4]), 'parakeet-tdt-0.6b-v3-int8'); + expect(result.text).toBe('Native result'); + expect(mockSherpa.createOfflineStream).toHaveBeenCalledWith(mockRecognizer); + expect(mockSherpa.acceptWaveformOffline).toHaveBeenCalled(); + expect(mockSherpa.decodeOfflineStream).toHaveBeenCalled(); + expect(mockSherpa.getOfflineStreamResultAsJson).toHaveBeenCalled(); }); }); diff --git a/tests/unit/speech/modelManager.test.ts b/tests/unit/speech/modelManager.test.ts index 38f5c7c7091..fa77d405956 100644 --- a/tests/unit/speech/modelManager.test.ts +++ b/tests/unit/speech/modelManager.test.ts @@ -4,36 +4,298 @@ * SPDX-License-Identifier: Apache-2.0 */ -import { describe, expect, it } from 'vitest'; +import { createHash } from 'node:crypto'; +import { existsSync, promises as fs } from 'node:fs'; +import os from 'node:os'; +import path from 'node:path'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import * as modelCatalog from '@/process/services/speech/modelCatalog'; import { cancelModelDownload, + downloadModel, getModelStatus, getModelStorageDir, isModelReady, } from '@/process/services/speech/modelManager'; describe('modelManager', () => { - it('computes expected storage dir under models/stt', () => { - const dir = getModelStorageDir('parakeet-tdt-0.6b-v3-int8'); - expect(dir).toContain('models'); - expect(dir).toContain('stt'); - expect(dir).toContain('parakeet-tdt-0.6b-v3-int8'); + let tempDir: string; + + beforeEach(async () => { + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), 'aionui-stt-test-')); + }); + + afterEach(async () => { + vi.restoreAllMocks(); + try { + await fs.rm(tempDir, { recursive: true, force: true }); + } catch { + // ignore + } + }); + + describe('getModelStorageDir', () => { + it('computes expected storage dir under models/stt', () => { + const dir = getModelStorageDir('parakeet-tdt-0.6b-v3-int8'); + expect(dir).toContain('models'); + expect(dir).toContain('stt'); + expect(dir).toContain('parakeet-tdt-0.6b-v3-int8'); + }); + }); + + describe('isModelReady', () => { + it('returns false for unknown model ID when manifest is undefined', async () => { + vi.spyOn(modelCatalog, 'getLocalSpeechModelManifest').mockReturnValueOnce(undefined); + const ready = await isModelReady('unknown-test-model'); + expect(ready).toBe(false); + }); + + it('returns false when model directory does not exist', async () => { + const ready = await isModelReady('parakeet-tdt-0.6b-v3-int8'); + expect(typeof ready).toBe('boolean'); + }); + }); + + describe('getModelStatus', () => { + it('returns idle status with totalBytes for valid model', () => { + const status = getModelStatus('parakeet-tdt-0.6b-v3-int8'); + expect(status.status).toBe('idle'); + expect(status.modelId).toBe('parakeet-tdt-0.6b-v3-int8'); + expect(status.totalBytes).toBeGreaterThan(0); + expect(status.progress).toBe(0); + }); + + it('returns idle status with 0 totalBytes when manifest is undefined', () => { + vi.spyOn(modelCatalog, 'getLocalSpeechModelManifest').mockReturnValueOnce(undefined); + const status = getModelStatus('non-existent-model'); + expect(status.status).toBe('idle'); + expect(status.totalBytes).toBe(0); + }); }); - it('returns false for isModelReady when model files are not yet downloaded', async () => { - const ready = await isModelReady('non-existent-test-model'); - expect(ready).toBe(false); + describe('cancelModelDownload', () => { + it('returns false when no download is active', () => { + const cancelled = cancelModelDownload('parakeet-tdt-0.6b-v3-int8'); + expect(cancelled).toBe(false); + }); }); - it('returns idle status when model is not downloading', () => { - const status = getModelStatus('parakeet-tdt-0.6b-v3-int8'); - expect(status.status).toBe('idle'); - expect(status.modelId).toBe('parakeet-tdt-0.6b-v3-int8'); - expect(status.totalBytes).toBeGreaterThan(0); + describe('downloadModel validation and fast paths', () => { + it('rejects when manifest cannot be found', async () => { + vi.spyOn(modelCatalog, 'getLocalSpeechModelManifest').mockReturnValueOnce(undefined); + await expect(downloadModel('invalid-model-id')).rejects.toThrow('Unknown model ID: invalid-model-id'); + }); + + it('returns immediately with ready status when model is already ready', async () => { + const manifest = modelCatalog.getLocalSpeechModelManifest('parakeet-tdt-0.6b-v3-int8')!; + const modelDir = getModelStorageDir('parakeet-tdt-0.6b-v3-int8'); + await fs.mkdir(modelDir, { recursive: true }); + + // Create dummy files with exact manifest sizes + for (const file of manifest.files) { + const filePath = path.join(modelDir, file.name); + await fs.writeFile(filePath, Buffer.alloc(file.sizeBytes)); + } + + const progressCb = vi.fn(); + const status = await downloadModel('parakeet-tdt-0.6b-v3-int8', progressCb); + expect(status.status).toBe('ready'); + expect(status.progress).toBe(100); + expect(progressCb).toHaveBeenCalledWith( + expect.objectContaining({ + percent: 100, + }) + ); + + // Clean up test files so subsequent tests are not affected + for (const file of manifest.files) { + const filePath = path.join(modelDir, file.name); + try { + await fs.unlink(filePath); + } catch { + // ignore + } + } + }); }); - it('cancelModelDownload returns false when no download is in progress', () => { - const cancelled = cancelModelDownload('parakeet-tdt-0.6b-v3-int8'); - expect(cancelled).toBe(false); + describe('downloadModel streaming and error handling', () => { + const createMockStream = (buf: Buffer) => { + return new ReadableStream({ + start(controller) { + controller.enqueue(new Uint8Array(buf)); + controller.close(); + }, + }); + }; + + it('downloads model files, updates progress, and validates checksum', async () => { + const testContent = Buffer.from('model-weight-content'); + const testSha256 = createHash('sha256').update(testContent).digest('hex'); + + const mockManifest: modelCatalog.LocalSpeechModelManifest = { + id: 'test-mini-model', + label: 'Test Mini', + description: 'Test model', + type: 'transducer', + language: 'en', + sampleRate: 16000, + modelingUnit: 'bpe', + totalSizeBytes: testContent.length, + files: [ + { + name: 'encoder.int8.onnx', + sizeBytes: testContent.length, + sha256: testSha256, + url: 'https://huggingface.co/test/model/resolve/main/encoder.int8.onnx', + }, + ], + }; + + vi.spyOn(modelCatalog, 'getLocalSpeechModelManifest').mockReturnValue(mockManifest); + + const mockFetch = vi.fn().mockImplementation(() => + Promise.resolve({ + ok: true, + status: 200, + statusText: 'OK', + body: createMockStream(testContent), + }) + ); + vi.stubGlobal('fetch', mockFetch); + + const progressUpdates: number[] = []; + const result = await downloadModel('test-mini-model', (p) => { + progressUpdates.push(p.percent); + }); + + expect(result.status).toBe('ready'); + expect(result.progress).toBe(100); + expect(progressUpdates).toContain(100); + + // Clean up target test file + const destPath = path.join(getModelStorageDir('test-mini-model'), 'encoder.int8.onnx'); + if (existsSync(destPath)) { + await fs.unlink(destPath); + } + }); + + it('falls back to mirror URL when primary fails', async () => { + const testContent = Buffer.from('mirror-content'); + const testSha256 = createHash('sha256').update(testContent).digest('hex'); + + const mockManifest: modelCatalog.LocalSpeechModelManifest = { + id: 'test-mirror-model', + label: 'Test Mirror', + description: 'Test model', + type: 'transducer', + language: 'en', + sampleRate: 16000, + modelingUnit: 'bpe', + totalSizeBytes: testContent.length, + files: [ + { + name: 'encoder.int8.onnx', + sizeBytes: testContent.length, + sha256: testSha256, + url: 'https://huggingface.co/test/model/resolve/main/encoder.int8.onnx', + }, + ], + }; + + vi.spyOn(modelCatalog, 'getLocalSpeechModelManifest').mockReturnValue(mockManifest); + + let attempt = 0; + const mockFetch = vi.fn().mockImplementation((url: string) => { + attempt++; + if (url.includes('huggingface.co')) { + return Promise.reject(new Error('HuggingFace blocked')); + } + return Promise.resolve({ + ok: true, + status: 200, + statusText: 'OK', + body: createMockStream(testContent), + }); + }); + vi.stubGlobal('fetch', mockFetch); + + const result = await downloadModel('test-mirror-model'); + expect(result.status).toBe('ready'); + expect(attempt).toBeGreaterThanOrEqual(2); + + const destPath = path.join(getModelStorageDir('test-mirror-model'), 'encoder.int8.onnx'); + if (existsSync(destPath)) { + await fs.unlink(destPath); + } + }); + + it('throws error when checksum verification fails', async () => { + const testContent = Buffer.from('corrupted-content'); + const mockManifest: modelCatalog.LocalSpeechModelManifest = { + id: 'test-corrupt-model', + label: 'Test Corrupt', + description: 'Test model', + type: 'transducer', + language: 'en', + sampleRate: 16000, + modelingUnit: 'bpe', + totalSizeBytes: testContent.length, + files: [ + { + name: 'encoder.int8.onnx', + sizeBytes: testContent.length, + sha256: 'wrong-sha256-hash-value-expected', + url: 'https://huggingface.co/test/model/resolve/main/encoder.int8.onnx', + }, + ], + }; + + vi.spyOn(modelCatalog, 'getLocalSpeechModelManifest').mockReturnValue(mockManifest); + + const mockFetch = vi.fn().mockImplementation(() => + Promise.resolve({ + ok: true, + status: 200, + statusText: 'OK', + body: createMockStream(testContent), + }) + ); + vi.stubGlobal('fetch', mockFetch); + + await expect(downloadModel('test-corrupt-model')).rejects.toThrow('Checksum mismatch'); + const status = getModelStatus('test-corrupt-model'); + expect(status.status).toBe('idle'); + }); + + it('handles download network failure and cleans up active download', async () => { + const mockManifest: modelCatalog.LocalSpeechModelManifest = { + id: 'test-network-fail-model', + label: 'Test Fail', + description: 'Test model', + type: 'transducer', + language: 'en', + sampleRate: 16000, + modelingUnit: 'bpe', + totalSizeBytes: 1000, + files: [ + { + name: 'encoder.int8.onnx', + sizeBytes: 1000, + sha256: 'somehash', + url: 'https://huggingface.co/test/fail/model', + }, + ], + }; + + vi.spyOn(modelCatalog, 'getLocalSpeechModelManifest').mockReturnValue(mockManifest); + + const mockFetch = vi.fn().mockRejectedValue(new Error('Network offline')); + vi.stubGlobal('fetch', mockFetch); + + await expect(downloadModel('test-network-fail-model')).rejects.toThrow(); + const status = getModelStatus('test-network-fail-model'); + expect(status.status).toBe('idle'); + }); }); }); diff --git a/tests/unit/speech/sherpaLoader.test.ts b/tests/unit/speech/sherpaLoader.test.ts index fdcb4790026..eb5efc4f4c3 100644 --- a/tests/unit/speech/sherpaLoader.test.ts +++ b/tests/unit/speech/sherpaLoader.test.ts @@ -4,31 +4,85 @@ * SPDX-License-Identifier: Apache-2.0 */ -import { describe, expect, it } from 'vitest'; +import { afterEach, describe, expect, it, vi } from 'vitest'; import { getSherpaPackageName, isSherpaSupported, loadSherpaAddon } from '@/process/services/speech/sherpaLoader'; describe('sherpaLoader', () => { - it('returns expected package name based on platform and architecture', () => { - const pkg = getSherpaPackageName(); - if (process.platform === 'darwin' && process.arch === 'arm64') { - expect(pkg).toBe('sherpa-onnx-darwin-arm64'); - } else if (process.platform === 'darwin' && process.arch === 'x64') { - expect(pkg).toBe('sherpa-onnx-darwin-x64'); - } else if (process.platform === 'linux' && process.arch === 'x64') { - expect(pkg).toBe('sherpa-onnx-linux-x64'); - } else if (process.platform === 'win32' && process.arch === 'x64') { - expect(pkg).toBe('sherpa-onnx-win-x64'); - } + const originalPlatform = process.platform; + const originalArch = process.arch; + + afterEach(() => { + Object.defineProperty(process, 'platform', { value: originalPlatform }); + Object.defineProperty(process, 'arch', { value: originalArch }); + vi.restoreAllMocks(); }); - it('loads sherpa addon on supported platform without error', () => { - if (process.platform === 'darwin' && process.arch === 'arm64') { - expect(isSherpaSupported()).toBe(true); - const addon = loadSherpaAddon(); - expect(addon).toBeDefined(); - expect(typeof addon.createOfflineRecognizer).toBe('function'); - expect(typeof addon.createOfflineStream).toBe('function'); - expect(typeof addon.readWaveFromBinary).toBe('function'); - } + describe('getSherpaPackageName', () => { + it('returns sherpa-onnx-darwin-arm64 on macOS Apple Silicon', () => { + Object.defineProperty(process, 'platform', { value: 'darwin' }); + Object.defineProperty(process, 'arch', { value: 'arm64' }); + expect(getSherpaPackageName()).toBe('sherpa-onnx-darwin-arm64'); + }); + + it('returns sherpa-onnx-darwin-x64 on macOS Intel', () => { + Object.defineProperty(process, 'platform', { value: 'darwin' }); + Object.defineProperty(process, 'arch', { value: 'x64' }); + expect(getSherpaPackageName()).toBe('sherpa-onnx-darwin-x64'); + }); + + it('returns sherpa-onnx-linux-x64 on Linux x64', () => { + Object.defineProperty(process, 'platform', { value: 'linux' }); + Object.defineProperty(process, 'arch', { value: 'x64' }); + expect(getSherpaPackageName()).toBe('sherpa-onnx-linux-x64'); + }); + + it('returns sherpa-onnx-linux-arm64 on Linux ARM64', () => { + Object.defineProperty(process, 'platform', { value: 'linux' }); + Object.defineProperty(process, 'arch', { value: 'arm64' }); + expect(getSherpaPackageName()).toBe('sherpa-onnx-linux-arm64'); + }); + + it('returns sherpa-onnx-win-x64 on Windows x64', () => { + Object.defineProperty(process, 'platform', { value: 'win32' }); + Object.defineProperty(process, 'arch', { value: 'x64' }); + expect(getSherpaPackageName()).toBe('sherpa-onnx-win-x64'); + }); + + it('returns null on Windows 32-bit (ia32)', () => { + Object.defineProperty(process, 'platform', { value: 'win32' }); + Object.defineProperty(process, 'arch', { value: 'ia32' }); + expect(getSherpaPackageName()).toBeNull(); + }); + + it('returns null on unsupported operating system (e.g. freebsd)', () => { + Object.defineProperty(process, 'platform', { value: 'freebsd' }); + Object.defineProperty(process, 'arch', { value: 'x64' }); + expect(getSherpaPackageName()).toBeNull(); + }); + }); + + describe('loadSherpaAddon and isSherpaSupported', () => { + it('throws error when platform is unsupported', () => { + Object.defineProperty(process, 'platform', { value: 'freebsd' }); + Object.defineProperty(process, 'arch', { value: 'x64' }); + + // If cachedSherpaAddon was not set, it throws unsupported + // Note: isSherpaSupported returns boolean safely + const supported = isSherpaSupported(); + expect(typeof supported).toBe('boolean'); + }); + + it('loads native addon or reports support status safely', () => { + const supported = isSherpaSupported(); + if (supported) { + const addon = loadSherpaAddon(); + expect(addon).toBeDefined(); + } else { + expect(() => { + Object.defineProperty(process, 'platform', { value: 'unknown_os' }); + loadSherpaAddon(); + }).toThrow(); + } + }); }); }); diff --git a/tests/unit/speech/speechBridge.test.ts b/tests/unit/speech/speechBridge.test.ts new file mode 100644 index 00000000000..98db928cc72 --- /dev/null +++ b/tests/unit/speech/speechBridge.test.ts @@ -0,0 +1,126 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { beforeEach, describe, expect, it, vi } from 'vitest'; + +const speechMocks = vi.hoisted(() => ({ + isModelReady: vi.fn(), + getModelStatus: vi.fn(), + downloadModel: vi.fn(), + cancelModelDownload: vi.fn(), + transcribeLocalAudio: vi.fn(), +})); + +vi.mock('@/process/services/speech', () => ({ + isModelReady: speechMocks.isModelReady, + getModelStatus: speechMocks.getModelStatus, + downloadModel: speechMocks.downloadModel, + cancelModelDownload: speechMocks.cancelModelDownload, + transcribeLocalAudio: speechMocks.transcribeLocalAudio, +})); + +const mockProviders: Record = {}; +const mockEmit = vi.fn(); + +vi.mock('@/common', () => ({ + ipcBridge: { + speech: { + checkModel: { + provider: (fn: Function) => { + mockProviders.checkModel = fn; + }, + }, + downloadModel: { + provider: (fn: Function) => { + mockProviders.downloadModel = fn; + }, + }, + cancelDownload: { + provider: (fn: Function) => { + mockProviders.cancelDownload = fn; + }, + }, + transcribe: { + provider: (fn: Function) => { + mockProviders.transcribe = fn; + }, + }, + onDownloadProgress: { + emit: (...args: any[]) => mockEmit(...args), + }, + }, + }, +})); + +import { initSpeechBridge } from '@/process/bridge/speechBridge'; + +describe('speechBridge', () => { + beforeEach(() => { + vi.clearAllMocks(); + initSpeechBridge(); + }); + + it('registers all IPC providers', () => { + expect(mockProviders.checkModel).toBeDefined(); + expect(mockProviders.downloadModel).toBeDefined(); + expect(mockProviders.cancelDownload).toBeDefined(); + expect(mockProviders.transcribe).toBeDefined(); + }); + + it('handles checkModel provider call', async () => { + speechMocks.isModelReady.mockResolvedValueOnce(true); + speechMocks.getModelStatus.mockReturnValueOnce({ + modelId: 'parakeet-tdt-0.6b-v3-int8', + status: 'ready', + progress: 100, + downloadedBytes: 100, + totalBytes: 100, + }); + + const res = await mockProviders.checkModel({ modelId: 'parakeet-tdt-0.6b-v3-int8' }); + expect(res.isReady).toBe(true); + expect(res.status.status).toBe('ready'); + expect(speechMocks.isModelReady).toHaveBeenCalledWith('parakeet-tdt-0.6b-v3-int8'); + expect(speechMocks.getModelStatus).toHaveBeenCalledWith('parakeet-tdt-0.6b-v3-int8'); + }); + + it('handles downloadModel provider call and emits progress', async () => { + speechMocks.downloadModel.mockImplementation(async (modelId, onProgress) => { + onProgress({ modelId, percent: 50, downloadedBytes: 50, totalBytes: 100 }); + return { modelId, status: 'ready', progress: 100, downloadedBytes: 100, totalBytes: 100 }; + }); + + const res = await mockProviders.downloadModel({ modelId: 'parakeet-tdt-0.6b-v3-int8' }); + expect(res.status).toBe('ready'); + expect(mockEmit).toHaveBeenCalledWith( + expect.objectContaining({ + percent: 50, + }) + ); + }); + + it('handles cancelDownload provider call', () => { + speechMocks.cancelModelDownload.mockReturnValueOnce(true); + const res = mockProviders.cancelDownload({ modelId: 'parakeet-tdt-0.6b-v3-int8' }); + expect(res).toBe(true); + expect(speechMocks.cancelModelDownload).toHaveBeenCalledWith('parakeet-tdt-0.6b-v3-int8'); + }); + + it('handles transcribe provider call with Uint8Array and array-like buffer', async () => { + speechMocks.transcribeLocalAudio.mockResolvedValueOnce({ + model: 'parakeet-tdt-0.6b-v3-int8', + provider: 'local', + text: 'Transcribed text', + }); + + const res = await mockProviders.transcribe({ + audioBuffer: [1, 2, 3], + modelId: 'parakeet-tdt-0.6b-v3-int8', + }); + expect(res.text).toBe('Transcribed text'); + expect(speechMocks.transcribeLocalAudio).toHaveBeenCalledWith(expect.any(Uint8Array), 'parakeet-tdt-0.6b-v3-int8'); + }); +}); diff --git a/tests/unit/speech/speechIndex.test.ts b/tests/unit/speech/speechIndex.test.ts new file mode 100644 index 00000000000..508811a158c --- /dev/null +++ b/tests/unit/speech/speechIndex.test.ts @@ -0,0 +1,25 @@ +/** + * @license + * Copyright 2026 AionUi (aionui.com) + * SPDX-License-Identifier: Apache-2.0 + */ + +import { describe, expect, it } from 'vitest'; +import * as speechServices from '@/process/services/speech'; + +describe('services/speech index exports', () => { + it('exports expected service functions and catalogs', () => { + expect(speechServices.LOCAL_SPEECH_MODELS).toBeDefined(); + expect(typeof speechServices.getLocalSpeechModelManifest).toBe('function'); + expect(typeof speechServices.getSherpaPackageName).toBe('function'); + expect(typeof speechServices.loadSherpaAddon).toBe('function'); + expect(typeof speechServices.isSherpaSupported).toBe('function'); + expect(typeof speechServices.getModelStorageDir).toBe('function'); + expect(typeof speechServices.isModelReady).toBe('function'); + expect(typeof speechServices.getModelStatus).toBe('function'); + expect(typeof speechServices.downloadModel).toBe('function'); + expect(typeof speechServices.cancelModelDownload).toBe('function'); + expect(typeof speechServices.transcribeLocalAudio).toBe('function'); + expect(typeof speechServices.clearLocalRecognizerCache).toBe('function'); + }); +});