diff --git a/src/transfer.ts b/src/transfer.ts index 7e499fe..a6073de 100644 --- a/src/transfer.ts +++ b/src/transfer.ts @@ -25,7 +25,7 @@ export function extractTransferables(args: unknown[]): { args: unknown[]; transf return { args: processedArgs, transfer: transferList }; } -function collectTransferables(value: unknown, list: Transferable[]): void { +export function collectTransferables(value: unknown, list: Transferable[]): void { if (!(typeof value === 'object' && value !== null)) return; // Directly transferable diff --git a/src/worker-entry.ts b/src/worker-entry.ts index 6bdab8b..4aafa70 100644 --- a/src/worker-entry.ts +++ b/src/worker-entry.ts @@ -1,7 +1,8 @@ import { parentPort } from 'node:worker_threads'; import { registry } from './registry.ts'; import { deserializeArg, serializeArg } from './sync/reconstruct.ts'; -import { extractTransferables } from './transfer.ts'; +import { collectTransferables } from './transfer.ts'; +import type { Transferable } from 'node:worker_threads'; const imported = new Set(); @@ -19,9 +20,10 @@ parentPort!.on('message', async (msg: { callId: number; id: string; args: unknow const deserializedArgs = args.map(deserializeArg); const value = await fn(...deserializedArgs); - const extracted = extractTransferables([value]); - const returnValue = serializeArg(extracted.args[0]); - parentPort!.postMessage({ callId, value: returnValue }, extracted.transfer); + const returnValue = serializeArg(value); + const transferList: Transferable[] = []; + collectTransferables(value, transferList); + parentPort!.postMessage({ callId, value: returnValue }, transferList); } catch (err) { const message = err instanceof Error ? err.message : String(err); parentPort!.postMessage({ callId, error: message }); diff --git a/test/fixtures/transfer.ts b/test/fixtures/transfer.ts index 7ea975e..9f535a4 100644 --- a/test/fixtures/transfer.ts +++ b/test/fixtures/transfer.ts @@ -16,3 +16,12 @@ export const sumUint8 = mo(import.meta, (arr: Uint8Array): number => { } return sum; }); + +export const makeBuffer = mo(import.meta, (size: number): ArrayBuffer => { + const buf = new ArrayBuffer(size); + const view = new Uint8Array(buf); + for (let i = 0; i < size; i++) { + view[i] = i + 1; + } + return buf; +}); diff --git a/test/transfer.test.ts b/test/transfer.test.ts index e42e6f5..fc4d489 100644 --- a/test/transfer.test.ts +++ b/test/transfer.test.ts @@ -1,7 +1,7 @@ import { describe, it } from 'node:test'; import assert from 'node:assert/strict'; import { workerPool, transfer } from 'moroutine'; -import { sumBuffer, sumUint8 } from './fixtures/transfer.ts'; +import { sumBuffer, sumUint8, makeBuffer } from './fixtures/transfer.ts'; describe('transfer', () => { it('transfers an ArrayBuffer to a worker (zero-copy)', async () => { @@ -35,6 +35,18 @@ describe('transfer', () => { } }); + it('auto-transfers return values from worker (zero-copy)', async () => { + const pool = workerPool(1); + try { + const buf = await pool(makeBuffer(4)); + const view = new Uint8Array(buf); + assert.equal(buf.byteLength, 4); + assert.deepEqual([...view], [1, 2, 3, 4]); + } finally { + pool[Symbol.dispose](); + } + }); + it('without transfer, buffer is copied (not detached)', async () => { const buf = new ArrayBuffer(4); const view = new Uint8Array(buf);