diff --git a/packages/cli/src/lib/parallel.test.ts b/packages/cli/src/lib/parallel.test.ts index 6c7e7387d7..824a83d335 100644 --- a/packages/cli/src/lib/parallel.test.ts +++ b/packages/cli/src/lib/parallel.test.ts @@ -169,7 +169,9 @@ describe('runWorkerQueueThreads', () => { ]); expect(results).toEqual([20, 21, 22, 23, 24, 25, 26, 27, 28, 29]); }); +}); +describe('runWorkerThreads', () => { it('should run a single thread without items', async () => { const [result] = await runWorkerThreads({ threadCount: 1, @@ -188,4 +190,22 @@ describe('runWorkerQueueThreads', () => { expect(results).toEqual(['foo', 'foo', 'foo', 'foo']); }); + + it('should send messages', async () => { + const messages = new Array(); + + await runWorkerThreads({ + threadCount: 2, + worker: async (_data, sendMessage) => { + sendMessage('foo'); + await new Promise(resolve => setTimeout(resolve, 50)); + sendMessage('bar'); + await new Promise(resolve => setTimeout(resolve, 50)); + sendMessage('baz'); + }, + onMessage: (message: string) => messages.push(message), + }); + + expect(messages).toEqual(['foo', 'foo', 'bar', 'bar', 'baz', 'baz']); + }); }); diff --git a/packages/cli/src/lib/parallel.ts b/packages/cli/src/lib/parallel.ts index 3e9479407d..246020fcc3 100644 --- a/packages/cli/src/lib/parallel.ts +++ b/packages/cli/src/lib/parallel.ts @@ -112,54 +112,13 @@ type WorkerThreadMessage = | { type: 'error'; error: ErrorLike; + } + | { + type: 'message'; + message: unknown; }; -function workerQueueThread( - workerFuncFactory: (data: unknown) => (item: unknown) => Promise, -) { - const { parentPort, workerData } = require('worker_threads'); - const workerFunc = workerFuncFactory(workerData); - - parentPort.on('message', async (message: WorkerThreadMessage) => { - if (message.type === 'done') { - parentPort.close(); - return; - } - if (message.type === 'item') { - try { - const result = await workerFunc(message.item); - parentPort.postMessage({ - type: 'result', - index: message.index, - result, - }); - } catch (error) { - parentPort.postMessage({ type: 'error', error }); - } - } - }); - - parentPort.postMessage({ type: 'start' }); -} - -function workerThread(workerFunc: (data: unknown) => Promise) { - const { parentPort, workerData } = require('worker_threads'); - - workerFunc(workerData).then( - result => { - parentPort.postMessage({ - type: 'result', - index: 0, - result, - }); - }, - error => { - parentPort.postMessage({ type: 'error', error }); - }, - ); -} - -type WorkerQueueThreadsOptions = { +export type WorkerQueueThreadsOptions = { /** The items to process */ items: Iterable; /** @@ -203,13 +162,10 @@ export async function runWorkerQueueThreads( Array(threadCount) .fill(0) .map(async () => { - const thread = new Worker( - `(${workerQueueThread})((${workerFactory}))`, - { - eval: true, - workerData, - }, - ); + const thread = new Worker(`(${workerQueueThread})(${workerFactory})`, { + eval: true, + workerData, + }); return new Promise((resolve, reject) => { thread.on('message', (message: WorkerThreadMessage) => { @@ -251,7 +207,35 @@ export async function runWorkerQueueThreads( return results; } -type WorkerThreadsOptions = { +function workerQueueThread( + workerFuncFactory: (data: unknown) => (item: unknown) => Promise, +) { + const { parentPort, workerData } = require('worker_threads'); + const workerFunc = workerFuncFactory(workerData); + + parentPort.on('message', async (message: WorkerThreadMessage) => { + if (message.type === 'done') { + parentPort.close(); + return; + } + if (message.type === 'item') { + try { + const result = await workerFunc(message.item); + parentPort.postMessage({ + type: 'result', + index: message.index, + result, + }); + } catch (error) { + parentPort.postMessage({ type: 'error', error }); + } + } + }); + + parentPort.postMessage({ type: 'start' }); +} + +export type WorkerThreadsOptions = { /** * A function that is called by each worker thread to produce a result. * @@ -264,26 +248,31 @@ type WorkerThreadsOptions = { * note that they are both copied by value into the worker thread, except for * types that are explicitly shareable across threads, such as `SharedArrayBuffer`. */ - worker: (data: TData) => Promise; + worker: ( + data: TData, + sendMessage: (message: TMessage) => void, + ) => Promise; /** Data supplied to each worker */ workerData?: TData; /** Number of threads, defaults to 1 */ threadCount?: number; + /** An optional handler for messages posted from the worker thread */ + onMessage?: (message: TMessage) => void; }; /** * Spawns one or more worker threads using the `worker_threads` module. */ -export async function runWorkerThreads( - options: WorkerThreadsOptions, +export async function runWorkerThreads( + options: WorkerThreadsOptions, ): Promise { - const { worker, workerData, threadCount = 1 } = options; + const { worker, workerData, threadCount = 1, onMessage } = options; return Promise.all( Array(threadCount) .fill(0) .map(async () => { - const thread = new Worker(`(${workerThread})((${worker}))`, { + const thread = new Worker(`(${workerThread})(${worker})`, { eval: true, workerData, }); @@ -294,6 +283,8 @@ export async function runWorkerThreads( resolve(message.result as TResult); } else if (message.type === 'error') { reject(message.error); + } else if (message.type === 'message') { + onMessage?.(message.message as TMessage); } }); @@ -307,3 +298,29 @@ export async function runWorkerThreads( }), ); } + +function workerThread( + workerFunc: ( + data: unknown, + sendMessage: (message: unknown) => void, + ) => Promise, +) { + const { parentPort, workerData } = require('worker_threads'); + + const sendMessage = (message: unknown) => { + parentPort.postMessage({ type: 'message', message }); + }; + + workerFunc(workerData, sendMessage).then( + result => { + parentPort.postMessage({ + type: 'result', + index: 0, + result, + }); + }, + error => { + parentPort.postMessage({ type: 'error', error }); + }, + ); +}