diff --git a/src/execute.ts b/src/execute.ts index 9a4060e..3480962 100644 --- a/src/execute.ts +++ b/src/execute.ts @@ -1,6 +1,7 @@ import type { Worker } from 'node:worker_threads'; import { freezeModule } from './registry.ts'; import { serializeArg, deserializeArg } from './sync/reconstruct.ts'; +import { extractTransferables } from './transfer.ts'; let nextCallId = 0; const pending = new Map void; reject: (reason: any) => void }>(); @@ -24,6 +25,8 @@ export function execute(worker: Worker, id: string, args: unknown[]): Promise const callId = nextCallId++; return new Promise((resolve, reject) => { pending.set(callId, { resolve, reject }); - worker.postMessage({ callId, id, args: args.map(serializeArg) }); + const extracted = extractTransferables(args); + const msg = { callId, id, args: extracted.args.map(serializeArg) }; + worker.postMessage(msg, extracted.transfer); }); } diff --git a/src/index.ts b/src/index.ts index 628fbd6..6015586 100644 --- a/src/index.ts +++ b/src/index.ts @@ -1,6 +1,7 @@ export { mo } from './mo.ts'; export { Task } from './task.ts'; export { workerPool } from './worker-pool.ts'; +export { transfer } from './transfer.ts'; export type { Runner } from './runner.ts'; export { AtomicBool, diff --git a/src/transfer.ts b/src/transfer.ts new file mode 100644 index 0000000..7e499fe --- /dev/null +++ b/src/transfer.ts @@ -0,0 +1,53 @@ +import type { Transferable } from 'node:worker_threads'; +import { MessagePort } from 'node:worker_threads'; + +const TRANSFER = Symbol.for('moroutine.transfer'); + +export interface Transferred { + readonly [TRANSFER]: true; + readonly value: T; +} + +export function transfer(value: T): Transferred { + return { [TRANSFER]: true as const, value }; +} + +export function extractTransferables(args: unknown[]): { args: unknown[]; transfer: Transferable[] } { + const transferList: Transferable[] = []; + const processedArgs = args.map((arg) => { + if (typeof arg === 'object' && arg !== null && TRANSFER in arg) { + const value = (arg as Transferred).value; + collectTransferables(value, transferList); + return value; + } + return arg; + }); + return { args: processedArgs, transfer: transferList }; +} + +function collectTransferables(value: unknown, list: Transferable[]): void { + if (!(typeof value === 'object' && value !== null)) return; + + // Directly transferable + if ( + value instanceof ArrayBuffer || + value instanceof MessagePort || + value instanceof ReadableStream || + value instanceof WritableStream + ) { + list.push(value); + return; + } + + // TypedArray or DataView — extract .buffer + if (ArrayBuffer.isView(value) && value.buffer instanceof ArrayBuffer) { + list.push(value.buffer); + return; + } + + // TransformStream — extract both streams + if (value instanceof TransformStream) { + list.push(value.readable, value.writable); + return; + } +} diff --git a/src/worker-entry.ts b/src/worker-entry.ts index 0d27f92..6bdab8b 100644 --- a/src/worker-entry.ts +++ b/src/worker-entry.ts @@ -1,6 +1,7 @@ import { parentPort } from 'node:worker_threads'; import { registry } from './registry.ts'; import { deserializeArg, serializeArg } from './sync/reconstruct.ts'; +import { extractTransferables } from './transfer.ts'; const imported = new Set(); @@ -18,7 +19,9 @@ parentPort!.on('message', async (msg: { callId: number; id: string; args: unknow const deserializedArgs = args.map(deserializeArg); const value = await fn(...deserializedArgs); - parentPort!.postMessage({ callId, value: serializeArg(value) }); + const extracted = extractTransferables([value]); + const returnValue = serializeArg(extracted.args[0]); + parentPort!.postMessage({ callId, value: returnValue }, extracted.transfer); } 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 new file mode 100644 index 0000000..7ea975e --- /dev/null +++ b/test/fixtures/transfer.ts @@ -0,0 +1,18 @@ +import { mo } from 'moroutine'; + +export const sumBuffer = mo(import.meta, (buf: ArrayBuffer): number => { + const view = new Uint8Array(buf); + let sum = 0; + for (let i = 0; i < view.length; i++) { + sum += view[i]; + } + return sum; +}); + +export const sumUint8 = mo(import.meta, (arr: Uint8Array): number => { + let sum = 0; + for (let i = 0; i < arr.length; i++) { + sum += arr[i]; + } + return sum; +}); diff --git a/test/transfer.test.ts b/test/transfer.test.ts new file mode 100644 index 0000000..e42e6f5 --- /dev/null +++ b/test/transfer.test.ts @@ -0,0 +1,53 @@ +import { describe, it } from 'node:test'; +import assert from 'node:assert/strict'; +import { workerPool, transfer } from 'moroutine'; +import { sumBuffer, sumUint8 } from './fixtures/transfer.ts'; + +describe('transfer', () => { + it('transfers an ArrayBuffer to a worker (zero-copy)', async () => { + const buf = new ArrayBuffer(4); + const view = new Uint8Array(buf); + view[0] = 1; view[1] = 2; view[2] = 3; view[3] = 4; + + const pool = workerPool(1); + try { + const sum = await pool(sumBuffer(transfer(buf))); + assert.equal(sum, 10); + // Original buffer should be detached after transfer + assert.equal(buf.byteLength, 0); + } finally { + pool[Symbol.dispose](); + } + }); + + it('transfers a Uint8Array to a worker via its buffer', async () => { + const arr = new Uint8Array([5, 10, 15, 20]); + const originalBuffer = arr.buffer; + + const pool = workerPool(1); + try { + const sum = await pool(sumUint8(transfer(arr))); + assert.equal(sum, 50); + // Underlying buffer should be detached after transfer + assert.equal(originalBuffer.byteLength, 0); + } finally { + pool[Symbol.dispose](); + } + }); + + it('without transfer, buffer is copied (not detached)', async () => { + const buf = new ArrayBuffer(4); + const view = new Uint8Array(buf); + view[0] = 1; view[1] = 2; view[2] = 3; view[3] = 4; + + const pool = workerPool(1); + try { + const sum = await pool(sumBuffer(buf)); + assert.equal(sum, 10); + // Original buffer should NOT be detached (it was copied) + assert.equal(buf.byteLength, 4); + } finally { + pool[Symbol.dispose](); + } + }); +});