diff --git a/packages/kernel-browser-runtime/src/internal-comms/internal-connections.test.ts b/packages/kernel-browser-runtime/src/internal-comms/internal-connections.test.ts index 0040658412..2c734c921e 100644 --- a/packages/kernel-browser-runtime/src/internal-comms/internal-connections.test.ts +++ b/packages/kernel-browser-runtime/src/internal-comms/internal-connections.test.ts @@ -39,7 +39,7 @@ vi.mock('@metamask/streams/browser', async () => { messageTarget: MockPostMessageTarget; constructor({ onEnd, messageTarget }: MockStreamOptions) { - super(() => undefined, { readerOnEnd: onEnd, writerOnEnd: onEnd }); + super(() => undefined, { onEnd }); MockStream.instances.push(this); this.messageTarget = messageTarget; this.messageTarget.onmessage = (event) => { diff --git a/packages/ocap-kernel/src/vats/VatSupervisor.test.ts b/packages/ocap-kernel/src/vats/VatSupervisor.test.ts index 2c103fd153..1f1ff2518b 100644 --- a/packages/ocap-kernel/src/vats/VatSupervisor.test.ts +++ b/packages/ocap-kernel/src/vats/VatSupervisor.test.ts @@ -40,7 +40,7 @@ const makeVatSupervisor = async ({ platformOptions, makeAllowedGlobals, fetchBlob, - writerOnEnd, + onEnd, }: { dispatch?: (input: unknown) => void | Promise; logger?: Logger; @@ -49,7 +49,7 @@ const makeVatSupervisor = async ({ platformOptions?: Record; makeAllowedGlobals?: (options: { logger: Logger }) => VatEndowments; fetchBlob?: FetchBlob; - writerOnEnd?: () => void; + onEnd?: () => void; } = {}): Promise<{ supervisor: VatSupervisor; stream: TestDuplexStream; @@ -59,7 +59,7 @@ const makeVatSupervisor = async ({ JsonRpcMessage >(dispatch ?? (() => undefined), { validateInput: isJsonRpcMessage, - writerOnEnd, + onEnd, }); // Provide a default makePlatform if none is specified @@ -160,20 +160,20 @@ describe('VatSupervisor', () => { it('calls the endowments teardown before closing the stream', async () => { // The stream is hardened, so we can't vi.spyOn(stream, 'end'). Instead, - // observe the writer's onEnd callback, which fires as part of stream.end(). + // observe the stream's onEnd callback, which fires as part of stream.end(). const teardown = vi.fn().mockResolvedValue(undefined); - const writerOnEnd = vi.fn(); + const onEnd = vi.fn(); const { supervisor } = await makeVatSupervisor({ makeAllowedGlobals: () => makeVatEndowments({}, teardown), - writerOnEnd, + onEnd, }); await supervisor.terminate(); expect(teardown).toHaveBeenCalledTimes(1); - expect(writerOnEnd).toHaveBeenCalledTimes(1); + expect(onEnd).toHaveBeenCalledTimes(1); expect(teardown.mock.invocationCallOrder[0]).toBeLessThan( - writerOnEnd.mock.invocationCallOrder[0] as number, + onEnd.mock.invocationCallOrder[0] as number, ); }); diff --git a/packages/streams/CHANGELOG.md b/packages/streams/CHANGELOG.md index 632a47a708..6bff5c2e6e 100644 --- a/packages/streams/CHANGELOG.md +++ b/packages/streams/CHANGELOG.md @@ -7,6 +7,20 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- Export `BaseReader` and `BaseWriter` for one-way streams over any transport ([#1138](https://github.com/Consensys-Incorporated/ocap-kernel/pull/1138)) + +### Changed + +- **BREAKING:** `NodePort` requires an `off` method, which `NodeWorkerDuplexStream` uses to remove its listener when the stream ends ([#1138](https://github.com/Consensys-Incorporated/ocap-kernel/pull/1138)) +- `split` accepts any number of predicates and narrows the type of each split ([#1138](https://github.com/Consensys-Incorporated/ocap-kernel/pull/1138)) +- The `ChromeRuntimeDuplexStream` constructor throws if `localTarget` and `remoteTarget` are the same, as `make()` already did ([#1138](https://github.com/Consensys-Incorporated/ocap-kernel/pull/1138)) + +### Removed + +- **BREAKING:** Remove the `MessagePortReader`, `MessagePortWriter`, `PostMessageReader`, `PostMessageWriter`, `ChromeRuntimeReader`, `ChromeRuntimeWriter`, `NodeWorkerReader`, and `NodeWorkerWriter` exports; use the corresponding duplex streams instead ([#1138](https://github.com/Consensys-Incorporated/ocap-kernel/pull/1138)) + ### Fixed - `PostMessageDuplexStream` calls `onEnd` once per stream, and writes after the remote side ends return a done result instead of throwing when `onEnd` closes the transport ([#1137](https://github.com/Consensys-Incorporated/ocap-kernel/pull/1137)) diff --git a/packages/streams/package.json b/packages/streams/package.json index d6c73fb323..ee90e89be5 100644 --- a/packages/streams/package.json +++ b/packages/streams/package.json @@ -69,7 +69,6 @@ }, "dependencies": { "@endo/promise-kit": "^1.2.1", - "@endo/stream": "^1.3.1", "@metamask/kernel-errors": "workspace:^", "@metamask/kernel-utils": "workspace:^", "@metamask/superstruct": "^3.4.1", diff --git a/packages/streams/src/BaseDuplexStream.test.ts b/packages/streams/src/BaseDuplexStream.test.ts index d3f7d4ede8..f0f9186bb2 100644 --- a/packages/streams/src/BaseDuplexStream.test.ts +++ b/packages/streams/src/BaseDuplexStream.test.ts @@ -167,8 +167,9 @@ describe('BaseDuplexStream', () => { throw new Error('foo'); }); - await expect(stream.synchronize()).rejects.toThrow('foo'); - await expect(stream.synchronize()).rejects.toThrow('foo'); + const message = 'TestDuplexStream experienced a dispatch failure'; + await expect(stream.synchronize()).rejects.toThrow(message); + await expect(stream.synchronize()).rejects.toThrow(message); }); }); @@ -343,53 +344,54 @@ describe('BaseDuplexStream', () => { await sink.return(); }); - it('return calls ends both the reader and writer', async () => { - const readerOnEnd = vi.fn(); - const writerOnEnd = vi.fn(); - const stream = await TestDuplexStream.make(() => undefined, { - readerOnEnd, - writerOnEnd, - }); - + it.each([ + ['returning', async (stream: TestDuplexStream) => stream.return()], + [ + 'throwing', + async (stream: TestDuplexStream) => stream.throw(new Error('foo')), + ], + [ + 'the remote ending', + async (stream: TestDuplexStream) => + stream.receiveInput(makeStreamDoneSignal()), + ], + ])('calls onEnd once after %s', async (_, endStream) => { + const onEnd = vi.fn(); + const stream = await TestDuplexStream.make(() => undefined, { onEnd }); + + await endStream(stream); await stream.return(); - expect(readerOnEnd).toHaveBeenCalledOnce(); - expect(writerOnEnd).toHaveBeenCalledOnce(); + expect(onEnd).toHaveBeenCalledOnce(); + expect(await stream.next()).toStrictEqual(makeDoneResult()); }); - it('throw calls throw on the writer but return on the reader', async () => { - const readerOnEnd = vi.fn(); - const writerOnEnd = vi.fn(); - const stream = await TestDuplexStream.make(() => undefined, { - readerOnEnd, - writerOnEnd, + it('ends and calls onEnd if a write fails', async () => { + const onDispatch = vi.fn(); + const onEnd = vi.fn(); + const stream = await TestDuplexStream.make(onDispatch, { onEnd }); + onDispatch.mockImplementation(() => { + throw new Error('foo'); }); - await stream.throw(new Error('foo')); - expect(readerOnEnd).toHaveBeenCalledOnce(); - expect(writerOnEnd).toHaveBeenCalledOnce(); + await expect(stream.write(42)).rejects.toThrow( + 'TestDuplexStream experienced a dispatch failure', + ); + expect(onEnd).toHaveBeenCalledOnce(); + expect(await stream.next()).toStrictEqual(makeDoneResult()); }); - it('ending the reader calls reader onEnd function', async () => { - const readerOnEnd = vi.fn(); - const stream = await TestDuplexStream.make(() => undefined, { - readerOnEnd, - }); + it('dispatches the done signal before calling onEnd', async () => { + const calls: string[] = []; + const stream = await TestDuplexStream.make( + (value) => { + calls.push(stringify(value)); + }, + { onEnd: () => calls.push('onEnd') }, + ); + calls.length = 0; await stream.receiveInput(makeStreamDoneSignal()); - expect(readerOnEnd).toHaveBeenCalledOnce(); - }); - - it('ending the writer calls writer onEnd function', async () => { - const onDispatch = vi.fn(() => { - throw new Error('foo'); - }); - const writerOnEnd = vi.fn(); - const stream = await TestDuplexStream.make(onDispatch, { - writerOnEnd, - }); - - await expect(stream.write(42)).rejects.toThrow('foo'); - expect(writerOnEnd).toHaveBeenCalledOnce(); + expect(calls).toStrictEqual([stringify(makeStreamDoneSignal()), 'onEnd']); }); describe('end', () => { diff --git a/packages/streams/src/BaseDuplexStream.ts b/packages/streams/src/BaseDuplexStream.ts index aed522fd3a..515f3c5d90 100644 --- a/packages/streams/src/BaseDuplexStream.ts +++ b/packages/streams/src/BaseDuplexStream.ts @@ -1,11 +1,12 @@ import type { PromiseKit } from '@endo/promise-kit'; import { makePromiseKit } from '@endo/promise-kit'; -import type { Reader } from '@endo/stream'; import { stringify } from '@metamask/kernel-utils'; import { is, literal, object } from '@metamask/superstruct'; import type { Infer } from '@metamask/superstruct'; -import type { BaseReader, BaseWriter, ValidateInput } from './BaseStream.ts'; +import { BaseReader, BaseWriter } from './BaseStream.ts'; +import type { Dispatch, Listen, OnEnd, ValidateInput } from './BaseStream.ts'; +import type { Reader } from './utils.ts'; import { makeDoneResult } from './utils.ts'; export const DuplexStreamSentinel = { @@ -46,19 +47,15 @@ export const isDuplexStreamSignal = ( ): value is DuplexStreamSignal => isSyn(value) || isAck(value); /** - * Make a validator for input to a duplex stream. Constructor helper for concrete - * duplex stream implementations. - * - * Validators passed in by consumers must be augmented such that errors aren't - * thrown for {@link DuplexStreamSignal} values. + * Augments a consumer-provided validator so that it accepts + * {@link DuplexStreamSignal} values. * * @param validateInput - The validator for the stream's input type. - * @returns A validator for the stream's input type, or `undefined` if no - * validation is desired. + * @returns The augmented validator, or `undefined` if none was provided. */ -export const makeDuplexStreamInputValidator = ( +const makeDuplexStreamInputValidator = ( validateInput?: ValidateInput, -): ((value: unknown) => value is Read) | undefined => +): ValidateInput | undefined => validateInput && ((value: unknown): value is Read => isDuplexStreamSignal(value) || validateInput(value)); @@ -77,26 +74,28 @@ const isEnded = (status: SynchronizationStatus): boolean => status === SynchronizationStatus.Complete || status === SynchronizationStatus.Failed; -/** - * The base of a duplex stream. Essentially a {@link BaseReader} with a `write()` method. - * Backed up by separate {@link BaseReader} and {@link BaseWriter} instances under the hood. - */ -export abstract class BaseDuplexStream< - Read, - ReadStream extends BaseReader, - Write = Read, - WriteStream extends BaseWriter = BaseWriter, -> implements Reader -{ +export type BaseDuplexStreamArgs = { + name: string; + listen: Listen; + onDispatch: Dispatch; + validateInput?: ValidateInput | undefined; /** - * The underlying reader for the duplex stream. + * Called once when the stream ends, after the final signal has been dispatched. + * For cleanup such as closing the transport. */ - readonly #reader: ReadStream; + onEnd?: OnEnd | undefined; +}; - /** - * The underlying writer for the duplex stream. - */ - readonly #writer: WriteStream; +/** + * The base of a duplex stream over some transport. Essentially a + * {@link BaseReader} with a `write()` method. Backed up by separate + * {@link BaseReader} and {@link BaseWriter} instances under the hood, each of + * which ends the other. + */ +export class BaseDuplexStream implements Reader { + readonly #reader: BaseReader; + + readonly #writer: BaseWriter; /** * The promise for the synchronization of the stream with its remote @@ -127,10 +126,39 @@ export abstract class BaseDuplexStream< /** * Constructs a new {@link BaseDuplexStream}. * - * @param reader - The underlying reader for the duplex stream. - * @param writer - The underlying writer for the duplex stream. + * @param options - Options bag for configuring the duplex stream. + * @param options.name - The name of the stream, for logging purposes. + * @param options.listen - Subscribes the stream to its transport. + * @param options.onDispatch - Dispatches messages over the transport. + * @param options.validateInput - A function that validates input from the transport. + * @param options.onEnd - A function that is called once when the stream ends. */ - constructor(reader: ReadStream, writer: WriteStream) { + constructor({ + name, + listen, + onDispatch, + validateInput, + onEnd, + }: BaseDuplexStreamArgs) { + // The writer ends last, so that its final signal is dispatched before onEnd. + const writer: BaseWriter = new BaseWriter({ + name, + onDispatch, + onEnd: async (error) => { + // eslint-disable-next-line @typescript-eslint/no-use-before-define + await reader.return(); + await onEnd?.(error); + }, + }); + const reader = new BaseReader({ + name, + listen, + validateInput: makeDuplexStreamInputValidator(validateInput), + onEnd: async () => { + await writer.return(); + }, + }); + // Set a catch handler to avoid unhandled rejection errors. The promise may // reject before reads or writes occur, in which case there are no handlers. this.#syncKit.promise.catch(() => undefined); @@ -348,7 +376,7 @@ harden(BaseDuplexStream); * A duplex stream. Essentially a {@link Reader} with a `write()` method. */ export type DuplexStream = Pick< - BaseDuplexStream, Write, BaseWriter>, + BaseDuplexStream, 'next' | 'write' | 'drain' | 'pipe' | 'return' | 'throw' | 'end' > & { [Symbol.asyncIterator]: () => DuplexStream; diff --git a/packages/streams/src/BaseStream.test.ts b/packages/streams/src/BaseStream.test.ts index a4d9adfbed..924143f72a 100644 --- a/packages/streams/src/BaseStream.test.ts +++ b/packages/streams/src/BaseStream.test.ts @@ -22,13 +22,6 @@ describe('BaseReader', () => { expect(reader[Symbol.asyncIterator]()).toBe(reader); }); - it('throws if getReceiveInput is called more than once', () => { - const reader = new TestReader(); - expect(() => reader.getReceiveInput()).toThrow( - 'TestReader received multiple calls to getReceiveInput()', - ); - }); - it('calls onEnd once when ending', async () => { const onEnd = vi.fn(); const reader = new TestReader({ onEnd }); @@ -329,32 +322,23 @@ describe('BaseWriter', () => { }); }); - it('handles repeated failures to dispatch messages', async () => { - const dispatchSpy = vi - .fn() - .mockImplementationOnce(() => { - throw new Error('foo'); - }) - .mockImplementationOnce(() => { - throw new Error('foo'); - }); - const writer = new TestWriter({ onDispatch: dispatchSpy }); + it('ends the stream if failing to dispatch the error signal', async () => { + const onEnd = vi.fn(); + const dispatchSpy = vi.fn(() => { + throw new Error('foo'); + }); + const writer = new TestWriter({ onDispatch: dispatchSpy, onEnd }); await expect(writer.next(42)).rejects.toThrow( - 'TestWriter experienced repeated dispatch failures.', - ); - expect(dispatchSpy).toHaveBeenCalledTimes(3); - expect(dispatchSpy).toHaveBeenNthCalledWith(1, 42); - expect(dispatchSpy).toHaveBeenNthCalledWith(2, { - [StreamSentinel.Error]: true, - error: makeErrorMatcher('foo'), - }); - expect(dispatchSpy).toHaveBeenNthCalledWith(3, { - [StreamSentinel.Error]: true, - error: makeErrorMatcher( - 'TestWriter experienced repeated dispatch failures.', + makeErrorMatcher( + new Error('TestWriter experienced a dispatch failure', { + cause: new Error('foo'), + }), ), - }); + ); + expect(dispatchSpy).toHaveBeenCalledTimes(2); + expect(onEnd).toHaveBeenCalledOnce(); + expect(await writer.next(43)).toStrictEqual(makeDoneResult()); }); }); diff --git a/packages/streams/src/BaseStream.ts b/packages/streams/src/BaseStream.ts index ab1f2ba9a7..b2185f6c5e 100644 --- a/packages/streams/src/BaseStream.ts +++ b/packages/streams/src/BaseStream.ts @@ -1,17 +1,15 @@ import { makePromiseKit } from '@endo/promise-kit'; -import type { Reader, Writer } from '@endo/stream'; import { stringify } from '@metamask/kernel-utils'; import type { PromiseCallbacks } from '@metamask/kernel-utils'; -import type { Dispatchable, Writable } from './utils.ts'; +import type { Dispatchable, Reader, Writer } from './utils.ts'; import { + isSignalLike, makeDoneResult, makePendingResult, makeStreamDoneSignal, makeStreamErrorSignal, - marshal, - StreamDoneSymbol, - unmarshal, + parseSignal, } from './utils.ts'; const makeStreamBuffer = < @@ -100,12 +98,21 @@ export type OnEnd = (error?: Error) => void | Promise; export type ValidateInput = (input: unknown) => input is Read; /** - * A function that receives input from a transport mechanism to a readable stream. - * Validates that the input is an {@link IteratorResult}, and throws if it is not. + * Forwards input from a transport to a reader. Never rejects; invalid input ends + * the reader with an error. */ export type ReceiveInput = (input: unknown) => Promise; +/** + * Subscribes a reader to its transport. May return a function that unsubscribes it, + * which is called when the reader ends. + */ +export type Listen = ( + receiveInput: (input: unknown) => void, +) => (() => void) | void; + export type BaseReaderArgs = { + listen: Listen; name?: string | undefined; onEnd?: OnEnd | undefined; validateInput?: ValidateInput | undefined; @@ -114,10 +121,6 @@ export type BaseReaderArgs = { /** * The base of a readable async iterator stream. * - * Subclasses must forward input received from the transport mechanism via the function - * returned by `getReceiveInput()`. Any cleanup required by subclasses should be performed - * in a callback passed to `setOnEnd()`. - * * The result of any value received before the stream ends is guaranteed to be observable * by the consumer. */ @@ -134,79 +137,53 @@ export class BaseReader implements Reader { #onEnd?: OnEnd | undefined; - #didExposeReceiveInput: boolean = false; - /** * Constructs a {@link BaseReader}. * * @param options - Options bag for configuring the reader. + * @param options.listen - Subscribes the reader to its transport. * @param options.name - The name of the stream, for logging purposes. Defaults to the class name. - * @param options.onEnd - A function that is called when the stream ends. For any cleanup that - * should happen when the stream ends, such as closing a message port. + * @param options.onEnd - A function that is called when the stream ends. * @param options.validateInput - A function that validates input from the transport. */ - constructor({ name, onEnd, validateInput }: BaseReaderArgs) { + constructor({ listen, name, onEnd, validateInput }: BaseReaderArgs) { this.#name = name ?? this.constructor.name; - this.#onEnd = onEnd; this.#validateInput = validateInput; + // eslint-disable-next-line @typescript-eslint/no-misused-promises -- Never rejects. + const unlisten = listen(this.#receiveInput); + this.#onEnd = async (error) => { + unlisten?.(); + await onEnd?.(error); + }; harden(this); } - /** - * Returns the `receiveInput()` method, which is used to receive input from the stream. - * Attempting to call this method more than once will throw an error. - * - * @returns The `receiveInput()` method. - */ - protected getReceiveInput(): ReceiveInput { - if (this.#didExposeReceiveInput) { - throw new Error( - `${this.#name} received multiple calls to getReceiveInput()`, - ); - } - this.#didExposeReceiveInput = true; - return this.#receiveInput.bind(this); - } - - readonly #receiveInput = async (input: unknown): Promise => { + readonly #receiveInput: ReceiveInput = async (input) => { // eslint-disable-next-line @typescript-eslint/await-thenable await null; - - const unmarshaled = unmarshal(input); - if (unmarshaled instanceof Error) { - await this.#handleInputError(unmarshaled); - return; - } - - if (unmarshaled === StreamDoneSymbol) { - await this.#end(); - return; - } - - if (this.#validateInput?.(unmarshaled) === false) { - await this.#handleInputError( - new Error( - `${this.#name}: Message failed type validation:\n${stringify(unmarshaled)}`, - ), - ); - return; + try { + if (isSignalLike(input)) { + const error = parseSignal(input); + if (error) { + throw error; + } + await this.#end(); + return; + } + if (this.#validateInput?.(input) === false) { + throw new Error( + `${this.#name}: Message failed type validation:\n${stringify(input)}`, + ); + } + this.#buffer.put(makePendingResult(input)); + } catch (error) { + if (!this.#buffer.hasPendingReads()) { + this.#buffer.put(error as Error); + } + await this.#end(error as Error).catch(() => undefined); } - - this.#buffer.put(makePendingResult(unmarshaled)); }; - /** - * Handles an input error by putting it into the buffer and ending the stream. - * - * @param error - The error to handle. - */ - async #handleInputError(error: Error): Promise { - if (!this.#buffer.hasPendingReads()) { - this.#buffer.put(error); - } - await this.#end(error); - } - /** * Ends the stream. Calls and then unsets the `#onEnd` method. * Idempotent. @@ -215,9 +192,9 @@ export class BaseReader implements Reader { */ async #end(error?: Error): Promise { this.#buffer.end(error); - const onEndP = this.#onEnd?.(error); + const onEnd = this.#onEnd; this.#onEnd = undefined; - await onEndP; + await onEnd?.(error); } /** @@ -288,7 +265,7 @@ export type BaseWriterArgs = { export class BaseWriter implements Writer { #isDone: boolean = false; - readonly #name: string = 'BaseWriter'; + readonly #name: string; readonly #onDispatch: Dispatch; @@ -300,8 +277,7 @@ export class BaseWriter implements Writer { * @param options - Options bag for configuring the writer. * @param options.onDispatch - A function that dispatches messages over the underlying transport mechanism. * @param options.name - The name of the stream, for logging purposes. Defaults to the class name. - * @param options.onEnd - A function that is called when the stream ends. For any cleanup that - * should happen when the stream ends, such as closing a message port. + * @param options.onEnd - A function that is called when the stream ends. */ constructor({ name, onDispatch, onEnd }: BaseWriterArgs) { this.#name = name ?? this.constructor.name; @@ -311,59 +287,25 @@ export class BaseWriter implements Writer { } /** - * Dispatches the value, via the dispatch function registered in the constructor. - * If dispatching fails, calls `#throw()`, and is therefore mutually recursive with - * that method. For this reason, includes a flag indicating past failure to dispatch - * a value, which is used to avoid infinite recursion. If dispatching succeeds, returns a - * `{ done: true }` result if the value was an {@link Error} or itself a `done` result, - * otherwise returns `{ done: false }`. - * - * @param value - The value to dispatch. - * @param hasFailed - Whether dispatching has failed previously. - * @returns The result of dispatching the value. - */ - async #dispatch( - value: Writable, - hasFailed = false, - ): Promise> { - try { - await this.#onDispatch(marshal(value)); - return value === StreamDoneSymbol || value instanceof Error - ? makeDoneResult() - : makePendingResult(undefined); - } catch (error) { - if (hasFailed) { - // Break out of repeated failure to dispatch an error. It is unclear how this would occur - // in practice, but it's the kind of failure mode where it's better to be sure. - const repeatedFailureError = new Error( - `${this.#name} experienced repeated dispatch failures.`, - { cause: error }, - ); - await this.#onDispatch(makeStreamErrorSignal(repeatedFailureError)); - throw repeatedFailureError; - } else { - await this.#throw( - /* istanbul ignore next: The ternary is mostly to please TypeScript */ - error instanceof Error ? error : new Error(String(error)), - true, - ); - throw new Error(`${this.#name} experienced a dispatch failure`, { - cause: error, - }); - } - } - } - - /** - * Ends the stream and calls the onEnd callback. Idempotent. + * Dispatches the final signal and calls `onEnd`. The writer ends even if + * either throws. Idempotent. * * @param error - The error to end the stream with. */ async #end(error?: Error): Promise { + if (this.#isDone) { + return; + } this.#isDone = true; - const onEndP = this.#onEnd?.(error); + const onEnd = this.#onEnd; this.#onEnd = undefined; - await onEndP; + try { + await this.#onDispatch( + error ? makeStreamErrorSignal(error) : makeStreamDoneSignal(), + ); + } finally { + await onEnd?.(error); + } } /** @@ -376,7 +318,8 @@ export class BaseWriter implements Writer { } /** - * Writes the next message to the transport. + * Writes the next message to the transport. If dispatching fails, forwards the + * failure to the transport (if possible) and ends the stream. * * @param value - The next message to write to the transport. * @returns The result of writing the message. @@ -385,24 +328,27 @@ export class BaseWriter implements Writer { if (this.#isDone) { return makeDoneResult(); } - return this.#dispatch(value); + try { + await this.#onDispatch(value); + } catch (cause) { + await this.#end( + /* istanbul ignore next: The ternary is mostly to please TypeScript */ + cause instanceof Error ? cause : new Error(String(cause)), + ).catch(() => undefined); + throw new Error(`${this.#name} experienced a dispatch failure`, { + cause, + }); + } + return makePendingResult(undefined); } /** * Closes the underlying transport and returns. Idempotent. - * The stream ends even if dispatching the done signal fails, in which case - * the dispatch error is rethrown. * * @returns The final result for this stream. */ async return(): Promise> { - if (!this.#isDone) { - try { - await this.#onDispatch(makeStreamDoneSignal()); - } finally { - await this.#end(); - } - } + await this.#end(); return makeDoneResult(); } @@ -413,9 +359,7 @@ export class BaseWriter implements Writer { * @returns The final result for this stream. */ async throw(error: Error): Promise> { - if (!this.#isDone) { - await this.#throw(error); - } + await this.#end(error); return makeDoneResult(); } @@ -428,25 +372,5 @@ export class BaseWriter implements Writer { async end(error?: Error): Promise> { return error ? this.throw(error) : this.return(); } - - /** - * Dispatches the error and calls `#end()`. Mutually recursive with `dispatch()`. - * For this reason, includes a flag indicating past failure, so that `dispatch()` - * can avoid infinite recursion. See `dispatch()` for more details. - * - * @param error - The error to forward. - * @param hasFailed - Whether dispatching has failed previously. - * @returns The final result for this stream. - */ - async #throw( - error: Error, - hasFailed = false, - ): Promise> { - const result = this.#dispatch(error, hasFailed); - if (!this.#isDone) { - await this.#end(error); - } - return result; - } } harden(BaseWriter); diff --git a/packages/streams/src/browser/ChromeRuntimeStream.test.ts b/packages/streams/src/browser/ChromeRuntimeStream.test.ts index 812352246d..77ad320556 100644 --- a/packages/streams/src/browser/ChromeRuntimeStream.test.ts +++ b/packages/streams/src/browser/ChromeRuntimeStream.test.ts @@ -6,11 +6,7 @@ import type { MessageEnvelope, ChromeRuntimeTarget, } from './ChromeRuntimeStream.ts'; -import { - ChromeRuntimeReader, - ChromeRuntimeWriter, - ChromeRuntimeDuplexStream, -} from './ChromeRuntimeStream.ts'; +import { ChromeRuntimeDuplexStream } from './ChromeRuntimeStream.ts'; import { makeAck } from '../BaseDuplexStream.ts'; import type { ValidateInput } from '../BaseStream.ts'; import { @@ -64,246 +60,9 @@ const asChromeRuntime = ( runtime: ReturnType['runtime'], ): ChromeRuntime => runtime as unknown as ChromeRuntime; -// TODO: Further investigation is needed to determine whether these tests -// can be run concurrently. -describe('ChromeRuntimeReader', () => { - it('constructs a ChromeRuntimeReader', () => { - const { runtime } = makeRuntime(); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - - expect(reader).toBeInstanceOf(ChromeRuntimeReader); - expect(reader[Symbol.asyncIterator]()).toBe(reader); - expect(runtime.onMessage.addListener).toHaveBeenCalledTimes(1); - }); - - it('emits messages received from runtime', async () => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - - const message = { foo: 'bar' }; - dispatchRuntimeMessage(message); - - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - }); - - it('calls validateInput with received input if specified', async () => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); - const validateInput = vi - .fn() - .mockReturnValue(true) as unknown as ValidateInput; - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - { validateInput }, - ); - - const message = { foo: 'bar' }; - dispatchRuntimeMessage(message); - - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - expect(validateInput).toHaveBeenCalledWith(message); - }); - - it('throws if validateInput throws', async () => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); - const validateInput = (() => { - throw new Error('foo'); - }) as unknown as ValidateInput; - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - { validateInput }, - ); - - const message = { foo: 'bar' }; - dispatchRuntimeMessage(message); - await expect(reader.next()).rejects.toThrow('foo'); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - }); - - it('ignores messages from other extensions', async () => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - - const nextP = reader.next(); - const message1 = { foo: 'bar' }; - const message2 = { fizz: 'buzz' }; - dispatchRuntimeMessage( - message1, - 'background', - 'offscreen', - 'other-extension-id', - ); - dispatchRuntimeMessage(message2); - - expect(await nextP).toStrictEqual(makePendingResult(message2)); - }); - - it('ignores messages that are not valid envelopes', async () => { - const { runtime, dispatchRuntimeMessage, listeners } = makeRuntime(); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - - const nextP = reader.next(); - - vi.spyOn(console, 'debug'); - listeners[0]?.({ not: 'an envelope' }, { id: EXTENSION_ID }); - - expect(console.debug).toHaveBeenCalledWith( - `ChromeRuntimeReader received unexpected message: ${stringify({ - not: 'an envelope', - })}`, - ); - - const message = { foo: 'bar' }; - dispatchRuntimeMessage(message); - expect(await nextP).toStrictEqual(makePendingResult(message)); - }); - - it('ignores messages for other targets', async () => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - - const nextP = reader.next(); - - vi.spyOn(console, 'warn'); - const message1 = { foo: 'bar' }; - dispatchRuntimeMessage( - message1, - // @ts-expect-error Intentional destructive testing - 'foo', - 'offscreen', - ); - - const message2 = { fizz: 'buzz' }; - dispatchRuntimeMessage(message2); - expect(await nextP).toStrictEqual(makePendingResult(message2)); - }); - - it('removes runtime.onMessage listener when done', async () => { - const { runtime, dispatchRuntimeMessage, listeners } = makeRuntime(); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - expect(listeners).toHaveLength(1); - - dispatchRuntimeMessage(makeStreamDoneSignal()); - - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(runtime.onMessage.removeListener).toHaveBeenCalledTimes(1); - expect(listeners).toHaveLength(0); - }); - - it('calls onEnd once when ending', async () => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); - const onEnd = vi.fn(); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - { onEnd }, - ); - - dispatchRuntimeMessage(makeStreamDoneSignal()); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); - - it('handles errors from onEnd function', async () => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); - const onEnd = vi.fn(() => { - throw new Error('foo'); - }); - const reader = new ChromeRuntimeReader( - asChromeRuntime(runtime), - 'background', - 'offscreen', - { onEnd }, - ); - - dispatchRuntimeMessage(makeStreamDoneSignal()); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - -describe.concurrent('ChromeRuntimeWriter', () => { - it('constructs a ChromeRuntimeWriter', () => { - const { runtime } = makeRuntime(); - const writer = new ChromeRuntimeWriter( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - - expect(writer).toBeInstanceOf(ChromeRuntimeWriter); - expect(writer[Symbol.asyncIterator]()).toBe(writer); - }); - - it('writes messages to runtime.sendMessage', async () => { - const { runtime } = makeRuntime(); - const writer = new ChromeRuntimeWriter( - asChromeRuntime(runtime), - 'background', - 'offscreen', - ); - - const message = { foo: 'bar' }; - const nextP = writer.next(message); - - expect(await nextP).toStrictEqual(makePendingResult(undefined)); - expect(runtime.sendMessage).toHaveBeenCalledWith( - makeEnvelope(message, 'background', 'offscreen'), - ); - }); - - it('calls onEnd once when ending', async () => { - const { runtime } = makeRuntime(); - const onEnd = vi.fn(); - const writer = new ChromeRuntimeWriter( - asChromeRuntime(runtime), - 'background', - 'offscreen', - { onEnd }, - ); - - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - describe.concurrent('ChromeRuntimeDuplexStream', () => { const makeDuplexStream = async (validateInput?: ValidateInput) => { - const { runtime, dispatchRuntimeMessage } = makeRuntime(); + const { runtime, dispatchRuntimeMessage, listeners } = makeRuntime(); const duplexStreamP = ChromeRuntimeDuplexStream.make( asChromeRuntime(runtime), 'background', @@ -312,7 +71,10 @@ describe.concurrent('ChromeRuntimeDuplexStream', () => { ); dispatchRuntimeMessage(makeAck()); - return [await duplexStreamP, { runtime, dispatchRuntimeMessage }] as const; + return [ + await duplexStreamP, + { runtime, dispatchRuntimeMessage, listeners }, + ] as const; }; it('throws an error when localTarget and remoteTarget are the same', async () => { @@ -348,6 +110,55 @@ describe.concurrent('ChromeRuntimeDuplexStream', () => { expect(validateInput).toHaveBeenCalledWith(message); }); + it('writes enveloped messages to runtime.sendMessage', async () => { + const [duplexStream, { runtime }] = await makeDuplexStream(); + + expect(await duplexStream.write(42)).toStrictEqual( + makePendingResult(undefined), + ); + expect(runtime.sendMessage).toHaveBeenLastCalledWith( + makeEnvelope(42, 'offscreen', 'background'), + ); + }); + + it('ignores messages from other extensions and for other targets', async () => { + const [duplexStream, { dispatchRuntimeMessage }] = await makeDuplexStream(); + + const nextP = duplexStream.next(); + dispatchRuntimeMessage(1, 'background', 'offscreen', 'other-extension-id'); + // @ts-expect-error Intentional destructive testing + dispatchRuntimeMessage(2, 'foo', 'offscreen'); + dispatchRuntimeMessage(3); + + expect(await nextP).toStrictEqual(makePendingResult(3)); + }); + + it('ignores messages that are not valid envelopes', async () => { + const [duplexStream, { dispatchRuntimeMessage, listeners }] = + await makeDuplexStream(); + const nextP = duplexStream.next(); + + vi.spyOn(console, 'debug'); + listeners[0]?.({ not: 'an envelope' }, { id: EXTENSION_ID }); + + expect(console.debug).toHaveBeenCalledWith( + `ChromeRuntimeDuplexStream received unexpected message: ${stringify({ + not: 'an envelope', + })}`, + ); + + dispatchRuntimeMessage(42); + expect(await nextP).toStrictEqual(makePendingResult(42)); + }); + + it('removes its runtime.onMessage listener when it ends', async () => { + const [duplexStream, { listeners }] = await makeDuplexStream(); + expect(listeners).toHaveLength(1); + + await duplexStream.return(); + expect(listeners).toHaveLength(0); + }); + it('ends the reader when the writer ends', async () => { const [duplexStream, { runtime }] = await makeDuplexStream(); runtime.sendMessage.mockImplementationOnce(() => { diff --git a/packages/streams/src/browser/ChromeRuntimeStream.ts b/packages/streams/src/browser/ChromeRuntimeStream.ts index 28174addad..f660d1cbb4 100644 --- a/packages/streams/src/browser/ChromeRuntimeStream.ts +++ b/packages/streams/src/browser/ChromeRuntimeStream.ts @@ -1,14 +1,13 @@ /** - * This module provides a pair of classes for creating readable and writable streams - * over the Chrome Extension Runtime messaging API. + * This module provides a duplex stream over the Chrome Extension Runtime messaging API. * - * These streams utilize `chrome.runtime.sendMessage` for sending data and + * The stream uses `chrome.runtime.sendMessage` for sending data and * `chrome.runtime.onMessage.addListener` for receiving data. This allows for * communication between different parts of a Chrome extension (e.g., background scripts, * content scripts, and popup pages). * * Note that unlike e.g. the `MessagePort` API, the Chrome Extension Runtime messaging API - * doesn't have a built-in way to close the connection. The streams will continue to operate + * doesn't have a built-in way to close the connection. The stream will continue to operate * as long as the extension is running, unless manually ended. * * @module ChromeRuntime streams @@ -18,18 +17,8 @@ import { stringify } from '@metamask/kernel-utils'; import type { Json } from '@metamask/utils'; import type { ChromeRuntime, ChromeMessageSender } from './chrome.d.ts'; -import { - BaseDuplexStream, - makeDuplexStreamInputValidator, -} from '../BaseDuplexStream.ts'; -import type { - BaseReaderArgs, - ValidateInput, - BaseWriterArgs, - ReceiveInput, -} from '../BaseStream.ts'; -import { BaseReader, BaseWriter } from '../BaseStream.ts'; -import type { Dispatchable } from '../utils.ts'; +import { BaseDuplexStream } from '../BaseDuplexStream.ts'; +import type { ValidateInput } from '../BaseStream.ts'; export type ChromeRuntimeTarget = 'background' | 'offscreen' | 'popup'; @@ -49,166 +38,14 @@ const isMessageEnvelope = ( 'payload' in message; /** - * A readable stream over the Chrome Extension Runtime messaging API. - * - * This class is a naive passthrough mechanism for data using `chrome.runtime.onMessage`. - * Expects exclusive read access to the messaging API. - * - * @see - * - {@link ChromeRuntimeWriter} for the corresponding writable stream. - * - The module-level documentation for more details. - */ -export class ChromeRuntimeReader extends BaseReader { - readonly #receiveInput: ReceiveInput; - - readonly #target: ChromeRuntimeTarget; - - readonly #source: ChromeRuntimeTarget; - - readonly #extensionId: string; - - /** - * Constructs a new {@link ChromeRuntimeReader}. - * - * @param runtime - The Chrome runtime API object. - * @param target - The target context (e.g., 'background', 'offscreen', 'popup'). - * @param source - The source context that messages are expected from. - * @param options - Options bag for configuring the reader. - * @param options.validateInput - A function that validates input from the transport. - * @param options.onEnd - A function that is called when the stream ends. - */ - constructor( - runtime: ChromeRuntime, - target: ChromeRuntimeTarget, - source: ChromeRuntimeTarget, - { validateInput, onEnd }: BaseReaderArgs = {}, - ) { - // eslint-disable-next-line prefer-const - let messageListener: ( - message: unknown, - sender: ChromeMessageSender, - ) => void; - - super({ - validateInput, - onEnd: async (error) => { - runtime.onMessage.removeListener(messageListener); - await onEnd?.(error); - }, - }); - - this.#receiveInput = super.getReceiveInput(); - this.#target = target; - this.#source = source; - this.#extensionId = runtime.id; - - messageListener = this.#onMessage.bind(this); - // Begin listening for messages from the Chrome runtime. - runtime.onMessage.addListener(messageListener); - - harden(this); - } - - /** - * Handles incoming messages from the Chrome runtime. - * - * @param message - The message received from the Chrome runtime. - * @param sender - The sender information for the message. - */ - #onMessage(message: unknown, sender: ChromeMessageSender): void { - if (sender.id !== this.#extensionId) { - return; - } - - if (!isMessageEnvelope(message)) { - // TODO(#562): Use logger instead. - // eslint-disable-next-line no-console - console.debug( - `ChromeRuntimeReader received unexpected message: ${stringify( - message, - )}`, - ); - return; - } - - if (message.target !== this.#target || message.source !== this.#source) { - // TODO(#562): Use logger instead. - // eslint-disable-next-line no-console - console.debug( - `ChromeRuntimeReader received message with incorrect target or source: ${stringify(message)}`, - `Expected target: ${this.#target}`, - `Expected source: ${this.#source}`, - ); - return; - } - - this.#receiveInput(message.payload).catch(async (error) => - this.throw(error), - ); - } -} -harden(ChromeRuntimeReader); - -/** - * A writable stream over the Chrome Extension Runtime messaging API. - * - * This class is a naive passthrough mechanism for data using `chrome.runtime.sendMessage`. - * - * @see - * - {@link ChromeRuntimeReader} for the corresponding readable stream. - * - The module-level documentation for more details. - */ -export class ChromeRuntimeWriter extends BaseWriter { - /** - * Constructs a new {@link ChromeRuntimeWriter}. - * - * @param runtime - The Chrome runtime API object. - * @param target - The target context to send messages to. - * @param source - The source context identifying where messages originate. - * @param options - Options bag for configuring the writer. - * @param options.name - The name of the stream, for logging purposes. - * @param options.onEnd - A function that is called when the stream ends. - */ - constructor( - runtime: ChromeRuntime, - target: ChromeRuntimeTarget, - source: ChromeRuntimeTarget, - { name, onEnd }: Omit, 'onDispatch'> = {}, - ) { - super({ - name, - onDispatch: async (value: Dispatchable) => { - await runtime.sendMessage({ - target, - source, - payload: value, - }); - }, - onEnd, - }); - harden(this); - } -} -harden(ChromeRuntimeWriter); - -/** - * A duplex stream over the Chrome Extension Runtime messaging API. - * - * This class is a naive passthrough mechanism for data using `chrome.runtime.onMessage`. - * - * @see - * - {@link ChromeRuntimeReader} for the corresponding readable stream. - * - {@link ChromeRuntimeWriter} for the corresponding writable stream. + * A duplex stream over the Chrome Extension Runtime messaging API. Reads only + * enveloped messages from this extension that are addressed from `remoteTarget` + * to `localTarget`. */ export class ChromeRuntimeDuplexStream< Read extends Json, Write extends Json = Read, -> extends BaseDuplexStream< - Read, - ChromeRuntimeReader, - Write, - ChromeRuntimeWriter -> { +> extends BaseDuplexStream { /** * Constructs a new {@link ChromeRuntimeDuplexStream}. * @@ -223,32 +60,46 @@ export class ChromeRuntimeDuplexStream< remoteTarget: ChromeRuntimeTarget, validateInput?: ValidateInput, ) { - let writer: ChromeRuntimeWriter; // eslint-disable-line prefer-const - const reader = new ChromeRuntimeReader( - runtime, - localTarget, - remoteTarget, - { - name: 'ChromeRuntimeDuplexStream', - validateInput: makeDuplexStreamInputValidator(validateInput), - onEnd: async () => { - await writer.return(); - }, + if (localTarget === remoteTarget) { + throw new Error('localTarget and remoteTarget must be different'); + } + + super({ + name: 'ChromeRuntimeDuplexStream', + validateInput, + listen: (receiveInput) => { + const onMessage = ( + message: unknown, + sender: ChromeMessageSender, + ): void => { + if (sender.id !== runtime.id) { + return; + } + if ( + isMessageEnvelope(message) && + message.target === localTarget && + message.source === remoteTarget + ) { + receiveInput(message.payload); + return; + } + // TODO(#562): Use logger instead. + // eslint-disable-next-line no-console + console.debug( + `ChromeRuntimeDuplexStream received unexpected message: ${stringify(message)}`, + ); + }; + runtime.onMessage.addListener(onMessage); + return () => runtime.onMessage.removeListener(onMessage); }, - ); - writer = new ChromeRuntimeWriter( - runtime, - remoteTarget, - localTarget, - { - name: 'ChromeRuntimeDuplexStream', - onEnd: async () => { - await reader.return(); - }, + onDispatch: async (payload) => { + await runtime.sendMessage({ + target: remoteTarget, + source: localTarget, + payload, + }); }, - ); - super(reader, writer); - harden(this); + }); } /** @@ -266,10 +117,6 @@ export class ChromeRuntimeDuplexStream< remoteTarget: ChromeRuntimeTarget, validateInput?: ValidateInput, ): Promise> { - if (localTarget === remoteTarget) { - throw new Error('localTarget and remoteTarget must be different'); - } - const stream = new ChromeRuntimeDuplexStream( runtime, localTarget, diff --git a/packages/streams/src/browser/MessagePortStream.test.ts b/packages/streams/src/browser/MessagePortStream.test.ts index f3667f5174..4611303507 100644 --- a/packages/streams/src/browser/MessagePortStream.test.ts +++ b/packages/streams/src/browser/MessagePortStream.test.ts @@ -1,11 +1,7 @@ import { delay } from '@metamask/kernel-utils'; import { describe, expect, it, vi } from 'vitest'; -import { - MessagePortDuplexStream, - MessagePortReader, - MessagePortWriter, -} from './MessagePortStream.ts'; +import { MessagePortDuplexStream } from './MessagePortStream.ts'; import { makeAck } from '../BaseDuplexStream.ts'; import type { ValidateInput } from '../BaseStream.ts'; import { @@ -14,144 +10,6 @@ import { makeStreamDoneSignal, } from '../utils.ts'; -describe('MessagePortReader', () => { - it('constructs a MessagePortReader', () => { - const { port1 } = new MessageChannel(); - const addListenerSpy = vi.spyOn(port1, 'addEventListener'); - const reader = new MessagePortReader(port1); - - expect(reader).toBeInstanceOf(MessagePortReader); - expect(reader[Symbol.asyncIterator]()).toBe(reader); - expect(addListenerSpy).toHaveBeenCalledOnce(); - }); - - it('emits messages received from port', async () => { - const { port1, port2 } = new MessageChannel(); - const reader = new MessagePortReader(port1); - - const message = { foo: 'bar' }; - port2.postMessage(message); - await delay(10); - - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - }); - - it('calls validateInput with received input if specified', async () => { - const validateInput = vi - .fn() - .mockReturnValue(true) as unknown as ValidateInput; - const { port1, port2 } = new MessageChannel(); - const reader = new MessagePortReader(port1, { validateInput }); - - const message = { foo: 'bar' }; - port2.postMessage(message); - await delay(10); - - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - expect(validateInput).toHaveBeenCalledWith(message); - }); - - it('throws if validateInput throws', async () => { - const validateInput = (() => { - throw new Error('foo'); - }) as unknown as ValidateInput; - const { port1, port2 } = new MessageChannel(); - const reader = new MessagePortReader(port1, { validateInput }); - - port2.postMessage(42); - await expect(reader.next()).rejects.toThrow('foo'); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - }); - - it('closes the port when done', async () => { - const { port1, port2 } = new MessageChannel(); - const closeSpy = vi.spyOn(port1, 'close'); - const addListenerSpy = vi.spyOn(port1, 'addEventListener'); - const removeListenerSpy = vi.spyOn(port1, 'removeEventListener'); - const reader = new MessagePortReader(port1); - expect(addListenerSpy).toHaveBeenCalledOnce(); - - port2.postMessage(makeStreamDoneSignal()); - await delay(10); - - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(closeSpy).toHaveBeenCalledOnce(); - expect(removeListenerSpy).toHaveBeenCalledOnce(); - }); - - it('ignores messages with ports', async () => { - const { port1, port2 } = new MessageChannel(); - const reader = new MessagePortReader(port1); - const { port1: otherPort } = new MessageChannel(); - - port2.postMessage(makeDoneResult(), [otherPort]); - port2.postMessage({ foo: 'bar' }); - await delay(10); - - expect(await reader.next()).toStrictEqual( - makePendingResult({ foo: 'bar' }), - ); - }); - - it('calls onEnd once when ending', async () => { - const { port1, port2 } = new MessageChannel(); - const onEnd = vi.fn(); - const reader = new MessagePortReader(port1, { onEnd }); - - port2.postMessage(makeStreamDoneSignal()); - await delay(10); - - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - -describe('MessagePortWriter', () => { - it('constructs a MessagePortWriter', () => { - const { port1 } = new MessageChannel(); - const writer = new MessagePortWriter(port1); - - expect(writer).toBeInstanceOf(MessagePortWriter); - expect(writer[Symbol.asyncIterator]()).toBe(writer); - }); - - it('writes messages to the port', async () => { - const { port1, port2 } = new MessageChannel(); - const writer = new MessagePortWriter(port1); - - const message = { foo: 'bar' }; - const messageP = new Promise((resolve) => { - port2.onmessage = (messageEvent): void => resolve(messageEvent.data); - }); - const nextP = writer.next(message); - - expect(await nextP).toStrictEqual(makePendingResult(undefined)); - expect(await messageP).toStrictEqual(message); - }); - - it('closes the port when done', async () => { - const { port1 } = new MessageChannel(); - const closeSpy = vi.spyOn(port1, 'close'); - const writer = new MessagePortWriter(port1); - - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(closeSpy).toHaveBeenCalledOnce(); - }); - - it('calls onEnd once when ending', async () => { - const { port1 } = new MessageChannel(); - const onEnd = vi.fn(); - const writer = new MessagePortWriter(port1, { onEnd }); - - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - describe('MessagePortDuplexStream', () => { const makeDuplexStream = async ( channel: MessageChannel = new MessageChannel(), @@ -174,6 +32,24 @@ describe('MessagePortDuplexStream', () => { expect(duplexStream[Symbol.asyncIterator]()).toBe(duplexStream); }); + it('reads messages from and writes messages to the port', async () => { + const channel = new MessageChannel(); + const duplexStream = await makeDuplexStream(channel); + + const messageP = new Promise((resolve) => { + // The port's queue also holds the SYN sent during synchronization. + channel.port2.onmessage = ({ data }): void => + data === 43 ? resolve(data) : undefined; + }); + channel.port2.postMessage(42); + + expect(await duplexStream.next()).toStrictEqual(makePendingResult(42)); + expect(await duplexStream.write(43)).toStrictEqual( + makePendingResult(undefined), + ); + expect(await messageP).toBe(43); + }); + it('calls validateInput with received input if specified', async () => { const validateInput = vi .fn() @@ -187,6 +63,17 @@ describe('MessagePortDuplexStream', () => { expect(validateInput).toHaveBeenCalledWith(42); }); + it('ignores messages with ports', async () => { + const channel = new MessageChannel(); + const duplexStream = await makeDuplexStream(channel); + const { port1: otherPort } = new MessageChannel(); + + channel.port2.postMessage(1, [otherPort]); + channel.port2.postMessage(2); + + expect(await duplexStream.next()).toStrictEqual(makePendingResult(2)); + }); + it('ends the reader when the writer ends', async () => { const { port1, port2 } = new MessageChannel(); vi.spyOn(port1, 'postMessage') @@ -203,8 +90,10 @@ describe('MessagePortDuplexStream', () => { expect(await duplexStream.next()).toStrictEqual(makeDoneResult()); }); - it('ends the writer when the reader ends', async () => { + it('ends the writer and closes the port when the reader ends', async () => { const { port1, port2 } = new MessageChannel(); + const closeSpy = vi.spyOn(port1, 'close'); + const removeListenerSpy = vi.spyOn(port1, 'removeEventListener'); const duplexStream = await makeDuplexStream({ port1, port2 }); const readP = duplexStream.next(); @@ -212,5 +101,7 @@ describe('MessagePortDuplexStream', () => { await delay(10); expect(await duplexStream.write(42)).toStrictEqual(makeDoneResult()); expect(await readP).toStrictEqual(makeDoneResult()); + expect(closeSpy).toHaveBeenCalledOnce(); + expect(removeListenerSpy).toHaveBeenCalledOnce(); }); }); diff --git a/packages/streams/src/browser/MessagePortStream.ts b/packages/streams/src/browser/MessagePortStream.ts index 2573e65e5e..b228912408 100644 --- a/packages/streams/src/browser/MessagePortStream.ts +++ b/packages/streams/src/browser/MessagePortStream.ts @@ -1,16 +1,13 @@ /** - * This module provides a pair of classes for creating readable and writable streams - * over a [MessagePort](https://developer.mozilla.org/en-US/docs/Web/API/MessagePort). - * The classes are naive passthrough mechanisms for data that assume exclusive access - * to their ports. The lifetime of the underlying message port is expected to be + * This module provides a duplex stream over a + * [MessagePort](https://developer.mozilla.org/en-US/docs/Web/API/MessagePort). + * The stream is a naive passthrough mechanism for data that assumes exclusive access + * to its port. The lifetime of the underlying message port is expected to be * coextensive with "the other side". * * At the time of writing, there is no ergonomic way to detect the closure of a port. For - * this reason, ports have to be ended manually via `.return()` or `.throw()`. Ending a - * {@link MessagePortWriter} will end any {@link MessagePortReader} reading from the - * remote port and close the entangled ports, but it will not affect any other streams - * connected to the remote or local port, which must also be ended manually. Use - * {@link MessagePortDuplexStream} to create a duplex stream over a single port. + * this reason, streams have to be ended manually via `.return()` or `.throw()`. Ending a + * stream ends the stream on the remote port and closes the entangled ports. * * Regarding limitations around detecting `MessagePort` closure, see: * - https://github.com/fergald/explainer-messageport-close @@ -20,115 +17,17 @@ */ import type { OnMessage } from './utils.ts'; -import { - BaseDuplexStream, - makeDuplexStreamInputValidator, -} from '../BaseDuplexStream.ts'; -import type { - BaseReaderArgs, - BaseWriterArgs, - ValidateInput, -} from '../BaseStream.ts'; -import { BaseReader, BaseWriter } from '../BaseStream.ts'; -import type { Dispatchable } from '../utils.ts'; +import { BaseDuplexStream } from '../BaseDuplexStream.ts'; +import type { ValidateInput } from '../BaseStream.ts'; /** - * A readable stream over a {@link MessagePort}. - * - * This class is a naive passthrough mechanism for data over a pair of linked message - * ports. Ignores message events dispatched on its port that contain ports, but - * otherwise expects {@link Dispatchable} values to be posted to its port. - * - * @see - * - {@link MessagePortWriter} for the corresponding writable stream. - * - The module-level documentation for more details. - */ -export class MessagePortReader extends BaseReader { - /** - * Constructs a new {@link MessagePortReader}. - * - * @param port - The message port to read from. - * @param options - Options bag for configuring the reader. - * @param options.validateInput - A function that validates input from the transport. - * @param options.onEnd - A function that is called when the stream ends. - */ - constructor( - port: MessagePort, - { validateInput, onEnd }: BaseReaderArgs = {}, - ) { - super({ - validateInput, - onEnd: async (error) => { - // eslint-disable-next-line @typescript-eslint/no-use-before-define - port.removeEventListener('message', onMessage); - port.close(); - await onEnd?.(error); - }, - }); - - const receiveInput = super.getReceiveInput(); - - const onMessage: OnMessage = (messageEvent) => { - if (messageEvent.ports.length > 0) { - return; - } - - receiveInput(messageEvent.data).catch(async (error) => this.throw(error)); - }; - port.addEventListener('message', onMessage); - port.start(); - - harden(this); - } -} -harden(MessagePortReader); - -/** - * A writable stream over a {@link MessagePort}. - * - * @see - * - {@link MessagePortReader} for the corresponding readable stream. - * - The module-level documentation for more details. - */ -export class MessagePortWriter extends BaseWriter { - /** - * Constructs a new {@link MessagePortWriter}. - * - * @param port - The message port to write to. - * @param options - Options bag for configuring the writer. - * @param options.name - The name of the stream, for logging purposes. - * @param options.onEnd - A function that is called when the stream ends. - */ - constructor( - port: MessagePort, - { name, onEnd }: Omit, 'onDispatch'> = {}, - ) { - super({ - name, - onDispatch: (value: Dispatchable) => port.postMessage(value), - onEnd: async (error) => { - port.close(); - await onEnd?.(error); - }, - }); - port.start(); - harden(this); - } -} -harden(MessagePortWriter); - -/** - * A duplex stream over a {@link MessagePort}. + * A duplex stream over a {@link MessagePort}. Ignores message events that + * transfer ports. */ export class MessagePortDuplexStream< Read, Write = Read, -> extends BaseDuplexStream< - Read, - MessagePortReader, - Write, - MessagePortWriter -> { +> extends BaseDuplexStream { /** * Constructs a new {@link MessagePortDuplexStream}. * @@ -136,21 +35,22 @@ export class MessagePortDuplexStream< * @param validateInput - A function that validates input from the transport. */ constructor(port: MessagePort, validateInput?: ValidateInput) { - let writer: MessagePortWriter; // eslint-disable-line prefer-const - const reader = new MessagePortReader(port, { - name: 'MessagePortDuplexStream', - validateInput: makeDuplexStreamInputValidator(validateInput), - onEnd: async () => { - await writer.return(); - }, - }); - writer = new MessagePortWriter(port, { + super({ name: 'MessagePortDuplexStream', - onEnd: async () => { - await reader.return(); + validateInput, + listen: (receiveInput) => { + const onMessage: OnMessage = (messageEvent) => { + if (messageEvent.ports.length === 0) { + receiveInput(messageEvent.data); + } + }; + port.addEventListener('message', onMessage); + port.start(); + return () => port.removeEventListener('message', onMessage); }, + onDispatch: (value) => port.postMessage(value), + onEnd: () => port.close(), }); - super(reader, writer); } /** diff --git a/packages/streams/src/browser/PostMessageStream.test.ts b/packages/streams/src/browser/PostMessageStream.test.ts index 7c3dbd132a..659cf19bf3 100644 --- a/packages/streams/src/browser/PostMessageStream.test.ts +++ b/packages/streams/src/browser/PostMessageStream.test.ts @@ -2,11 +2,7 @@ import { delay } from '@metamask/kernel-utils'; import { makeMockMessageTarget } from '@ocap/repo-tools/test-utils'; import { describe, it, expect, vi } from 'vitest'; -import { - PostMessageDuplexStream, - PostMessageReader, - PostMessageWriter, -} from './PostMessageStream.ts'; +import { PostMessageDuplexStream } from './PostMessageStream.ts'; import type { PostMessageTarget } from './PostMessageStream.ts'; import type { PostMessage } from './utils.ts'; import { makeAck } from '../BaseDuplexStream.ts'; @@ -18,169 +14,19 @@ import { makeStreamErrorSignal, } from '../utils.ts'; -describe('PostMessageReader', () => { - it('constructs a PostMessageReader', () => { - const reader = new PostMessageReader({ - messageTarget: makeMockMessageTarget(), - }); - expect(reader).toBeInstanceOf(PostMessageReader); - }); - - it('emits messages received from postMessage', async () => { - const messageTarget = makeMockMessageTarget(); - const reader = new PostMessageReader({ - messageTarget, - }); - - const message = { foo: 'bar' }; - - messageTarget.postMessage(message); - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - }); - - it('can yield MessageEvents directly', async () => { - const messageTarget = makeMockMessageTarget(); - const reader = new PostMessageReader({ - messageTarget, - messageEventMode: 'event', - }); - - const message = new MessageEvent('message', { data: 'bar' }); - - messageTarget.postMessage(message); - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - }); - - it('handles stream done signals normally when yielding MessageEvents', async () => { - const messageTarget = makeMockMessageTarget(); - const reader = new PostMessageReader({ - messageTarget, - messageEventMode: 'event', - }); - - messageTarget.postMessage( - new MessageEvent('message', { data: makeStreamDoneSignal() }), - ); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - }); - - it('handles stream error signals normally when yielding MessageEvents', async () => { - const messageTarget = makeMockMessageTarget(); - const reader = new PostMessageReader({ - messageTarget, - messageEventMode: 'event', - }); - - const nextP = reader.next(); - - messageTarget.postMessage( - new MessageEvent('message', { - data: makeStreamErrorSignal(new Error('foo')), - }), - ); - await expect(nextP).rejects.toThrow('foo'); - }); - - it('calls validateInput with received input if specified', async () => { - const validateInput = vi - .fn() - .mockReturnValue(true) as unknown as ValidateInput; - const messageTarget = makeMockMessageTarget(); - const reader = new PostMessageReader({ - messageTarget, - validateInput, - }); - - const message = { foo: 'bar' }; - messageTarget.postMessage(message); - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - expect(validateInput).toHaveBeenCalledWith(message); - }); - - it('throws if validateInput throws', async () => { - const messageTarget = makeMockMessageTarget(); - const validateInput = (() => { - throw new Error('foo'); - }) as unknown as ValidateInput; - const reader = new PostMessageReader({ - messageTarget, - validateInput, - }); - - messageTarget.postMessage(42); - await expect(reader.next()).rejects.toThrow('foo'); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - }); - - it('removes its listener when it ends', async () => { - const messageTarget = makeMockMessageTarget(); - const reader = new PostMessageReader({ - messageTarget, - }); - expect(messageTarget.listeners).toHaveLength(1); - - const message = makeStreamDoneSignal(); - messageTarget.postMessage(message); - - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(messageTarget.removeEventListener).toHaveBeenCalled(); - expect(messageTarget.listeners).toHaveLength(0); - }); - - it('calls onEnd once when ending', async () => { - const messageTarget = makeMockMessageTarget(); - const onEnd = vi.fn(); - const reader = new PostMessageReader({ - messageTarget, - onEnd, - }); - - messageTarget.postMessage(makeStreamDoneSignal()); - - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - -describe('PostMessageWriter', () => { - it('constructs a PostMessageWriter', () => { - const writer = new PostMessageWriter(makeMockMessageTarget()); - expect(writer).toBeInstanceOf(PostMessageWriter); - }); - - it('writes messages to postMessage', async () => { - const messageTarget = makeMockMessageTarget(); - const writer = new PostMessageWriter(messageTarget); - const message = { foo: 'bar' }; - await writer.next({ payload: message, transfer: [] }); - expect(messageTarget.postMessage).toHaveBeenCalledWith(message, []); - }); - - it('calls onEnd once when ending', async () => { - const messageTarget = makeMockMessageTarget(); - const onEnd = vi.fn(); - const writer = new PostMessageWriter(messageTarget, { onEnd }); - - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - describe('PostMessageDuplexStream', () => { const makeDuplexStream = async ({ messageTarget = makeMockMessageTarget(), postRemoteMessage = vi.fn(), validateInput, onEnd, + messageEventMode, }: { - messageTarget?: PostMessageTarget; + messageTarget?: ReturnType; postRemoteMessage?: PostMessage; validateInput?: ValidateInput; onEnd?: () => Promise; + messageEventMode?: 'data' | 'event'; } = {}) => { const postLocalMessage = messageTarget.postMessage; // @ts-expect-error In reality you have to be explicit about `messageEventMode` @@ -188,6 +34,7 @@ describe('PostMessageDuplexStream', () => { messageTarget: { ...messageTarget, postMessage: postRemoteMessage }, validateInput, onEnd, + messageEventMode, }); postLocalMessage(makeAck()); await delay(10); @@ -222,6 +69,59 @@ describe('PostMessageDuplexStream', () => { expect(validateInput).toHaveBeenCalledWith(42); }); + it('can yield MessageEvents directly', async () => { + const { duplexStream, postLocalMessage } = await makeDuplexStream< + MessageEvent, + unknown + >({ messageEventMode: 'event' }); + + const message = new MessageEvent('message', { data: 'bar' }); + postLocalMessage(message); + expect(await duplexStream.next()).toStrictEqual(makePendingResult(message)); + }); + + it('reads done signals as data when yielding MessageEvents', async () => { + const { duplexStream, postLocalMessage } = await makeDuplexStream< + MessageEvent, + unknown + >({ messageEventMode: 'event' }); + + postLocalMessage(makeStreamDoneSignal()); + expect(await duplexStream.next()).toStrictEqual(makeDoneResult()); + }); + + it('reads error signals as data when yielding MessageEvents', async () => { + const { duplexStream, postLocalMessage } = await makeDuplexStream< + MessageEvent, + unknown + >({ messageEventMode: 'event' }); + + const nextP = duplexStream.next(); + postLocalMessage(makeStreamErrorSignal(new Error('foo'))); + await expect(nextP).rejects.toThrow('foo'); + }); + + it('removes its listener when it ends', async () => { + const { duplexStream, messageTarget } = await makeDuplexStream(); + expect(messageTarget.listeners).toHaveLength(1); + + await duplexStream.return(); + expect(messageTarget.listeners).toHaveLength(0); + }); + + it('ends with an error if validateInput throws', async () => { + const validateInput = (() => { + throw new Error('foo'); + }) as unknown as ValidateInput; + const { duplexStream, postLocalMessage } = await makeDuplexStream({ + validateInput, + }); + + postLocalMessage(42); + await expect(duplexStream.next()).rejects.toThrow('foo'); + expect(await duplexStream.next()).toStrictEqual(makeDoneResult()); + }); + it('calls onEnd when ending if specified', async () => { const onEnd = vi.fn(); const { duplexStream } = await makeDuplexStream({ diff --git a/packages/streams/src/browser/PostMessageStream.ts b/packages/streams/src/browser/PostMessageStream.ts index 9892f943dd..c50a3e33c8 100644 --- a/packages/streams/src/browser/PostMessageStream.ts +++ b/packages/streams/src/browser/PostMessageStream.ts @@ -1,6 +1,6 @@ /** - * This module provides a pair of classes for creating readable and writable streams - * over a [postMessage](https://developer.mozilla.org/en-US/docs/Web/API/Window/postMessage). + * This module provides a duplex stream over a + * [postMessage](https://developer.mozilla.org/en-US/docs/Web/API/Window/postMessage) * function. * * @module PostMessage streams @@ -9,15 +9,9 @@ import { isObject } from '@metamask/utils'; import type { OnMessage, PostMessage } from './utils.ts'; -import { - BaseDuplexStream, - isDuplexStreamSignal, - makeDuplexStreamInputValidator, -} from '../BaseDuplexStream.ts'; -import type { BaseReaderArgs, BaseWriterArgs } from '../BaseStream.ts'; -import { BaseReader, BaseWriter } from '../BaseStream.ts'; +import { BaseDuplexStream, isDuplexStreamSignal } from '../BaseDuplexStream.ts'; +import type { OnEnd, ValidateInput } from '../BaseStream.ts'; import { isSignalLike } from '../utils.ts'; -import type { Dispatchable } from '../utils.ts'; export type PostMessageTarget = { addEventListener: (type: 'message', listener: OnMessage) => void; @@ -25,68 +19,18 @@ export type PostMessageTarget = { postMessage: PostMessage; }; -type PostMessageReaderArgs = BaseReaderArgs & { +type PostMessageDuplexStreamArgs = { messageTarget: PostMessageTarget; + validateInput?: ValidateInput | undefined; + onEnd?: OnEnd | undefined; } & (Read extends MessageEvent - ? { - messageEventMode: 'event'; - } - : { - messageEventMode?: 'data' | undefined; - }); - -/** - * A readable stream over a {@link PostMessage} function. - * - * Ignores message events dispatched on its port that contain ports, but otherwise - * expects {@link Dispatchable} values to be posted to its port. - * - * @see {@link PostMessageWriter} for the corresponding writable stream. - */ -export class PostMessageReader extends BaseReader { - /** - * Constructs a new {@link PostMessageReader}. - * - * @param options - Options bag for configuring the reader. - * @param options.messageTarget - The target to listen for messages on. - * @param options.validateInput - A function that validates input from the transport. - * @param options.onEnd - A function that is called when the stream ends. - * @param options.messageEventMode - Whether to pass the message event or just the data to the stream. - */ - constructor({ - validateInput, - onEnd, - messageTarget, - messageEventMode = 'data', - }: PostMessageReaderArgs) { - // eslint-disable-next-line prefer-const - let onMessage: OnMessage; - - super({ - validateInput, - onEnd: async (error) => { - messageTarget.removeEventListener('message', onMessage); - await onEnd?.(error); - }, + ? { + messageEventMode: 'event'; + } + : { + messageEventMode?: 'data' | undefined; }); - const receiveInput = super.getReceiveInput(); - onMessage = (messageEvent) => { - const value = - messageEventMode === 'data' || - isSignalLike(messageEvent.data) || - isDuplexStreamSignal(messageEvent.data) - ? messageEvent.data - : messageEvent; - receiveInput(value).catch(async (error) => this.throw(error)); - }; - messageTarget.addEventListener('message', onMessage); - - harden(this); - } -} -harden(PostMessageReader); - export type PostMessageEnvelope = { payload: Write; transfer: Transferable[]; @@ -106,97 +50,49 @@ const isPostMessageEnvelope = ( Array.isArray(value.transfer); /** - * A writable stream over a {@link PostMessage} function. - * - * @see {@link PostMessageReader} for the corresponding readable stream. - */ -export class PostMessageWriter extends BaseWriter { - /** - * Constructs a new {@link PostMessageWriter}. - * - * @param messageTarget - The target to post messages to. - * @param options - Options bag for configuring the writer. - * @param options.name - The name of the stream, for logging purposes. - * @param options.onEnd - A function that is called when the stream ends. - */ - constructor( - messageTarget: PostMessageTarget, - { name, onEnd }: Omit, 'onDispatch'> = {}, - ) { - super({ - name, - onDispatch: (value: Dispatchable) => { - return isPostMessageEnvelope(value) - ? messageTarget.postMessage(value.payload, value.transfer) - : messageTarget.postMessage(value); - }, - onEnd: async (error) => { - await onEnd?.(error); - }, - }); - harden(this); - } -} -harden(PostMessageWriter); - -type PostMessageDuplexStreamArgs = PostMessageReaderArgs; - -/** - * A duplex stream over a {@link PostMessage} function. - * - * @see {@link PostMessageReader} for the corresponding readable stream. - * @see {@link PostMessageWriter} for the corresponding writable stream. + * A duplex stream over a {@link PostMessage} function. Writes of a + * {@link PostMessageEnvelope} post its payload with its transfer list. */ export class PostMessageDuplexStream< Read, Write = Read, -> extends BaseDuplexStream< - Read, - PostMessageReader, - Write, - PostMessageWriter -> { +> extends BaseDuplexStream { /** * Constructs a new {@link PostMessageDuplexStream}. * * @param options - Options bag for configuring the duplex stream. * @param options.messageTarget - The target for sending and receiving messages. * @param options.validateInput - A function that validates input from the transport. - * @param options.onEnd - A function that is called when the stream ends. + * @param options.onEnd - A function that is called once when the stream ends. + * @param options.messageEventMode - Whether to read whole message events or just their data. */ constructor({ messageTarget, validateInput, onEnd, - ...args + messageEventMode = 'data', }: PostMessageDuplexStreamArgs) { - let didCallOnEnd = false; - const callOnEndOnce = async (): Promise => { - if (!didCallOnEnd) { - didCallOnEnd = true; - await onEnd?.(); - } - }; - - let writer: PostMessageWriter; // eslint-disable-line prefer-const - const reader = new PostMessageReader({ - ...args, - messageTarget, - validateInput: makeDuplexStreamInputValidator(validateInput), - // End the writer first, since onEnd may close the transport. - onEnd: async () => { - await writer.return(); - await callOnEndOnce(); - }, - } as PostMessageReaderArgs); - writer = new PostMessageWriter(messageTarget, { + super({ name: 'PostMessageDuplexStream', - onEnd: async () => { - await reader.return(); - await callOnEndOnce(); + validateInput, + onEnd, + listen: (receiveInput) => { + const onMessage: OnMessage = (messageEvent) => + receiveInput( + messageEventMode === 'data' || + isSignalLike(messageEvent.data) || + isDuplexStreamSignal(messageEvent.data) + ? messageEvent.data + : messageEvent, + ); + messageTarget.addEventListener('message', onMessage); + return () => messageTarget.removeEventListener('message', onMessage); }, + onDispatch: (value) => + isPostMessageEnvelope(value) + ? messageTarget.postMessage(value.payload, value.transfer) + : messageTarget.postMessage(value), }); - super(reader, writer); } /** diff --git a/packages/streams/src/browser/index.test.ts b/packages/streams/src/browser/index.test.ts index 6233c78a26..ee65312c9b 100644 --- a/packages/streams/src/browser/index.test.ts +++ b/packages/streams/src/browser/index.test.ts @@ -6,14 +6,8 @@ describe('index', () => { it('has the expected exports', () => { expect(Object.keys(indexModule).sort()).toStrictEqual([ 'ChromeRuntimeDuplexStream', - 'ChromeRuntimeReader', - 'ChromeRuntimeWriter', 'MessagePortDuplexStream', - 'MessagePortReader', - 'MessagePortWriter', 'PostMessageDuplexStream', - 'PostMessageReader', - 'PostMessageWriter', 'initializeMessageChannel', 'receiveMessagePort', 'split', diff --git a/packages/streams/src/browser/index.ts b/packages/streams/src/browser/index.ts index 1414d54d3f..18f5ef4c81 100644 --- a/packages/streams/src/browser/index.ts +++ b/packages/streams/src/browser/index.ts @@ -2,23 +2,11 @@ export { initializeMessageChannel, receiveMessagePort, } from './message-channel.ts'; -export { - MessagePortDuplexStream, - MessagePortReader, - MessagePortWriter, -} from './MessagePortStream.ts'; +export { MessagePortDuplexStream } from './MessagePortStream.ts'; export type { ChromeRuntime, ChromeMessageSender } from './chrome.d.ts'; export type { ChromeRuntimeTarget } from './ChromeRuntimeStream.ts'; -export { - ChromeRuntimeDuplexStream, - ChromeRuntimeReader, - ChromeRuntimeWriter, -} from './ChromeRuntimeStream.ts'; -export { - PostMessageDuplexStream, - PostMessageReader, - PostMessageWriter, -} from './PostMessageStream.ts'; +export { ChromeRuntimeDuplexStream } from './ChromeRuntimeStream.ts'; +export { PostMessageDuplexStream } from './PostMessageStream.ts'; export type { PostMessageEnvelope, PostMessageTarget, diff --git a/packages/streams/src/index.test.ts b/packages/streams/src/index.test.ts index 06e1774037..45f842ae4b 100644 --- a/packages/streams/src/index.test.ts +++ b/packages/streams/src/index.test.ts @@ -5,9 +5,9 @@ import * as indexModule from './index.ts'; describe('index', () => { it('has the expected exports', () => { expect(Object.keys(indexModule).sort()).toStrictEqual([ + 'BaseReader', + 'BaseWriter', 'NodeWorkerDuplexStream', - 'NodeWorkerReader', - 'NodeWorkerWriter', 'split', ]); }); diff --git a/packages/streams/src/index.ts b/packages/streams/src/index.ts index 65bf2c0528..2f42b33340 100644 --- a/packages/streams/src/index.ts +++ b/packages/streams/src/index.ts @@ -1,8 +1,5 @@ export type { Reader, Writer } from './utils.ts'; export type { DuplexStream } from './BaseDuplexStream.ts'; -export { - NodeWorkerReader, - NodeWorkerWriter, - NodeWorkerDuplexStream, -} from './node/NodeWorkerStream.ts'; +export { BaseReader, BaseWriter } from './BaseStream.ts'; +export { NodeWorkerDuplexStream } from './node/NodeWorkerStream.ts'; export { split } from './split.ts'; diff --git a/packages/streams/src/node/NodeWorkerStream.test.ts b/packages/streams/src/node/NodeWorkerStream.test.ts index 3cc9e5a5ea..9f2b2c8f12 100644 --- a/packages/streams/src/node/NodeWorkerStream.test.ts +++ b/packages/streams/src/node/NodeWorkerStream.test.ts @@ -2,11 +2,7 @@ import { delay } from '@metamask/kernel-utils'; import { describe, it, expect, vi } from 'vitest'; import type { Mocked } from 'vitest'; -import { - NodeWorkerDuplexStream, - NodeWorkerReader, - NodeWorkerWriter, -} from './NodeWorkerStream.ts'; +import { NodeWorkerDuplexStream } from './NodeWorkerStream.ts'; import type { NodePort, OnMessage } from './NodeWorkerStream.ts'; import { makeAck } from '../BaseDuplexStream.ts'; import type { ValidateInput } from '../BaseStream.ts'; @@ -23,104 +19,15 @@ const makeMockNodePort = (): Mocked & { on: vi.fn((_event, listener) => { port.messageHandler = listener; }), + off: vi.fn(() => { + port.messageHandler = undefined; + }), postMessage: vi.fn(), - messageHandler: undefined, + messageHandler: undefined as OnMessage | undefined, }; return port; }; -describe('NodeWorkerReader', () => { - it('constructs a NodeWorkerReader', () => { - const port = makeMockNodePort(); - const reader = new NodeWorkerReader(port); - - expect(reader).toBeInstanceOf(NodeWorkerReader); - expect(reader[Symbol.asyncIterator]()).toBe(reader); - expect(port.on).toHaveBeenCalledOnce(); - }); - - it('emits messages received from port', async () => { - const port = makeMockNodePort(); - const reader = new NodeWorkerReader(port); - - const message = { foo: 'bar' }; - port.messageHandler?.(message); - - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - }); - - it('calls validateInput with received input if specified', async () => { - const port = makeMockNodePort(); - const validateInput = vi - .fn() - .mockReturnValue(true) as unknown as ValidateInput; - const reader = new NodeWorkerReader(port, { validateInput }); - - const message = { foo: 'bar' }; - port.messageHandler?.(message); - - expect(await reader.next()).toStrictEqual(makePendingResult(message)); - expect(validateInput).toHaveBeenCalledWith(message); - }); - - it('throws if validateInput throws', async () => { - const port = makeMockNodePort(); - const validateInput = (() => { - throw new Error('foo'); - }) as unknown as ValidateInput; - const reader = new NodeWorkerReader(port, { validateInput }); - - port.messageHandler?.(42); - await expect(reader.next()).rejects.toThrow('foo'); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - }); - - it('calls onEnd once when ending', async () => { - const port = makeMockNodePort(); - const onEnd = vi.fn(); - const reader = new NodeWorkerReader(port, { onEnd }); - - port.messageHandler?.(makeStreamDoneSignal()); - - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await reader.next()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - -describe('NodeWorkerWriter', () => { - it('constructs a NodeWorkerWriter', () => { - const port = makeMockNodePort(); - const writer = new NodeWorkerWriter(port); - - expect(writer).toBeInstanceOf(NodeWorkerWriter); - expect(writer[Symbol.asyncIterator]()).toBe(writer); - }); - - it('writes messages to the port', async () => { - const port = makeMockNodePort(); - const writer = new NodeWorkerWriter(port); - - const message = { foo: 'bar' }; - const nextP = writer.next(message); - - expect(await nextP).toStrictEqual(makePendingResult(undefined)); - expect(port.postMessage).toHaveBeenCalledWith(message); - }); - - it('calls onEnd once when ending', async () => { - const port = makeMockNodePort(); - const onEnd = vi.fn(); - const writer = new NodeWorkerWriter(port, { onEnd }); - - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - expect(await writer.return()).toStrictEqual(makeDoneResult()); - expect(onEnd).toHaveBeenCalledTimes(1); - }); -}); - describe('NodeWorkerDuplexStream', () => { const makeDuplexStream = async ( port = makeMockNodePort(), @@ -135,10 +42,24 @@ describe('NodeWorkerDuplexStream', () => { }; it('constructs a NodeWorkerDuplexStream', async () => { - const duplexStream = await makeDuplexStream(); + const port = makeMockNodePort(); + const duplexStream = await makeDuplexStream(port); expect(duplexStream).toBeInstanceOf(NodeWorkerDuplexStream); expect(duplexStream[Symbol.asyncIterator]()).toBe(duplexStream); + expect(port.on).toHaveBeenCalledOnce(); + }); + + it('reads messages from and writes messages to the port', async () => { + const port = makeMockNodePort(); + const duplexStream = await makeDuplexStream(port); + + port.messageHandler?.(42); + expect(await duplexStream.next()).toStrictEqual(makePendingResult(42)); + expect(await duplexStream.write(43)).toStrictEqual( + makePendingResult(undefined), + ); + expect(port.postMessage).toHaveBeenLastCalledWith(43); }); it('calls validateInput with received input if specified', async () => { @@ -170,14 +91,16 @@ describe('NodeWorkerDuplexStream', () => { expect(await duplexStream.next()).toStrictEqual(makeDoneResult()); }); - it('ends the writer when the reader ends', async () => { + it('ends the writer and removes its listener when the reader ends', async () => { const port = makeMockNodePort(); const duplexStream = await makeDuplexStream(port); + const listener = port.messageHandler; const readP = duplexStream.next(); port.messageHandler?.(makeStreamDoneSignal()); await delay(10); expect(await duplexStream.write(42)).toStrictEqual(makeDoneResult()); expect(await readP).toStrictEqual(makeDoneResult()); + expect(port.off).toHaveBeenCalledWith('message', listener); }); }); diff --git a/packages/streams/src/node/NodeWorkerStream.ts b/packages/streams/src/node/NodeWorkerStream.ts index 97533cd6cc..12baff2e32 100644 --- a/packages/streams/src/node/NodeWorkerStream.ts +++ b/packages/streams/src/node/NodeWorkerStream.ts @@ -2,103 +2,25 @@ * @module Node Worker streams */ -import { - BaseDuplexStream, - makeDuplexStreamInputValidator, -} from '../BaseDuplexStream.ts'; -import type { - BaseReaderArgs, - BaseWriterArgs, - ValidateInput, -} from '../BaseStream.ts'; -import { BaseReader, BaseWriter } from '../BaseStream.ts'; -import type { Dispatchable } from '../utils.ts'; +import { BaseDuplexStream } from '../BaseDuplexStream.ts'; +import type { ValidateInput } from '../BaseStream.ts'; export type OnMessage = (message: unknown) => void; export type NodePort = { on: (event: 'message', listener: OnMessage) => void; + off: (event: 'message', listener: OnMessage) => void; postMessage: (message: unknown) => void; }; /** - * A readable stream over a {@link NodePort}. - * - * @see - * - {@link NodeWorkerWriter} for the corresponding writable stream. - * - The module-level documentation for more details. - */ -export class NodeWorkerReader extends BaseReader { - /** - * Constructs a new {@link NodeWorkerReader}. - * - * @param port - The node worker port to read from. - * @param options - Options bag for configuring the reader. - * @param options.validateInput - A function that validates input from the transport. - * @param options.onEnd - A function that is called when the stream ends. - */ - constructor( - port: NodePort, - { validateInput, onEnd }: BaseReaderArgs = {}, - ) { - super({ - validateInput, - onEnd: async () => await onEnd?.(), - }); - - const receiveInput = super.getReceiveInput(); - port.on('message', (data) => { - receiveInput(data).catch(async (error) => this.throw(error)); - }); - harden(this); - } -} -harden(NodeWorkerReader); - -/** - * A writable stream over a {@link NodeWorker}. - * - * @see - * - {@link NodeWorkerReader} for the corresponding readable stream. - * - The module-level documentation for more details. - */ -export class NodeWorkerWriter extends BaseWriter { - /** - * Constructs a new {@link NodeWorkerWriter}. - * - * @param port - The node worker port to write to. - * @param options - Options bag for configuring the writer. - * @param options.name - The name of the stream, for logging purposes. - * @param options.onEnd - A function that is called when the stream ends. - */ - constructor( - port: NodePort, - { name, onEnd }: Omit, 'onDispatch'> = {}, - ) { - super({ - name, - onDispatch: (value: Dispatchable) => port.postMessage(value), - onEnd: async () => { - await onEnd?.(); - }, - }); - harden(this); - } -} -harden(NodeWorkerWriter); - -/** - * A duplex stream over a Node worker port. + * A duplex stream over a Node worker port, i.e. a `Worker` or a + * `worker_threads` `MessagePort`. */ export class NodeWorkerDuplexStream< Read, Write = Read, -> extends BaseDuplexStream< - Read, - NodeWorkerReader, - Write, - NodeWorkerWriter -> { +> extends BaseDuplexStream { /** * Constructs a new {@link NodeWorkerDuplexStream}. * @@ -106,21 +28,15 @@ export class NodeWorkerDuplexStream< * @param validateInput - A function that validates input from the transport. */ constructor(port: NodePort, validateInput?: ValidateInput) { - let writer: NodeWorkerWriter; // eslint-disable-line prefer-const - const reader = new NodeWorkerReader(port, { - name: 'NodeWorkerDuplexStream', - validateInput: makeDuplexStreamInputValidator(validateInput), - onEnd: async () => { - await writer.return(); - }, - }); - writer = new NodeWorkerWriter(port, { + super({ name: 'NodeWorkerDuplexStream', - onEnd: async () => { - await reader.return(); + validateInput, + listen: (receiveInput) => { + port.on('message', receiveInput); + return () => port.off('message', receiveInput); }, + onDispatch: (value) => port.postMessage(value), }); - super(reader, writer); } /** diff --git a/packages/streams/src/split.ts b/packages/streams/src/split.ts index 893577dfe3..7751fc0e38 100644 --- a/packages/streams/src/split.ts +++ b/packages/streams/src/split.ts @@ -2,85 +2,29 @@ import { stringify } from '@metamask/kernel-utils'; import type { DuplexStream } from './BaseDuplexStream.ts'; import { BaseReader } from './BaseStream.ts'; -import type { BaseReaderArgs, ReceiveInput } from './BaseStream.ts'; +import type { ReceiveInput } from './BaseStream.ts'; /** - * A reader for use within {@link split} that reads from a reader and forwards - * writes to a parent. The reader should output a subset of the parent stream's values based - * on some predicate. + * A {@link DuplexStream} for use within {@link split} that reads a subset of its + * parent's values and forwards writes to its parent. */ -class SplitReader extends BaseReader { - /** - * Constructs a new {@link SplitReader}. - * - * @param args - The arguments to pass to the base reader. - */ - // eslint-disable-next-line no-restricted-syntax - private constructor(args: BaseReaderArgs) { - super(args); - } +class SplitStream implements DuplexStream { + readonly #parent: DuplexStream; - /** - * Creates a new {@link SplitReader}. - * - * @param args - The arguments to pass to the base reader. - * @returns A new {@link SplitReader} and the receive input function. - */ - static make( - args: BaseReaderArgs, - ): [SplitReader, ReceiveInput] { - const reader = new SplitReader(args); - return [reader, reader.getReceiveInput()] as const; - } -} -harden(SplitReader); - -/** - * A {@link DuplexStream} for use within {@link split} that reads from a reader and forwards - * writes to a parent. The reader should output a subset of the parent stream's values based - * on some predicate. - */ -class SplitStream - implements DuplexStream -{ - readonly #parent: DuplexStream; - - readonly #reader: SplitReader; + readonly #reader: BaseReader; /** * Constructs a new {@link SplitStream}. * - * @param parent - The parent stream to read from. - * @param reader - The reader to use to read from the parent stream. + * @param parent - The parent stream. + * @param reader - The reader that receives this split's subset of the parent's values. */ - constructor( - parent: DuplexStream, - reader: SplitReader, - ) { + constructor(parent: DuplexStream, reader: BaseReader) { this.#parent = parent; this.#reader = reader; harden(this); } - /** - * Constructs a new {@link SplitStream}. - * - * @param parent - The parent stream to read from. - * @returns A new {@link SplitStream} and the receive input function. - */ - static make( - parent: DuplexStream, - ): { - stream: SplitStream; - receiveInput: ReceiveInput; - } { - const [reader, receiveInput] = SplitReader.make({ - name: this.constructor.name, - }); - const stream = new SplitStream(parent, reader); - return { stream, receiveInput }; - } - /** * Reads the next value from the stream. * @@ -91,9 +35,9 @@ class SplitStream } /** - * Writes a value to the stream. + * Writes a value to the parent stream. * - * @param value - The value to write to the stream. + * @param value - The value to write. * @returns The result of writing the value. */ async write(value: Write): Promise> { @@ -123,28 +67,26 @@ class SplitStream } /** - * Closes the stream. Idempotent. + * Closes the stream and its parent. Idempotent. * * @returns The final result for this stream. */ async return(): Promise> { - await this.#parent.return(); - return this.#reader.return(); + return this.end(); } /** - * Closes the stream with an error. Idempotent. + * Closes the stream and its parent with an error. Idempotent. * * @param error - The error to close the stream with. * @returns The final result for this stream. */ async throw(error: Error): Promise> { - await this.#parent.throw(error); - return this.#reader.throw(error); + return this.end(error); } /** - * Closes the stream. Syntactic sugar for `throw(error)` or `return()`. Idempotent. + * Closes the stream and its parent. Idempotent. * * @param error - The error to close the stream with. * @returns The final result for this stream. @@ -165,97 +107,67 @@ class SplitStream } harden(SplitStream); -// There's no reason to do this but we leave it in for the sake of completeness. -export function split( - stream: DuplexStream, - predicateA: (value: Read) => value is ReadA, -): [DuplexStream]; - -export function split( - stream: DuplexStream, - predicateA: (value: Read) => value is ReadA, - predicateB: (value: Read) => value is ReadB, -): [DuplexStream, DuplexStream]; - -export function split< - Read, - Write, - ReadA extends Read, - ReadB extends Read, - ReadC extends Read, ->( - stream: DuplexStream, - predicateA: (value: Read) => value is ReadA, - predicateB: (value: Read) => value is ReadB, - predicateC: (value: Read) => value is ReadC, -): [ - DuplexStream, - DuplexStream, - DuplexStream, -]; - -export function split< - Read, - Write, - ReadA extends Read, - ReadB extends Read, - ReadC extends Read, - ReadD extends Read, ->( - stream: DuplexStream, - predicateA: (value: Read) => value is ReadA, - predicateB: (value: Read) => value is ReadB, - predicateC: (value: Read) => value is ReadC, - predicateD: (value: Read) => value is ReadD, -): [ - DuplexStream, - DuplexStream, - DuplexStream, - DuplexStream, -]; +type Splits = { + [Index in keyof Predicates]: DuplexStream< + Predicates[Index] extends (( + value: Read, + ) => value is infer Narrowed extends Read) + ? Narrowed + : Read, + Write + >; +}; /** - * Splits a stream into multiple streams based on a list of predicates. - * Supports up to 4 predicates with type checking, and any number without! + * Splits a stream into one stream per predicate. Each value read from the parent + * goes to the first split whose predicate it matches; a value that matches none + * ends all splits with an error. Writes to any split go to the parent, and ending + * any split ends the parent and therefore all splits. * * @param parentStream - The stream to split. * @param predicates - The predicates to use to split the stream. * @returns An array of "splits" of the parent stream. */ -export function split( +export function split< + Read, + Write, + Predicates extends ((value: Read) => boolean)[], +>( parentStream: DuplexStream, - ...predicates: ((value: Read) => boolean)[] -): DuplexStream[] { - const splits = predicates.map( - (predicate) => [predicate, SplitStream.make(parentStream)] as const, - ); + ...predicates: Predicates +): Splits { + const splits = predicates.map((predicate) => { + let receiveInput!: ReceiveInput; + const reader = new BaseReader({ + name: 'SplitStream', + listen: (receive) => { + receiveInput = receive as ReceiveInput; + }, + }); + const stream = new SplitStream(parentStream, reader); + return { predicate, receiveInput, stream }; + }); // eslint-disable-next-line no-void void (async () => { let error: Error | undefined; try { for await (const value of parentStream) { - let matched = false; - for (const [predicate, { receiveInput }] of splits) { - if (predicate(value)) { - matched = true; - await receiveInput(value); - break; - } - } - - if (!matched) { + const match = splits.find(({ predicate }) => predicate(value)); + if (!match) { throw new Error( `Failed to match any predicate for value: ${stringify(value)}`, ); } + // Awaited so that every value is received before the splits end. + await match.receiveInput(value); } } catch (caughtError) { error = caughtError as Error; } - await Promise.all(splits.map(async ([, { stream }]) => stream.end(error))); + await Promise.all(splits.map(async ({ stream }) => stream.end(error))); })(); - return splits.map(([, { stream }]) => stream); + return splits.map(({ stream }) => stream) as Splits; } diff --git a/packages/streams/src/utils.test.ts b/packages/streams/src/utils.test.ts index 84bdb82294..503050e3e9 100644 --- a/packages/streams/src/utils.test.ts +++ b/packages/streams/src/utils.test.ts @@ -2,67 +2,31 @@ import { stringify } from '@metamask/kernel-utils'; import { makeErrorMatcherFactory } from '@ocap/repo-tools/test-utils'; import { describe, expect, it } from 'vitest'; -import type { Dispatchable, Writable } from './utils.ts'; import { makeDoneResult, makePendingResult, makeStreamDoneSignal, makeStreamErrorSignal, - marshal, - StreamDoneSymbol, + parseSignal, StreamSentinel, - unmarshal, } from './utils.ts'; const makeErrorMatcher = makeErrorMatcherFactory(expect); -describe('marshal', () => { - it.each([ - ['StreamDoneSymbol', StreamDoneSymbol, makeStreamDoneSignal()], - [ - 'Error', - new Error('foo'), - { - [StreamSentinel.Error]: true, - error: makeErrorMatcher('foo'), - }, - ], - ['number', 42], - ['string', 'foo'], - ['object', { foo: 'bar' }], - ['array', [1, 2, 3]], - ['null', null], - ['Symbol', Symbol('foo')], - ] as [string, Writable, Dispatchable | undefined][])( - 'should marshal a %s value', - (_, value, expected) => { - const marshaledValue = marshal(value); - expect(marshaledValue).toStrictEqual(expected ?? value); - }, - ); -}); +describe('parseSignal', () => { + it('returns undefined for a done signal', () => { + expect(parseSignal(makeStreamDoneSignal())).toBeUndefined(); + }); -describe('unmarshal', () => { - it.each([ - ['StreamDoneSignal', makeStreamDoneSignal(), StreamDoneSymbol], - ['Error', makeStreamErrorSignal(new Error('foo')), new Error('foo')], - ['number', 42], - ['string', 'foo'], - ['object', { foo: 'bar' }], - ['array', [1, 2, 3]], - ['null', null], - ['Symbol', Symbol('foo')], - ] as [string, Dispatchable, Writable | undefined][])( - 'should unmarshal a %s value', - (_, value, expected) => { - const unmarshaledValue = unmarshal(value); - expect(unmarshaledValue).toStrictEqual(expected ?? value); - }, - ); + it('returns the error carried by an error signal', () => { + expect(parseSignal(makeStreamErrorSignal(new Error('foo')))).toStrictEqual( + makeErrorMatcher('foo'), + ); + }); it('throws if the value is not a valid stream signal', () => { - const badSignal = { [StreamSentinel.Error]: true, error: 'foo' }; - expect(() => unmarshal(badSignal)).toThrow( + const badSignal = { [StreamSentinel.Error]: true, error: 'foo' } as const; + expect(() => parseSignal(badSignal)).toThrow( `Invalid stream signal: ${stringify(badSignal)}`, ); }); diff --git a/packages/streams/src/utils.ts b/packages/streams/src/utils.ts index 0493901898..81444c2c9f 100644 --- a/packages/streams/src/utils.ts +++ b/packages/streams/src/utils.ts @@ -1,4 +1,3 @@ -import type { Reader, Writer } from '@endo/stream'; import { isMarshaledError, marshalError, @@ -14,15 +13,26 @@ import { UnsafeJsonStruct, } from '@metamask/utils'; -export type { Reader, Writer }; +/** + * An async iterator that does not conflate its read and write types. Matches the + * `Stream` type of `@endo/stream`. + */ +type Stream = { + next(value: Write): Promise>; + return(): Promise>; + throw(error: Error): Promise>; + [Symbol.asyncIterator](): Stream; +}; + +export type Reader = Stream; + +export type Writer = Stream; export const StreamSentinel = { Error: '@@StreamError', Done: '@@StreamDone', } as const; -export const StreamDoneSymbol = Symbol('StreamDone'); - const StreamDoneStruct = object({ [StreamSentinel.Done]: literal(true), }); @@ -53,11 +63,21 @@ export const makeStreamDoneSignal = (): StreamDone => ({ }); /** - * A value that can be written to a stream. + * Parses a stream signal. * - * @template Yield - The type of the values yielded by the iterator. + * @param signal - The signal to parse. + * @returns The error carried by an error signal, or `undefined` for a done signal. + * @throws If the value is not a valid stream signal. */ -export type Writable = Yield | Error | typeof StreamDoneSymbol; +export const parseSignal = (signal: StreamSignal): Error | undefined => { + if (is(signal, StreamDoneStruct)) { + return undefined; + } + if (is(signal, StreamErrorStruct) && isMarshaledError(signal.error)) { + return unmarshalError(signal.error); + } + throw new Error(`Invalid stream signal: ${stringify(signal)}`); +}; /** * A value that can be dispatched to the internal transport mechanism of a stream. @@ -66,44 +86,6 @@ export type Writable = Yield | Error | typeof StreamDoneSymbol; */ export type Dispatchable = Yield | StreamSignal; -/** - * Marshals a {@link Writable} into a {@link Dispatchable}. - * - * @param value - The value to marshal. - * @returns The marshaled value. - */ -export function marshal(value: Writable): Dispatchable { - if (value === StreamDoneSymbol) { - return { [StreamSentinel.Done]: true }; - } - if (value instanceof Error) { - return { - [StreamSentinel.Error]: true, - error: marshalError(value), - }; - } - return value; -} - -/** - * Unmarshals a {@link Dispatchable} into a {@link Writable}. - * - * @param value - The value to unmarshal. - * @returns The unmarshaled value. - */ -export function unmarshal(value: Dispatchable): Writable { - if (isSignalLike(value)) { - if (is(value, StreamDoneStruct)) { - return StreamDoneSymbol; - } - if (is(value, StreamErrorStruct) && isMarshaledError(value.error)) { - return unmarshalError(value.error); - } - throw new Error(`Invalid stream signal: ${stringify(value)}`); - } - return value; -} - /** * Creates a {@link IteratorResult} with `{ done: true, value: undefined }`. * diff --git a/packages/streams/test/stream-mocks.ts b/packages/streams/test/stream-mocks.ts index b844a40845..2ee0651935 100644 --- a/packages/streams/test/stream-mocks.ts +++ b/packages/streams/test/stream-mocks.ts @@ -1,19 +1,15 @@ -import { - BaseDuplexStream, - makeAck, - makeDuplexStreamInputValidator, -} from '../src/BaseDuplexStream.ts'; +import { BaseDuplexStream, makeAck } from '../src/BaseDuplexStream.ts'; import type { Dispatch, ReceiveInput, BaseReaderArgs, + OnEnd, ValidateInput, - BaseWriterArgs, } from '../src/BaseStream.ts'; import { BaseReader, BaseWriter } from '../src/BaseStream.ts'; /** - * A test reader that exposes the receiveInput method for testing purposes. + * A test reader that exposes its receiveInput function for testing purposes. */ export class TestReader extends BaseReader { readonly #receiveInput: ReceiveInput; @@ -32,92 +28,34 @@ export class TestReader extends BaseReader { * * @param args - Options bag for configuring the reader. */ - constructor(args: BaseReaderArgs = {}) { - super(args); - this.#receiveInput = super.getReceiveInput(); - } - - /** - * Gets the receive input function. Overrides the protected method for testing. - * - * @returns The receive input function. - */ - getReceiveInput(): ReceiveInput { - return super.getReceiveInput(); - } - - /** - * Closes the underlying transport and returns. - * - * @returns The final result for this stream. - */ - async return(): Promise> { - return super.return(); - } - - /** - * Closes the stream with an error. - * - * @param error - The error to close the stream with. - * @returns The final result for this stream. - */ - async throw(error: Error): Promise> { - return super.throw(error); + constructor(args: Omit, 'listen'> = {}) { + let receiveInput!: ReceiveInput; + super({ + ...args, + listen: (receive) => { + receiveInput = receive as ReceiveInput; + }, + }); + this.#receiveInput = receiveInput; } } -/** - * A test writer that exposes the onDispatch function for testing purposes. - */ -export class TestWriter extends BaseWriter { - readonly #onDispatch: Dispatch; - - /** - * Gets the dispatch function for this writer. - * - * @returns The dispatch function. - */ - get onDispatch(): Dispatch { - return this.#onDispatch; - } - - /** - * Constructs a new {@link TestWriter}. - * - * @param args - Options bag for configuring the writer. - */ - constructor(args: BaseWriterArgs) { - super(args); - this.#onDispatch = args.onDispatch; - } -} +export class TestWriter extends BaseWriter {} type TestDuplexStreamOptions = { validateInput?: ValidateInput | undefined; - readerOnEnd?: () => void; - writerOnEnd?: () => void; + onEnd?: OnEnd | undefined; }; /** - * A test duplex stream that exposes internal methods for testing purposes. + * A test duplex stream that exposes its receiveInput function for testing purposes. */ export class TestDuplexStream< Read = number, Write = Read, -> extends BaseDuplexStream, Write, TestWriter> { - readonly #onDispatch: Dispatch; - +> extends BaseDuplexStream { readonly #receiveInput: ReceiveInput; - /** - * Gets the dispatch function for the underlying writer. - * - * @returns The dispatch function. - */ - get onDispatch(): Dispatch { - return this.#onDispatch; - } - /** * Gets the receive input function for the underlying reader. * @@ -133,41 +71,23 @@ export class TestDuplexStream< * @param onDispatch - The dispatch function to use for writing. * @param options - Options bag for configuring the stream. * @param options.validateInput - A function that validates input from the transport. - * @param options.readerOnEnd - A function that is called when the reader ends. - * @param options.writerOnEnd - A function that is called when the writer ends. + * @param options.onEnd - A function that is called once when the stream ends. */ constructor( onDispatch: Dispatch, - { - validateInput, - readerOnEnd, - writerOnEnd, - }: TestDuplexStreamOptions = {}, + { validateInput, onEnd }: TestDuplexStreamOptions = {}, ) { - const reader = new TestReader({ + let receiveInput!: ReceiveInput; + super({ name: 'TestDuplexStream', - onEnd: readerOnEnd, - validateInput: makeDuplexStreamInputValidator(validateInput), + validateInput, + onEnd, + onDispatch, + listen: (receive) => { + receiveInput = receive as ReceiveInput; + }, }); - super( - reader, - new TestWriter({ - name: 'TestDuplexStream', - onDispatch, - onEnd: writerOnEnd, - }), - ); - this.#onDispatch = onDispatch; - this.#receiveInput = reader.receiveInput; - } - - /** - * Synchronizes the stream with its remote counterpart. - * - * @returns A promise that resolves when the stream is synchronized. - */ - async synchronize(): Promise { - return super.synchronize(); + this.#receiveInput = receiveInput; } /** diff --git a/yarn.lock b/yarn.lock index a94d8d9126..4be4a9cf52 100644 --- a/yarn.lock +++ b/yarn.lock @@ -3616,7 +3616,6 @@ __metadata: dependencies: "@arethetypeswrong/cli": "npm:^0.17.4" "@endo/promise-kit": "npm:^1.2.1" - "@endo/stream": "npm:^1.3.1" "@metamask/auto-changelog": "npm:^6.2.1" "@metamask/eslint-config": "npm:^15.0.1" "@metamask/eslint-config-nodejs": "npm:^15.0.1"