diff --git a/packages/shared/src/iterator.test.ts b/packages/shared/src/iterator.test.ts index dff50c030..7cfafebee 100644 --- a/packages/shared/src/iterator.test.ts +++ b/packages/shared/src/iterator.test.ts @@ -263,8 +263,10 @@ describe('replicateAsyncIterator', async () => { expect(iterators.length).toBe(3) - expect(await iterators[0]!.next()).toEqual({ done: false, value: 1 }) - expect(await iterators[1]!.next()).toEqual({ done: false, value: 1 }) + await Promise.all([ + expect(iterators[0]!.next()).resolves.toEqual({ done: false, value: 1 }), + expect(iterators[1]!.next()).resolves.toEqual({ done: false, value: 1 }), + ]) expect(await iterators[0]!.next()).toEqual({ done: false, value: 2 }) expect(await iterators[1]!.next()).toEqual({ done: false, value: 2 }) @@ -295,9 +297,7 @@ describe('replicateAsyncIterator', async () => { const gen = async function* () { yield 1 - await new Promise(resolve => setTimeout(resolve, 1)) - yield 2 - yield 3 + await new Promise(resolve => setTimeout(resolve, 10)) throw error } @@ -308,28 +308,18 @@ describe('replicateAsyncIterator', async () => { expect(await iterators[0]!.next()).toEqual({ done: false, value: 1 }) expect(await iterators[1]!.next()).toEqual({ done: false, value: 1 }) - expect(await iterators[0]!.next()).toEqual({ done: false, value: 2 }) - expect(await iterators[1]!.next()).toEqual({ done: false, value: 2 }) - - expect(await iterators[0]!.next()).toEqual({ done: false, value: 3 }) - expect(await iterators[1]!.next()).toEqual({ done: false, value: 3 }) - expect(await iterators[2]!.next()).toEqual({ done: false, value: 1 }) - - await expect(iterators[0]!.next()).rejects.toThrow(error) - await expect(iterators[1]!.next()).rejects.toThrow(error) - expect(await iterators[2]!.next()).toEqual({ done: false, value: 2 }) + await Promise.all([ + expect(iterators[0]!.next()).rejects.toThrow(error), + expect(iterators[1]!.next()).rejects.toThrow(error), + ]) expect(await iterators[0]!.next()).toEqual({ done: true, value: undefined }) expect(await iterators[1]!.next()).toEqual({ done: true, value: undefined }) - expect(await iterators[2]!.next()).toEqual({ done: false, value: 3 }) + expect(await iterators[2]!.next()).toEqual({ done: false, value: 1 }) expect(await iterators[0]!.next()).toEqual({ done: true, value: undefined }) expect(await iterators[1]!.next()).toEqual({ done: true, value: undefined }) await expect(iterators[2]!.next()).rejects.toThrow(error) - - expect(await iterators[0]!.next()).toEqual({ done: true, value: undefined }) - expect(await iterators[1]!.next()).toEqual({ done: true, value: undefined }) - expect(await iterators[2]!.next()).toEqual({ done: true, value: undefined }) }) it('on manual close', async () => { diff --git a/packages/shared/src/iterator.ts b/packages/shared/src/iterator.ts index 64e146dee..8902d4314 100644 --- a/packages/shared/src/iterator.ts +++ b/packages/shared/src/iterator.ts @@ -1,5 +1,5 @@ import type { SetSpanErrorOptions } from './otel' -import { defer, once, sequential } from './function' +import { once, sequential } from './function' import { runInSpanContext, setSpanError, startSpan } from './otel' import { AsyncIdQueue } from './queue' @@ -112,9 +112,11 @@ export function replicateAsyncIterator( while (true) { const item = await source.next() - for (let id = 0; id < count; id++) { - if (queue.isOpen(id.toString())) { - queue.push(id.toString(), item) + for (let i = 0; i < count; i++) { + const id = i.toString() + + if (queue.isOpen(id)) { + queue.push(id, item) } } @@ -123,34 +125,39 @@ export function replicateAsyncIterator( } } } - catch (e) { - error = { value: e } + catch (reason) { + error = { value: reason } + + queue.waiterIds.forEach((id) => { + queue.close({ id, reason }) + }) } }) - for (let id = 0; id < count; id++) { - queue.open(id.toString()) + for (let i = 0; i < count; i++) { + const id = i.toString() + + queue.open(id) replicated.push(new AsyncIteratorClass( () => { start() return new Promise((resolve, reject) => { - queue.pull(id.toString()) - .then(resolve) - .catch(reject) - - defer(() => { - if (error) { - reject(error.value) - } - }) + if (!error || queue.hasBufferedItems(id)) { + queue.pull(id) + .then(resolve) + .catch(reject) + } + else { + reject(error.value) + } }) }, async (reason) => { - queue.close({ id: id.toString() }) + queue.close({ id }) if (reason !== 'next') { - if (replicated.every((_, id) => !queue.isOpen(id.toString()))) { + if (!queue.length) { await source?.return?.() } } diff --git a/packages/shared/src/queue.test.ts b/packages/shared/src/queue.test.ts index 1c99a6c16..5090f0c6e 100644 --- a/packages/shared/src/queue.test.ts +++ b/packages/shared/src/queue.test.ts @@ -132,4 +132,44 @@ describe('asyncIdQueue', () => { expect(queue.isOpen('2')).toBe(false) expect(queue.isOpen('3')).toBe(false) }) + + it('waiterIds', async () => { + queue.open('1') + queue.open('2') + + const p1 = queue.pull('1') + const p2 = queue.pull('2') + + expect(queue.waiterIds).toEqual(['1', '2']) + + queue.push('1', 'item1') + queue.push('2', 'item2') + + await expect(p1).resolves.toBe('item1') + await expect(p2).resolves.toBe('item2') + + expect(queue.waiterIds).toEqual([]) + }) + + it('hasBufferedItems', async () => { + queue.open('1') + queue.open('2') + + expect(queue.hasBufferedItems('1')).toBe(false) + expect(queue.hasBufferedItems('2')).toBe(false) + + queue.push('1', 'item1') + queue.push('2', 'item2') + + expect(queue.hasBufferedItems('1')).toBe(true) + expect(queue.hasBufferedItems('2')).toBe(true) + + await queue.pull('1') + expect(queue.hasBufferedItems('1')).toBe(false) + expect(queue.hasBufferedItems('2')).toBe(true) + + await queue.pull('2') + expect(queue.hasBufferedItems('1')).toBe(false) + expect(queue.hasBufferedItems('2')).toBe(false) + }) }) diff --git a/packages/shared/src/queue.ts b/packages/shared/src/queue.ts index 8d7cee15e..4f16c2566 100644 --- a/packages/shared/src/queue.ts +++ b/packages/shared/src/queue.ts @@ -5,13 +5,21 @@ export interface AsyncIdQueueCloseOptions { export class AsyncIdQueue { private readonly openIds = new Set() - private readonly items = new Map() - private readonly pendingPulls = new Map void, reject: (err: unknown) => void])[]>() + private readonly queues = new Map() + private readonly waiters = new Map void, reject: (err: unknown) => void])[]>() get length(): number { return this.openIds.size } + get waiterIds(): string[] { + return Array.from(this.waiters.keys()) + } + + hasBufferedItems(id: string): boolean { + return Boolean(this.queues.get(id)?.length) + } + open(id: string): void { this.openIds.add(id) } @@ -23,23 +31,23 @@ export class AsyncIdQueue { push(id: string, item: T): void { this.assertOpen(id) - const pending = this.pendingPulls.get(id) + const pending = this.waiters.get(id) if (pending?.length) { pending.shift()![0](item) if (pending.length === 0) { - this.pendingPulls.delete(id) + this.waiters.delete(id) } } else { - const items = this.items.get(id) + const items = this.queues.get(id) if (items) { items.push(item) } else { - this.items.set(id, [item]) + this.queues.set(id, [item]) } } } @@ -47,20 +55,20 @@ export class AsyncIdQueue { async pull(id: string): Promise { this.assertOpen(id) - const items = this.items.get(id) + const items = this.queues.get(id) if (items?.length) { const item = items.shift()! if (items.length === 0) { - this.items.delete(id) + this.queues.delete(id) } return item } return new Promise((resolve, reject) => { - const waitingPulls = this.pendingPulls.get(id) + const waitingPulls = this.waiters.get(id) const pending = [resolve, reject] as const @@ -68,32 +76,32 @@ export class AsyncIdQueue { waitingPulls.push(pending) } else { - this.pendingPulls.set(id, [pending]) + this.waiters.set(id, [pending]) } }) } close({ id, reason }: AsyncIdQueueCloseOptions = {}): void { if (id === undefined) { - this.pendingPulls.forEach((pendingPulls, id) => { + this.waiters.forEach((pendingPulls, id) => { pendingPulls.forEach(([, reject]) => { reject(reason ?? new Error(`[AsyncIdQueue] Queue[${id}] was closed or aborted while waiting for pulling.`)) }) }) - this.pendingPulls.clear() + this.waiters.clear() this.openIds.clear() - this.items.clear() + this.queues.clear() return } - this.pendingPulls.get(id)?.forEach(([, reject]) => { + this.waiters.get(id)?.forEach(([, reject]) => { reject(reason ?? new Error(`[AsyncIdQueue] Queue[${id}] was closed or aborted while waiting for pulling.`)) }) - this.pendingPulls.delete(id) + this.waiters.delete(id) this.openIds.delete(id) - this.items.delete(id) + this.queues.delete(id) } assertOpen(id: string): void { diff --git a/packages/standard-server/src/utils.test.ts b/packages/standard-server/src/utils.test.ts index 1bae68bc7..f3d402aa9 100644 --- a/packages/standard-server/src/utils.test.ts +++ b/packages/standard-server/src/utils.test.ts @@ -111,10 +111,15 @@ describe('replicateStandardLazyResponse', () => { replicateAsyncIteratorSpy.mockReturnValueOnce([1, 2, 3] as any) - expect(await replicated[0]!.body()).toBe(1) + // parallel test is important + await Promise.all([ + expect(replicated[0]!.body()).resolves.toEqual(1), + expect(replicated[1]!.body()).resolves.toEqual(2), + ]) + expect(await replicated[0]!.body()).toBe(1) // make sure cached - expect(await replicated[1]!.body()).toBe(2) expect(await replicated[1]!.body()).toBe(2) // make sure cached + expect(await replicated[2]!.body()).toBe(3) expect(await replicated[2]!.body()).toBe(3) // make sure cached diff --git a/packages/standard-server/src/utils.ts b/packages/standard-server/src/utils.ts index 7ea1dc1ea..8e667fc56 100644 --- a/packages/standard-server/src/utils.ts +++ b/packages/standard-server/src/utils.ts @@ -73,17 +73,13 @@ export function replicateStandardLazyResponse( replicated.push({ ...response, body: once(async () => { - if (replicatedAsyncIteratorObjects) { - return replicatedAsyncIteratorObjects.shift() - } - const body = await (bodyPromise ??= response.body()) if (!isAsyncIteratorObject(body)) { return body } - replicatedAsyncIteratorObjects = replicateAsyncIterator(body, count) + replicatedAsyncIteratorObjects ??= replicateAsyncIterator(body, count) return replicatedAsyncIteratorObjects.shift() }), })