diff --git a/src/active-counts.ts b/src/active-counts.ts index 3d6a9dd..b7942d2 100644 --- a/src/active-counts.ts +++ b/src/active-counts.ts @@ -10,9 +10,18 @@ export class ActiveCounts { readonly #view: Int32Array; constructor(sizeOrBuffer: number | SharedArrayBuffer) { - this.#view = new Int32Array( - typeof sizeOrBuffer === 'number' ? new SharedArrayBuffer(sizeOrBuffer * STRIDE * 4) : sizeOrBuffer, - ); + if (typeof sizeOrBuffer === 'number') { + this.#view = new Int32Array(new SharedArrayBuffer(sizeOrBuffer * STRIDE * Int32Array.BYTES_PER_ELEMENT)); + return; + } + if (!(sizeOrBuffer instanceof SharedArrayBuffer)) { + throw new TypeError('ActiveCounts buffer must be a SharedArrayBuffer'); + } + const slotBytes = STRIDE * Int32Array.BYTES_PER_ELEMENT; + if (sizeOrBuffer.byteLength % slotBytes !== 0) { + throw new RangeError(`ActiveCounts buffer byteLength must be a multiple of ${slotBytes}`); + } + this.#view = new Int32Array(sizeOrBuffer); } get buffer(): SharedArrayBuffer { diff --git a/src/worker-pool.ts b/src/worker-pool.ts index 9568848..c9954e0 100644 --- a/src/worker-pool.ts +++ b/src/worker-pool.ts @@ -94,19 +94,20 @@ export function workers(sizeOrOpts?: number | WorkerOptions, opts?: WorkerOption } } - function resolveWorkerAndHandle(task: Task): { worker: Worker; handle: WorkerHandle; idx: number } { + function resolveWorker(task: Task): { worker: Worker; idx: number } { if (task.worker != null) { const idx = workerHandles.indexOf(task.worker); - if (idx !== -1) return { worker: pool[idx], handle: workerHandles[idx], idx }; + if (idx !== -1) return { worker: pool[idx], idx }; } const handle = balancer.select(workerHandles, task); const idx = workerHandles.indexOf(handle); - return { worker: pool[idx], handle, idx }; + if (idx === -1) throw new Error('Balancer returned a handle that is not in this pool'); + return { worker: pool[idx], idx }; } function dispatch(task: Task): Promise { if (disposed) return Promise.reject(new Error('Worker pool is disposed')); - const { worker, idx } = resolveWorkerAndHandle(task); + const { worker, idx } = resolveWorker(task); return trackValue(idx, execute(worker, task.id, task.args)); } @@ -137,7 +138,7 @@ export function workers(sizeOrOpts?: number | WorkerOptions, opts?: WorkerOption (taskOrTasks: Task | Task[] | (Task & AsyncIterable), channelOpts?: ChannelOptions): any => { if (taskOrTasks instanceof AsyncIterableTask) { if (disposed) throw new Error('Worker pool is disposed'); - const { worker, idx } = resolveWorkerAndHandle(taskOrTasks); + const { worker, idx } = resolveWorker(taskOrTasks); const { iterable, done } = dispatchStream(worker, taskOrTasks.id, taskOrTasks.args, channelOpts); trackStream(idx, done); return iterable; diff --git a/test/active-counts.test.ts b/test/active-counts.test.ts index 1ebd4dc..7985e9a 100644 --- a/test/active-counts.test.ts +++ b/test/active-counts.test.ts @@ -28,4 +28,19 @@ describe('ActiveCounts', () => { const counts = new ActiveCounts(2); assert.equal(counts.buffer.byteLength, 2 * 64); }); + + it('throws RangeError for an out-of-range index', () => { + const counts = new ActiveCounts(2); + assert.throws(() => counts.inc(99), RangeError); + }); + + it('throws RangeError for a buffer not a multiple of the 64-byte stride', () => { + const bad = new SharedArrayBuffer(63); + assert.throws(() => new ActiveCounts(bad), RangeError); + }); + + it('throws TypeError for a plain ArrayBuffer', () => { + const bad = new ArrayBuffer(64); + assert.throws(() => new ActiveCounts(bad as unknown as SharedArrayBuffer), TypeError); + }); }); diff --git a/test/load-balancing.test.ts b/test/load-balancing.test.ts index be25496..e87af5a 100644 --- a/test/load-balancing.test.ts +++ b/test/load-balancing.test.ts @@ -107,6 +107,20 @@ describe('load balancing', () => { assert.ok(disposed); }); + it('throws when a custom balancer returns a handle not in the pool', () => { + const foreign: Balancer = { + select() { + return { index: 0, activeCount: 0, exec: {} as any }; + }, + }; + const run = workers(1, { balance: foreign }); + try { + assert.throws(() => run(identity(42)), /not in this pool/); + } finally { + run[Symbol.dispose](); + } + }); + it('pinned tasks bypass balancer', async () => { let called = false; const custom: Balancer = {