From 277080a004c00d68f0ceb7b50c4e58c1b1e50c8d Mon Sep 17 00:00:00 2001 From: Mohamed Boudra Date: Fri, 6 Feb 2026 12:43:10 +0700 Subject: [PATCH] Refactor speech providers to stream-first local/openai architecture --- packages/cli/src/cli.ts | 4 + packages/cli/src/commands/speech/download.ts | 68 ++++ packages/cli/src/commands/speech/index.ts | 29 ++ packages/cli/src/commands/speech/models.ts | 72 ++++ .../server/scripts/download-speech-models.ts | 4 +- packages/server/scripts/list-speech-models.ts | 3 +- packages/server/src/client/daemon-client.ts | 53 +++ .../src/server/agent/stt-manager.test.ts | 56 ++- .../server/src/server/agent/stt-manager.ts | 122 +++++-- .../server/src/server/agent/stt-openai.ts | 139 ------- packages/server/src/server/bootstrap.ts | 103 ++++-- packages/server/src/server/config.ts | 20 +- .../src/server/daemon-client.e2e.test.ts | 6 +- .../dictation-stream-manager.test.ts | 86 ++--- .../dictation/dictation-stream-manager.ts | 338 ++++-------------- .../server/src/server/persisted-config.ts | 20 +- packages/server/src/server/session.ts | 131 ++++++- .../local}/pocket/pocket-tts-onnx.ts | 6 +- .../local}/sherpa/model-catalog.ts | 0 .../local}/sherpa/model-downloader.test.ts | 0 .../local}/sherpa/model-downloader.ts | 0 .../sherpa/sherpa-offline-recognizer.ts | 0 .../local}/sherpa/sherpa-online-recognizer.ts | 0 .../local}/sherpa/sherpa-onnx-loader.ts | 0 .../local}/sherpa/sherpa-onnx-node-loader.ts | 0 .../sherpa-parakeet-realtime-session.ts | 43 +-- .../local}/sherpa/sherpa-parakeet-stt.ts | 87 ++++- .../local}/sherpa/sherpa-realtime-session.ts | 45 +-- .../local}/sherpa/sherpa-stt.ts | 86 ++++- .../local}/sherpa/sherpa-tts.ts | 4 +- .../local}/sherpa/speech-download.e2e.test.ts | 10 +- .../openai/realtime-transcription-session.ts} | 18 +- .../src/server/speech/providers/openai/stt.ts | 269 ++++++++++++++ .../providers/openai/tts.ts} | 4 +- .../src/server/speech/speech-provider.ts | 48 ++- .../server/src/server/websocket-server.ts | 7 +- packages/server/src/shared/messages.ts | 49 +++ scripts/speech/download-sherpa-models.sh | 95 ----- 38 files changed, 1294 insertions(+), 731 deletions(-) create mode 100644 packages/cli/src/commands/speech/download.ts create mode 100644 packages/cli/src/commands/speech/index.ts create mode 100644 packages/cli/src/commands/speech/models.ts delete mode 100644 packages/server/src/server/agent/stt-openai.ts rename packages/server/src/server/speech/{ => providers/local}/pocket/pocket-tts-onnx.ts (99%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/model-catalog.ts (100%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/model-downloader.test.ts (100%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/model-downloader.ts (100%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-offline-recognizer.ts (100%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-online-recognizer.ts (100%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-onnx-loader.ts (100%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-onnx-node-loader.ts (100%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-parakeet-realtime-session.ts (73%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-parakeet-stt.ts (50%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-realtime-session.ts (70%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-stt.ts (55%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/sherpa-tts.ts (97%) rename packages/server/src/server/speech/{ => providers/local}/sherpa/speech-download.e2e.test.ts (97%) rename packages/server/src/server/{agent/openai-realtime-transcription.ts => speech/providers/openai/realtime-transcription-session.ts} (91%) create mode 100644 packages/server/src/server/speech/providers/openai/stt.ts rename packages/server/src/server/{agent/tts-openai.ts => speech/providers/openai/tts.ts} (93%) delete mode 100755 scripts/speech/download-sherpa-models.sh diff --git a/packages/cli/src/cli.ts b/packages/cli/src/cli.ts index cee6d72f6..b3222c7c4 100644 --- a/packages/cli/src/cli.ts +++ b/packages/cli/src/cli.ts @@ -3,6 +3,7 @@ import { createAgentCommand } from './commands/agent/index.js' import { createDaemonCommand } from './commands/daemon/index.js' import { createPermitCommand } from './commands/permit/index.js' import { createProviderCommand } from './commands/provider/index.js' +import { createSpeechCommand } from './commands/speech/index.js' import { createWorktreeCommand } from './commands/worktree/index.js' import { runLsCommand } from './commands/agent/ls.js' import { runRunCommand } from './commands/agent/run.js' @@ -134,6 +135,9 @@ export function createCli(): Command { // Provider commands program.addCommand(createProviderCommand()) + // Speech model commands + program.addCommand(createSpeechCommand()) + // Worktree commands program.addCommand(createWorktreeCommand()) diff --git a/packages/cli/src/commands/speech/download.ts b/packages/cli/src/commands/speech/download.ts new file mode 100644 index 000000000..c37b84b18 --- /dev/null +++ b/packages/cli/src/commands/speech/download.ts @@ -0,0 +1,68 @@ +import type { Command } from "commander"; +import type { + CommandError, + CommandOptions, + ListResult, + OutputSchema, +} from "../../output/index.js"; +import { connectToDaemon } from "../../utils/client.js"; + +interface SpeechDownloadRow { + modelId: string; + status: "downloaded"; +} + +const speechDownloadSchema: OutputSchema = { + idField: "modelId", + columns: [ + { header: "MODEL", field: "modelId", width: 36 }, + { header: "STATUS", field: "status", width: 12, color: () => "green" }, + ], +}; + +export type SpeechDownloadResult = ListResult; + +export interface SpeechDownloadOptions extends CommandOptions { + host?: string; + model?: string[]; +} + +export async function runSpeechDownloadCommand( + options: SpeechDownloadOptions, + _command: Command +): Promise { + const client = await connectToDaemon({ host: options.host }); + try { + const response = await client.downloadSpeechModels({ + modelIds: options.model && options.model.length > 0 ? options.model : undefined, + }); + if (response.error) { + const commandError: CommandError = { + code: "SPEECH_MODELS_DOWNLOAD_FAILED", + message: response.error, + }; + throw commandError; + } + + return { + type: "list", + data: response.downloadedModelIds.map((modelId) => ({ + modelId, + status: "downloaded" as const, + })), + schema: speechDownloadSchema, + }; + } catch (error) { + if (typeof error === "object" && error && "code" in error && "message" in error) { + throw error; + } + const message = error instanceof Error ? error.message : String(error); + const commandError: CommandError = { + code: "SPEECH_MODELS_DOWNLOAD_FAILED", + message: `Failed to download speech models: ${message}`, + }; + throw commandError; + } finally { + await client.close().catch(() => {}); + } +} diff --git a/packages/cli/src/commands/speech/index.ts b/packages/cli/src/commands/speech/index.ts new file mode 100644 index 000000000..6661313a5 --- /dev/null +++ b/packages/cli/src/commands/speech/index.ts @@ -0,0 +1,29 @@ +import { Command } from "commander"; +import { withOutput } from "../../output/index.js"; +import { runSpeechModelsCommand } from "./models.js"; +import { runSpeechDownloadCommand } from "./download.js"; + +function collectMultiple(value: string, previous: string[]): string[] { + return previous.concat([value]); +} + +export function createSpeechCommand(): Command { + const speech = new Command("speech").description("Manage local speech models"); + + speech + .command("models") + .description("List local speech model download status") + .option("--json", "Output in JSON format") + .option("--host ", "Daemon host:port (default: localhost:6767)") + .action(withOutput(runSpeechModelsCommand)); + + speech + .command("download") + .description("Download local speech models") + .option("--model ", "Model ID to download (repeatable)", collectMultiple, []) + .option("--json", "Output in JSON format") + .option("--host ", "Daemon host:port (default: localhost:6767)") + .action(withOutput(runSpeechDownloadCommand)); + + return speech; +} diff --git a/packages/cli/src/commands/speech/models.ts b/packages/cli/src/commands/speech/models.ts new file mode 100644 index 000000000..eca33edd6 --- /dev/null +++ b/packages/cli/src/commands/speech/models.ts @@ -0,0 +1,72 @@ +import type { Command } from "commander"; +import type { + CommandError, + CommandOptions, + ListResult, + OutputSchema, +} from "../../output/index.js"; +import { connectToDaemon } from "../../utils/client.js"; + +interface SpeechModelListItem { + id: string; + kind: string; + status: "downloaded" | "missing"; + modelDir: string; + missingFiles: string; +} + +const speechModelsSchema: OutputSchema = { + idField: "id", + columns: [ + { header: "MODEL", field: "id", width: 36 }, + { header: "KIND", field: "kind", width: 12 }, + { + header: "STATUS", + field: "status", + width: 12, + color: (value) => (value === "downloaded" ? "green" : "yellow"), + }, + { header: "MODEL DIR", field: "modelDir", width: 44 }, + { header: "MISSING FILES", field: "missingFiles", width: 40 }, + ], +}; + +export type SpeechModelsResult = ListResult; + +export interface SpeechModelsOptions extends CommandOptions { + host?: string; +} + +export async function runSpeechModelsCommand( + options: SpeechModelsOptions, + _command: Command +): Promise { + const client = await connectToDaemon({ host: options.host }); + try { + const response = await client.listSpeechModels(); + const rows: SpeechModelListItem[] = response.models + .slice() + .sort((a, b) => a.kind.localeCompare(b.kind) || a.id.localeCompare(b.id)) + .map((model) => ({ + id: model.id, + kind: model.kind, + status: model.isDownloaded ? "downloaded" : "missing", + modelDir: model.modelDir, + missingFiles: model.missingFiles?.join(", ") ?? "", + })); + return { + type: "list", + data: rows, + schema: speechModelsSchema, + }; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + const commandError: CommandError = { + code: "SPEECH_MODELS_LIST_FAILED", + message: `Failed to list speech models: ${message}`, + }; + throw commandError; + } finally { + await client.close().catch(() => {}); + } +} diff --git a/packages/server/scripts/download-speech-models.ts b/packages/server/scripts/download-speech-models.ts index 978cc3e85..a531c5942 100644 --- a/packages/server/scripts/download-speech-models.ts +++ b/packages/server/scripts/download-speech-models.ts @@ -1,7 +1,7 @@ import { resolvePaseoHome } from "../src/server/paseo-home.js"; import { createRootLogger } from "../src/server/logger.js"; -import { ensureSherpaOnnxModels } from "../src/server/speech/sherpa/model-downloader.js"; -import type { SherpaOnnxModelId } from "../src/server/speech/sherpa/model-catalog.js"; +import { ensureSherpaOnnxModels } from "../src/server/speech/providers/local/sherpa/model-downloader.js"; +import type { SherpaOnnxModelId } from "../src/server/speech/providers/local/sherpa/model-catalog.js"; function parseArgs(argv: string[]): { modelsDir: string; modelIds: SherpaOnnxModelId[] } { const home = resolvePaseoHome(); diff --git a/packages/server/scripts/list-speech-models.ts b/packages/server/scripts/list-speech-models.ts index d907ab4eb..8462dfc27 100644 --- a/packages/server/scripts/list-speech-models.ts +++ b/packages/server/scripts/list-speech-models.ts @@ -1,4 +1,4 @@ -import { listSherpaOnnxModels } from "../src/server/speech/sherpa/model-catalog.js"; +import { listSherpaOnnxModels } from "../src/server/speech/providers/local/sherpa/model-catalog.js"; const models = listSherpaOnnxModels() .slice() @@ -8,4 +8,3 @@ for (const m of models) { // eslint-disable-next-line no-console console.log(`${m.kind}\t${m.id}\t${m.description}`); } - diff --git a/packages/server/src/client/daemon-client.ts b/packages/server/src/client/daemon-client.ts index 9ced8cff3..8994e9407 100644 --- a/packages/server/src/client/daemon-client.ts +++ b/packages/server/src/client/daemon-client.ts @@ -36,6 +36,8 @@ import type { ExecuteCommandResponse, ListVoiceConversationsResponseMessage, ListProviderModelsResponseMessage, + SpeechModelsListResponse, + SpeechModelsDownloadResponse, ListTerminalsResponse, CreateTerminalResponse, SubscribeTerminalResponse, @@ -201,6 +203,8 @@ type PaseoWorktreeArchivePayload = PaseoWorktreeArchiveResponse["payload"]; type FileExplorerPayload = FileExplorerResponse["payload"]; type FileDownloadTokenPayload = FileDownloadTokenResponse["payload"]; type ListProviderModelsPayload = ListProviderModelsResponseMessage["payload"]; +type SpeechModelsListPayload = SpeechModelsListResponse["payload"]; +type SpeechModelsDownloadPayload = SpeechModelsDownloadResponse["payload"]; type ListCommandsPayload = ListCommandsResponse["payload"]; type ExecuteCommandPayload = ExecuteCommandResponse["payload"]; type AgentPermissionResolvedPayload = AgentPermissionResolvedMessage["payload"]; @@ -2014,6 +2018,55 @@ export class DaemonClient { }); } + async listSpeechModels(requestId?: string): Promise { + const resolvedRequestId = this.createRequestId(requestId); + const message = SessionInboundMessageSchema.parse({ + type: "speech_models_list_request", + requestId: resolvedRequestId, + }); + return this.sendRequest({ + requestId: resolvedRequestId, + message, + timeout: 30000, + options: { skipQueue: true }, + select: (msg) => { + if (msg.type !== "speech_models_list_response") { + return null; + } + if (msg.payload.requestId !== resolvedRequestId) { + return null; + } + return msg.payload; + }, + }); + } + + async downloadSpeechModels( + options?: { modelIds?: string[]; requestId?: string } + ): Promise { + const resolvedRequestId = this.createRequestId(options?.requestId); + const message = SessionInboundMessageSchema.parse({ + type: "speech_models_download_request", + modelIds: options?.modelIds, + requestId: resolvedRequestId, + }); + return this.sendRequest({ + requestId: resolvedRequestId, + message, + timeout: 30 * 60 * 1000, + options: { skipQueue: true }, + select: (msg) => { + if (msg.type !== "speech_models_download_response") { + return null; + } + if (msg.payload.requestId !== resolvedRequestId) { + return null; + } + return msg.payload; + }, + }); + } + async listCommands( agentId: string, requestId?: string diff --git a/packages/server/src/server/agent/stt-manager.test.ts b/packages/server/src/server/agent/stt-manager.test.ts index df88f3bea..8b8f91a68 100644 --- a/packages/server/src/server/agent/stt-manager.test.ts +++ b/packages/server/src/server/agent/stt-manager.test.ts @@ -1,14 +1,53 @@ import { describe, expect, it } from "vitest"; import pino from "pino"; +import { EventEmitter } from "node:events"; import { STTManager } from "./stt-manager.js"; -import type { SpeechToTextProvider, TranscriptionResult } from "../speech/speech-provider.js"; +import type { + SpeechToTextProvider, + StreamingTranscriptionSession, + TranscriptionResult, +} from "../speech/speech-provider.js"; class FakeStt implements SpeechToTextProvider { + public readonly id = "fake"; constructor(private readonly result: TranscriptionResult) {} - async transcribeAudio(): Promise { - return this.result; + createSession(_params: { + logger: any; + language?: string; + prompt?: string; + }): StreamingTranscriptionSession { + const emitter = new EventEmitter(); + const result = this.result; + let segmentId = "seg-1"; + let previousSegmentId: string | null = null; + + return { + requiredSampleRate: 24000, + async connect() {}, + appendPcm16() {}, + commit() { + (emitter as any).emit("committed", { segmentId, previousSegmentId }); + (emitter as any).emit("transcript", { + segmentId, + transcript: result.text, + isFinal: true, + language: result.language, + logprobs: result.logprobs, + avgLogprob: result.avgLogprob, + isLowConfidence: result.isLowConfidence, + }); + previousSegmentId = segmentId; + segmentId = "seg-2"; + }, + clear() {}, + close() {}, + on(event: any, handler: any) { + emitter.on(event, handler); + return undefined; + }, + }; } } @@ -20,10 +59,12 @@ describe("STTManager", () => { new FakeStt({ text: "um", isLowConfidence: true, avgLogprob: -10 }) ); - const result = await manager.transcribe(Buffer.from("x"), "audio/wav", { label: "t" }); + const result = await manager.transcribe(Buffer.alloc(2), "audio/pcm;rate=24000", { + label: "t", + }); expect(result.text).toBe(""); expect(result.isLowConfidence).toBe(true); - expect(result.byteLength).toBe(1); + expect(result.byteLength).toBe(2); }); it("passes through normal transcriptions", async () => { @@ -33,10 +74,9 @@ describe("STTManager", () => { new FakeStt({ text: "hello world", language: "en", isLowConfidence: false }) ); - const result = await manager.transcribe(Buffer.from("abc"), "audio/wav"); + const result = await manager.transcribe(Buffer.alloc(4), "audio/pcm;rate=24000"); expect(result.text).toBe("hello world"); expect(result.language).toBe("en"); - expect(result.byteLength).toBe(3); + expect(result.byteLength).toBe(4); }); }); - diff --git a/packages/server/src/server/agent/stt-manager.ts b/packages/server/src/server/agent/stt-manager.ts index cb454dc3b..6cdd35ea8 100644 --- a/packages/server/src/server/agent/stt-manager.ts +++ b/packages/server/src/server/agent/stt-manager.ts @@ -1,6 +1,8 @@ import type pino from "pino"; import type { SpeechToTextProvider, TranscriptionResult } from "../speech/speech-provider.js"; import { maybePersistDebugAudio } from "./stt-debug.js"; +import { parsePcm16MonoWav, parsePcmRateFromFormat } from "../speech/audio.js"; +import { Pcm16MonoResampler } from "./pcm16-resampler.js"; interface TranscriptionMetadata { agentId?: string; @@ -63,36 +65,104 @@ export class STTManager { this.logger.warn({ err: error }, "Failed to persist debug audio"); } - const result = await this.stt.transcribeAudio(audio, format); + const session = this.stt.createSession({ + logger: this.logger.child({ component: "stt-session" }), + language: "en", + }); - // Filter out low-confidence transcriptions (non-speech sounds) - if (result.isLowConfidence) { - this.logger.debug( - { text: result.text, avgLogprob: result.avgLogprob }, - "Filtered low-confidence transcription (likely non-speech)" - ); - - // Return empty text to ignore this transcription - return { - ...result, - text: "", - byteLength: audio.length, - format, - debugRecordingPath: debugRecordingPath ?? undefined, - }; + let inputRate: number; + let pcm16: Buffer; + if (format.toLowerCase().includes("audio/wav")) { + const parsed = parsePcm16MonoWav(audio); + inputRate = parsed.sampleRate; + pcm16 = parsed.pcm16; + } else if (format.toLowerCase().includes("audio/pcm")) { + inputRate = + parsePcmRateFromFormat(format, session.requiredSampleRate) ?? + session.requiredSampleRate; + pcm16 = audio; + } else { + throw new Error(`Unsupported audio format for STT: ${format}`); } - this.logger.debug( - { text: result.text, avgLogprob: result.avgLogprob }, - "Transcription complete" - ); + let pcmForModel = pcm16; + if (inputRate !== session.requiredSampleRate) { + const resampler = new Pcm16MonoResampler({ + inputRate, + outputRate: session.requiredSampleRate, + }); + pcmForModel = resampler.processChunk(pcm16); + inputRate = session.requiredSampleRate; + } - return { - ...result, - debugRecordingPath: debugRecordingPath ?? undefined, - byteLength: audio.length, - format, - }; + try { + const startedAt = Date.now(); + const finalEventPromise = new Promise<{ + transcript: string; + language?: string; + logprobs?: TranscriptionResult["logprobs"]; + avgLogprob?: number; + isLowConfidence?: boolean; + }>((resolve, reject) => { + session.on("error", reject); + session.on("transcript", (payload) => { + if (!payload.isFinal) { + return; + } + resolve({ + transcript: payload.transcript, + language: payload.language, + logprobs: payload.logprobs, + avgLogprob: payload.avgLogprob, + isLowConfidence: payload.isLowConfidence, + }); + }); + }); + + await session.connect(); + session.appendPcm16(pcmForModel); + session.commit(); + const finalEvent = await finalEventPromise; + const result: TranscriptionResult = { + text: finalEvent.transcript, + language: finalEvent.language, + logprobs: finalEvent.logprobs, + avgLogprob: finalEvent.avgLogprob, + isLowConfidence: finalEvent.isLowConfidence, + duration: Date.now() - startedAt, + }; + + // Filter out low-confidence transcriptions (non-speech sounds) + if (result.isLowConfidence) { + this.logger.debug( + { text: result.text, avgLogprob: result.avgLogprob }, + "Filtered low-confidence transcription (likely non-speech)" + ); + + // Return empty text to ignore this transcription + return { + ...result, + text: "", + byteLength: audio.length, + format, + debugRecordingPath: debugRecordingPath ?? undefined, + }; + } + + this.logger.debug( + { text: result.text, avgLogprob: result.avgLogprob }, + "Transcription complete" + ); + + return { + ...result, + debugRecordingPath: debugRecordingPath ?? undefined, + byteLength: audio.length, + format, + }; + } finally { + session.close(); + } } /** diff --git a/packages/server/src/server/agent/stt-openai.ts b/packages/server/src/server/agent/stt-openai.ts deleted file mode 100644 index 4326f6c68..000000000 --- a/packages/server/src/server/agent/stt-openai.ts +++ /dev/null @@ -1,139 +0,0 @@ -import type pino from "pino"; -import OpenAI from "openai"; -import { writeFile, unlink } from "fs/promises"; -import { join } from "path"; -import { tmpdir } from "os"; -import { v4 } from "uuid"; -import { inferAudioExtension } from "./audio-utils.js"; -import type { LogprobToken, TranscriptionResult } from "../speech/speech-provider.js"; - -export type { LogprobToken, TranscriptionResult }; - -export interface STTConfig { - apiKey: string; - model?: "whisper-1" | "gpt-4o-transcribe" | "gpt-4o-mini-transcribe" | (string & {}); - confidenceThreshold?: number; // Default: -3.0 -} - -function isObject(value: unknown): value is { [key: string]: unknown } { - return typeof value === "object" && value !== null; -} - -function isLogprobToken(value: unknown): value is LogprobToken { - if (!isObject(value)) { - return false; - } - if (typeof value.token !== "string") { - return false; - } - if (typeof value.logprob !== "number") { - return false; - } - if (value.bytes === undefined) { - return true; - } - return Array.isArray(value.bytes) && value.bytes.every((entry) => typeof entry === "number"); -} - -function isLogprobTokenArray(value: unknown): value is LogprobToken[] { - return Array.isArray(value) && value.every((entry) => isLogprobToken(entry)); -} - -export class OpenAISTT { - private readonly openaiClient: OpenAI; - private readonly config: STTConfig; - private readonly logger: pino.Logger; - - constructor(sttConfig: STTConfig, parentLogger: pino.Logger) { - this.config = sttConfig; - this.logger = parentLogger.child({ module: "agent", provider: "openai", component: "stt" }); - this.openaiClient = new OpenAI({ - apiKey: sttConfig.apiKey, - }); - this.logger.info({ model: sttConfig.model || "whisper-1" }, "STT (OpenAI Whisper) initialized"); - } - - public async transcribeAudio(audioBuffer: Buffer, format: string): Promise { - const startTime = Date.now(); - let tempFilePath: string | null = null; - - try { - const ext = inferAudioExtension(format); - tempFilePath = join(tmpdir(), `audio-${v4()}.${ext}`); - await writeFile(tempFilePath, audioBuffer); - - this.logger.debug( - { tempFilePath, bytes: audioBuffer.length }, - "Transcribing audio file" - ); - - const modelToUse = this.config.model ?? "whisper-1"; - const supportsLogprobs = - modelToUse === "gpt-4o-transcribe" || modelToUse === "gpt-4o-mini-transcribe"; - const includeLogprobs: ["logprobs"] = ["logprobs"]; - - const response = await this.openaiClient.audio.transcriptions.create({ - file: await import("fs").then((fs) => fs.createReadStream(tempFilePath!)), - language: "en", - model: modelToUse, - ...(supportsLogprobs ? { include: includeLogprobs } : {}), - response_format: "json", - }); - - const duration = Date.now() - startTime; - const confidenceThreshold = this.config.confidenceThreshold ?? -3.0; - - let avgLogprob: number | undefined; - let isLowConfidence = false; - const logprobs = - supportsLogprobs && - isObject(response) && - isLogprobTokenArray(response.logprobs) - ? response.logprobs - : undefined; - - if (logprobs && logprobs.length > 0) { - const totalLogprob = logprobs.reduce((sum, token) => sum + token.logprob, 0); - avgLogprob = totalLogprob / logprobs.length; - isLowConfidence = avgLogprob < confidenceThreshold; - - if (isLowConfidence) { - this.logger.debug( - { - avgLogprob, - threshold: confidenceThreshold, - text: response.text, - tokenLogprobs: logprobs.map((t) => `${t.token}:${t.logprob.toFixed(2)}`).join(", "), - }, - "Low confidence transcription detected" - ); - } - } - - this.logger.debug({ duration, text: response.text, avgLogprob }, "Transcription complete"); - - return { - text: response.text, - duration: duration, - logprobs: logprobs, - avgLogprob: avgLogprob, - isLowConfidence: isLowConfidence, - language: - isObject(response) && typeof response.language === "string" - ? response.language - : undefined, - }; - } catch (error: any) { - this.logger.error({ err: error }, "Transcription error"); - throw new Error(`STT transcription failed: ${error.message}`); - } finally { - if (tempFilePath) { - try { - await unlink(tempFilePath); - } catch (cleanupError) { - this.logger.warn({ tempFilePath }, "Failed to clean up temp file"); - } - } - } - } -} diff --git a/packages/server/src/server/bootstrap.ts b/packages/server/src/server/bootstrap.ts index b47deb61f..5dcfd16d3 100644 --- a/packages/server/src/server/bootstrap.ts +++ b/packages/server/src/server/bootstrap.ts @@ -40,20 +40,20 @@ function parseListenString(listen: string): ListenTarget { import { VoiceAssistantWebSocketServer } from "./websocket-server.js"; import { DownloadTokenStore } from "./file-download/token-store.js"; -import { OpenAISTT, type STTConfig } from "./agent/stt-openai.js"; -import { OpenAITTS, type TTSConfig } from "./agent/tts-openai.js"; +import { OpenAISTT, type STTConfig } from "./speech/providers/openai/stt.js"; +import { OpenAITTS, type TTSConfig } from "./speech/providers/openai/tts.js"; +import { OpenAIRealtimeTranscriptionSession } from "./speech/providers/openai/realtime-transcription-session.js"; import type { SpeechToTextProvider, TextToSpeechProvider } from "./speech/speech-provider.js"; -import type { RealtimeTranscriptionSessionFactory } from "./dictation/dictation-stream-manager.js"; -import { SherpaOnlineRecognizerEngine } from "./speech/sherpa/sherpa-online-recognizer.js"; -import { SherpaOfflineRecognizerEngine } from "./speech/sherpa/sherpa-offline-recognizer.js"; -import { SherpaOnnxSTT } from "./speech/sherpa/sherpa-stt.js"; -import { SherpaOnnxParakeetSTT } from "./speech/sherpa/sherpa-parakeet-stt.js"; -import { SherpaOnnxTTS } from "./speech/sherpa/sherpa-tts.js"; -import { SherpaRealtimeTranscriptionSession } from "./speech/sherpa/sherpa-realtime-session.js"; -import { SherpaParakeetRealtimeTranscriptionSession } from "./speech/sherpa/sherpa-parakeet-realtime-session.js"; -import { ensureSherpaOnnxModels, getSherpaOnnxModelDir } from "./speech/sherpa/model-downloader.js"; -import type { SherpaOnnxModelId } from "./speech/sherpa/model-catalog.js"; -import { PocketTtsOnnxTTS } from "./speech/pocket/pocket-tts-onnx.js"; +import { SherpaOnlineRecognizerEngine } from "./speech/providers/local/sherpa/sherpa-online-recognizer.js"; +import { SherpaOfflineRecognizerEngine } from "./speech/providers/local/sherpa/sherpa-offline-recognizer.js"; +import { SherpaOnnxSTT } from "./speech/providers/local/sherpa/sherpa-stt.js"; +import { SherpaOnnxParakeetSTT } from "./speech/providers/local/sherpa/sherpa-parakeet-stt.js"; +import { SherpaOnnxTTS } from "./speech/providers/local/sherpa/sherpa-tts.js"; +import { SherpaRealtimeTranscriptionSession } from "./speech/providers/local/sherpa/sherpa-realtime-session.js"; +import { SherpaParakeetRealtimeTranscriptionSession } from "./speech/providers/local/sherpa/sherpa-parakeet-realtime-session.js"; +import { ensureSherpaOnnxModels, getSherpaOnnxModelDir } from "./speech/providers/local/sherpa/model-downloader.js"; +import type { SherpaOnnxModelId } from "./speech/providers/local/sherpa/model-catalog.js"; +import { PocketTtsOnnxTTS } from "./speech/providers/local/pocket/pocket-tts-onnx.js"; import { AgentManager } from "./agent/agent-manager.js"; import { AgentStorage } from "./agent/agent-storage.js"; import { attachAgentStoragePersistence } from "./persistence-hooks.js"; @@ -98,9 +98,9 @@ export type PaseoSherpaOnnxConfig = { }; export type PaseoSpeechConfig = { - dictationSttProvider?: "openai" | "sherpa"; - voiceSttProvider?: "openai" | "sherpa"; - voiceTtsProvider?: "openai" | "sherpa"; + dictationSttProvider?: "openai" | "local"; + voiceSttProvider?: "openai" | "local"; + voiceTtsProvider?: "openai" | "local"; sherpaOnnx?: PaseoSherpaOnnxConfig; }; @@ -417,7 +417,7 @@ export async function createPaseoDaemon( let sttService: SpeechToTextProvider | null = null; let ttsService: TextToSpeechProvider | null = null; - let dictationSessionFactory: RealtimeTranscriptionSessionFactory | undefined; + let dictationSttService: SpeechToTextProvider | null = null; let sherpaOnline: SherpaOnlineRecognizerEngine | null = null; let sherpaOffline: SherpaOfflineRecognizerEngine | null = null; @@ -427,11 +427,11 @@ export async function createPaseoDaemon( const speechConfig = config.speech ?? null; const sherpaConfig = speechConfig?.sherpaOnnx ?? null; - const wantsSherpaDictation = (speechConfig?.dictationSttProvider ?? "openai") === "sherpa"; - const wantsSherpaVoiceStt = (speechConfig?.voiceSttProvider ?? "openai") === "sherpa"; - const wantsSherpaVoiceTts = (speechConfig?.voiceTtsProvider ?? "openai") === "sherpa"; + const wantsLocalDictation = (speechConfig?.dictationSttProvider ?? "openai") === "local"; + const wantsLocalVoiceStt = (speechConfig?.voiceSttProvider ?? "openai") === "local"; + const wantsLocalVoiceTts = (speechConfig?.voiceTtsProvider ?? "openai") === "local"; - if ((wantsSherpaDictation || wantsSherpaVoiceStt || wantsSherpaVoiceTts) && sherpaConfig) { + if ((wantsLocalDictation || wantsLocalVoiceStt || wantsLocalVoiceTts) && sherpaConfig) { const autoDownload = sherpaConfig.autoDownload ?? (process.env.VITEST ? false : true); let sttPreset = (sherpaConfig.stt?.preset ?? "zipformer-bilingual-zh-en-2023-02-20").trim(); if ( @@ -460,10 +460,10 @@ export async function createPaseoDaemon( } const modelIds: SherpaOnnxModelId[] = []; - if (wantsSherpaDictation || wantsSherpaVoiceStt) { + if (wantsLocalDictation || wantsLocalVoiceStt) { modelIds.push(sttPreset as SherpaOnnxModelId); } - if (wantsSherpaVoiceTts) { + if (wantsLocalVoiceTts) { modelIds.push(ttsPreset as SherpaOnnxModelId); } @@ -489,7 +489,7 @@ export async function createPaseoDaemon( } } - if ((wantsSherpaDictation || wantsSherpaVoiceStt) && sherpaConfig) { + if ((wantsLocalDictation || wantsLocalVoiceStt) && sherpaConfig) { let preset = (sherpaConfig.stt?.preset ?? "zipformer-bilingual-zh-en-2023-02-20").trim(); if ( preset !== "zipformer-bilingual-zh-en-2023-02-20" && @@ -561,14 +561,14 @@ export async function createPaseoDaemon( sherpaOnline = null; sherpaOffline = null; } - } else if (wantsSherpaDictation || wantsSherpaVoiceStt) { + } else if (wantsLocalDictation || wantsLocalVoiceStt) { logger.warn( { configured: Boolean(sherpaConfig) }, "Sherpa STT selected but no sherpaOnnx config found; STT will be unavailable" ); } - if (wantsSherpaVoiceTts && sherpaConfig) { + if (wantsLocalVoiceTts && sherpaConfig) { let preset = (sherpaConfig.tts?.preset ?? "pocket-tts-onnx-int8").trim(); if ( preset !== "kitten-nano-en-v0_1-fp16" && @@ -615,27 +615,34 @@ export async function createPaseoDaemon( ); sherpaTts = null; } - } else if (wantsSherpaVoiceTts) { + } else if (wantsLocalVoiceTts) { logger.warn( { configured: Boolean(sherpaConfig) }, "Sherpa TTS selected but no sherpaOnnx config found; TTS will be unavailable" ); } - if (wantsSherpaVoiceStt && sherpaOffline) { + if (wantsLocalVoiceStt && sherpaOffline) { sttService = new SherpaOnnxParakeetSTT({ engine: sherpaOffline }, logger); - } else if (wantsSherpaVoiceStt && sherpaOnline) { + } else if (wantsLocalVoiceStt && sherpaOnline) { sttService = new SherpaOnnxSTT({ engine: sherpaOnline }, logger); } - if (wantsSherpaVoiceTts && sherpaTts) { + if (wantsLocalVoiceTts && sherpaTts) { ttsService = sherpaTts; } - if (wantsSherpaDictation && sherpaOnline) { - dictationSessionFactory = () => new SherpaRealtimeTranscriptionSession({ engine: sherpaOnline! }); - } else if (wantsSherpaDictation && sherpaOffline) { - dictationSessionFactory = () => new SherpaParakeetRealtimeTranscriptionSession({ engine: sherpaOffline! }); + if (wantsLocalDictation && sherpaOnline) { + dictationSttService = { + id: "local", + createSession: () => new SherpaRealtimeTranscriptionSession({ engine: sherpaOnline! }), + }; + } else if (wantsLocalDictation && sherpaOffline) { + dictationSttService = { + id: "local", + createSession: () => + new SherpaParakeetRealtimeTranscriptionSession({ engine: sherpaOffline! }), + }; } const voiceSttProvider = speechConfig?.voiceSttProvider ?? "openai"; @@ -645,10 +652,10 @@ export async function createPaseoDaemon( const needsOpenAiStt = !sttService && voiceSttProvider === "openai"; const needsOpenAiTts = !ttsService && voiceTtsProvider === "openai"; const needsOpenAiDictation = - dictationSttProvider === "openai" || (dictationSttProvider === "sherpa" && !dictationSessionFactory); + dictationSttProvider === "openai" || (dictationSttProvider === "local" && !dictationSttService); - const fallbackOpenAiStt = !sttService && voiceSttProvider === "sherpa" && Boolean(openaiApiKey); - const fallbackOpenAiTts = !ttsService && voiceTtsProvider === "sherpa" && Boolean(openaiApiKey); + const fallbackOpenAiStt = !sttService && voiceSttProvider === "local" && Boolean(openaiApiKey); + const fallbackOpenAiTts = !ttsService && voiceTtsProvider === "local" && Boolean(openaiApiKey); if ((needsOpenAiStt || needsOpenAiTts || needsOpenAiDictation || fallbackOpenAiStt || fallbackOpenAiTts) && openaiApiKey) { logger.info("OpenAI client initialized"); @@ -689,6 +696,25 @@ export async function createPaseoDaemon( ); } } + + if (needsOpenAiDictation) { + const dictationApiKey = config.openai?.apiKey ?? openaiApiKey; + const transcriptionModel = + process.env.OPENAI_REALTIME_TRANSCRIPTION_MODEL ?? "gpt-4o-transcribe"; + + dictationSttService = { + id: "openai", + createSession: ({ logger: sessionLogger, language, prompt }) => + new OpenAIRealtimeTranscriptionSession({ + apiKey: dictationApiKey, + logger: sessionLogger, + transcriptionModel, + ...(language ? { language } : {}), + ...(prompt ? { prompt } : {}), + turnDetection: null, + }), + }; + } } else if (needsOpenAiStt || needsOpenAiTts || needsOpenAiDictation || fallbackOpenAiStt || fallbackOpenAiTts) { logger.warn("OPENAI_API_KEY not set - OpenAI STT/TTS/dictation fallback is unavailable"); } @@ -710,9 +736,8 @@ export async function createPaseoDaemon( voiceLlmModel: config.voiceLlmModel ?? null, }, { - openaiApiKey: config.openai?.apiKey ?? null, finalTimeoutMs: config.dictationFinalTimeoutMs, - ...(dictationSessionFactory ? { sessionFactory: dictationSessionFactory } : {}), + stt: dictationSttService, } ); diff --git a/packages/server/src/server/config.ts b/packages/server/src/server/config.ts index fca13a681..84f4a10b4 100644 --- a/packages/server/src/server/config.ts +++ b/packages/server/src/server/config.ts @@ -1,8 +1,8 @@ import path from "node:path"; import type { PaseoDaemonConfig } from "./bootstrap.js"; -import type { STTConfig } from "./agent/stt-openai.js"; -import type { TTSConfig } from "./agent/tts-openai.js"; +import type { STTConfig } from "./speech/providers/openai/stt.js"; +import type { TTSConfig } from "./speech/providers/openai/tts.js"; import { loadPersistedConfig } from "./persisted-config.js"; import { mergeAllowedHosts, @@ -76,7 +76,7 @@ function parseOpenAIConfig( }; } -function parseSpeechProviderId(value: unknown): "openai" | "sherpa" | null { +function parseSpeechProviderId(value: unknown): "openai" | "local" | null { if (typeof value !== "string") { return null; } @@ -86,7 +86,7 @@ function parseSpeechProviderId(value: unknown): "openai" | "sherpa" | null { } if (normalized === "openai") return "openai"; if (normalized === "sherpa" || normalized === "sherpa-onnx" || normalized === "local") { - return "sherpa"; + return "local"; } return null; } @@ -185,22 +185,22 @@ export function loadConfig( const dictationSttProvider = parseSpeechProviderId(env.PASEO_DICTATION_STT_PROVIDER) ?? parseSpeechProviderId(persisted.features?.dictation?.stt?.provider) ?? - "sherpa"; + "local"; const voiceSttProvider = parseSpeechProviderId(env.PASEO_VOICE_STT_PROVIDER) ?? parseSpeechProviderId(persisted.features?.voiceMode?.stt?.provider) ?? - "sherpa"; + "local"; const voiceTtsProvider = parseSpeechProviderId(env.PASEO_VOICE_TTS_PROVIDER) ?? parseSpeechProviderId(persisted.features?.voiceMode?.tts?.provider) ?? - "sherpa"; + "local"; const shouldConfigureSherpa = - dictationSttProvider === "sherpa" || - voiceSttProvider === "sherpa" || - voiceTtsProvider === "sherpa" || + dictationSttProvider === "local" || + voiceSttProvider === "local" || + voiceTtsProvider === "local" || typeof env.PASEO_SHERPA_ONNX_MODELS_DIR === "string" || Boolean(persisted.providers?.sherpaOnnx); diff --git a/packages/server/src/server/daemon-client.e2e.test.ts b/packages/server/src/server/daemon-client.e2e.test.ts index 5a253f513..d1119f402 100644 --- a/packages/server/src/server/daemon-client.e2e.test.ts +++ b/packages/server/src/server/daemon-client.e2e.test.ts @@ -96,9 +96,9 @@ describe("daemon client E2E", () => { dictationFinalTimeoutMs: 5000, ...(openaiApiKey ? { openai: { apiKey: openaiApiKey } } : {}), speech: { - dictationSttProvider: "sherpa", - voiceSttProvider: "sherpa", - voiceTtsProvider: "sherpa", + dictationSttProvider: "local", + voiceSttProvider: "local", + voiceTtsProvider: "local", sherpaOnnx: { modelsDir: sherpaModelsDir, stt: { diff --git a/packages/server/src/server/dictation/dictation-stream-manager.test.ts b/packages/server/src/server/dictation/dictation-stream-manager.test.ts index 830739d74..2e8ae0251 100644 --- a/packages/server/src/server/dictation/dictation-stream-manager.test.ts +++ b/packages/server/src/server/dictation/dictation-stream-manager.test.ts @@ -4,23 +4,26 @@ import pino from "pino"; import { DictationStreamManager, - type RealtimeTranscriptionSession, - type RealtimeTranscriptionSessionFactory, } from "./dictation-stream-manager.js"; +import type { + SpeechToTextProvider, + StreamingTranscriptionSession, +} from "../speech/speech-provider.js"; -class FakeRealtimeSession extends EventEmitter implements RealtimeTranscriptionSession { +class FakeRealtimeSession extends EventEmitter implements StreamingTranscriptionSession { connected = false; - appended: string[] = []; + appended: Buffer[] = []; commitCalls = 0; clearCalls = 0; closed = false; + requiredSampleRate = 24000; async connect(): Promise { this.connected = true; } - appendPcm16Base64(base64Audio: string): void { - this.appended.push(base64Audio); + appendPcm16(pcm16le: Buffer): void { + this.appended.push(pcm16le); } commit(): void { @@ -35,12 +38,12 @@ class FakeRealtimeSession extends EventEmitter implements RealtimeTranscriptionS this.closed = true; } - emitCommitted(itemId: string): void { - this.emit("committed", { itemId, previousItemId: null }); + emitCommitted(segmentId: string): void { + this.emit("committed", { segmentId, previousSegmentId: null }); } - emitTranscript(itemId: string, transcript: string, isFinal: boolean): void { - this.emit("transcript", { itemId, transcript, isFinal }); + emitTranscript(segmentId: string, transcript: string, isFinal: boolean): void { + this.emit("transcript", { segmentId, transcript, isFinal }); } emitError(message: string): void { @@ -48,6 +51,14 @@ class FakeRealtimeSession extends EventEmitter implements RealtimeTranscriptionS } } +class FakeSttProvider implements SpeechToTextProvider { + public readonly id = "fake"; + constructor(private readonly session: FakeRealtimeSession) {} + createSession(_params: { logger: any; language?: string; prompt?: string }): StreamingTranscriptionSession { + return this.session; + } +} + const buildPcmBase64 = (sampleValue: number, sampleCount: number): string => { const samples = new Int16Array(sampleCount); samples.fill(sampleValue); @@ -59,34 +70,29 @@ const tick = async (): Promise => { await Promise.resolve(); }; -describe("DictationStreamManager (semantic VAD grace fallback)", () => { +describe("DictationStreamManager (finish buffer-too-small tolerance)", () => { const env = { - turnDetection: process.env.OPENAI_REALTIME_DICTATION_TURN_DETECTION, dictationDebug: process.env.PASEO_DICTATION_DEBUG, }; beforeEach(() => { vi.useFakeTimers(); - process.env.OPENAI_REALTIME_DICTATION_TURN_DETECTION = "semantic_vad"; process.env.PASEO_DICTATION_DEBUG = "false"; }); afterEach(() => { vi.useRealTimers(); - process.env.OPENAI_REALTIME_DICTATION_TURN_DETECTION = env.turnDetection; process.env.PASEO_DICTATION_DEBUG = env.dictationDebug; }); it("treats buffer-too-small as benign and finalizes with existing transcripts", async () => { const session = new FakeRealtimeSession(); - const factory: RealtimeTranscriptionSessionFactory = () => session; const emitted: Array<{ type: string; payload: any }> = []; const manager = new DictationStreamManager({ logger: pino({ level: "silent" }), emit: (msg) => emitted.push(msg), sessionId: "s1", - openaiApiKey: "k", - sessionFactory: factory, + stt: new FakeSttProvider(session), finalTimeoutMs: 5000, }); @@ -98,14 +104,11 @@ describe("DictationStreamManager (semantic VAD grace fallback)", () => { format: "audio/pcm;rate=24000;bits=16", }); - session.emitTranscript("i1", "hello world", true); + session.emitTranscript("seg-1", "hello world", true); await manager.handleFinish("d1", 0); await tick(); - vi.advanceTimersByTime(2000); - await tick(); - session.emitError( "Error committing input audio buffer: buffer too small. Expected at least 100ms of audio, but buffer only has 0.00ms of audio." ); @@ -117,56 +120,21 @@ describe("DictationStreamManager (semantic VAD grace fallback)", () => { expect(final?.payload.text).toBe("hello world"); expect(session.closed).toBe(true); }); - - it("does not fallback-commit if committed event arrives during grace window", async () => { - const session = new FakeRealtimeSession(); - const factory: RealtimeTranscriptionSessionFactory = () => session; - const emitted: Array<{ type: string; payload: any }> = []; - const manager = new DictationStreamManager({ - logger: pino({ level: "silent" }), - emit: (msg) => emitted.push(msg), - sessionId: "s1", - openaiApiKey: "k", - sessionFactory: factory, - finalTimeoutMs: 5000, - }); - - await manager.handleStart("d1", "audio/pcm;rate=24000;bits=16"); - await manager.handleChunk({ - dictationId: "d1", - seq: 0, - audioBase64: buildPcmBase64(2000, 2400), - format: "audio/pcm;rate=24000;bits=16", - }); - - await manager.handleFinish("d1", 0); - session.emitCommitted("i1"); - session.emitTranscript("i1", "hi there", true); - - vi.advanceTimersByTime(2000); - await tick(); - - expect(session.commitCalls).toBe(0); - const final = emitted.find((msg) => msg.type === "dictation_stream_final"); - expect(final?.payload.text).toBe("hi there"); - }); }); -describe("DictationStreamManager (provider-agnostic session factory)", () => { - it("can start with a custom session factory even without OPENAI_API_KEY", async () => { +describe("DictationStreamManager (provider-agnostic provider)", () => { + it("does not require OPENAI_API_KEY", async () => { const original = process.env.OPENAI_API_KEY; delete process.env.OPENAI_API_KEY; try { const session = new FakeRealtimeSession(); - const factory: RealtimeTranscriptionSessionFactory = () => session; const emitted: Array<{ type: string; payload: any }> = []; const manager = new DictationStreamManager({ logger: pino({ level: "silent" }), emit: (msg) => emitted.push(msg), sessionId: "s1", - openaiApiKey: null, - sessionFactory: factory, + stt: new FakeSttProvider(session), }); await manager.handleStart("d-local", "audio/pcm;rate=16000;bits=16"); diff --git a/packages/server/src/server/dictation/dictation-stream-manager.ts b/packages/server/src/server/dictation/dictation-stream-manager.ts index ebad8281b..0a6f21b38 100644 --- a/packages/server/src/server/dictation/dictation-stream-manager.ts +++ b/packages/server/src/server/dictation/dictation-stream-manager.ts @@ -7,112 +7,19 @@ import { } from "../agent/dictation-debug.js"; import { isPaseoDictationDebugEnabled } from "../agent/recordings-debug.js"; import { Pcm16MonoResampler } from "../agent/pcm16-resampler.js"; -import { OpenAIRealtimeTranscriptionSession } from "../agent/openai-realtime-transcription.js"; +import type { + SpeechToTextProvider, + StreamingTranscriptionSession, +} from "../speech/speech-provider.js"; +import { parsePcmRateFromFormat, pcm16lePeakAbs } from "../speech/audio.js"; const PCM_CHANNELS = 1; const PCM_BITS_PER_SAMPLE = 16; -const OPENAI_DICTATION_PCM_OUTPUT_RATE = 24000; const DEFAULT_DICTATION_FINAL_TIMEOUT_MS = 10000; -const DICTATION_VAD_GRACE_TIMEOUT_MS = Number.parseInt( - process.env.OPENAI_REALTIME_DICTATION_VAD_GRACE_TIMEOUT_MS ?? "2000", - 10 -); const DICTATION_SILENCE_PEAK_THRESHOLD = Number.parseInt( - process.env.OPENAI_REALTIME_DICTATION_SILENCE_PEAK_THRESHOLD ?? "300", + process.env.PASEO_DICTATION_SILENCE_PEAK_THRESHOLD ?? "300", 10 ); -const DICTATION_TURN_DETECTION = ( - process.env.OPENAI_REALTIME_DICTATION_TURN_DETECTION ?? "semantic_vad" -).trim(); -const DICTATION_SEMANTIC_VAD_EAGERNESS = ( - process.env.OPENAI_REALTIME_DICTATION_SEMANTIC_VAD_EAGERNESS ?? "medium" -).trim(); -const DICTATION_FLUSH_SILENCE_MS = Number.parseInt( - process.env.OPENAI_REALTIME_DICTATION_FLUSH_SILENCE_MS ?? "800", - 10 -); - -type OpenAITurnDetection = - | null - | { - type: "server_vad"; - create_response: false; - threshold?: number; - prefix_padding_ms?: number; - silence_duration_ms?: number; - } - | { type: "semantic_vad"; create_response: false; eagerness?: "low" | "medium" | "high" }; - -function pcm16lePeakAbs(pcm16le: Buffer): number { - if (pcm16le.length === 0) { - return 0; - } - if (pcm16le.length % 2 !== 0) { - throw new Error(`PCM16 chunk byteLength must be even, got ${pcm16le.length}`); - } - const samples = new Int16Array( - pcm16le.buffer, - pcm16le.byteOffset, - pcm16le.byteLength / 2 - ); - let peak = 0; - for (let i = 0; i < samples.length; i += 1) { - const v = samples[i]!; - const abs = v < 0 ? -v : v; - if (abs > peak) { - peak = abs; - if (peak >= 32767) { - break; - } - } - } - return peak; -} - -function parseDictationTurnDetection(): OpenAITurnDetection { - if ( - !DICTATION_TURN_DETECTION || - DICTATION_TURN_DETECTION === "none" || - DICTATION_TURN_DETECTION === "null" - ) { - return null; - } - if (DICTATION_TURN_DETECTION === "server_vad") { - return { type: "server_vad", create_response: false }; - } - const eagerness = - DICTATION_SEMANTIC_VAD_EAGERNESS === "low" || - DICTATION_SEMANTIC_VAD_EAGERNESS === "high" - ? (DICTATION_SEMANTIC_VAD_EAGERNESS as "low" | "high") - : ("medium" as const); - return { type: "semantic_vad", create_response: false, eagerness }; -} - -export type RealtimeTranscriptionSession = { - connect(): Promise; - appendPcm16Base64(base64Audio: string): void; - commit(): void; - clear(): void; - close(): void; - on( - event: "committed", - handler: (payload: { itemId: string; previousItemId: string | null }) => void - ): unknown; - on( - event: "transcript", - handler: (payload: { itemId: string; transcript: string; isFinal: boolean }) => void - ): unknown; - on(event: "error", handler: (err: unknown) => void): unknown; -}; - -export type RealtimeTranscriptionSessionFactory = (params: { - apiKey?: string | null; - logger: pino.Logger; - transcriptionModel: string; - language?: string; - prompt?: string; - turnDetection: OpenAITurnDetection; -}) => RealtimeTranscriptionSession; function convertPCMToWavBuffer( pcmBuffer: Buffer, @@ -147,7 +54,7 @@ type DictationStreamState = { dictationId: string; sessionId: string; inputFormat: string; - openai: RealtimeTranscriptionSession; + stt: StreamingTranscriptionSession; inputRate: number; outputRate: number; resampler: Pcm16MonoResampler | null; @@ -159,17 +66,14 @@ type DictationStreamState = { ackSeq: number; bytesSinceCommit: number; peakSinceCommit: number; - committedItemIds: string[]; - transcriptsByItemId: Map; - finalTranscriptItemIds: Set; + committedSegmentIds: string[]; + transcriptsBySegmentId: Map; + finalTranscriptSegmentIds: Set; awaitingFinalCommit: boolean; - vadGraceTimeout: ReturnType | null; - fallbackCommitAttempted: boolean; finishRequested: boolean; finishSealed: boolean; finalSeq: number | null; finalTimeout: ReturnType | null; - isSemanticVad: boolean; }; export type DictationStreamOutboundMessage = @@ -192,37 +96,22 @@ export class DictationStreamManager { private readonly logger: pino.Logger; private readonly emit: (msg: DictationStreamOutboundMessage) => void; private readonly sessionId: string; - private readonly openaiApiKey: string | null; + private readonly stt: SpeechToTextProvider | null; private readonly finalTimeoutMs: number; - private readonly createSession: RealtimeTranscriptionSessionFactory; - private readonly requiresApiKey: boolean; private readonly streams = new Map(); constructor(params: { logger: pino.Logger; emit: (msg: DictationStreamOutboundMessage) => void; sessionId: string; - openaiApiKey?: string | null; + stt: SpeechToTextProvider | null; finalTimeoutMs?: number; - sessionFactory?: RealtimeTranscriptionSessionFactory; }) { this.logger = params.logger.child({ component: "dictation-stream-manager" }); this.emit = params.emit; this.sessionId = params.sessionId; - this.openaiApiKey = params.openaiApiKey ?? null; + this.stt = params.stt; this.finalTimeoutMs = params.finalTimeoutMs ?? DEFAULT_DICTATION_FINAL_TIMEOUT_MS; - this.requiresApiKey = !params.sessionFactory; - this.createSession = - params.sessionFactory ?? - ((factoryParams) => { - if (!factoryParams.apiKey) { - throw new Error("OPENAI_API_KEY not set"); - } - return new OpenAIRealtimeTranscriptionSession({ - ...factoryParams, - apiKey: factoryParams.apiKey, - }); - }); } public cleanupAll(): void { @@ -234,39 +123,30 @@ export class DictationStreamManager { public async handleStart(dictationId: string, format: string): Promise { this.cleanupDictationStream(dictationId); - const apiKey = this.openaiApiKey ?? process.env.OPENAI_API_KEY; - if (this.requiresApiKey && !apiKey) { - this.failDictationStream(dictationId, "OPENAI_API_KEY not set", false); + if (!this.stt) { + this.failDictationStream(dictationId, "Dictation STT not configured", false); return; } - const transcriptionModel = - process.env.OPENAI_REALTIME_TRANSCRIPTION_MODEL ?? "gpt-4o-transcribe"; const transcriptionPrompt = - process.env.OPENAI_REALTIME_DICTATION_TRANSCRIPTION_PROMPT ?? + process.env.PASEO_DICTATION_TRANSCRIPTION_PROMPT ?? "Transcribe only what the speaker says. Do not add words. Preserve punctuation and casing. If the audio is silence or non-speech noise, return an empty transcript."; - const turnDetection = parseDictationTurnDetection(); - const openai = this.createSession({ - apiKey: apiKey ?? null, + const stt = this.stt.createSession({ logger: this.logger.child({ dictationId }), - transcriptionModel, language: "en", prompt: transcriptionPrompt, - turnDetection, }); - openai.on("committed", ({ itemId }: { itemId: string }) => { + stt.on("committed", ({ segmentId }) => { const state = this.streams.get(dictationId); if (!state) { return; } - this.clearVadGraceTimeout(state); - state.committedItemIds.push(itemId); + state.committedSegmentIds.push(segmentId); state.bytesSinceCommit = 0; state.peakSinceCommit = 0; - // When finishing, we require at least one commit after finish if we flushed pending audio. if (state.finishRequested && state.awaitingFinalCommit) { state.awaitingFinalCommit = false; } @@ -274,52 +154,37 @@ export class DictationStreamManager { this.maybeFinalizeDictationStream(dictationId); }); - openai.on( - "transcript", - ({ - itemId, - transcript, - isFinal, - }: { - itemId: string; - transcript: string; - isFinal: boolean; - }) => { - const state = this.streams.get(dictationId); - if (!state) { - return; - } - state.transcriptsByItemId.set(itemId, transcript); - if (isFinal) { - state.finalTranscriptItemIds.add(itemId); - } - - // If we triggered a finish commit but OpenAI doesn't emit committed events (or they arrive late), - // allow final transcripts to unblock finalization. - if (state.finishRequested && state.awaitingFinalCommit && isFinal) { - this.clearVadGraceTimeout(state); - state.awaitingFinalCommit = false; - } - - const orderedIds = state.committedItemIds.includes(itemId) - ? state.committedItemIds - : [...state.committedItemIds, itemId]; - const partialText = orderedIds - .map((id) => state.transcriptsByItemId.get(id) ?? "") - .join(" ") - .trim(); - this.emitDictationPartial(dictationId, partialText); - - this.maybeSealDictationStreamFinish(dictationId); - this.maybeFinalizeDictationStream(dictationId); + stt.on("transcript", ({ segmentId, transcript, isFinal }) => { + const state = this.streams.get(dictationId); + if (!state) { + return; + } + state.transcriptsBySegmentId.set(segmentId, transcript); + if (isFinal) { + state.finalTranscriptSegmentIds.add(segmentId); } - ); - openai.on("error", (err) => { + if (state.finishRequested && state.awaitingFinalCommit && isFinal) { + state.awaitingFinalCommit = false; + } + + const orderedIds = state.committedSegmentIds.includes(segmentId) + ? state.committedSegmentIds + : [...state.committedSegmentIds, segmentId]; + const partialText = orderedIds + .map((id) => state.transcriptsBySegmentId.get(id) ?? "") + .join(" ") + .trim(); + this.emitDictationPartial(dictationId, partialText); + + this.maybeSealDictationStreamFinish(dictationId); + this.maybeFinalizeDictationStream(dictationId); + }); + + stt.on("error", (err) => { const message = err instanceof Error ? err.message : String(err); const state = this.streams.get(dictationId); if (state && state.finishRequested && isBufferTooSmallError(message)) { - this.clearVadGraceTimeout(state); if (state.awaitingFinalCommit) { state.awaitingFinalCommit = false; } @@ -329,14 +194,13 @@ export class DictationStreamManager { void this.failAndCleanupDictationStream(dictationId, message, true); }); - await openai.connect(); + await stt.connect(); - const rateMatch = /(?:^|[;,\s])rate\s*=\s*(\d+)(?:$|[;,\s])/i.exec(format); - const inputRate = rateMatch ? Number.parseInt(rateMatch[1]!, 10) : 16000; + const inputRate = parsePcmRateFromFormat(format, 16000) ?? 16000; if (!Number.isFinite(inputRate) || inputRate <= 0) { this.failDictationStream(dictationId, `Invalid dictation input rate in format: ${format}`, false); try { - openai.close(); + stt.close(); } catch { // no-op } @@ -348,13 +212,13 @@ export class DictationStreamManager { this.logger ); - const outputRate = this.requiresApiKey ? OPENAI_DICTATION_PCM_OUTPUT_RATE : inputRate; + const outputRate = stt.requiredSampleRate; this.streams.set(dictationId, { dictationId, sessionId: this.sessionId, inputFormat: format, - openai, + stt, inputRate, outputRate, resampler: @@ -372,17 +236,14 @@ export class DictationStreamManager { ackSeq: -1, bytesSinceCommit: 0, peakSinceCommit: 0, - committedItemIds: [], - transcriptsByItemId: new Map(), - finalTranscriptItemIds: new Set(), + committedSegmentIds: [], + transcriptsBySegmentId: new Map(), + finalTranscriptSegmentIds: new Set(), awaitingFinalCommit: false, - vadGraceTimeout: null, - fallbackCommitAttempted: false, finishRequested: false, finishSealed: false, finalSeq: null, finalTimeout: null, - isSemanticVad: turnDetection?.type === "semantic_vad", }); this.emitDictationAck(dictationId, -1); @@ -425,7 +286,7 @@ export class DictationStreamManager { const resampled = state.resampler ? state.resampler.processChunk(pcm16) : pcm16; if (resampled.length > 0) { - state.openai.appendPcm16Base64(resampled.toString("base64")); + state.stt.appendPcm16(resampled); state.debugAudioChunks.push(resampled); state.bytesSinceCommit += resampled.length; state.peakSinceCommit = Math.max(state.peakSinceCommit, pcm16lePeakAbs(resampled)); @@ -579,12 +440,11 @@ export class DictationStreamManager { if (!state) { return; } - this.clearVadGraceTimeout(state); if (state.finalTimeout) { clearTimeout(state.finalTimeout); } try { - state.openai.close(); + state.stt.close(); } catch { // no-op } @@ -616,37 +476,18 @@ export class DictationStreamManager { }, "Dictation finish: clearing silence-only tail (skip final commit)" ); - state.openai.clear(); + state.stt.clear(); state.bytesSinceCommit = 0; state.peakSinceCommit = 0; state.awaitingFinalCommit = false; } else { - const silenceBytes = Math.max( - 0, - Math.round((state.outputRate * 2 * DICTATION_FLUSH_SILENCE_MS) / 1000) - ); - if (silenceBytes > 0) { - this.logger.debug( - { dictationId, silenceMs: DICTATION_FLUSH_SILENCE_MS, silenceBytes }, - "Dictation finish: appending silence tail for semantic VAD flush" - ); - const silence = Buffer.alloc(silenceBytes); - state.openai.appendPcm16Base64(silence.toString("base64")); - state.debugAudioChunks.push(silence); - state.bytesSinceCommit += silenceBytes; - } - state.awaitingFinalCommit = true; - if (state.isSemanticVad) { - this.startVadGraceTimeout(state); - } else { - try { - state.openai.commit(); - } catch (error) { - const message = error instanceof Error ? error.message : String(error); - void this.failAndCleanupDictationStream(dictationId, message, true); - return; - } + try { + state.stt.commit(); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + void this.failAndCleanupDictationStream(dictationId, message, true); + return; } } } else { @@ -672,15 +513,15 @@ export class DictationStreamManager { return; } - const committedSet = new Set(state.committedItemIds); - const orderedItemIds: string[] = [...state.committedItemIds]; - for (const itemId of state.transcriptsByItemId.keys()) { - if (!committedSet.has(itemId)) { - orderedItemIds.push(itemId); + const committedSet = new Set(state.committedSegmentIds); + const orderedSegmentIds: string[] = [...state.committedSegmentIds]; + for (const segmentId of state.transcriptsBySegmentId.keys()) { + if (!committedSet.has(segmentId)) { + orderedSegmentIds.push(segmentId); } } - if (orderedItemIds.length === 0) { + if (orderedSegmentIds.length === 0) { void (async () => { const debugRecordingPath = await this.maybePersistDictationStreamAudio(dictationId); this.emit({ @@ -708,15 +549,15 @@ export class DictationStreamManager { return; } - const allTranscriptsReady = orderedItemIds.every((itemId) => - state.finalTranscriptItemIds.has(itemId) + const allTranscriptsReady = orderedSegmentIds.every((segmentId) => + state.finalTranscriptSegmentIds.has(segmentId) ); if (!allTranscriptsReady) { return; } - const orderedText = orderedItemIds - .map((itemId) => state.transcriptsByItemId.get(itemId) ?? "") + const orderedText = orderedSegmentIds + .map((segmentId) => state.transcriptsBySegmentId.get(segmentId) ?? "") .join(" ") .trim(); @@ -745,41 +586,6 @@ export class DictationStreamManager { this.cleanupDictationStream(dictationId); })(); } - - private startVadGraceTimeout(state: DictationStreamState): void { - if (state.vadGraceTimeout || DICTATION_VAD_GRACE_TIMEOUT_MS <= 0) { - return; - } - state.vadGraceTimeout = setTimeout(() => { - state.vadGraceTimeout = null; - if (!state.finishRequested || !state.awaitingFinalCommit) { - return; - } - if (state.bytesSinceCommit <= 0 || state.fallbackCommitAttempted) { - return; - } - state.fallbackCommitAttempted = true; - try { - state.openai.commit(); - } catch (error) { - const message = error instanceof Error ? error.message : String(error); - if (isBufferTooSmallError(message)) { - state.awaitingFinalCommit = false; - this.maybeFinalizeDictationStream(state.dictationId); - return; - } - void this.failAndCleanupDictationStream(state.dictationId, message, true); - } - }, DICTATION_VAD_GRACE_TIMEOUT_MS); - } - - private clearVadGraceTimeout(state: DictationStreamState): void { - if (!state.vadGraceTimeout) { - return; - } - clearTimeout(state.vadGraceTimeout); - state.vadGraceTimeout = null; - } } function isBufferTooSmallError(message: string): boolean { diff --git a/packages/server/src/server/persisted-config.ts b/packages/server/src/server/persisted-config.ts index 0c3ad2792..59469bc3a 100644 --- a/packages/server/src/server/persisted-config.ts +++ b/packages/server/src/server/persisted-config.ts @@ -46,11 +46,25 @@ const ProvidersSchema = z }) .strict(); +const SpeechProviderIdSchema = z.preprocess( + (value) => { + if (typeof value !== "string") { + return value; + } + const normalized = value.trim().toLowerCase(); + if (normalized === "sherpa" || normalized === "sherpa-onnx") { + return "local"; + } + return normalized; + }, + z.enum(["openai", "local"]) +); + const FeatureDictationSchema = z .object({ stt: z .object({ - provider: z.enum(["openai", "sherpa"]).optional(), + provider: SpeechProviderIdSchema.optional(), model: z.string().min(1).optional(), preset: z.string().min(1).optional(), confidenceThreshold: z.number().optional(), @@ -71,7 +85,7 @@ const FeatureVoiceModeSchema = z .optional(), stt: z .object({ - provider: z.enum(["openai", "sherpa"]).optional(), + provider: SpeechProviderIdSchema.optional(), model: z.string().min(1).optional(), preset: z.string().min(1).optional(), }) @@ -79,7 +93,7 @@ const FeatureVoiceModeSchema = z .optional(), tts: z .object({ - provider: z.enum(["openai", "sherpa"]).optional(), + provider: SpeechProviderIdSchema.optional(), model: z.enum(["tts-1", "tts-1-hd"]).optional(), voice: z.enum(["alloy", "echo", "fable", "onyx", "nova", "shimmer"]).optional(), preset: z.string().min(1).optional(), diff --git a/packages/server/src/server/session.ts b/packages/server/src/server/session.ts index d6a6885de..ae5417017 100644 --- a/packages/server/src/server/session.ts +++ b/packages/server/src/server/session.ts @@ -38,7 +38,6 @@ import { maybePersistTtsDebugAudio } from "./agent/tts-debug.js"; import { isPaseoDictationDebugEnabled } from "./agent/recordings-debug.js"; import { DictationStreamManager, - type RealtimeTranscriptionSessionFactory, } from "./dictation/dictation-stream-manager.js"; import type { VoiceConversationStore } from "./voice-conversation-store.js"; import { @@ -106,6 +105,14 @@ import { } from "../utils/checkout-git.js"; import { getProjectIcon } from "../utils/project-icon.js"; import { expandTilde } from "../utils/path.js"; +import { + ensureSherpaOnnxModels, + getSherpaOnnxModelDir, +} from "./speech/providers/local/sherpa/model-downloader.js"; +import { + listSherpaOnnxModels, + type SherpaOnnxModelId, +} from "./speech/providers/local/sherpa/model-catalog.js"; import type pino from "pino"; const execAsync = promisify(exec); @@ -334,9 +341,8 @@ export class Session { voiceLlmModel?: string | null; }, dictation?: { - openaiApiKey?: string | null; finalTimeoutMs?: number; - sessionFactory?: RealtimeTranscriptionSessionFactory; + stt?: SpeechToTextProvider | null; } ) { this.clientId = clientId; @@ -367,9 +373,8 @@ export class Session { logger: this.sessionLogger, sessionId: this.sessionId, emit: (msg) => this.emit(msg as unknown as SessionOutboundMessage), - openaiApiKey: dictation?.openaiApiKey ?? null, + stt: dictation?.stt ?? null, finalTimeoutMs: dictation?.finalTimeoutMs, - ...(dictation?.sessionFactory ? { sessionFactory: dictation.sessionFactory } : {}), }); // Initialize agent MCP client asynchronously @@ -975,6 +980,14 @@ export class Session { await this.handleListProviderModelsRequest(msg); break; + case "speech_models_list_request": + await this.handleSpeechModelsListRequest(msg); + break; + + case "speech_models_download_request": + await this.handleSpeechModelsDownloadRequest(msg); + break; + case "clear_agent_attention": await this.handleClearAgentAttention(msg.agentId); break; @@ -1910,6 +1923,114 @@ export class Session { } } + private async handleSpeechModelsListRequest( + msg: Extract + ): Promise { + const modelsDir = + process.env.PASEO_SHERPA_ONNX_MODELS_DIR?.trim() || + join(this.paseoHome, "models", "sherpa-onnx"); + + const models = await Promise.all( + listSherpaOnnxModels().map(async (model) => { + const modelDir = getSherpaOnnxModelDir(modelsDir, model.id); + const missingFiles: string[] = []; + for (const rel of model.requiredFiles) { + const filePath = join(modelDir, rel); + try { + const fileStat = await stat(filePath); + if (fileStat.isDirectory()) { + continue; + } + if (!fileStat.isFile() || fileStat.size <= 0) { + missingFiles.push(rel); + } + } catch { + missingFiles.push(rel); + } + } + + return { + id: model.id, + kind: model.kind, + description: model.description, + modelDir, + isDownloaded: missingFiles.length === 0, + ...(missingFiles.length > 0 ? { missingFiles } : {}), + }; + }) + ); + + this.emit({ + type: "speech_models_list_response", + payload: { + modelsDir, + models, + requestId: msg.requestId, + }, + }); + } + + private async handleSpeechModelsDownloadRequest( + msg: Extract + ): Promise { + const modelsDir = + process.env.PASEO_SHERPA_ONNX_MODELS_DIR?.trim() || + join(this.paseoHome, "models", "sherpa-onnx"); + + const modelIdsRaw = + msg.modelIds && msg.modelIds.length > 0 + ? msg.modelIds + : [ + process.env.PASEO_SHERPA_STT_PRESET ?? "zipformer-bilingual-zh-en-2023-02-20", + process.env.PASEO_SHERPA_TTS_PRESET ?? "pocket-tts-onnx-int8", + ]; + + const allModelIds = new Set(listSherpaOnnxModels().map((m) => m.id)); + const invalid = modelIdsRaw.filter((id) => !allModelIds.has(id as SherpaOnnxModelId)); + if (invalid.length > 0) { + this.emit({ + type: "speech_models_download_response", + payload: { + modelsDir, + downloadedModelIds: [], + error: `Unknown speech model id(s): ${invalid.join(", ")}`, + requestId: msg.requestId, + }, + }); + return; + } + + const modelIds = modelIdsRaw as SherpaOnnxModelId[]; + try { + await ensureSherpaOnnxModels({ + modelsDir, + modelIds, + autoDownload: true, + logger: this.sessionLogger, + }); + this.emit({ + type: "speech_models_download_response", + payload: { + modelsDir, + downloadedModelIds: modelIds, + error: null, + requestId: msg.requestId, + }, + }); + } catch (error) { + this.sessionLogger.error({ err: error, modelIds }, "Failed to download speech models"); + this.emit({ + type: "speech_models_download_response", + payload: { + modelsDir, + downloadedModelIds: [], + error: error instanceof Error ? error.message : String(error), + requestId: msg.requestId, + }, + }); + } + } + private normalizeGitOptions( gitOptions?: GitSetupOptions, legacyWorktreeName?: string diff --git a/packages/server/src/server/speech/pocket/pocket-tts-onnx.ts b/packages/server/src/server/speech/providers/local/pocket/pocket-tts-onnx.ts similarity index 99% rename from packages/server/src/server/speech/pocket/pocket-tts-onnx.ts rename to packages/server/src/server/speech/providers/local/pocket/pocket-tts-onnx.ts index 58fd16e96..847edf13a 100644 --- a/packages/server/src/server/speech/pocket/pocket-tts-onnx.ts +++ b/packages/server/src/server/speech/providers/local/pocket/pocket-tts-onnx.ts @@ -3,9 +3,9 @@ import { readFile } from "node:fs/promises"; import { Readable } from "node:stream"; import type pino from "pino"; -import type { SpeechStreamResult, TextToSpeechProvider } from "../speech-provider.js"; -import { chunkBuffer, float32ToPcm16le, parsePcm16MonoWav, pcm16leToFloat32 } from "../audio.js"; -import { Pcm16MonoResampler } from "../../agent/pcm16-resampler.js"; +import type { SpeechStreamResult, TextToSpeechProvider } from "../../../speech-provider.js"; +import { chunkBuffer, float32ToPcm16le, parsePcm16MonoWav, pcm16leToFloat32 } from "../../../audio.js"; +import { Pcm16MonoResampler } from "../../../../agent/pcm16-resampler.js"; type OrtModule = typeof import("onnxruntime-node"); type OrtSession = import("onnxruntime-node").InferenceSession; diff --git a/packages/server/src/server/speech/sherpa/model-catalog.ts b/packages/server/src/server/speech/providers/local/sherpa/model-catalog.ts similarity index 100% rename from packages/server/src/server/speech/sherpa/model-catalog.ts rename to packages/server/src/server/speech/providers/local/sherpa/model-catalog.ts diff --git a/packages/server/src/server/speech/sherpa/model-downloader.test.ts b/packages/server/src/server/speech/providers/local/sherpa/model-downloader.test.ts similarity index 100% rename from packages/server/src/server/speech/sherpa/model-downloader.test.ts rename to packages/server/src/server/speech/providers/local/sherpa/model-downloader.test.ts diff --git a/packages/server/src/server/speech/sherpa/model-downloader.ts b/packages/server/src/server/speech/providers/local/sherpa/model-downloader.ts similarity index 100% rename from packages/server/src/server/speech/sherpa/model-downloader.ts rename to packages/server/src/server/speech/providers/local/sherpa/model-downloader.ts diff --git a/packages/server/src/server/speech/sherpa/sherpa-offline-recognizer.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-offline-recognizer.ts similarity index 100% rename from packages/server/src/server/speech/sherpa/sherpa-offline-recognizer.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-offline-recognizer.ts diff --git a/packages/server/src/server/speech/sherpa/sherpa-online-recognizer.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-online-recognizer.ts similarity index 100% rename from packages/server/src/server/speech/sherpa/sherpa-online-recognizer.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-online-recognizer.ts diff --git a/packages/server/src/server/speech/sherpa/sherpa-onnx-loader.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-onnx-loader.ts similarity index 100% rename from packages/server/src/server/speech/sherpa/sherpa-onnx-loader.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-onnx-loader.ts diff --git a/packages/server/src/server/speech/sherpa/sherpa-onnx-node-loader.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-onnx-node-loader.ts similarity index 100% rename from packages/server/src/server/speech/sherpa/sherpa-onnx-node-loader.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-onnx-node-loader.ts diff --git a/packages/server/src/server/speech/sherpa/sherpa-parakeet-realtime-session.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-parakeet-realtime-session.ts similarity index 73% rename from packages/server/src/server/speech/sherpa/sherpa-parakeet-realtime-session.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-parakeet-realtime-session.ts index 6a912bd69..52651543c 100644 --- a/packages/server/src/server/speech/sherpa/sherpa-parakeet-realtime-session.ts +++ b/packages/server/src/server/speech/providers/local/sherpa/sherpa-parakeet-realtime-session.ts @@ -1,22 +1,23 @@ import { EventEmitter } from "node:events"; import { v4 as uuidv4 } from "uuid"; -import type { RealtimeTranscriptionSession } from "../../dictation/dictation-stream-manager.js"; -import { pcm16lePeakAbs, pcm16leToFloat32 } from "../audio.js"; +import type { StreamingTranscriptionSession } from "../../../speech-provider.js"; +import { pcm16lePeakAbs, pcm16leToFloat32 } from "../../../audio.js"; import { SherpaOfflineRecognizerEngine } from "./sherpa-offline-recognizer.js"; export class SherpaParakeetRealtimeTranscriptionSession extends EventEmitter - implements RealtimeTranscriptionSession + implements StreamingTranscriptionSession { private readonly engine: SherpaOfflineRecognizerEngine; private connected = false; - private currentItemId: string | null = null; - private previousItemId: string | null = null; + public readonly requiredSampleRate: number; + private currentSegmentId: string | null = null; + private previousSegmentId: string | null = null; private lastPartialText = ""; - private pcm16 = Buffer.alloc(0); + private pcm16: Buffer = Buffer.alloc(0); private lastDecodeAt = 0; private decoding = false; private pendingDecode = false; @@ -25,6 +26,7 @@ export class SherpaParakeetRealtimeTranscriptionSession constructor(params: { engine: SherpaOfflineRecognizerEngine; minDecodeIntervalMs?: number }) { super(); this.engine = params.engine; + this.requiredSampleRate = this.engine.sampleRate; this.minDecodeIntervalMs = params.minDecodeIntervalMs ?? 350; } @@ -32,18 +34,17 @@ export class SherpaParakeetRealtimeTranscriptionSession if (this.connected) { return; } - this.currentItemId = uuidv4(); + this.currentSegmentId = uuidv4(); this.connected = true; } - appendPcm16Base64(base64Audio: string): void { - if (!this.connected || !this.currentItemId) { + appendPcm16(chunk: Buffer): void { + if (!this.connected || !this.currentSegmentId) { this.emit("error", new Error("Parakeet realtime session not connected")); return; } try { - const chunk = Buffer.from(base64Audio, "base64"); this.pcm16 = this.pcm16.length === 0 ? chunk : Buffer.concat([this.pcm16, chunk]); void this.maybeDecode(false); } catch (err) { @@ -52,7 +53,7 @@ export class SherpaParakeetRealtimeTranscriptionSession } commit(): void { - if (!this.connected || !this.currentItemId) { + if (!this.connected || !this.currentSegmentId) { this.emit("error", new Error("Parakeet realtime session not connected")); return; } @@ -61,14 +62,14 @@ export class SherpaParakeetRealtimeTranscriptionSession try { await this.maybeDecode(true); const finalText = this.lastPartialText; - const itemId = this.currentItemId!; - const previousItemId = this.previousItemId; + const segmentId = this.currentSegmentId!; + const previousSegmentId = this.previousSegmentId; - this.emit("committed", { itemId, previousItemId }); - this.emit("transcript", { itemId, transcript: finalText, isFinal: true }); + this.emit("committed", { segmentId, previousSegmentId }); + this.emit("transcript", { segmentId, transcript: finalText, isFinal: true }); - this.previousItemId = itemId; - this.currentItemId = uuidv4(); + this.previousSegmentId = segmentId; + this.currentSegmentId = uuidv4(); this.lastPartialText = ""; this.pcm16 = Buffer.alloc(0); } catch (err) { @@ -82,18 +83,18 @@ export class SherpaParakeetRealtimeTranscriptionSession return; } this.pcm16 = Buffer.alloc(0); - this.currentItemId = uuidv4(); + this.currentSegmentId = uuidv4(); this.lastPartialText = ""; } close(): void { this.connected = false; - this.currentItemId = null; + this.currentSegmentId = null; this.pcm16 = Buffer.alloc(0); } private async maybeDecode(force: boolean): Promise { - if (!this.connected || !this.currentItemId) { + if (!this.connected || !this.currentSegmentId) { return; } @@ -113,7 +114,7 @@ export class SherpaParakeetRealtimeTranscriptionSession this.lastDecodeAt = Date.now(); if (text !== this.lastPartialText) { this.lastPartialText = text; - this.emit("transcript", { itemId: this.currentItemId, transcript: text, isFinal: false }); + this.emit("transcript", { segmentId: this.currentSegmentId, transcript: text, isFinal: false }); } } finally { this.decoding = false; diff --git a/packages/server/src/server/speech/sherpa/sherpa-parakeet-stt.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-parakeet-stt.ts similarity index 50% rename from packages/server/src/server/speech/sherpa/sherpa-parakeet-stt.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-parakeet-stt.ts index d3b0c151f..3b23fa8e1 100644 --- a/packages/server/src/server/speech/sherpa/sherpa-parakeet-stt.ts +++ b/packages/server/src/server/speech/providers/local/sherpa/sherpa-parakeet-stt.ts @@ -1,8 +1,14 @@ +import { EventEmitter } from "node:events"; +import { v4 as uuidv4 } from "uuid"; import type pino from "pino"; -import type { SpeechToTextProvider, TranscriptionResult } from "../speech-provider.js"; -import { Pcm16MonoResampler } from "../../agent/pcm16-resampler.js"; -import { parsePcm16MonoWav, parsePcmRateFromFormat, pcm16lePeakAbs, pcm16leToFloat32 } from "../audio.js"; +import type { + SpeechToTextProvider, + StreamingTranscriptionSession, + TranscriptionResult, +} from "../../../speech-provider.js"; +import { Pcm16MonoResampler } from "../../../../agent/pcm16-resampler.js"; +import { parsePcm16MonoWav, parsePcmRateFromFormat, pcm16lePeakAbs, pcm16leToFloat32 } from "../../../audio.js"; import { SherpaOfflineRecognizerEngine } from "./sherpa-offline-recognizer.js"; export type SherpaParakeetSttConfig = { @@ -14,6 +20,7 @@ export class SherpaOnnxParakeetSTT implements SpeechToTextProvider { private readonly engine: SherpaOfflineRecognizerEngine; private readonly silencePeakThreshold: number; private readonly logger: pino.Logger; + public readonly id = "local" as const; constructor(config: SherpaParakeetSttConfig, logger: pino.Logger) { this.engine = config.engine; @@ -21,7 +28,79 @@ export class SherpaOnnxParakeetSTT implements SpeechToTextProvider { this.logger = logger.child({ module: "speech", provider: "sherpa-onnx", component: "parakeet-stt" }); } - async transcribeAudio(audioBuffer: Buffer, format: string): Promise { + public createSession(params: { + logger: pino.Logger; + language?: string; + prompt?: string; + }): StreamingTranscriptionSession { + const emitter = new EventEmitter(); + const logger = params.logger.child({ provider: "local", component: "parakeet-stt-session" }); + const requiredSampleRate = this.engine.sampleRate; + let connected = false; + let segmentId = uuidv4(); + let previousSegmentId: string | null = null; + let pcm16: Buffer = Buffer.alloc(0); + + return { + requiredSampleRate, + async connect() { + connected = true; + }, + appendPcm16(chunk: Buffer) { + if (!connected) { + (emitter as any).emit("error", new Error("STT session not connected")); + return; + } + pcm16 = pcm16.length === 0 ? chunk : Buffer.concat([pcm16, chunk]); + }, + commit: () => { + if (!connected) { + (emitter as any).emit("error", new Error("STT session not connected")); + return; + } + + const committedId = segmentId; + const prev = previousSegmentId; + (emitter as any).emit("committed", { segmentId: committedId, previousSegmentId: prev }); + + void (async () => { + try { + const rt = await this.transcribeAudio(pcm16, `audio/pcm;rate=${requiredSampleRate}`); + (emitter as any).emit("transcript", { + segmentId: committedId, + transcript: rt.text, + isFinal: true, + language: rt.language, + logprobs: rt.logprobs, + avgLogprob: rt.avgLogprob, + isLowConfidence: rt.isLowConfidence, + }); + } catch (err) { + (emitter as any).emit("error", err); + } finally { + previousSegmentId = committedId; + segmentId = uuidv4(); + pcm16 = Buffer.alloc(0); + logger.debug({ bytes: pcm16.length }, "Parakeet session reset"); + } + })(); + }, + clear() { + pcm16 = Buffer.alloc(0); + segmentId = uuidv4(); + }, + close() { + connected = false; + pcm16 = Buffer.alloc(0); + }, + on(event: any, handler: any) { + emitter.on(event, handler); + return undefined; + }, + }; + } + + public async transcribeAudio(audioBuffer: Buffer, format: string): Promise { const start = Date.now(); let inputRate: number; diff --git a/packages/server/src/server/speech/sherpa/sherpa-realtime-session.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-realtime-session.ts similarity index 70% rename from packages/server/src/server/speech/sherpa/sherpa-realtime-session.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-realtime-session.ts index dec4fc297..6c6a251f9 100644 --- a/packages/server/src/server/speech/sherpa/sherpa-realtime-session.ts +++ b/packages/server/src/server/speech/providers/local/sherpa/sherpa-realtime-session.ts @@ -1,26 +1,28 @@ import { EventEmitter } from "node:events"; import { v4 as uuidv4 } from "uuid"; -import type { RealtimeTranscriptionSession } from "../../dictation/dictation-stream-manager.js"; -import { pcm16lePeakAbs, pcm16leToFloat32 } from "../audio.js"; +import type { StreamingTranscriptionSession } from "../../../speech-provider.js"; +import { pcm16lePeakAbs, pcm16leToFloat32 } from "../../../audio.js"; import { SherpaOnlineRecognizerEngine } from "./sherpa-online-recognizer.js"; export class SherpaRealtimeTranscriptionSession extends EventEmitter - implements RealtimeTranscriptionSession + implements StreamingTranscriptionSession { private readonly engine: SherpaOnlineRecognizerEngine; private stream: any | null = null; private connected = false; - private currentItemId: string | null = null; - private previousItemId: string | null = null; + public readonly requiredSampleRate: number; + private currentSegmentId: string | null = null; + private previousSegmentId: string | null = null; private lastPartialText = ""; private readonly tailPaddingMs: number; constructor(params: { engine: SherpaOnlineRecognizerEngine; tailPaddingMs?: number }) { super(); this.engine = params.engine; + this.requiredSampleRate = this.engine.sampleRate; this.tailPaddingMs = params.tailPaddingMs ?? 500; } @@ -29,19 +31,18 @@ export class SherpaRealtimeTranscriptionSession return; } this.stream = this.engine.createStream(); - this.currentItemId = uuidv4(); + this.currentSegmentId = uuidv4(); this.connected = true; } - appendPcm16Base64(base64Audio: string): void { - if (!this.connected || !this.stream || !this.currentItemId) { + appendPcm16(pcm16le: Buffer): void { + if (!this.connected || !this.stream || !this.currentSegmentId) { this.emit("error", new Error("Sherpa realtime session not connected")); return; } try { - const pcm16 = Buffer.from(base64Audio, "base64"); - const peak = pcm16lePeakAbs(pcm16); + const peak = pcm16lePeakAbs(pcm16le); const peakFloat = peak / 32768.0; const targetPeak = 0.6; const maxGain = 50; @@ -49,7 +50,7 @@ export class SherpaRealtimeTranscriptionSession peakFloat > 0 && peakFloat < targetPeak ? Math.min(maxGain, targetPeak / peakFloat) : 1; - const floatSamples = pcm16leToFloat32(pcm16, gain); + const floatSamples = pcm16leToFloat32(pcm16le, gain); this.stream.acceptWaveform(this.engine.sampleRate, floatSamples); while (this.engine.recognizer.isReady(this.stream)) { @@ -59,7 +60,11 @@ export class SherpaRealtimeTranscriptionSession const text = String(this.engine.recognizer.getResult(this.stream)?.text ?? "").trim(); if (text !== this.lastPartialText) { this.lastPartialText = text; - this.emit("transcript", { itemId: this.currentItemId, transcript: text, isFinal: false }); + this.emit("transcript", { + segmentId: this.currentSegmentId, + transcript: text, + isFinal: false, + }); } } catch (err) { this.emit("error", err instanceof Error ? err : new Error(String(err))); @@ -67,7 +72,7 @@ export class SherpaRealtimeTranscriptionSession } commit(): void { - if (!this.connected || !this.stream || !this.currentItemId) { + if (!this.connected || !this.stream || !this.currentSegmentId) { this.emit("error", new Error("Sherpa realtime session not connected")); return; } @@ -83,14 +88,14 @@ export class SherpaRealtimeTranscriptionSession } const finalText = String(this.engine.recognizer.getResult(this.stream)?.text ?? "").trim(); - const itemId = this.currentItemId; - const previousItemId = this.previousItemId; + const segmentId = this.currentSegmentId; + const previousSegmentId = this.previousSegmentId; - this.emit("committed", { itemId, previousItemId }); - this.emit("transcript", { itemId, transcript: finalText, isFinal: true }); + this.emit("committed", { segmentId, previousSegmentId }); + this.emit("transcript", { segmentId, transcript: finalText, isFinal: true }); - this.previousItemId = itemId; - this.currentItemId = uuidv4(); + this.previousSegmentId = segmentId; + this.currentSegmentId = uuidv4(); this.lastPartialText = ""; this.engine.recognizer.reset(this.stream); } catch (err) { @@ -104,7 +109,7 @@ export class SherpaRealtimeTranscriptionSession } try { this.engine.recognizer.reset(this.stream); - this.currentItemId = uuidv4(); + this.currentSegmentId = uuidv4(); this.lastPartialText = ""; } catch (err) { this.emit("error", err instanceof Error ? err : new Error(String(err))); diff --git a/packages/server/src/server/speech/sherpa/sherpa-stt.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-stt.ts similarity index 55% rename from packages/server/src/server/speech/sherpa/sherpa-stt.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-stt.ts index 863b4270d..ff308f51d 100644 --- a/packages/server/src/server/speech/sherpa/sherpa-stt.ts +++ b/packages/server/src/server/speech/providers/local/sherpa/sherpa-stt.ts @@ -1,8 +1,14 @@ +import { EventEmitter } from "node:events"; +import { v4 as uuidv4 } from "uuid"; import type pino from "pino"; -import type { SpeechToTextProvider, TranscriptionResult } from "../speech-provider.js"; -import { Pcm16MonoResampler } from "../../agent/pcm16-resampler.js"; -import { parsePcm16MonoWav, parsePcmRateFromFormat, pcm16lePeakAbs, pcm16leToFloat32 } from "../audio.js"; +import type { + SpeechToTextProvider, + StreamingTranscriptionSession, + TranscriptionResult, +} from "../../../speech-provider.js"; +import { Pcm16MonoResampler } from "../../../../agent/pcm16-resampler.js"; +import { parsePcm16MonoWav, parsePcmRateFromFormat, pcm16lePeakAbs, pcm16leToFloat32 } from "../../../audio.js"; import { SherpaOnlineRecognizerEngine } from "./sherpa-online-recognizer.js"; export type SherpaSttConfig = { @@ -16,6 +22,7 @@ export class SherpaOnnxSTT implements SpeechToTextProvider { private readonly silencePeakThreshold: number; private readonly tailPaddingMs: number; private readonly logger: pino.Logger; + public readonly id = "local" as const; constructor(config: SherpaSttConfig, logger: pino.Logger) { this.engine = config.engine; @@ -24,7 +31,78 @@ export class SherpaOnnxSTT implements SpeechToTextProvider { this.logger = logger.child({ module: "speech", provider: "sherpa-onnx", component: "stt" }); } - async transcribeAudio(audioBuffer: Buffer, format: string): Promise { + public createSession(params: { + logger: pino.Logger; + language?: string; + prompt?: string; + }): StreamingTranscriptionSession { + const emitter = new EventEmitter(); + void params; + const requiredSampleRate = this.engine.sampleRate; + let connected = false; + let segmentId = uuidv4(); + let previousSegmentId: string | null = null; + let pcm16: Buffer = Buffer.alloc(0); + + return { + requiredSampleRate, + async connect() { + connected = true; + }, + appendPcm16(chunk: Buffer) { + if (!connected) { + (emitter as any).emit("error", new Error("STT session not connected")); + return; + } + pcm16 = pcm16.length === 0 ? chunk : Buffer.concat([pcm16, chunk]); + }, + commit: () => { + if (!connected) { + (emitter as any).emit("error", new Error("STT session not connected")); + return; + } + + const committedId = segmentId; + const prev = previousSegmentId; + (emitter as any).emit("committed", { segmentId: committedId, previousSegmentId: prev }); + + void (async () => { + try { + const rt = await this.transcribeAudio(pcm16, `audio/pcm;rate=${requiredSampleRate}`); + (emitter as any).emit("transcript", { + segmentId: committedId, + transcript: rt.text, + isFinal: true, + language: rt.language, + logprobs: rt.logprobs, + avgLogprob: rt.avgLogprob, + isLowConfidence: rt.isLowConfidence, + }); + } catch (err) { + (emitter as any).emit("error", err); + } finally { + previousSegmentId = committedId; + segmentId = uuidv4(); + pcm16 = Buffer.alloc(0); + } + })(); + }, + clear() { + pcm16 = Buffer.alloc(0); + segmentId = uuidv4(); + }, + close() { + connected = false; + pcm16 = Buffer.alloc(0); + }, + on(event: any, handler: any) { + emitter.on(event, handler); + return undefined; + }, + }; + } + + public async transcribeAudio(audioBuffer: Buffer, format: string): Promise { const start = Date.now(); let inputRate: number; diff --git a/packages/server/src/server/speech/sherpa/sherpa-tts.ts b/packages/server/src/server/speech/providers/local/sherpa/sherpa-tts.ts similarity index 97% rename from packages/server/src/server/speech/sherpa/sherpa-tts.ts rename to packages/server/src/server/speech/providers/local/sherpa/sherpa-tts.ts index 2b5482248..99d9dab6a 100644 --- a/packages/server/src/server/speech/sherpa/sherpa-tts.ts +++ b/packages/server/src/server/speech/providers/local/sherpa/sherpa-tts.ts @@ -2,8 +2,8 @@ import type pino from "pino"; import { Readable } from "node:stream"; import { existsSync } from "node:fs"; -import type { SpeechStreamResult, TextToSpeechProvider } from "../speech-provider.js"; -import { chunkBuffer, float32ToPcm16le } from "../audio.js"; +import type { SpeechStreamResult, TextToSpeechProvider } from "../../../speech-provider.js"; +import { chunkBuffer, float32ToPcm16le } from "../../../audio.js"; import { loadSherpaOnnx } from "./sherpa-onnx-loader.js"; export type SherpaTtsPreset = "kokoro-en-v0_19" | "kitten-nano-en-v0_1-fp16"; diff --git a/packages/server/src/server/speech/sherpa/speech-download.e2e.test.ts b/packages/server/src/server/speech/providers/local/sherpa/speech-download.e2e.test.ts similarity index 97% rename from packages/server/src/server/speech/sherpa/speech-download.e2e.test.ts rename to packages/server/src/server/speech/providers/local/sherpa/speech-download.e2e.test.ts index da828a29d..8d66efdc5 100644 --- a/packages/server/src/server/speech/sherpa/speech-download.e2e.test.ts +++ b/packages/server/src/server/speech/providers/local/sherpa/speech-download.e2e.test.ts @@ -6,8 +6,8 @@ import pino from "pino"; import { ensureSherpaOnnxModels, getSherpaOnnxModelDir } from "./model-downloader.js"; import type { SherpaOnnxModelId } from "./model-catalog.js"; -import { createDaemonTestContext } from "../../test-utils/index.js"; -import { parsePcm16MonoWav, wordSimilarity } from "../../test-utils/dictation-e2e.js"; +import { createDaemonTestContext } from "../../../../test-utils/index.js"; +import { parsePcm16MonoWav, wordSimilarity } from "../../../../test-utils/dictation-e2e.js"; import { SherpaOnnxTTS } from "./sherpa-tts.js"; import { PocketTtsOnnxTTS } from "../pocket/pocket-tts-onnx.js"; import { SherpaOnlineRecognizerEngine } from "./sherpa-online-recognizer.js"; @@ -114,9 +114,9 @@ describe("speech models (download E2E)", () => { paseoHomeRoot, dictationFinalTimeoutMs: 8000, speech: { - dictationSttProvider: "sherpa", - voiceSttProvider: "sherpa", - voiceTtsProvider: "sherpa", + dictationSttProvider: "local", + voiceSttProvider: "local", + voiceTtsProvider: "local", sherpaOnnx: { modelsDir, autoDownload: false, diff --git a/packages/server/src/server/agent/openai-realtime-transcription.ts b/packages/server/src/server/speech/providers/openai/realtime-transcription-session.ts similarity index 91% rename from packages/server/src/server/agent/openai-realtime-transcription.ts rename to packages/server/src/server/speech/providers/openai/realtime-transcription-session.ts index a7beaced8..c4854f256 100644 --- a/packages/server/src/server/agent/openai-realtime-transcription.ts +++ b/packages/server/src/server/speech/providers/openai/realtime-transcription-session.ts @@ -1,6 +1,7 @@ import type pino from "pino"; import WebSocket from "ws"; import { EventEmitter } from "node:events"; +import type { StreamingTranscriptionSession } from "../../speech-provider.js"; type OpenAITurnDetection = | null @@ -60,7 +61,11 @@ type OpenAIServerEvent = } | { type: "error"; error?: { message?: string } }; -export class OpenAIRealtimeTranscriptionSession extends EventEmitter { +export class OpenAIRealtimeTranscriptionSession + extends EventEmitter + implements StreamingTranscriptionSession +{ + public readonly requiredSampleRate = 24000; private readonly apiKey: string; private readonly logger: pino.Logger; private readonly transcriptionModel: string; @@ -161,8 +166,8 @@ export class OpenAIRealtimeTranscriptionSession extends EventEmitter { if (event.type === "input_audio_buffer.committed") { this.emit("committed", { - itemId: event.item_id, - previousItemId: event.previous_item_id, + segmentId: event.item_id, + previousSegmentId: event.previous_item_id, }); return; } @@ -182,13 +187,13 @@ export class OpenAIRealtimeTranscriptionSession extends EventEmitter { const prev = this.partialByItemId.get(event.item_id) ?? ""; const next = replaceDelta ? event.delta : prev + event.delta; this.partialByItemId.set(event.item_id, next); - this.emit("transcript", { itemId: event.item_id, transcript: next, isFinal: false }); + this.emit("transcript", { segmentId: event.item_id, transcript: next, isFinal: false }); return; } if (event.type === "conversation.item.input_audio_transcription.completed") { this.partialByItemId.set(event.item_id, event.transcript); - this.emit("transcript", { itemId: event.item_id, transcript: event.transcript, isFinal: true }); + this.emit("transcript", { segmentId: event.item_id, transcript: event.transcript, isFinal: true }); return; } @@ -218,10 +223,11 @@ export class OpenAIRealtimeTranscriptionSession extends EventEmitter { return this.ready; } - public appendPcm16Base64(base64Audio: string): void { + public appendPcm16(pcm16le: Buffer): void { if (!this.ws || this.ws.readyState !== WebSocket.OPEN) { throw new Error("OpenAI realtime websocket not connected"); } + const base64Audio = pcm16le.toString("base64"); const event: OpenAIClientEvent = { type: "input_audio_buffer.append", audio: base64Audio }; this.ws.send(JSON.stringify(event)); } diff --git a/packages/server/src/server/speech/providers/openai/stt.ts b/packages/server/src/server/speech/providers/openai/stt.ts new file mode 100644 index 000000000..f9ffab062 --- /dev/null +++ b/packages/server/src/server/speech/providers/openai/stt.ts @@ -0,0 +1,269 @@ +import { EventEmitter } from "node:events"; +import type pino from "pino"; +import OpenAI from "openai"; +import { writeFile, unlink } from "fs/promises"; +import { join } from "path"; +import { tmpdir } from "os"; +import { v4 } from "uuid"; +import { inferAudioExtension } from "../../../agent/audio-utils.js"; +import type { + LogprobToken, + SpeechToTextProvider, + StreamingTranscriptionSession, + TranscriptionResult, +} from "../../speech-provider.js"; + +export type { LogprobToken, TranscriptionResult }; + +export interface STTConfig { + apiKey: string; + model?: "whisper-1" | "gpt-4o-transcribe" | "gpt-4o-mini-transcribe" | (string & {}); + confidenceThreshold?: number; // Default: -3.0 +} + +function isObject(value: unknown): value is { [key: string]: unknown } { + return typeof value === "object" && value !== null; +} + +function isLogprobToken(value: unknown): value is LogprobToken { + if (!isObject(value)) { + return false; + } + if (typeof value.token !== "string") { + return false; + } + if (typeof value.logprob !== "number") { + return false; + } + if (value.bytes === undefined) { + return true; + } + return Array.isArray(value.bytes) && value.bytes.every((entry) => typeof entry === "number"); +} + +function isLogprobTokenArray(value: unknown): value is LogprobToken[] { + return Array.isArray(value) && value.every((entry) => isLogprobToken(entry)); +} + +export class OpenAISTT implements SpeechToTextProvider { + private readonly openaiClient: OpenAI; + private readonly config: STTConfig; + private readonly logger: pino.Logger; + public readonly id = "openai" as const; + + constructor(sttConfig: STTConfig, parentLogger: pino.Logger) { + this.config = sttConfig; + this.logger = parentLogger.child({ module: "agent", provider: "openai", component: "stt" }); + this.openaiClient = new OpenAI({ + apiKey: sttConfig.apiKey, + }); + this.logger.info({ model: sttConfig.model || "whisper-1" }, "STT (OpenAI Whisper) initialized"); + } + + public createSession(params: { + logger: pino.Logger; + language?: string; + prompt?: string; + }): StreamingTranscriptionSession { + const emitter = new EventEmitter(); + const logger = params.logger.child({ provider: "openai", component: "stt-session" }); + const requiredSampleRate = 24000; + + let connected = false; + let segmentId = v4(); + let previousSegmentId: string | null = null; + let pcm16: Buffer = Buffer.alloc(0); + const transcribeAudio = this.transcribeAudioInternal.bind(this); + + const convertPCMToWavBuffer = (pcmBuffer: Buffer): Buffer => { + const headerSize = 44; + const channels = 1; + const bitsPerSample = 16; + const sampleRate = requiredSampleRate; + const wavBuffer = Buffer.alloc(headerSize + pcmBuffer.length); + const byteRate = (sampleRate * channels * bitsPerSample) / 8; + const blockAlign = (channels * bitsPerSample) / 8; + + wavBuffer.write("RIFF", 0); + wavBuffer.writeUInt32LE(36 + pcmBuffer.length, 4); + wavBuffer.write("WAVE", 8); + wavBuffer.write("fmt ", 12); + wavBuffer.writeUInt32LE(16, 16); + wavBuffer.writeUInt16LE(1, 20); + wavBuffer.writeUInt16LE(channels, 22); + wavBuffer.writeUInt32LE(sampleRate, 24); + wavBuffer.writeUInt32LE(byteRate, 28); + wavBuffer.writeUInt16LE(blockAlign, 32); + wavBuffer.writeUInt16LE(bitsPerSample, 34); + wavBuffer.write("data", 36); + wavBuffer.writeUInt32LE(pcmBuffer.length, 40); + pcmBuffer.copy(wavBuffer, 44); + + return wavBuffer; + }; + + return { + requiredSampleRate, + async connect() { + connected = true; + }, + appendPcm16(chunk: Buffer) { + if (!connected) { + (emitter as any).emit("error", new Error("STT session not connected")); + return; + } + pcm16 = pcm16.length === 0 ? chunk : Buffer.concat([pcm16, chunk]); + }, + commit() { + if (!connected) { + (emitter as any).emit("error", new Error("STT session not connected")); + return; + } + + const committedId = segmentId; + const prev = previousSegmentId; + (emitter as any).emit("committed", { segmentId: committedId, previousSegmentId: prev }); + + void (async () => { + try { + if (pcm16.length === 0) { + (emitter as any).emit("transcript", { + segmentId: committedId, + transcript: "", + isFinal: true, + language: params.language, + isLowConfidence: true, + }); + return; + } + + const wav = convertPCMToWavBuffer(pcm16); + const result = await transcribeAudio( + wav, + "audio/wav", + params.language ?? "en", + logger + ); + + (emitter as any).emit("transcript", { + segmentId: committedId, + transcript: result.text, + isFinal: true, + language: result.language, + logprobs: result.logprobs, + avgLogprob: result.avgLogprob, + isLowConfidence: result.isLowConfidence, + }); + } catch (err) { + (emitter as any).emit("error", err); + } finally { + previousSegmentId = committedId; + segmentId = v4(); + pcm16 = Buffer.alloc(0); + } + })(); + }, + clear() { + pcm16 = Buffer.alloc(0); + segmentId = v4(); + }, + close() { + connected = false; + pcm16 = Buffer.alloc(0); + }, + on(event: any, handler: any) { + emitter.on(event, handler); + return undefined; + }, + }; + } + + private async transcribeAudioInternal( + audioBuffer: Buffer, + format: string, + language: string, + logger: pino.Logger + ): Promise { + const startTime = Date.now(); + let tempFilePath: string | null = null; + + try { + const ext = inferAudioExtension(format); + tempFilePath = join(tmpdir(), `audio-${v4()}.${ext}`); + await writeFile(tempFilePath, audioBuffer); + + logger.debug( + { tempFilePath, bytes: audioBuffer.length }, + "Transcribing audio file" + ); + + const modelToUse = this.config.model ?? "whisper-1"; + const supportsLogprobs = + modelToUse === "gpt-4o-transcribe" || modelToUse === "gpt-4o-mini-transcribe"; + const includeLogprobs: ["logprobs"] = ["logprobs"]; + + const response = await this.openaiClient.audio.transcriptions.create({ + file: await import("fs").then((fs) => fs.createReadStream(tempFilePath!)), + language, + model: modelToUse, + ...(supportsLogprobs ? { include: includeLogprobs } : {}), + response_format: "json", + }); + + const duration = Date.now() - startTime; + const confidenceThreshold = this.config.confidenceThreshold ?? -3.0; + + let avgLogprob: number | undefined; + let isLowConfidence = false; + const logprobs = + supportsLogprobs && + isObject(response) && + isLogprobTokenArray(response.logprobs) + ? response.logprobs + : undefined; + + if (logprobs && logprobs.length > 0) { + const totalLogprob = logprobs.reduce((sum, token) => sum + token.logprob, 0); + avgLogprob = totalLogprob / logprobs.length; + isLowConfidence = avgLogprob < confidenceThreshold; + + if (isLowConfidence) { + logger.debug( + { + avgLogprob, + threshold: confidenceThreshold, + text: response.text, + tokenLogprobs: logprobs.map((t) => `${t.token}:${t.logprob.toFixed(2)}`).join(", "), + }, + "Low confidence transcription detected" + ); + } + } + + logger.debug({ duration, text: response.text, avgLogprob }, "Transcription complete"); + + return { + text: response.text, + duration: duration, + logprobs: logprobs, + avgLogprob: avgLogprob, + isLowConfidence: isLowConfidence, + language: + isObject(response) && typeof response.language === "string" + ? response.language + : undefined, + }; + } catch (error: any) { + logger.error({ err: error }, "Transcription error"); + throw new Error(`STT transcription failed: ${error.message}`); + } finally { + if (tempFilePath) { + try { + await unlink(tempFilePath); + } catch (cleanupError) { + logger.warn({ tempFilePath }, "Failed to clean up temp file"); + } + } + } + } +} diff --git a/packages/server/src/server/agent/tts-openai.ts b/packages/server/src/server/speech/providers/openai/tts.ts similarity index 93% rename from packages/server/src/server/agent/tts-openai.ts rename to packages/server/src/server/speech/providers/openai/tts.ts index bd96a80b4..7333f3126 100644 --- a/packages/server/src/server/agent/tts-openai.ts +++ b/packages/server/src/server/speech/providers/openai/tts.ts @@ -1,7 +1,7 @@ import type pino from "pino"; import OpenAI from "openai"; import { Readable } from "node:stream"; -import type { SpeechStreamResult } from "../speech/speech-provider.js"; +import type { SpeechStreamResult, TextToSpeechProvider } from "../../speech-provider.js"; export type { SpeechStreamResult }; @@ -12,7 +12,7 @@ export interface TTSConfig { responseFormat?: "mp3" | "opus" | "aac" | "flac" | "wav" | "pcm"; } -export class OpenAITTS { +export class OpenAITTS implements TextToSpeechProvider { private readonly openaiClient: OpenAI; private readonly config: TTSConfig; private readonly logger: pino.Logger; diff --git a/packages/server/src/server/speech/speech-provider.ts b/packages/server/src/server/speech/speech-provider.ts index 3ed103666..fb58ba910 100644 --- a/packages/server/src/server/speech/speech-provider.ts +++ b/packages/server/src/server/speech/speech-provider.ts @@ -1,3 +1,4 @@ +import type pino from "pino"; import type { Readable } from "node:stream"; export interface LogprobToken { @@ -15,8 +16,52 @@ export interface TranscriptionResult { isLowConfidence?: boolean; } +export interface StreamingTranscriptionCommittedEvent { + segmentId: string; + previousSegmentId: string | null; +} + +export interface StreamingTranscriptionEvent { + segmentId: string; + transcript: string; + isFinal: boolean; + language?: string; + logprobs?: LogprobToken[]; + avgLogprob?: number; + isLowConfidence?: boolean; +} + +export type StreamingTranscriptionSession = { + /** + * Required PCM16LE sample rate for `appendPcm16()`. + * Callers are responsible for resampling before appending. + */ + requiredSampleRate: number; + + connect(): Promise; + appendPcm16(pcm16le: Buffer): void; + commit(): void; + clear(): void; + close(): void; + + on( + event: "committed", + handler: (payload: StreamingTranscriptionCommittedEvent) => void + ): unknown; + on( + event: "transcript", + handler: (payload: StreamingTranscriptionEvent) => void + ): unknown; + on(event: "error", handler: (err: unknown) => void): unknown; +}; + export interface SpeechToTextProvider { - transcribeAudio(audioBuffer: Buffer, format: string): Promise; + id: "openai" | "local" | (string & {}); + createSession(params: { + logger: pino.Logger; + language?: string; + prompt?: string; + }): StreamingTranscriptionSession; } export interface SpeechStreamResult { @@ -27,4 +72,3 @@ export interface SpeechStreamResult { export interface TextToSpeechProvider { synthesizeSpeech(text: string): Promise; } - diff --git a/packages/server/src/server/websocket-server.ts b/packages/server/src/server/websocket-server.ts index b7e84f4bf..4d347694e 100644 --- a/packages/server/src/server/websocket-server.ts +++ b/packages/server/src/server/websocket-server.ts @@ -21,7 +21,6 @@ import { PushTokenStore } from "./push/token-store.js"; import { PushService } from "./push/push-service.js"; import { VoiceConversationStore } from "./voice-conversation-store.js"; import type { SpeechToTextProvider, TextToSpeechProvider } from "./speech/speech-provider.js"; -import type { RealtimeTranscriptionSessionFactory } from "./dictation/dictation-stream-manager.js"; export type AgentMcpTransportFactory = () => Promise; @@ -72,9 +71,8 @@ export class VoiceAssistantWebSocketServer { private readonly terminalManager: TerminalManager | null; private readonly voiceConversationStore: VoiceConversationStore; private readonly dictation: { - openaiApiKey?: string | null; finalTimeoutMs?: number; - sessionFactory?: RealtimeTranscriptionSessionFactory; + stt?: SpeechToTextProvider | null; } | null; private readonly voice: { openrouterApiKey?: string | null; @@ -98,9 +96,8 @@ export class VoiceAssistantWebSocketServer { voiceLlmModel?: string | null; }, dictation?: { - openaiApiKey?: string | null; finalTimeoutMs?: number; - sessionFactory?: RealtimeTranscriptionSessionFactory; + stt?: SpeechToTextProvider | null; } ) { this.logger = logger.child({ module: "websocket-server" }); diff --git a/packages/server/src/shared/messages.ts b/packages/server/src/shared/messages.ts index f53783f8b..a8f2a3d03 100644 --- a/packages/server/src/shared/messages.ts +++ b/packages/server/src/shared/messages.ts @@ -489,6 +489,17 @@ export const ListProviderModelsRequestMessageSchema = z.object({ requestId: z.string(), }); +export const SpeechModelsListRequestSchema = z.object({ + type: z.literal("speech_models_list_request"), + requestId: z.string(), +}); + +export const SpeechModelsDownloadRequestSchema = z.object({ + type: z.literal("speech_models_download_request"), + modelIds: z.array(z.string()).optional(), + requestId: z.string(), +}); + export const ResumeAgentRequestMessageSchema = z.object({ type: z.literal("resume_agent_request"), handle: AgentPersistenceHandleSchema, @@ -887,6 +898,8 @@ export const SessionInboundMessageSchema = z.discriminatedUnion("type", [ DictationStreamCancelMessageSchema, CreateAgentRequestMessageSchema, ListProviderModelsRequestMessageSchema, + SpeechModelsListRequestSchema, + SpeechModelsDownloadRequestSchema, ResumeAgentRequestMessageSchema, RefreshAgentRequestMessageSchema, CancelAgentRequestMessageSchema, @@ -1518,6 +1531,34 @@ export const ListProviderModelsResponseMessageSchema = z.object({ }), }); +export const SpeechModelsListResponseSchema = z.object({ + type: z.literal("speech_models_list_response"), + payload: z.object({ + modelsDir: z.string(), + models: z.array( + z.object({ + id: z.string(), + kind: z.string(), + description: z.string(), + modelDir: z.string(), + isDownloaded: z.boolean(), + missingFiles: z.array(z.string()).optional(), + }) + ), + requestId: z.string(), + }), +}); + +export const SpeechModelsDownloadResponseSchema = z.object({ + type: z.literal("speech_models_download_response"), + payload: z.object({ + modelsDir: z.string(), + downloadedModelIds: z.array(z.string()), + error: z.string().nullable(), + requestId: z.string(), + }), +}); + const AgentSlashCommandSchema = z.object({ name: z.string(), description: z.string(), @@ -1673,6 +1714,8 @@ export const SessionOutboundMessageSchema = z.discriminatedUnion("type", [ ProjectIconResponseSchema, FileDownloadTokenResponseSchema, ListProviderModelsResponseMessageSchema, + SpeechModelsListResponseSchema, + SpeechModelsDownloadResponseSchema, ListCommandsResponseSchema, ExecuteCommandResponseSchema, ListTerminalsResponseSchema, @@ -1727,6 +1770,8 @@ export type AgentDeletedMessage = z.infer; export type ListProviderModelsResponseMessage = z.infer< typeof ListProviderModelsResponseMessageSchema >; +export type SpeechModelsListResponse = z.infer; +export type SpeechModelsDownloadResponse = z.infer; export type InitializeAgentResponseMessage = z.infer; // Type exports for payload types @@ -1747,6 +1792,10 @@ export type CreateAgentRequestMessage = z.infer; +export type SpeechModelsListRequestMessage = z.infer; +export type SpeechModelsDownloadRequestMessage = z.infer< + typeof SpeechModelsDownloadRequestSchema +>; export type ResumeAgentRequestMessage = z.infer; export type DeleteAgentRequestMessage = z.infer; export type InitializeAgentRequestMessage = z.infer; diff --git a/scripts/speech/download-sherpa-models.sh b/scripts/speech/download-sherpa-models.sh deleted file mode 100755 index 5a3ccdc5c..000000000 --- a/scripts/speech/download-sherpa-models.sh +++ /dev/null @@ -1,95 +0,0 @@ -#!/usr/bin/env bash -set -euo pipefail - -usage() { - cat <<'EOF' -Download local speech models for Paseo (sherpa-onnx). - -Defaults: - - STT: sherpa-onnx-streaming-zipformer-bilingual-zh-en-2023-02-20 - - TTS: kitten-nano-en-v0_1-fp16 - -Usage: - scripts/speech/download-sherpa-models.sh [--models-dir DIR] [--with-kokoro] [--with-paraformer] - -Preferred: - npm run speech:download --workspace=@getpaseo/server - -Notes: - - Models are downloaded from the sherpa-onnx GitHub releases. - - Pocket TTS is downloaded by the Node script (`npm run speech:download --workspace=@getpaseo/server`) - because it is a file-based HuggingFace model (not a single tarball). - - Set PASEO_SHERPA_ONNX_MODELS_DIR to override where the daemon looks. -EOF -} - -MODELS_DIR="" -WITH_KOKORO=0 -WITH_PARAFORMER=0 - -while [[ $# -gt 0 ]]; do - case "$1" in - --models-dir) - MODELS_DIR="${2:-}" - shift 2 - ;; - --with-kokoro) - WITH_KOKORO=1 - shift 1 - ;; - --with-paraformer) - WITH_PARAFORMER=1 - shift 1 - ;; - -h|--help) - usage - exit 0 - ;; - *) - echo "Unknown arg: $1" >&2 - usage >&2 - exit 2 - ;; - esac -done - -if [[ -z "${MODELS_DIR}" ]]; then - if [[ -n "${PASEO_SHERPA_ONNX_MODELS_DIR:-}" ]]; then - MODELS_DIR="${PASEO_SHERPA_ONNX_MODELS_DIR}" - elif [[ -n "${PASEO_HOME:-}" ]]; then - MODELS_DIR="${PASEO_HOME}/models/sherpa-onnx" - else - MODELS_DIR="${HOME}/.paseo/models/sherpa-onnx" - fi -fi - -mkdir -p "${MODELS_DIR}" -cd "${MODELS_DIR}" - -download_and_extract() { - local url="$1" - local filename - filename="$(basename "$url")" - - echo "Downloading ${filename}..." - curl -fsSL -O "${url}" - echo "Extracting ${filename}..." - tar xf "${filename}" - rm -f "${filename}" -} - -echo "NOTE: This script is deprecated. Prefer: npm run speech:download --workspace=@getpaseo/server" >&2 - -download_and_extract "https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-streaming-zipformer-bilingual-zh-en-2023-02-20.tar.bz2" -download_and_extract "https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/kitten-nano-en-v0_1-fp16.tar.bz2" - -if [[ "${WITH_PARAFORMER}" -eq 1 ]]; then - download_and_extract "https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-streaming-paraformer-bilingual-zh-en.tar.bz2" -fi - -if [[ "${WITH_KOKORO}" -eq 1 ]]; then - download_and_extract "https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/kokoro-en-v0_19.tar.bz2" -fi - -echo "Done." -echo "Models dir: ${MODELS_DIR}"