Skip to content

Commit 34eaed6

Browse files
tegefaulkesCMCDragonkai
authored andcommitted
feat: added reverse response message to raw duplex streams
1 parent 0479d04 commit 34eaed6

8 files changed

Lines changed: 302 additions & 182 deletions

File tree

src/rpc/RPCClient.ts

Lines changed: 68 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,13 @@
11
import type { WritableStream, ReadableStream } from 'stream/web';
2-
import type { ContextTimed, ContextTimedInput } from '@matrixai/contexts';
2+
import type { ContextTimedInput } from '@matrixai/contexts';
33
import type {
44
HandlerType,
55
JSONRPCRequestMessage,
66
StreamFactory,
77
ClientManifest,
8-
RPCStream, JSONRPCResponseResult
9-
} from "./types";
8+
RPCStream,
9+
JSONRPCResponseResult,
10+
} from './types';
1011
import type { JSONValue } from '../types';
1112
import type {
1213
JSONRPCRequest,
@@ -20,8 +21,7 @@ import { Timer } from '@matrixai/timer';
2021
import * as rpcUtilsMiddleware from './utils/middleware';
2122
import * as rpcErrors from './errors';
2223
import * as rpcUtils from './utils/utils';
23-
import { never, promise } from "../utils";
24-
import { parseJSONRPCResponse } from "./utils/utils";
24+
import { never, promise } from '../utils';
2525

2626
const timerCleanupReasonSymbol = Symbol('timerCleanUpReasonSymbol');
2727

@@ -255,15 +255,18 @@ class RPCClient<M extends ClientManifest> {
255255
method: string,
256256
ctx: Partial<ContextTimedInput> = {},
257257
): Promise<RPCStream<O, I>> {
258+
// Setting up abort signal and timer
258259
const abortController = new AbortController();
259260
const signal = abortController.signal;
260261
// A promise that will reject if there is an abort signal or timeout
261262
const abortRaceProm = promise<never>();
262263
// Prevent unhandled rejection when we're done with the promise
263264
abortRaceProm.p.catch(() => {});
264-
signal.addEventListener('abort', () => {
265+
const abortRacePromHandler = () => {
265266
abortRaceProm.rejectP(signal.reason);
266-
}, { once: true});
267+
};
268+
signal.addEventListener('abort', abortRacePromHandler);
269+
267270
let abortHandler: () => void;
268271
if (ctx.signal != null) {
269272
// Propagate signal events
@@ -284,7 +287,10 @@ class RPCClient<M extends ClientManifest> {
284287
const cleanUp = () => {
285288
// Clean up the timer and signal
286289
if (ctx.timer == null) timer.cancel(timerCleanupReasonSymbol);
287-
signal.removeEventListener('abort', abortHandler);
290+
if (ctx.signal != null) {
291+
ctx.signal.removeEventListener('abort', abortHandler);
292+
}
293+
signal.addEventListener('abort', abortRacePromHandler);
288294
};
289295
// Setting up abort events for timeout
290296
const timeoutError = new rpcErrors.ErrorRPCTimedOut();
@@ -294,6 +300,7 @@ class RPCClient<M extends ClientManifest> {
294300
},
295301
() => {}, // Ignore cancellation error
296302
);
303+
297304
// Hooking up agnostic stream side
298305
let rpcStream: RPCStream<Uint8Array, Uint8Array>;
299306
const streamFactoryProm = this.streamFactory({ signal, timer });
@@ -302,19 +309,10 @@ class RPCClient<M extends ClientManifest> {
302309
} catch (e) {
303310
cleanUp();
304311
void streamFactoryProm.then((stream) =>
305-
stream.cancel('stream timed out early'),
312+
stream.cancel(Error('TMP stream timed out early')),
306313
);
307314
throw e;
308315
}
309-
const cancelStream = () => {
310-
rpcStream.cancel(signal.reason);
311-
};
312-
if (signal.aborted) {
313-
cancelStream();
314-
} else {
315-
signal.addEventListener('abort', cancelStream);
316-
}
317-
// Setting up event for stream timeout
318316
void timer.then(
319317
() => {
320318
rpcStream.cancel(new rpcErrors.ErrorRPCTimedOut());
@@ -356,7 +354,6 @@ class RPCClient<M extends ClientManifest> {
356354
.catch(() => {}),
357355
]).finally(() => {
358356
cleanUp();
359-
signal.removeEventListener('abort', cancelStream);
360357
});
361358

362359
// Returning interface
@@ -387,9 +384,25 @@ class RPCClient<M extends ClientManifest> {
387384
method: string,
388385
headerParams: JSONValue,
389386
ctx: Partial<ContextTimedInput> = {},
390-
): Promise<RPCStream<Uint8Array, Uint8Array, Record<string, JSONValue> & {result: JSONValue, command: string}>> {
387+
): Promise<
388+
RPCStream<
389+
Uint8Array,
390+
Uint8Array,
391+
Record<string, JSONValue> & { result: JSONValue; command: string }
392+
>
393+
> {
394+
// Setting up abort signal and timer
391395
const abortController = new AbortController();
392396
const signal = abortController.signal;
397+
// A promise that will reject if there is an abort signal or timeout
398+
const abortRaceProm = promise<never>();
399+
// Prevent unhandled rejection when we're done with the promise
400+
abortRaceProm.p.catch(() => {});
401+
const abortRacePromHandler = () => {
402+
abortRaceProm.rejectP(signal.reason);
403+
};
404+
signal.addEventListener('abort', abortRacePromHandler);
405+
393406
let abortHandler: () => void;
394407
if (ctx.signal != null) {
395408
// Propagate signal events
@@ -410,24 +423,34 @@ class RPCClient<M extends ClientManifest> {
410423
const cleanUp = () => {
411424
// Clean up the timer and signal
412425
if (ctx.timer == null) timer.cancel(timerCleanupReasonSymbol);
413-
signal.removeEventListener('abort', abortHandler);
426+
if (ctx.signal != null) {
427+
ctx.signal.removeEventListener('abort', abortHandler);
428+
}
429+
signal.addEventListener('abort', abortRacePromHandler);
414430
};
431+
// Setting up abort events for timeout
415432
const timeoutError = new rpcErrors.ErrorRPCTimedOut();
416433
void timer.then(
417434
() => {
418435
abortController.abort(timeoutError);
419436
},
420-
() => {},
437+
() => {}, // Ignore cancellation error
421438
);
422-
let streamCreation: [JSONValue, RPCStream<Uint8Array, Uint8Array>];
423-
const setupStream = async (): Promise<[JSONValue, RPCStream<Uint8Array, Uint8Array>]> => {
439+
440+
const setupStream = async (): Promise<
441+
[JSONValue, RPCStream<Uint8Array, Uint8Array>]
442+
> => {
424443
if (signal.aborted) throw signal.reason;
425444
const abortProm = promise<never>();
426-
// ignore error if orphaned
445+
// Ignore error if orphaned
427446
void abortProm.p.catch(() => {});
428-
signal.addEventListener('abort', () => {
429-
abortProm.rejectP(signal.reason);
430-
}, {once: true});
447+
signal.addEventListener(
448+
'abort',
449+
() => {
450+
abortProm.rejectP(signal.reason);
451+
},
452+
{ once: true },
453+
);
431454
const rpcStream = await Promise.race([
432455
this.streamFactory({ signal, timer }),
433456
abortProm.p,
@@ -441,48 +464,48 @@ class RPCClient<M extends ClientManifest> {
441464
};
442465
await tempWriter.write(Buffer.from(JSON.stringify(header)));
443466
tempWriter.releaseLock();
444-
const headTransformStream = rpcUtilsMiddleware.binaryToJsonMessageStream(
467+
const headTransformStream = rpcUtils.parseHeadStream(
445468
rpcUtils.parseJSONRPCResponse,
446469
);
447470
void rpcStream.readable
448471
// Allow us to re-use the readable after reading the first message
449-
.pipeTo(headTransformStream.writable, {
450-
preventClose: true,
451-
preventCancel: true,
452-
})
472+
.pipeTo(headTransformStream.writable)
453473
// Ignore any errors here, we only care that it ended
454474
.catch(() => {});
455475
const tempReader = headTransformStream.readable.getReader();
456476
let leadingMessage: JSONRPCResponseResult;
457477
try {
458-
const message = await Promise.race([
459-
tempReader.read(),
460-
abortProm.p,
461-
]);
478+
const message = await Promise.race([tempReader.read(), abortProm.p]);
479+
const messageValue = message.value as JSONRPCResponse;
462480
if (message.done) never();
463-
if ('error'in message.value) {
481+
if ('error' in messageValue) {
464482
const metadata = {
465483
...(rpcStream.meta ?? {}),
466484
command: method,
467485
};
468-
throw rpcUtils.toError(message.value.error.data, metadata);
486+
throw rpcUtils.toError(messageValue.error.data, metadata);
469487
}
470-
leadingMessage = message.value;
488+
leadingMessage = messageValue;
471489
} catch (e) {
472-
await tempReader.cancel();
473490
rpcStream.cancel(Error('TMP received error in leading response'));
474491
throw e;
475492
}
476-
// Downgrade back to the raw stream
477-
await tempReader.cancel();
478-
return [leadingMessage.result, rpcStream];
493+
tempReader.releaseLock();
494+
const newRpcStream: RPCStream<Uint8Array, Uint8Array> = {
495+
writable: rpcStream.writable,
496+
readable: headTransformStream.readable as ReadableStream<Uint8Array>,
497+
cancel: rpcStream.cancel,
498+
meta: rpcStream.meta,
499+
};
500+
return [leadingMessage.result, newRpcStream];
479501
};
502+
let streamCreation: [JSONValue, RPCStream<Uint8Array, Uint8Array>];
480503
try {
481504
streamCreation = await setupStream();
482505
} finally {
483506
cleanUp();
484507
}
485-
const [result, rpcStream] = streamCreation
508+
const [result, rpcStream] = streamCreation;
486509
const metadata = {
487510
...(rpcStream.meta ?? {}),
488511
result,

src/rpc/RPCServer.ts

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,6 @@ import * as rpcEvents from './events';
3131
import * as rpcUtils from './utils/utils';
3232
import * as rpcErrors from './errors';
3333
import * as rpcUtilsMiddleware from './utils/middleware';
34-
import * as utils from '../utils/utils';
3534
import { never } from '../utils/utils';
3635
import { sysexits } from '../errors';
3736

@@ -586,8 +585,8 @@ class RPCServer extends EventTarget {
586585
error: rpcError,
587586
id: null,
588587
};
589-
await headerWriter.write(Buffer.from(JSON.stringify(rpcErrorMessage)))
590-
// clean up and return
588+
await headerWriter.write(Buffer.from(JSON.stringify(rpcErrorMessage)));
589+
// Clean up and return
591590
timer.cancel(cleanupReason);
592591
abortController.signal.removeEventListener('abort', handleAbort);
593592
graceTimer?.cancel(cleanupReason);
@@ -597,12 +596,12 @@ class RPCServer extends EventTarget {
597596
}
598597
const [leadingResult, outputStream] = handlerResult;
599598

600-
if (leadingResult !== undefined){
601-
// writing leading metadata
599+
if (leadingResult !== undefined) {
600+
// Writing leading metadata
602601
const leadingMessage: JSONRPCResponseResult = {
603-
jsonrpc: "2.0",
602+
jsonrpc: '2.0',
604603
result: leadingResult,
605-
id: null
604+
id: null,
606605
};
607606
await headerWriter.write(Buffer.from(JSON.stringify(leadingMessage)));
608607
}

src/rpc/types.ts

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -264,7 +264,13 @@ type DuplexCallerImplementation<
264264
type RawCallerImplementation = (
265265
headerParams: JSONValue,
266266
ctx?: Partial<ContextTimedInput>,
267-
) => Promise<RPCStream<Uint8Array, Uint8Array, Record<string, JSONValue> & {result: JSONValue, command: string}>>;
267+
) => Promise<
268+
RPCStream<
269+
Uint8Array,
270+
Uint8Array,
271+
Record<string, JSONValue> & { result: JSONValue; command: string }
272+
>
273+
>;
268274

269275
type ConvertDuplexCaller<T> = T extends DuplexCaller<infer I, infer O>
270276
? DuplexCallerImplementation<I, O>

src/rpc/utils/utils.ts

Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@ import type {
1313
import type { JSONValue } from '../../types';
1414
import type { Timer } from '@matrixai/timer';
1515
import { TransformStream } from 'stream/web';
16+
import { JSONParser } from '@streamparser/json';
1617
import { AbstractError } from '@matrixai/errors';
1718
import * as rpcErrors from '../errors';
1819
import * as utils from '../../utils';
@@ -429,6 +430,72 @@ function getHandlerTypes(
429430
return out;
430431
}
431432

433+
/**
434+
* This function is a factory to create a TransformStream that will
435+
* transform a `Uint8Array` stream to a JSONRPC message stream.
436+
* The parsed messages will be validated with the provided messageParser, this
437+
* also infers the type of the stream output.
438+
* @param messageParser - Validates the JSONRPC messages, so you can select for a
439+
* specific type of message
440+
* @param bufferByteLimit - sets the number of bytes buffered before throwing an
441+
* error. This is used to avoid infinitely buffering the input.
442+
*/
443+
function parseHeadStream<T extends JSONRPCMessage>(
444+
messageParser: (message: unknown) => T,
445+
bufferByteLimit: number = 1024 * 1024,
446+
): TransformStream<Uint8Array, T | Uint8Array> {
447+
const parser = new JSONParser({
448+
separator: '',
449+
paths: ['$'],
450+
});
451+
let bytesWritten: number = 0;
452+
let parsing = true;
453+
let ended = false;
454+
455+
const endP = utils.promise();
456+
parser.onEnd = () => endP.resolveP();
457+
458+
return new TransformStream<Uint8Array, T | Uint8Array>(
459+
{
460+
flush: async () => {
461+
if (!parser.isEnded) parser.end();
462+
await endP.p;
463+
},
464+
start: (controller) => {
465+
parser.onValue = async (value) => {
466+
const jsonMessage = messageParser(value.value);
467+
controller.enqueue(jsonMessage);
468+
bytesWritten = 0;
469+
parsing = false;
470+
};
471+
},
472+
transform: async (chunk, controller) => {
473+
if (parsing) {
474+
try {
475+
bytesWritten += chunk.byteLength;
476+
parser.write(chunk);
477+
} catch (e) {
478+
throw new rpcErrors.ErrorRPCParse(undefined, { cause: e });
479+
}
480+
if (bytesWritten > bufferByteLimit) {
481+
throw new rpcErrors.ErrorRPCMessageLength();
482+
}
483+
} else {
484+
// Wait for parser to end
485+
if (!ended) {
486+
parser.end();
487+
await endP.p;
488+
ended = true;
489+
}
490+
// Pass through normal chunks
491+
controller.enqueue(chunk);
492+
}
493+
},
494+
},
495+
{ highWaterMark: 1 },
496+
);
497+
}
498+
432499
export {
433500
parseJSONRPCRequest,
434501
parseJSONRPCRequestMessage,
@@ -442,4 +509,5 @@ export {
442509
clientInputTransformStream,
443510
clientOutputTransformStream,
444511
getHandlerTypes,
512+
parseHeadStream,
445513
};

0 commit comments

Comments
 (0)