diff --git a/doc/api/stream_iter.md b/doc/api/stream_iter.md index 9fafdc4b62c3..146b09fd0900 100644 --- a/doc/api/stream_iter.md +++ b/doc/api/stream_iter.md @@ -431,14 +431,16 @@ the write. Use [`ondrain()`][] to wait for capacity rather than polling. the pending `end()` call; it does not fail the writer itself. * Returns: {Promise} Fulfills with the total number of bytes written. -Signal that no more data will be written. +Signals that no more data will be written and waits for buffered data to drain. #### `writer.endSync()` -* Returns: {number} Total bytes written, or `-1` if the writer is not open. +* Returns: {number} Total bytes written, or `-1` if ending cannot complete + synchronously. -Synchronous variant of `writer.end()`. Returns `-1` if the writer is already -closed or errored. Can be used as a try-fallback pattern: +Synchronous variant of `writer.end()`. A return value of `-1` means closing has +started but requires asynchronous draining. Use the try-fallback pattern to +await completion: ```cjs const result = writer.endSync(); diff --git a/lib/internal/streams/iter/broadcast.js b/lib/internal/streams/iter/broadcast.js index a1384c9e4d51..d8cba9b0141b 100644 --- a/lib/internal/streams/iter/broadcast.js +++ b/lib/internal/streams/iter/broadcast.js @@ -14,6 +14,8 @@ const { PromiseReject, PromiseResolve, PromiseWithResolvers, + SafePromisePrototypeFinally, + SafePromiseRace, SafeSet, Symbol, SymbolAsyncDispose, @@ -77,6 +79,22 @@ const kEnd = Symbol('kEnd'); const kAbort = Symbol('kAbort'); const kCanWrite = Symbol('kCanWrite'); const kOnBufferDrained = Symbol('kOnBufferDrained'); +const kOnEndDrained = Symbol('kOnEndDrained'); +const kPendingWriteRemoved = Symbol('kPendingWriteRemoved'); + +function raceEndWithSignal(promise, signal) { + if (!signal) return promise; + + const { promise: aborted, reject } = PromiseWithResolvers(); + const onAbort = () => reject(signal.reason); + signal.addEventListener('abort', onAbort, { __proto__: null, once: true }); + if (signal.aborted) onAbort(); + + return SafePromisePrototypeFinally( + SafePromiseRace([promise, aborted]), + () => signal.removeEventListener('abort', onAbort), + ); +} // ============================================================================= // Broadcast Implementation @@ -100,6 +118,7 @@ class BroadcastImpl { constructor(options) { this.#options = options; this[kOnBufferDrained] = null; + this[kOnEndDrained] = null; } setWriter(writer) { @@ -168,6 +187,7 @@ class BroadcastImpl { if (self.#deleteConsumer(state)) { self.#tryTrimBuffer(); } + self.#notifyEndDrained(); } return { @@ -343,10 +363,11 @@ class BroadcastImpl { } } } + this.#notifyEndDrained(); } [kAbort](reason) { - if (this.#ended || this.#error !== undefined) return; + if (this.#error !== undefined) return; this.#error = reason; this.#ended = true; @@ -381,6 +402,12 @@ class BroadcastImpl { // Private methods + #notifyEndDrained() { + if (this.#ended && this.#consumers.size === 0) { + this[kOnEndDrained]?.(); + } + } + #recomputeMinCursor() { const { minCursor, minCursorConsumers } = getMinCursor( this.#consumers, this.#bufferStart + this.#buffer.length); @@ -501,8 +528,9 @@ let getBroadcastPendingWrites; class BroadcastWriter { #broadcast; #totalBytes = 0; - #closed; - #aborted = false; + #state = 'open'; + #error; + #pendingEnd; #pendingWrites = new RingBuffer(); #pendingDrains = []; @@ -517,8 +545,11 @@ class BroadcastWriter { this.#broadcast[kOnBufferDrained] = () => { this.#resolvePendingWrites(); - this.#resolvePendingDrains(true); + if (this.#state === 'open') { + this.#resolvePendingDrains(true); + } }; + this.#broadcast[kOnEndDrained] = () => this.#endDrained(); } // The drainable protocol works with Stream.ondrain to provide a notification @@ -532,20 +563,12 @@ class BroadcastWriter { return promise; } - #isClosed() { - return this.#closed !== undefined; - } - - #isClosedOrAborted() { - return this.#isClosed() || this.#aborted; - } - get canWrite() { - return this.#isClosedOrAborted() ? null : this.#broadcast[kCanWrite](); + return this.#state === 'open' ? this.#broadcast[kCanWrite]() : null; } #canUseWriteFastPath(signal) { - return !signal && !this.#isClosed() && !this.#aborted && + return !signal && this.#state === 'open' && this.#broadcast[kCanWrite](); } @@ -577,13 +600,15 @@ class BroadcastWriter { } async #writevSlow(chunks, signal) { - // Check for pre-aborted - signal?.throwIfAborted(); - - if (this.#isClosedOrAborted()) { + if (this.#state === 'errored') { + throw this.#error; + } + if (this.#state !== 'open') { throw new ERR_INVALID_STATE.TypeError('Writer is closed'); } + signal?.throwIfAborted(); + const converted = convertChunks(chunks); if (this.#broadcast[kWrite](converted)) { @@ -609,7 +634,7 @@ class BroadcastWriter { } writeSync(chunk) { - if (this.#isClosedOrAborted()) return false; + if (this.#state !== 'open') return false; if (!this.#broadcast[kCanWrite]()) return false; const converted = toUint8Array(chunk); @@ -622,7 +647,7 @@ class BroadcastWriter { writevSync(chunks) { validateArray(chunks, 'chunks'); - if (this.#isClosedOrAborted()) return false; + if (this.#state !== 'open') return false; if (!this.#broadcast[kCanWrite]()) return false; const converted = convertChunks(chunks); if (this.#broadcast[kWrite](converted)) { @@ -636,34 +661,43 @@ class BroadcastWriter { end(options) { const signal = getWriterSignal(options); + if (this.#state === 'errored') return PromiseReject(this.#error); + if (this.#state === 'closed') return PromiseResolve(this.#totalBytes); if (signal?.aborted) return PromiseReject(signal.reason); - if (this.#isClosed()) return this.#closed; - this.#closed = PromiseResolve(this.#totalBytes); - this.#broadcast[kEnd](); - this.#resolvePendingDrains(false); - return this.#closed; + const endPromise = this.#getEndPromise(); + if (this.#state === 'open') { + this.#state = 'closing'; + this.#resolvePendingDrains(false); + this.#finishEndIfReady(); + } + + return raceEndWithSignal(endPromise, signal); } endSync() { - if (this.#closed) return this.#totalBytes; - this.#closed = PromiseResolve(this.#totalBytes); - this.#broadcast[kEnd](); + if (this.#state === 'closed') return this.#totalBytes; + if (this.#state === 'errored' || this.#state === 'closing') return -1; + + this.#state = 'closing'; this.#resolvePendingDrains(false); - return this.#totalBytes; + this.#finishEndIfReady(); + return this.#state === 'closed' ? this.#totalBytes : -1; } fail(reason) { - if (this.#isClosedOrAborted()) return; - this.#aborted = true; - this.#closed = PromiseResolve(this.#totalBytes); + if (this.#state === 'errored' || this.#state === 'closed') return; + this.#state = 'errored'; const error = reason ?? new ERR_INVALID_STATE.TypeError('Failed'); + this.#error = error; this.#rejectPendingWrites(error); this.#rejectPendingDrains(error); + this.#pendingEnd?.reject(error); this.#broadcast[kAbort](error); } [SymbolAsyncDispose]() { + if (this.#state === 'closing') return this.#getEndPromise(); this.fail(); return PromiseResolve(); } @@ -673,11 +707,33 @@ class BroadcastWriter { } [kCancelWriter]() { - if (this.#isClosed()) return; - this.#closed = PromiseResolve(this.#totalBytes); + if (this.#state === 'closed' || this.#state === 'errored') return; + this.#state = 'closed'; this.#rejectPendingWrites( lazyDOMException('Broadcast cancelled', 'AbortError')); this.#resolvePendingDrains(false); + this.#pendingEnd?.resolve(this.#totalBytes); + } + + #getEndPromise() { + this.#pendingEnd ??= PromiseWithResolvers(); + return this.#pendingEnd.promise; + } + + #finishEndIfReady() { + if (this.#state === 'closing' && this.#pendingWrites.length === 0) { + this.#broadcast[kEnd](); + } + } + + #endDrained() { + if (this.#state !== 'closing') return; + this.#state = 'closed'; + this.#pendingEnd?.resolve(this.#totalBytes); + } + + [kPendingWriteRemoved]() { + this.#finishEndIfReady(); } /** @@ -709,6 +765,7 @@ class BroadcastWriter { break; } } + this.#finishEndIfReady(); } #rejectPendingWrites(error) { @@ -741,6 +798,7 @@ function wireBroadcastWriteSignal(entry, signal, resolve, reject, self) { if (idx !== -1) pendingWrites.removeAt(idx); entry.chunk = null; reject(signal.reason ?? lazyDOMException('Aborted', 'AbortError')); + if (idx !== -1) self[kPendingWriteRemoved](); }; entry.resolve = function() { signal.removeEventListener('abort', onAbort); diff --git a/test/parallel/test-stream-iter-broadcast-backpressure.js b/test/parallel/test-stream-iter-broadcast-backpressure.js index 35a0cf0e238c..698f1821aeea 100644 --- a/test/parallel/test-stream-iter-broadcast-backpressure.js +++ b/test/parallel/test-stream-iter-broadcast-backpressure.js @@ -135,15 +135,99 @@ async function testStrictBackpressureOverflow() { }); } +async function testEndDrainsPendingWrite() { + const chunk1 = new Uint8Array(16384).fill(65); // 'A' + const chunk2 = Uint8Array.of(66); // 'B' + const { writer, broadcast: bc } = broadcast({ + budget: 16384, + backpressure: 'unbounded', + }); + const iter = bc.push()[Symbol.asyncIterator](); + + await writer.write(chunk1); + const pendingWrite = writer.write(chunk2); + const endPromise = writer.end(); + + assert.strictEqual(writer.canWrite, null); + assert.strictEqual(writer.writeSync('late'), false); + await assert.rejects(writer.write('late'), { + code: 'ERR_INVALID_STATE', + }); + + const first = await iter.next(); + assert.strictEqual(first.done, false); + assert.strictEqual(first.value[0][0], 65); + await pendingWrite; + + const second = await iter.next(); + assert.strictEqual(second.done, false); + assert.strictEqual(second.value[0][0], 66); + + let endResolved = false; + endPromise.then(common.mustCall(() => { endResolved = true; })); + await new Promise(setImmediate); + assert.strictEqual(endResolved, false); + + assert.strictEqual((await iter.next()).done, true); + assert.strictEqual(await endPromise, 16385); +} + +async function testEndSyncDrainsPendingWrite() { + const chunk1 = new Uint8Array(16384).fill(65); // 'A' + const chunk2 = Uint8Array.of(66); // 'B' + const { writer, broadcast: bc } = broadcast({ + budget: 16384, + backpressure: 'unbounded', + }); + const iter = bc.push()[Symbol.asyncIterator](); + + await writer.write(chunk1); + const pendingWrite = writer.write(chunk2); + assert.strictEqual(writer.endSync(), -1); + const endPromise = writer.end(); + + assert.strictEqual((await iter.next()).value[0][0], 65); + await pendingWrite; + assert.strictEqual((await iter.next()).value[0][0], 66); + assert.strictEqual((await iter.next()).done, true); + assert.strictEqual(await endPromise, 16385); + assert.strictEqual(writer.endSync(), 16385); +} + +async function testAbortedPendingWriteAllowsEnd() { + const ac = new AbortController(); + const reason = new Error('write aborted'); + const { writer, broadcast: bc } = broadcast({ + budget: 16384, + backpressure: 'unbounded', + }); + const iter = bc.push()[Symbol.asyncIterator](); + + await writer.write(new Uint8Array(16384)); + const pendingWrite = writer.write('blocked', { signal: ac.signal }); + const writeRejected = assert.rejects( + pendingWrite, + (error) => error === reason, + ); + const endPromise = writer.end(); + + ac.abort(reason); + await writeRejected; + assert.strictEqual((await iter.next()).done, false); + assert.strictEqual((await iter.next()).done, true); + assert.strictEqual(await endPromise, 16384); +} + // Writev async path async function testWritevAsync() { const { writer, broadcast: bc } = broadcast({ budget: 16384 }); const consumer = bc.push(); await writer.writev(['hello', ' ', 'world']); + const dataPromise = text(consumer); await writer.end(); - const data = await text(consumer); + const data = await dataPromise; assert.strictEqual(data, 'hello world'); } @@ -167,15 +251,19 @@ async function testZeroByteWrites() { assert.strictEqual(entries, 0); } -// endSync returns the total byte count +// endSync falls back to end() when consumers still need to drain. async function testEndSyncReturnValue() { const { writer, broadcast: bc } = broadcast({ budget: 16384 }); - bc.push(); // Need a consumer to write to + const consumer = bc.push(); writer.writeSync('hello'); // 5 bytes writer.writeSync(' world'); // 6 bytes - const total = writer.endSync(); - assert.strictEqual(total, 11); + assert.strictEqual(writer.endSync(), -1); + + const dataPromise = text(consumer); + assert.strictEqual(await writer.end(), 11); + assert.strictEqual(await dataPromise, 'hello world'); + assert.strictEqual(writer.endSync(), 11); } Promise.all([ @@ -184,6 +272,9 @@ Promise.all([ testBlockBackpressure(), testBlockBackpressureContent(), testStrictBackpressureOverflow(), + testEndDrainsPendingWrite(), + testEndSyncDrainsPendingWrite(), + testAbortedPendingWriteAllowsEnd(), testWritevAsync(), testZeroByteWrites(), testEndSyncReturnValue(), diff --git a/test/parallel/test-stream-iter-broadcast-basic.js b/test/parallel/test-stream-iter-broadcast-basic.js index 125f386210c9..d534c9acf3a7 100644 --- a/test/parallel/test-stream-iter-broadcast-basic.js +++ b/test/parallel/test-stream-iter-broadcast-basic.js @@ -19,13 +19,14 @@ async function testBasicBroadcast() { assert.strictEqual(bc.consumerCount, 2); - await writer.write('hello'); - await writer.end(); - - const [data1, data2] = await Promise.all([ + const dataPromise = Promise.all([ text(consumer1), text(consumer2), ]); + await writer.write('hello'); + await writer.end(); + + const [data1, data2] = await dataPromise; assert.strictEqual(data1, 'hello'); assert.strictEqual(data2, 'hello'); @@ -39,9 +40,10 @@ async function testMultipleWrites() { await writer.write('a'); await writer.write('b'); await writer.write('c'); + const dataPromise = text(consumer); await writer.end(); - const data = await text(consumer); + const data = await dataPromise; assert.strictEqual(data, 'abc'); } @@ -104,10 +106,11 @@ async function testWriterEnd() { const consumer = bc.push(); await writer.write('data'); + const dataPromise = text(consumer); const totalBytes = await writer.end(); assert.strictEqual(totalBytes, 4); // 'data' = 4 UTF-8 bytes - const data = await text(consumer); + const data = await dataPromise; assert.strictEqual(data, 'data'); } @@ -123,8 +126,59 @@ async function testWriterEndWithPreAbortedSignal() { // A rejected end must leave the writer open. await writer.write('data'); + const dataPromise = text(consumer); + assert.strictEqual(await writer.end(), 4); + assert.strictEqual(await dataPromise, 'data'); +} + +async function testWriterEndWaitsForAllConsumers() { + const { writer, broadcast: bc } = broadcast(); + const iter1 = bc.push()[Symbol.asyncIterator](); + const iter2 = bc.push()[Symbol.asyncIterator](); + + await writer.write('data'); + const endPromise = writer.end(); + let endResolved = false; + endPromise.then(common.mustCall(() => { endResolved = true; })); + + assert.strictEqual((await iter1.next()).done, false); + assert.strictEqual((await iter2.next()).done, false); + assert.strictEqual((await iter1.next()).done, true); + assert.strictEqual(endResolved, false); + + assert.strictEqual((await iter2.next()).done, true); + assert.strictEqual(await endPromise, 4); +} + +async function testWriterEndSignalDoesNotFailWriter() { + const { writer, broadcast: bc } = broadcast(); + const consumer = bc.push(); + const ac = new AbortController(); + const reason = new Error('end aborted'); + + await writer.write('data'); + const signaledEnd = writer.end({ signal: ac.signal }); + const rejected = assert.rejects(signaledEnd, (error) => error === reason); + ac.abort(reason); + await rejected; + + const dataPromise = text(consumer); assert.strictEqual(await writer.end(), 4); - assert.strictEqual(await text(consumer), 'data'); + assert.strictEqual(await dataPromise, 'data'); +} + +async function testWriterFailWhileClosing() { + const { writer, broadcast: bc } = broadcast(); + const iter = bc.push()[Symbol.asyncIterator](); + const reason = new Error('writer failed while closing'); + + await writer.write('data'); + const endPromise = writer.end(); + const endRejected = assert.rejects(endPromise, (error) => error === reason); + writer.fail(reason); + + await endRejected; + await assert.rejects(iter.next(), (error) => error === reason); } async function testWriterFail() { @@ -325,6 +379,9 @@ Promise.all([ testWritevSync(), testWriterEnd(), testWriterEndWithPreAbortedSignal(), + testWriterEndWaitsForAllConsumers(), + testWriterEndSignalDoesNotFailWriter(), + testWriterFailWhileClosing(), testWriterFail(), testCancelWithoutReason(), testCancelWithReason(),