fix(security): isolate audio recorder generations (#11456)

This commit is contained in:
Jasper
2026-09-01 04:39:08 +00:00
committed by GitHub
parent 34898f1652
commit 9a5319a6e8
2 changed files with 554 additions and 90 deletions
@@ -0,0 +1,348 @@
import { act, renderHook, waitFor } from '@testing-library/react';
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
const mocks = vi.hoisted(() => ({
config: {},
getDictationConfig: vi.fn(),
onError: vi.fn(),
onTranscription: vi.fn(),
read: vi.fn(),
transcribeDictation: vi.fn(),
}));
vi.mock('../components/ConfigContext', () => ({
useConfig: () => ({ config: mocks.config, read: mocks.read }),
}));
vi.mock('../acp/dictation', () => ({
getDictationConfig: mocks.getDictationConfig,
transcribeDictation: mocks.transcribeDictation,
}));
import { useAudioRecorder } from './useAudioRecorder';
type Deferred<T> = {
promise: Promise<T>;
resolve: (value: T) => void;
};
type FakeStream = MediaStream & {
track: { stop: ReturnType<typeof vi.fn> };
};
const deferred = <T,>(): Deferred<T> => {
let resolve = (_value: T) => {};
const promise = new Promise<T>((resolvePromise) => {
resolve = resolvePromise;
});
return { promise, resolve };
};
const createStream = (): FakeStream => {
const track = { stop: vi.fn() };
return {
track,
getTracks: () => [track],
} as unknown as FakeStream;
};
let moduleLoads: Deferred<void>[];
let audioContexts: MockAudioContext[];
let worklets: MockAudioWorkletNode[];
class MockAudioContext {
audioWorklet: { addModule: ReturnType<typeof vi.fn> };
close = vi.fn(() => Promise.resolve());
createGain = vi.fn(() => ({
connect: vi.fn(),
gain: { value: 1 },
}));
createMediaStreamSource = vi.fn(() => ({ connect: vi.fn() }));
destination = {};
constructor() {
const moduleLoad = moduleLoads.shift();
this.audioWorklet = {
addModule: vi.fn(() => moduleLoad?.promise ?? Promise.resolve()),
};
audioContexts.push(this);
}
}
class MockAudioWorkletNode {
connect = vi.fn();
disconnect = vi.fn();
port: { onmessage: ((event: MessageEvent<Float32Array>) => void) | null } = {
onmessage: null,
};
constructor() {
worklets.push(this);
}
emit(samples: Float32Array) {
this.port.onmessage?.({ data: samples } as MessageEvent<Float32Array>);
}
}
const renderRecorder = async () => {
const hook = renderHook(() =>
useAudioRecorder({
onError: mocks.onError,
onTranscription: mocks.onTranscription,
})
);
await waitFor(() => expect(hook.result.current.isEnabled).toBe(true));
return hook;
};
const startRecorder = async (startRecording: () => Promise<void>) => {
await act(async () => {
await startRecording();
});
};
const emitSpeechChunk = (worklet: MockAudioWorkletNode) => {
let now = 100;
vi.spyOn(Date, 'now').mockImplementation(() => now);
const speech = new Float32Array(3200).fill(0.1);
const silence = new Float32Array(3200);
act(() => {
worklet.emit(speech);
now = 400;
worklet.emit(silence);
now = 1301;
worklet.emit(silence);
});
};
describe('useAudioRecorder lifecycle', () => {
let getUserMedia: ReturnType<typeof vi.fn>;
beforeEach(() => {
moduleLoads = [];
audioContexts = [];
worklets = [];
getUserMedia = vi.fn();
mocks.getDictationConfig.mockReset().mockResolvedValue({
openai: { configured: true },
});
mocks.onError.mockReset();
mocks.onTranscription.mockReset();
mocks.read
.mockReset()
.mockImplementation((key: string) =>
Promise.resolve(key === 'voice_dictation_provider' ? 'openai' : null)
);
mocks.transcribeDictation.mockReset();
vi.stubGlobal('AudioContext', MockAudioContext);
vi.stubGlobal('AudioWorkletNode', MockAudioWorkletNode);
Object.defineProperty(navigator, 'mediaDevices', {
configurable: true,
value: { getUserMedia },
});
});
afterEach(() => {
vi.restoreAllMocks();
vi.unstubAllGlobals();
});
it('cleans the superseded generation when starts overlap', async () => {
const firstStream = createStream();
const secondStream = createStream();
const firstModuleLoad = deferred<void>();
moduleLoads.push(firstModuleLoad);
getUserMedia.mockResolvedValueOnce(firstStream).mockResolvedValueOnce(secondStream);
const { result } = await renderRecorder();
let firstStart = Promise.resolve();
act(() => {
firstStart = result.current.startRecording();
});
await waitFor(() => expect(audioContexts).toHaveLength(1));
await startRecorder(result.current.startRecording);
expect(firstStream.track.stop).toHaveBeenCalledOnce();
expect(audioContexts[0].close).toHaveBeenCalledOnce();
expect(worklets).toHaveLength(1);
expect(result.current.isRecording).toBe(true);
await act(async () => {
firstModuleLoad.resolve();
await firstStart;
});
expect(worklets).toHaveLength(1);
act(() => result.current.stopRecording());
expect(secondStream.track.stop).toHaveBeenCalledOnce();
expect(audioContexts[1].close).toHaveBeenCalledOnce();
expect(worklets[0].disconnect).toHaveBeenCalledOnce();
});
it('stops a stream that arrives after recording is stopped', async () => {
const pendingStream = deferred<MediaStream>();
const stream = createStream();
getUserMedia.mockReturnValueOnce(pendingStream.promise);
const { result } = await renderRecorder();
let start = Promise.resolve();
act(() => {
start = result.current.startRecording();
});
await waitFor(() => expect(getUserMedia).toHaveBeenCalledOnce());
act(() => result.current.stopRecording());
await act(async () => {
pendingStream.resolve(stream);
await start;
});
expect(stream.track.stop).toHaveBeenCalledOnce();
expect(audioContexts).toHaveLength(0);
expect(worklets).toHaveLength(0);
expect(result.current.isRecording).toBe(false);
expect(mocks.onError).not.toHaveBeenCalled();
});
it('does not finish startup after unmount', async () => {
const stream = createStream();
const moduleLoad = deferred<void>();
moduleLoads.push(moduleLoad);
getUserMedia.mockResolvedValueOnce(stream);
const { result, unmount } = await renderRecorder();
let start = Promise.resolve();
act(() => {
start = result.current.startRecording();
});
await waitFor(() => expect(audioContexts).toHaveLength(1));
unmount();
expect(stream.track.stop).toHaveBeenCalledOnce();
expect(audioContexts[0].close).toHaveBeenCalledOnce();
await act(async () => {
moduleLoad.resolve();
await start;
});
expect(worklets).toHaveLength(0);
expect(mocks.onError).not.toHaveBeenCalled();
});
it('suppresses an in-flight transcription after stop', async () => {
const stream = createStream();
const transcription = deferred<string>();
getUserMedia.mockResolvedValueOnce(stream);
mocks.transcribeDictation.mockReturnValueOnce(transcription.promise);
const { result } = await renderRecorder();
await startRecorder(result.current.startRecording);
emitSpeechChunk(worklets[0]);
await waitFor(() => expect(mocks.transcribeDictation).toHaveBeenCalledOnce());
act(() => result.current.stopRecording());
await act(async () => {
transcription.resolve('stale transcription');
await transcription.promise;
await Promise.resolve();
});
expect(mocks.onTranscription).not.toHaveBeenCalled();
});
it('transcribes the buffered final phrase when recording is stopped', async () => {
const stream = createStream();
getUserMedia.mockResolvedValueOnce(stream);
mocks.transcribeDictation.mockResolvedValueOnce('final phrase');
const { result } = await renderRecorder();
await startRecorder(result.current.startRecording);
act(() => {
worklets[0].emit(new Float32Array(3200).fill(0.1));
result.current.stopRecording();
});
await waitFor(() => expect(mocks.onTranscription).toHaveBeenCalledWith('final phrase'));
await waitFor(() => expect(result.current.isTranscribing).toBe(false));
expect(stream.track.stop).toHaveBeenCalledOnce();
expect(audioContexts[0].close).toHaveBeenCalledOnce();
expect(worklets[0].disconnect).toHaveBeenCalledOnce();
});
it('cancels final-phrase transcription when a new recording starts', async () => {
const firstStream = createStream();
const secondStream = createStream();
const transcription = deferred<string>();
getUserMedia.mockResolvedValueOnce(firstStream).mockResolvedValueOnce(secondStream);
mocks.transcribeDictation.mockReturnValueOnce(transcription.promise);
const { result } = await renderRecorder();
await startRecorder(result.current.startRecording);
act(() => {
worklets[0].emit(new Float32Array(3200).fill(0.1));
result.current.stopRecording();
});
await waitFor(() => expect(mocks.transcribeDictation).toHaveBeenCalledOnce());
await startRecorder(result.current.startRecording);
await act(async () => {
transcription.resolve('stale final phrase');
await transcription.promise;
await Promise.resolve();
});
expect(mocks.onTranscription).not.toHaveBeenCalled();
expect(result.current.isRecording).toBe(true);
act(() => result.current.stopRecording());
});
it('cancels final-phrase transcription when unmounted', async () => {
const stream = createStream();
const transcription = deferred<string>();
getUserMedia.mockResolvedValueOnce(stream);
mocks.transcribeDictation.mockReturnValueOnce(transcription.promise);
const { result, unmount } = await renderRecorder();
await startRecorder(result.current.startRecording);
act(() => {
worklets[0].emit(new Float32Array(3200).fill(0.1));
result.current.stopRecording();
});
await waitFor(() => expect(mocks.transcribeDictation).toHaveBeenCalledOnce());
unmount();
await act(async () => {
transcription.resolve('stale final phrase');
await transcription.promise;
await Promise.resolve();
});
expect(mocks.onTranscription).not.toHaveBeenCalled();
});
it('records and transcribes a normal generation', async () => {
const stream = createStream();
getUserMedia.mockResolvedValueOnce(stream);
mocks.transcribeDictation.mockResolvedValueOnce('hello world');
const { result } = await renderRecorder();
await startRecorder(result.current.startRecording);
expect(result.current.isRecording).toBe(true);
emitSpeechChunk(worklets[0]);
await waitFor(() => expect(mocks.onTranscription).toHaveBeenCalledWith('hello world'));
await waitFor(() => expect(result.current.isTranscribing).toBe(false));
act(() => result.current.stopRecording());
expect(result.current.isRecording).toBe(false);
expect(stream.track.stop).toHaveBeenCalledOnce();
expect(audioContexts[0].close).toHaveBeenCalledOnce();
expect(worklets[0].disconnect).toHaveBeenCalledOnce();
});
});
+206 -90
View File
@@ -17,6 +17,36 @@ const MIN_SPEECH_MS = 200;
// without clipping early speech onsets. Determined empirically for 16kHz mono input.
const RMS_THRESHOLD = 0.015;
interface RecorderGeneration {
audioContext: AudioContext | null;
cancelled: boolean;
completesAfterTranscription: boolean;
pendingTranscriptions: number;
stream: MediaStream | null;
worklet: AudioWorkletNode | null;
}
function cleanupGeneration(generation: RecorderGeneration) {
generation.cancelled = true;
const worklet = generation.worklet;
generation.worklet = null;
if (worklet) {
worklet.port.onmessage = null;
worklet.disconnect();
}
const audioContext = generation.audioContext;
generation.audioContext = null;
if (audioContext) {
void audioContext.close();
}
const stream = generation.stream;
generation.stream = null;
stream?.getTracks().forEach((track) => track.stop());
}
// Resolve worklet URL at runtime from window.location so it works under both
// the dev server (http://localhost) and packaged builds (file://).
const WORKLET_URL = new URL('audio-capture-worklet.js', window.location.href.split('#')[0]).href;
@@ -55,6 +85,17 @@ function rms(samples: Float32Array): number {
return Math.sqrt(sum / samples.length);
}
function mergeSamples(chunks: Float32Array[]): Float32Array {
const total = chunks.reduce((length, chunk) => length + chunk.length, 0);
const merged = new Float32Array(total);
let offset = 0;
for (const chunk of chunks) {
merged.set(chunk, offset);
offset += chunk.length;
}
return merged;
}
function blobToBase64(blob: Blob): Promise<string> {
return new Promise((resolve, reject) => {
const r = new FileReader();
@@ -72,15 +113,14 @@ export const useAudioRecorder = ({ onTranscription, onError }: UseAudioRecorderO
const { read, config } = useConfig();
const audioContextRef = useRef<AudioContext | null>(null);
const streamRef = useRef<MediaStream | null>(null);
const activeGenerationRef = useRef<RecorderGeneration | null>(null);
const mountedRef = useRef(true);
// VAD state (all refs to avoid re-render/stale closure issues)
const samplesRef = useRef<Float32Array[]>([]);
const isSpeakingRef = useRef(false);
const silenceStartRef = useRef(0);
const speechStartRef = useRef(0);
const pendingTranscriptions = useRef(0);
const providerRef = useRef(provider);
providerRef.current = provider;
@@ -112,95 +152,149 @@ export const useAudioRecorder = ({ onTranscription, onError }: UseAudioRecorderO
check();
}, [read, config]);
const transcribeChunk = useCallback(async (samples: Float32Array) => {
const prov = providerRef.current;
if (!prov) return;
pendingTranscriptions.current++;
setIsTranscribing(true);
try {
const wav = new Blob([encodeWav(samples, SAMPLE_RATE)], { type: 'audio/wav' });
const base64 = await blobToBase64(wav);
const text = await transcribeDictation(base64, 'audio/wav', prov);
if (text) {
onTranscriptionRef.current(text);
}
} catch (error) {
onErrorRef.current(errorMessage(error));
} finally {
pendingTranscriptions.current--;
if (pendingTranscriptions.current === 0) setIsTranscribing(false);
}
}, []);
const flush = useCallback(() => {
const chunks = samplesRef.current;
if (chunks.length === 0) return;
const total = chunks.reduce((n, c) => n + c.length, 0);
const merged = new Float32Array(total);
let off = 0;
for (const c of chunks) {
merged.set(c, off);
off += c.length;
}
const resetSpeech = useCallback(() => {
samplesRef.current = [];
transcribeChunk(merged);
}, [transcribeChunk]);
const flushRef = useRef(flush);
flushRef.current = flush;
const handleSamples = useCallback((samples: Float32Array) => {
const now = Date.now();
if (rms(samples) > RMS_THRESHOLD) {
if (!isSpeakingRef.current) {
isSpeakingRef.current = true;
speechStartRef.current = now;
}
silenceStartRef.current = 0;
samplesRef.current.push(new Float32Array(samples));
} else if (isSpeakingRef.current) {
samplesRef.current.push(new Float32Array(samples));
if (silenceStartRef.current === 0) {
silenceStartRef.current = now;
} else if (now - silenceStartRef.current > SILENCE_MS) {
if (now - speechStartRef.current > MIN_SPEECH_MS) {
flushRef.current();
} else {
samplesRef.current = [];
}
isSpeakingRef.current = false;
silenceStartRef.current = 0;
}
}
}, []);
const stopRecording = useCallback(() => {
if (isSpeakingRef.current && samplesRef.current.length > 0) {
flushRef.current();
}
isSpeakingRef.current = false;
silenceStartRef.current = 0;
audioContextRef.current?.close();
audioContextRef.current = null;
streamRef.current?.getTracks().forEach((t) => t.stop());
streamRef.current = null;
setIsRecording(false);
speechStartRef.current = 0;
}, []);
const isActiveGeneration = useCallback(
(generation: RecorderGeneration) =>
mountedRef.current && activeGenerationRef.current === generation && !generation.cancelled,
[]
);
const transcribeChunk = useCallback(
async (samples: Float32Array, generation: RecorderGeneration) => {
const prov = providerRef.current;
if (!prov || !isActiveGeneration(generation)) return;
generation.pendingTranscriptions++;
setIsTranscribing(true);
try {
const wav = new Blob([encodeWav(samples, SAMPLE_RATE)], { type: 'audio/wav' });
const base64 = await blobToBase64(wav);
if (!isActiveGeneration(generation)) return;
const text = await transcribeDictation(base64, 'audio/wav', prov);
if (text && isActiveGeneration(generation)) {
onTranscriptionRef.current(text);
}
} catch (error) {
if (isActiveGeneration(generation)) {
onErrorRef.current(errorMessage(error));
}
} finally {
generation.pendingTranscriptions--;
if (generation.pendingTranscriptions === 0 && isActiveGeneration(generation)) {
setIsTranscribing(false);
if (generation.completesAfterTranscription) {
activeGenerationRef.current = null;
generation.cancelled = true;
}
}
}
},
[isActiveGeneration]
);
const flush = useCallback(
(generation: RecorderGeneration) => {
if (!isActiveGeneration(generation)) return;
const chunks = samplesRef.current;
if (chunks.length === 0) return;
samplesRef.current = [];
void transcribeChunk(mergeSamples(chunks), generation);
},
[isActiveGeneration, transcribeChunk]
);
const handleSamples = useCallback(
(samples: Float32Array, generation: RecorderGeneration) => {
if (!isActiveGeneration(generation)) return;
const now = Date.now();
if (rms(samples) > RMS_THRESHOLD) {
if (!isSpeakingRef.current) {
isSpeakingRef.current = true;
speechStartRef.current = now;
}
silenceStartRef.current = 0;
samplesRef.current.push(new Float32Array(samples));
} else if (isSpeakingRef.current) {
samplesRef.current.push(new Float32Array(samples));
if (silenceStartRef.current === 0) {
silenceStartRef.current = now;
} else if (now - silenceStartRef.current > SILENCE_MS) {
if (now - speechStartRef.current > MIN_SPEECH_MS) {
flush(generation);
} else {
samplesRef.current = [];
}
isSpeakingRef.current = false;
silenceStartRef.current = 0;
}
}
},
[flush, isActiveGeneration]
);
const cancelActiveGeneration = useCallback(() => {
const generation = activeGenerationRef.current;
activeGenerationRef.current = null;
if (generation) cleanupGeneration(generation);
resetSpeech();
if (mountedRef.current) {
setIsRecording(false);
setIsTranscribing(false);
}
}, [resetSpeech]);
const stopRecording = useCallback(() => {
const finalChunks = isSpeakingRef.current ? samplesRef.current : [];
cancelActiveGeneration();
if (mountedRef.current && finalChunks.length > 0) {
const finalGeneration: RecorderGeneration = {
audioContext: null,
cancelled: false,
completesAfterTranscription: true,
pendingTranscriptions: 0,
stream: null,
worklet: null,
};
activeGenerationRef.current = finalGeneration;
void transcribeChunk(mergeSamples(finalChunks), finalGeneration);
}
}, [cancelActiveGeneration, transcribeChunk]);
const startRecording = useCallback(async () => {
if (!isEnabled) {
onError('Voice dictation is not enabled');
onErrorRef.current('Voice dictation is not enabled');
return;
}
cancelActiveGeneration();
const generation: RecorderGeneration = {
audioContext: null,
cancelled: false,
completesAfterTranscription: false,
pendingTranscriptions: 0,
stream: null,
worklet: null,
};
activeGenerationRef.current = generation;
try {
const preferredMic = await read('voice_dictation_preferred_mic', false);
if (!isActiveGeneration(generation)) return;
const audioConstraints: MediaTrackConstraints = {
echoCancellation: true,
@@ -215,6 +309,7 @@ export const useAudioRecorder = ({ onTranscription, onError }: UseAudioRecorderO
try {
stream = await navigator.mediaDevices.getUserMedia({ audio: audioConstraints });
} catch (e) {
if (!isActiveGeneration(generation)) return;
if (
preferredMic &&
e instanceof DOMException &&
@@ -226,17 +321,26 @@ export const useAudioRecorder = ({ onTranscription, onError }: UseAudioRecorderO
throw e;
}
}
streamRef.current = stream;
generation.stream = stream;
if (!isActiveGeneration(generation)) {
cleanupGeneration(generation);
return;
}
const ctx = new AudioContext({ sampleRate: SAMPLE_RATE });
audioContextRef.current = ctx;
generation.audioContext = ctx;
await ctx.audioWorklet.addModule(WORKLET_URL);
if (!isActiveGeneration(generation)) {
cleanupGeneration(generation);
return;
}
const source = ctx.createMediaStreamSource(stream);
const worklet = new AudioWorkletNode(ctx, 'audio-capture');
generation.worklet = worklet;
worklet.port.onmessage = (e: MessageEvent<Float32Array>) => handleSamples(e.data);
worklet.port.onmessage = (e: MessageEvent<Float32Array>) => handleSamples(e.data, generation);
// Connect through silent gain to keep worklet processing alive
const silence = ctx.createGain();
@@ -245,19 +349,31 @@ export const useAudioRecorder = ({ onTranscription, onError }: UseAudioRecorderO
worklet.connect(silence);
silence.connect(ctx.destination);
setIsRecording(true);
if (isActiveGeneration(generation)) {
setIsRecording(true);
}
} catch (error) {
stopRecording();
onError(errorMessage(error));
const isCurrent = isActiveGeneration(generation);
if (activeGenerationRef.current === generation) {
activeGenerationRef.current = null;
}
cleanupGeneration(generation);
if (isCurrent) {
resetSpeech();
setIsRecording(false);
setIsTranscribing(false);
onErrorRef.current(errorMessage(error));
}
}
}, [isEnabled, onError, handleSamples, stopRecording, read]);
}, [cancelActiveGeneration, handleSamples, isActiveGeneration, isEnabled, read, resetSpeech]);
useEffect(() => {
mountedRef.current = true;
return () => {
audioContextRef.current?.close();
streamRef.current?.getTracks().forEach((t) => t.stop());
mountedRef.current = false;
cancelActiveGeneration();
};
}, []);
}, [cancelActiveGeneration]);
return {
isEnabled,