diff --git a/src/worker-entry.ts b/src/worker-entry.ts index 4e5b6b1..7072312 100644 --- a/src/worker-entry.ts +++ b/src/worker-entry.ts @@ -6,11 +6,26 @@ import { collectTransferables } from './transfer.ts'; const imported = new Set(); const taskCache = new Map(); +const fnCache = new Map unknown>(); function isTaskArg(arg: unknown): arg is { __task__: number; id: string; args: unknown[] } { return typeof arg === 'object' && arg !== null && '__task__' in arg; } +async function resolveFn(id: string): Promise<(...args: unknown[]) => unknown> { + const cached = fnCache.get(id); + if (cached) return cached; + const url = id.slice(0, id.lastIndexOf('#')); + if (!imported.has(url)) { + await import(url); + imported.add(url); + } + const fn = registry.get(id) as ((...args: unknown[]) => unknown) | undefined; + if (!fn) throw new Error(`Moroutine not found: ${id}`); + fnCache.set(id, fn); + return fn; +} + function portToAsyncIterable(port: MessagePort): AsyncIterable { const queue: T[] = []; let done = false; @@ -78,6 +93,14 @@ function portToAsyncIterable(port: MessagePort): AsyncIterable { }; } +function needsAsyncResolve(arg: unknown): boolean { + return arg instanceof MessagePort || isTaskArg(arg); +} + +function resolveArgs(args: unknown[]): unknown[] | Promise { + return args.some(needsAsyncResolve) ? Promise.all(args.map(resolveArg)) : args.map(deserializeArg); +} + async function resolveArg(arg: unknown): Promise { if (arg instanceof MessagePort) { return portToAsyncIterable(arg); @@ -86,16 +109,8 @@ async function resolveArg(arg: unknown): Promise { if (taskCache.has(arg.__task__)) { return taskCache.get(arg.__task__); } - // Resolve the task's own args recursively - const resolvedArgs = await Promise.all(arg.args.map(resolveArg)); - // Import the module and run the function - const url = arg.id.slice(0, arg.id.lastIndexOf('#')); - if (!imported.has(url)) { - await import(url); - imported.add(url); - } - const fn = registry.get(arg.id); - if (!fn) throw new Error(`Moroutine not found: ${arg.id}`); + const resolvedArgs = await resolveArgs(arg.args); + const fn = await resolveFn(arg.id); const value = await fn(...resolvedArgs); taskCache.set(arg.__task__, value); return value; @@ -106,16 +121,9 @@ async function resolveArg(arg: unknown): Promise { parentPort!.on('message', async (msg: { callId?: number; id: string; args: unknown[]; port?: MessagePort }) => { const { id, args, port } = msg; try { - const url = id.slice(0, id.lastIndexOf('#')); - if (!imported.has(url)) { - await import(url); - imported.add(url); - } - - const fn = registry.get(id); - if (!fn) throw new Error(`Moroutine not found: ${id}`); + const fn = await resolveFn(id); - const resolvedArgs = await Promise.all(args.map(resolveArg)); + const resolvedArgs = await resolveArgs(args); if (port) { let paused = false;