Skip to content

Commit 0479d04

Browse files
tegefaulkesCMCDragonkai
authored andcommitted
wip adding leading response message to raw RPC calls
1 parent c3fc105 commit 0479d04

8 files changed

Lines changed: 232 additions & 52 deletions

File tree

src/rpc/RPCClient.ts

Lines changed: 67 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,8 @@ import type {
55
JSONRPCRequestMessage,
66
StreamFactory,
77
ClientManifest,
8-
RPCStream,
9-
} from './types';
8+
RPCStream, JSONRPCResponseResult
9+
} from "./types";
1010
import type { JSONValue } from '../types';
1111
import type {
1212
JSONRPCRequest,
@@ -20,7 +20,8 @@ import { Timer } from '@matrixai/timer';
2020
import * as rpcUtilsMiddleware from './utils/middleware';
2121
import * as rpcErrors from './errors';
2222
import * as rpcUtils from './utils/utils';
23-
import { promise } from '../utils';
23+
import { never, promise } from "../utils";
24+
import { parseJSONRPCResponse } from "./utils/utils";
2425

2526
const timerCleanupReasonSymbol = Symbol('timerCleanUpReasonSymbol');
2627

@@ -260,12 +261,14 @@ class RPCClient<M extends ClientManifest> {
260261
const abortRaceProm = promise<never>();
261262
// Prevent unhandled rejection when we're done with the promise
262263
abortRaceProm.p.catch(() => {});
264+
signal.addEventListener('abort', () => {
265+
abortRaceProm.rejectP(signal.reason);
266+
}, { once: true});
263267
let abortHandler: () => void;
264268
if (ctx.signal != null) {
265269
// Propagate signal events
266270
abortHandler = () => {
267271
abortController.abort(ctx.signal?.reason);
268-
abortRaceProm.rejectP(ctx.signal?.reason);
269272
};
270273
if (ctx.signal.aborted) abortHandler();
271274
ctx.signal.addEventListener('abort', abortHandler);
@@ -288,7 +291,6 @@ class RPCClient<M extends ClientManifest> {
288291
void timer.then(
289292
() => {
290293
abortController.abort(timeoutError);
291-
abortRaceProm.rejectP(timeoutError);
292294
},
293295
() => {}, // Ignore cancellation error
294296
);
@@ -384,29 +386,27 @@ class RPCClient<M extends ClientManifest> {
384386
public async rawStreamCaller(
385387
method: string,
386388
headerParams: JSONValue,
387-
ctx: Partial<ContextTimed> = {},
388-
): Promise<RPCStream<Uint8Array, Uint8Array>> {
389+
ctx: Partial<ContextTimedInput> = {},
390+
): Promise<RPCStream<Uint8Array, Uint8Array, Record<string, JSONValue> & {result: JSONValue, command: string}>> {
389391
const abortController = new AbortController();
390392
const signal = abortController.signal;
391-
// A promise that will reject if there is an abort signal or timeout
392-
const abortRaceProm = promise<never>();
393-
// Prevent unhandled rejection when we're done with the promise
394-
abortRaceProm.p.catch(() => {});
395393
let abortHandler: () => void;
396394
if (ctx.signal != null) {
397395
// Propagate signal events
398396
abortHandler = () => {
399397
abortController.abort(ctx.signal?.reason);
400-
abortRaceProm.rejectP(ctx.signal?.reason);
401398
};
402399
if (ctx.signal.aborted) abortHandler();
403400
ctx.signal.addEventListener('abort', abortHandler);
404401
}
405-
const timer =
406-
ctx.timer ??
407-
new Timer({
408-
delay: this.streamKeepAliveTimeoutTime,
402+
let timer: Timer;
403+
if (!(ctx.timer instanceof Timer)) {
404+
timer = new Timer({
405+
delay: ctx.timer ?? this.streamKeepAliveTimeoutTime,
409406
});
407+
} else {
408+
timer = ctx.timer;
409+
}
410410
const cleanUp = () => {
411411
// Clean up the timer and signal
412412
if (ctx.timer == null) timer.cancel(timerCleanupReasonSymbol);
@@ -416,13 +416,22 @@ class RPCClient<M extends ClientManifest> {
416416
void timer.then(
417417
() => {
418418
abortController.abort(timeoutError);
419-
abortRaceProm.rejectP(timeoutError);
420419
},
421420
() => {},
422421
);
423-
let rpcStream: RPCStream<Uint8Array, Uint8Array>;
424-
const setupStream = async () => {
425-
const rpcStream = await this.streamFactory({ signal, timer });
422+
let streamCreation: [JSONValue, RPCStream<Uint8Array, Uint8Array>];
423+
const setupStream = async (): Promise<[JSONValue, RPCStream<Uint8Array, Uint8Array>]> => {
424+
if (signal.aborted) throw signal.reason;
425+
const abortProm = promise<never>();
426+
// ignore error if orphaned
427+
void abortProm.p.catch(() => {});
428+
signal.addEventListener('abort', () => {
429+
abortProm.rejectP(signal.reason);
430+
}, {once: true});
431+
const rpcStream = await Promise.race([
432+
this.streamFactory({ signal, timer }),
433+
abortProm.p,
434+
]);
426435
const tempWriter = rpcStream.writable.getWriter();
427436
const header: JSONRPCRequestMessage = {
428437
jsonrpc: '2.0',
@@ -432,15 +441,51 @@ class RPCClient<M extends ClientManifest> {
432441
};
433442
await tempWriter.write(Buffer.from(JSON.stringify(header)));
434443
tempWriter.releaseLock();
435-
return rpcStream;
444+
const headTransformStream = rpcUtilsMiddleware.binaryToJsonMessageStream(
445+
rpcUtils.parseJSONRPCResponse,
446+
);
447+
void rpcStream.readable
448+
// Allow us to re-use the readable after reading the first message
449+
.pipeTo(headTransformStream.writable, {
450+
preventClose: true,
451+
preventCancel: true,
452+
})
453+
// Ignore any errors here, we only care that it ended
454+
.catch(() => {});
455+
const tempReader = headTransformStream.readable.getReader();
456+
let leadingMessage: JSONRPCResponseResult;
457+
try {
458+
const message = await Promise.race([
459+
tempReader.read(),
460+
abortProm.p,
461+
]);
462+
if (message.done) never();
463+
if ('error'in message.value) {
464+
const metadata = {
465+
...(rpcStream.meta ?? {}),
466+
command: method,
467+
};
468+
throw rpcUtils.toError(message.value.error.data, metadata);
469+
}
470+
leadingMessage = message.value;
471+
} catch (e) {
472+
await tempReader.cancel();
473+
rpcStream.cancel(Error('TMP received error in leading response'));
474+
throw e;
475+
}
476+
// Downgrade back to the raw stream
477+
await tempReader.cancel();
478+
return [leadingMessage.result, rpcStream];
436479
};
437480
try {
438-
rpcStream = await Promise.race([setupStream(), abortRaceProm.p]);
481+
streamCreation = await setupStream();
439482
} finally {
440483
cleanUp();
441484
}
485+
const [result, rpcStream] = streamCreation
442486
const metadata = {
443487
...(rpcStream.meta ?? {}),
488+
result,
444489
command: method,
445490
};
446491
return {

src/rpc/RPCServer.ts

Lines changed: 43 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ 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';
3435
import { never } from '../utils/utils';
3536
import { sysexits } from '../errors';
3637

@@ -344,7 +345,7 @@ class RPCServer extends EventTarget {
344345
});
345346
// Ignore any errors here, it should propagate to the ends of the stream
346347
void reverseMiddlewareStream.pipeTo(reverseStream).catch(() => {});
347-
return middleware.reverse.readable;
348+
return [undefined, middleware.reverse.readable];
348349
};
349350
this.registerRawStreamHandler(method, rawSteamHandler, timeout);
350351
}
@@ -565,12 +566,47 @@ class RPCServer extends EventTarget {
565566
timer.refresh();
566567
}
567568
this.logger.info(`Handling stream with method (${method})`);
568-
const outputStream = handler(
569-
[headerMessage.value, inputStream],
570-
rpcStream.cancel,
571-
rpcStream.meta,
572-
{ signal: abortController.signal, timer },
573-
);
569+
let handlerResult: [JSONValue | undefined, ReadableStream<Uint8Array>];
570+
const headerWriter = rpcStream.writable.getWriter();
571+
try {
572+
handlerResult = handler(
573+
[headerMessage.value, inputStream],
574+
rpcStream.cancel,
575+
rpcStream.meta,
576+
{ signal: abortController.signal, timer },
577+
);
578+
} catch (e) {
579+
const rpcError: JSONRPCError = {
580+
code: e.exitCode ?? sysexits.UNKNOWN,
581+
message: e.description ?? '',
582+
data: rpcUtils.fromError(e, this.sensitive),
583+
};
584+
const rpcErrorMessage: JSONRPCResponseError = {
585+
jsonrpc: '2.0',
586+
error: rpcError,
587+
id: null,
588+
};
589+
await headerWriter.write(Buffer.from(JSON.stringify(rpcErrorMessage)))
590+
// clean up and return
591+
timer.cancel(cleanupReason);
592+
abortController.signal.removeEventListener('abort', handleAbort);
593+
graceTimer?.cancel(cleanupReason);
594+
abortController.abort(new rpcErrors.ErrorRPCStreamEnded());
595+
rpcStream.cancel(Error('TMP header message was an error'));
596+
return;
597+
}
598+
const [leadingResult, outputStream] = handlerResult;
599+
600+
if (leadingResult !== undefined){
601+
// writing leading metadata
602+
const leadingMessage: JSONRPCResponseResult = {
603+
jsonrpc: "2.0",
604+
result: leadingResult,
605+
id: null
606+
};
607+
await headerWriter.write(Buffer.from(JSON.stringify(leadingMessage)));
608+
}
609+
headerWriter.releaseLock();
574610
const outputStreamEndProm = outputStream
575611
.pipeTo(rpcStream.writable)
576612
.catch(() => {}); // Ignore any errors, we only care that it finished

src/rpc/handlers.ts

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ abstract class RawHandler<
2929
cancel: (reason?: any) => void,
3030
meta: Record<string, JSONValue> | undefined,
3131
ctx: ContextTimed,
32-
): ReadableStream<Uint8Array>;
32+
): [JSONValue, ReadableStream<Uint8Array>];
3333
}
3434

3535
abstract class DuplexHandler<

src/rpc/types.ts

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -164,7 +164,7 @@ type HandlerImplementation<I, O> = (
164164

165165
type RawHandlerImplementation = HandlerImplementation<
166166
[JSONRPCRequest, ReadableStream<Uint8Array>],
167-
ReadableStream<Uint8Array>
167+
[JSONValue | undefined, ReadableStream<Uint8Array>]
168168
>;
169169

170170
type DuplexHandlerImplementation<
@@ -264,7 +264,7 @@ type DuplexCallerImplementation<
264264
type RawCallerImplementation = (
265265
headerParams: JSONValue,
266266
ctx?: Partial<ContextTimedInput>,
267-
) => Promise<RPCStream<Uint8Array, Uint8Array>>;
267+
) => Promise<RPCStream<Uint8Array, Uint8Array, Record<string, JSONValue> & {result: JSONValue, command: string}>>;
268268

269269
type ConvertDuplexCaller<T> = T extends DuplexCaller<infer I, infer O>
270270
? DuplexCallerImplementation<I, O>

tests/rpc/RPC.test.ts

Lines changed: 86 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ import type { JSONValue } from '@/types';
44
import { TransformStream } from 'stream/web';
55
import { fc, testProp } from '@fast-check/jest';
66
import Logger, { LogLevel, StreamHandler } from '@matrixai/logger';
7-
import { sleep } from 'ix/asynciterable/_sleep';
7+
import * as utils from '@/utils';
88
import RPCServer from '@/rpc/RPCServer';
99
import RPCClient from '@/rpc/RPCClient';
1010
import {
@@ -43,10 +43,10 @@ describe('RPC', () => {
4343
class TestMethod extends RawHandler {
4444
public handle(
4545
input: [JSONRPCRequest, ReadableStream<Uint8Array>],
46-
): ReadableStream<Uint8Array> {
46+
): [JSONValue, ReadableStream<Uint8Array>] {
4747
const [header_, stream] = input;
4848
header = header_;
49-
return stream;
49+
return ['some leading data', stream];
5050
}
5151
}
5252
const rpcServer = await RPCServer.createRPCServer({
@@ -89,12 +89,94 @@ describe('RPC', () => {
8989
id: null,
9090
};
9191
expect(header).toStrictEqual(expectedHeader);
92+
expect(callerInterface.meta?.result).toBe('some leading data');
9293
expect(await outputResult).toStrictEqual(inputData);
9394
await pipeProm;
9495
await rpcServer.destroy();
9596
await rpcClient.destroy();
9697
},
9798
);
99+
test(
100+
'RPC communication with raw stream times out waiting for leading message',
101+
async () => {
102+
const { clientPair, serverPair } = rpcTestUtils.createTapPairs<
103+
Uint8Array,
104+
Uint8Array
105+
>();
106+
void (async () => {
107+
for await (const _ of serverPair.readable) {
108+
// just consume
109+
}
110+
})();
111+
112+
const rpcClient = await RPCClient.createRPCClient({
113+
manifest: {
114+
testMethod: new RawCaller(),
115+
},
116+
streamFactory: async () => {
117+
return {
118+
...clientPair,
119+
cancel: () => {},
120+
};
121+
},
122+
logger,
123+
});
124+
125+
await expect(rpcClient.methods.testMethod({
126+
hello: 'world',
127+
},
128+
{ timer: 100 },
129+
)).rejects.toThrow(rpcErrors.ErrorRPCTimedOut);
130+
await rpcClient.destroy();
131+
},
132+
);
133+
test(
134+
'RPC communication with raw stream, raw handler throws',
135+
async () => {
136+
const { clientPair, serverPair } = rpcTestUtils.createTapPairs<
137+
Uint8Array,
138+
Uint8Array
139+
>();
140+
141+
class TestMethod extends RawHandler {
142+
public handle(
143+
input: [JSONRPCRequest, ReadableStream<Uint8Array>],
144+
): [JSONValue, ReadableStream<Uint8Array>] {
145+
throw Error('some error');
146+
}
147+
}
148+
const rpcServer = await RPCServer.createRPCServer({
149+
manifest: {
150+
testMethod: new TestMethod({}),
151+
},
152+
logger,
153+
});
154+
rpcServer.handleStream({
155+
...serverPair,
156+
cancel: () => {},
157+
});
158+
159+
const rpcClient = await RPCClient.createRPCClient({
160+
manifest: {
161+
testMethod: new RawCaller(),
162+
},
163+
streamFactory: async () => {
164+
return {
165+
...clientPair,
166+
cancel: () => {},
167+
};
168+
},
169+
logger,
170+
});
171+
172+
await expect(rpcClient.methods.testMethod({
173+
hello: 'world',
174+
})).rejects.toThrow(rpcErrors.ErrorPolykeyRemote)
175+
176+
await rpcServer.destroy();
177+
await rpcClient.destroy();
178+
},
179+
);
98180
testProp(
99181
'RPC communication with duplex stream',
100182
[fc.array(rpcTestUtils.safeJsonValueArb, { minLength: 1 })],
@@ -466,7 +548,7 @@ describe('RPC', () => {
466548
const writer = callerInterface.writable.getWriter();
467549
await writer.write({});
468550
// Allow time to process buffer
469-
await sleep(0);
551+
await utils.sleep(0);
470552
await expect(writer.write({})).toReject();
471553
const reader = callerInterface.readable.getReader();
472554
await expect(reader.read()).toReject();

0 commit comments

Comments
 (0)