Skip to content

Commit 70bd656

Browse files
tegefaulkesCMCDragonkai
authored andcommitted
refactor: combined the websocket client and server stream implementation
* Related #540 [ci skip]
1 parent b73d10a commit 70bd656

5 files changed

Lines changed: 272 additions & 461 deletions

File tree

src/PolykeyAgent.ts

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,6 @@ type NetworkConfig = {
5555
clientHost?: string;
5656
clientPort?: number;
5757
// Websocket server config
58-
maxReadableStreamBytes?: number;
5958
maxIdleTimeout?: number;
6059
pingIntervalTime?: number;
6160
pingTimeoutTimeTime?: number;
@@ -496,11 +495,9 @@ class PolykeyAgent {
496495
(await WebSocketServer.createWebSocketServer({
497496
connectionCallback: (rpcStream) =>
498497
rpcServerClient!.handleStream(rpcStream),
499-
fs,
500498
host: networkConfig_.clientHost,
501499
port: networkConfig_.clientPort,
502500
tlsConfig,
503-
maxReadableStreamBytes: networkConfig_.maxReadableStreamBytes,
504501
maxIdleTimeout: networkConfig_.maxIdleTimeout,
505502
pingIntervalTime: networkConfig_.pingIntervalTime,
506503
pingTimeoutTimeTime: networkConfig_.pingTimeoutTimeTime,

src/config.ts

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -98,7 +98,6 @@ const config = {
9898
clientHost: '127.0.0.1',
9999
clientPort: 0,
100100
// Websocket server config
101-
maxReadableStreamBytes: 1_000_000_000, // About 1 GB
102101
maxIdleTimeout: 120, // 2 minutes
103102
pingIntervalTime: 1_000, // 1 second
104103
pingTimeoutTimeTime: 10_000, // 10 seconds

src/websockets/WebSocketClient.ts

Lines changed: 24 additions & 238 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,6 @@
11
import type { TLSSocket } from 'tls';
2-
import type {
3-
ReadableStreamController,
4-
WritableStreamDefaultController,
5-
} from 'stream/web';
62
import type { ContextTimed } from '@matrixai/contexts';
73
import type { NodeId, NodeIdEncoded } from '../ids';
8-
import type { JSONValue } from '../types';
9-
import { WritableStream, ReadableStream } from 'stream/web';
104
import { createDestroy } from '@matrixai/async-init';
115
import Logger from '@matrixai/logger';
126
import WebSocket from 'ws';
@@ -33,7 +27,6 @@ class WebSocketClient {
3327
* Default is 1,000 milliseconds.
3428
* @param obj.pingTimeoutTimeTime - Time before connection is cleaned up after no ping responses.
3529
* Default is 10,000 milliseconds.
36-
* @param obj.maxReadableStreamBytes - The number of bytes the readable stream will buffer until pausing.
3730
* @param obj.logger
3831
*/
3932
static async createWebSocketClient({
@@ -43,7 +36,6 @@ class WebSocketClient {
4336
connectionTimeoutTime = Infinity,
4437
pingIntervalTime = 1_000,
4538
pingTimeoutTimeTime = 10_000,
46-
maxReadableStreamBytes = 1_000, // About 1kB
4739
logger = new Logger(this.name),
4840
}: {
4941
host: string;
@@ -52,15 +44,13 @@ class WebSocketClient {
5244
connectionTimeoutTime?: number;
5345
pingIntervalTime?: number;
5446
pingTimeoutTimeTime?: number;
55-
maxReadableStreamBytes?: number;
5647
logger?: Logger;
5748
}): Promise<WebSocketClient> {
5849
logger.info(`Creating ${this.name}`);
5950
const clientClient = new this(
6051
logger,
6152
host,
6253
port,
63-
maxReadableStreamBytes,
6454
expectedNodeIds,
6555
connectionTimeoutTime,
6656
pingIntervalTime,
@@ -77,7 +67,6 @@ class WebSocketClient {
7767
protected logger: Logger,
7868
host: string,
7969
protected port: number,
80-
protected maxReadableStreamBytes: number,
8170
protected expectedNodeIds: Array<NodeId>,
8271
protected connectionTimeoutTime: number,
8372
protected pingIntervalTime: number,
@@ -126,7 +115,7 @@ class WebSocketClient {
126115
@createDestroy.ready(new webSocketErrors.ErrorClientDestroyed())
127116
public async startConnection(
128117
ctx: Partial<ContextTimed> = {},
129-
): Promise<WebSocketStreamClientInternal> {
118+
): Promise<WebSocketStream> {
130119
// Setting up abort/cancellation logic
131120
const abortRaceProm = promise<never>();
132121
// Ignore unhandled rejection
@@ -161,7 +150,13 @@ class WebSocketClient {
161150
const address = `wss://${this.host}:${this.port}`;
162151
this.logger.info(`Connecting to ${address}`);
163152
const connectProm = promise<void>();
164-
const authenticateProm = promise<NodeId>();
153+
const authenticateProm = promise<{
154+
nodeId: NodeIdEncoded;
155+
localHost: string;
156+
localPort: number;
157+
remoteHost: string;
158+
remotePort: number;
159+
}>();
165160
const ws = new WebSocket(address, {
166161
rejectUnauthorized: false,
167162
});
@@ -178,12 +173,21 @@ class WebSocketClient {
178173
ws.once('upgrade', async (request) => {
179174
const tlsSocket = request.socket as TLSSocket;
180175
const peerCert = tlsSocket.getPeerCertificate(true);
181-
webSocketUtils
182-
.verifyServerCertificateChain(
176+
try {
177+
const nodeId = await webSocketUtils.verifyServerCertificateChain(
183178
this.expectedNodeIds,
184179
webSocketUtils.detailedToCertChain(peerCert),
185-
)
186-
.then(authenticateProm.resolveP, authenticateProm.rejectP);
180+
);
181+
authenticateProm.resolveP({
182+
nodeId: nodesUtils.encodeNodeId(nodeId),
183+
localHost: request.connection.localAddress ?? '',
184+
localPort: request.connection.localPort ?? 0,
185+
remoteHost: request.connection.remoteAddress ?? '',
186+
remotePort: request.connection.remotePort ?? 0,
187+
});
188+
} catch (e) {
189+
authenticateProm.rejectP(e);
190+
}
187191
});
188192
ws.once('open', () => {
189193
this.logger.info('starting connection');
@@ -222,17 +226,14 @@ class WebSocketClient {
222226

223227
// Constructing the `ReadableWritablePair`, the lifecycle is handed off to
224228
// the webSocketStream at this point.
225-
const webSocketStreamClient = new WebSocketStreamClientInternal(
229+
const webSocketStreamClient = new WebSocketStream(
226230
ws,
227-
this.maxReadableStreamBytes,
228231
this.pingIntervalTime,
229232
this.pingTimeoutTimeTime,
230233
{
231-
host: this.host,
232-
nodeId: nodesUtils.encodeNodeId(await authenticateProm.p),
233-
port: this.port,
234+
...(await authenticateProm.p),
234235
},
235-
this.logger,
236+
this.logger.getChild(WebSocketStream.name),
236237
);
237238
const abortStream = () => {
238239
webSocketStreamClient.cancel(
@@ -258,219 +259,4 @@ class WebSocketClient {
258259
}
259260

260261
// This is the internal implementation of the client's stream pair.
261-
class WebSocketStreamClientInternal extends WebSocketStream {
262-
protected readableController:
263-
| ReadableStreamController<Uint8Array>
264-
| undefined;
265-
protected writableController: WritableStreamDefaultController | undefined;
266-
267-
constructor(
268-
protected ws: WebSocket,
269-
maxReadableStreamBytes: number,
270-
pingInterval: number,
271-
pingTimeoutTime: number,
272-
protected clientMetadata: {
273-
nodeId: NodeIdEncoded;
274-
host: string;
275-
port: number;
276-
},
277-
logger: Logger,
278-
) {
279-
super();
280-
const readableLogger = logger.getChild('readable');
281-
const writableLogger = logger.getChild('writable');
282-
283-
this.readable = new ReadableStream<Uint8Array>(
284-
{
285-
start: (controller) => {
286-
this.readableController = controller;
287-
readableLogger.info('Starting');
288-
const messageHandler = (data) => {
289-
readableLogger.debug(`Received ${data.toString()}`);
290-
if (controller.desiredSize == null) {
291-
controller.error(Error('NEVER'));
292-
return;
293-
}
294-
if (controller.desiredSize < 0) {
295-
readableLogger.debug('Applying readable backpressure');
296-
ws.pause();
297-
}
298-
const message = data as Buffer;
299-
if (message.length === 0) {
300-
readableLogger.debug('Null message received');
301-
ws.removeListener('message', messageHandler);
302-
if (!this._readableEnded) {
303-
this.signalReadableEnd();
304-
readableLogger.debug('Closing');
305-
controller.close();
306-
}
307-
if (this._writableEnded) {
308-
logger.debug('Closing socket');
309-
ws.close();
310-
}
311-
return;
312-
}
313-
controller.enqueue(message);
314-
};
315-
readableLogger.debug('Registering socket message handler');
316-
ws.on('message', messageHandler);
317-
ws.once('close', (code, reason) => {
318-
logger.info('Socket closed');
319-
ws.removeListener('message', messageHandler);
320-
if (!this._readableEnded) {
321-
readableLogger.debug(
322-
`Closed early, ${code}, ${reason.toString()}`,
323-
);
324-
const e = new webSocketErrors.ErrorClientConnectionEndedEarly();
325-
this.signalReadableEnd(e);
326-
controller.error(e);
327-
}
328-
});
329-
ws.once('error', (e) => {
330-
if (!this._readableEnded) {
331-
readableLogger.error(e);
332-
this.signalReadableEnd(e);
333-
controller.error(e);
334-
}
335-
});
336-
},
337-
cancel: (reason) => {
338-
readableLogger.debug('Cancelled');
339-
this.signalReadableEnd(reason);
340-
if (!this._writableEnded) {
341-
readableLogger.debug('Closing socket');
342-
this.signalWritableEnd(reason);
343-
ws.close();
344-
}
345-
},
346-
pull: () => {
347-
readableLogger.debug('Releasing backpressure');
348-
ws.resume();
349-
},
350-
},
351-
{
352-
highWaterMark: maxReadableStreamBytes,
353-
size: (chunk) => chunk?.byteLength ?? 0,
354-
},
355-
);
356-
this.writable = new WritableStream<Uint8Array>({
357-
start: (controller) => {
358-
this.writableController = controller;
359-
writableLogger.info('Starting');
360-
ws.once('error', (e) => {
361-
if (!this._writableEnded) {
362-
writableLogger.error(e);
363-
this.signalWritableEnd(e);
364-
controller.error(e);
365-
}
366-
});
367-
ws.once('close', (code, reason) => {
368-
if (!this._writableEnded) {
369-
writableLogger.debug(`Closed early, ${code}, ${reason.toString()}`);
370-
const e = new webSocketErrors.ErrorClientConnectionEndedEarly();
371-
this.signalWritableEnd(e);
372-
controller.error(e);
373-
}
374-
});
375-
},
376-
close: () => {
377-
writableLogger.debug('Closing, sending null message');
378-
ws.send(Buffer.from([]));
379-
this.signalWritableEnd();
380-
if (this._readableEnded) {
381-
writableLogger.debug('Closing socket');
382-
ws.close();
383-
}
384-
},
385-
abort: (reason) => {
386-
writableLogger.debug('Aborted');
387-
this.signalWritableEnd(reason);
388-
if (this._readableEnded) {
389-
writableLogger.debug('Closing socket');
390-
ws.close();
391-
}
392-
},
393-
write: async (chunk, controller) => {
394-
if (this._writableEnded) return;
395-
writableLogger.debug(`Sending ${chunk?.toString()}`);
396-
const wait = promise<void>();
397-
ws.send(chunk, (e) => {
398-
if (e != null && !this._writableEnded) {
399-
// Opting to debug message here and not log an error, sending
400-
// failure is common if we send before the close event.
401-
writableLogger.debug('failed to send');
402-
const err = new webSocketErrors.ErrorClientConnectionEndedEarly(
403-
undefined,
404-
{
405-
cause: e,
406-
},
407-
);
408-
this.signalWritableEnd(err);
409-
controller.error(err);
410-
}
411-
wait.resolveP();
412-
});
413-
await wait.p;
414-
},
415-
});
416-
417-
// Setting up heartbeat
418-
const pingTimer = setInterval(() => {
419-
ws.ping();
420-
}, pingInterval);
421-
const pingTimeoutTimeTimer = setTimeout(() => {
422-
logger.debug('Ping timed out');
423-
ws.close(4002, 'Timed out');
424-
}, pingTimeoutTime);
425-
ws.on('ping', () => {
426-
logger.debug('Received ping');
427-
ws.pong();
428-
});
429-
ws.on('pong', () => {
430-
logger.debug('Received pong');
431-
pingTimeoutTimeTimer.refresh();
432-
});
433-
ws.once('close', (code, reason) => {
434-
logger.debug('WebSocket closed');
435-
const err =
436-
code !== 1000
437-
? new webSocketErrors.ErrorClientConnectionEndedEarly(
438-
`ended with code ${code}, ${reason.toString()}`,
439-
)
440-
: undefined;
441-
this.signalWebSocketEnd(err);
442-
logger.debug('Cleaning up timers');
443-
// Clean up timers
444-
clearTimeout(pingTimer);
445-
clearTimeout(pingTimeoutTimeTimer);
446-
});
447-
}
448-
449-
get meta(): Record<string, JSONValue> {
450-
// Spreading to avoid modifying the data
451-
return {
452-
...this.clientMetadata,
453-
};
454-
}
455-
456-
cancel(reason?: any): void {
457-
// Default error
458-
const err = reason ?? new webSocketErrors.ErrorClientConnectionEndedEarly();
459-
// Close the streams with the given error,
460-
if (!this._readableEnded) {
461-
this.readableController?.error(err);
462-
this.signalReadableEnd(err);
463-
}
464-
if (!this._writableEnded) {
465-
this.writableController?.error(err);
466-
this.signalWritableEnd(err);
467-
}
468-
// Then close the websocket
469-
if (!this._webSocketEnded) {
470-
this.ws.close(4000, 'Ending connection');
471-
this.signalWebSocketEnd(err);
472-
}
473-
}
474-
}
475-
476262
export default WebSocketClient;

0 commit comments

Comments
 (0)