@nestjs/websockets
Version:
Nest - modern, fast, powerful node.js web framework (@websockets)
283 lines (282 loc) • 13.3 kB
JavaScript
import { from as fromPromise, isObservable, of, } from 'rxjs';
import { distinctUntilChanged, mergeAll } from 'rxjs/operators';
import { GATEWAY_OPTIONS, PORT_METADATA } from './constants.js';
import { InvalidSocketPortException } from './errors/invalid-socket-port.exception.js';
import { GatewayMetadataExplorer, } from './gateway-metadata-explorer.js';
import { compareElementAt } from './utils/compare-element.util.js';
import { Logger } from '@nestjs/common';
import { ContextIdFactory, MetadataScanner, } from '@nestjs/core';
import { ExecutionContextHost, REQUEST_CONTEXT_ID, STATIC_CONTEXT, } from '@nestjs/core/internal';
export class WebSocketsController {
socketServerProvider;
config;
contextCreator;
container;
injector;
exceptionFiltersContext;
graphInspector;
appOptions;
logger = new Logger(WebSocketsController.name, {
timestamp: true,
});
metadataExplorer = new GatewayMetadataExplorer(new MetadataScanner());
exceptionFiltersCache = new WeakMap();
constructor(socketServerProvider, config, contextCreator, container, injector, exceptionFiltersContext, graphInspector, appOptions = {}) {
this.socketServerProvider = socketServerProvider;
this.config = config;
this.contextCreator = contextCreator;
this.container = container;
this.injector = injector;
this.exceptionFiltersContext = exceptionFiltersContext;
this.graphInspector = graphInspector;
this.appOptions = appOptions;
}
connectGatewayToServer(instanceOrWrapper, metatypeOrModuleKey, moduleKey, instanceWrapperId) {
const isInstanceWrapper = typeof metatypeOrModuleKey === 'string' &&
'instance' in instanceOrWrapper;
const instance = isInstanceWrapper
? instanceOrWrapper.instance
: instanceOrWrapper;
const metatype = isInstanceWrapper
? instanceOrWrapper.metatype
: metatypeOrModuleKey;
const targetModuleKey = isInstanceWrapper
? metatypeOrModuleKey
: moduleKey;
const targetInstanceWrapperId = isInstanceWrapper
? instanceOrWrapper.id
: instanceWrapperId;
const instanceWrapper = isInstanceWrapper
? instanceOrWrapper
: {
instance,
metatype,
id: targetInstanceWrapperId,
isDependencyTreeStatic: () => true,
isDependencyTreeDurable: () => false,
};
const gatewayMetatype = metatype ?? instance.constructor;
const options = Reflect.getMetadata(GATEWAY_OPTIONS, gatewayMetatype) || {};
const port = Reflect.getMetadata(PORT_METADATA, gatewayMetatype) || 0;
if (!Number.isInteger(port)) {
throw new InvalidSocketPortException(port, gatewayMetatype);
}
this.subscribeToServerEvents(instanceWrapper, options, port, targetModuleKey, targetInstanceWrapperId);
}
subscribeToServerEvents(instanceOrWrapper, options, port, moduleKey, instanceWrapperId) {
const instanceWrapper = 'instance' in instanceOrWrapper
? instanceOrWrapper
: {
instance: instanceOrWrapper,
metatype: instanceOrWrapper.constructor,
id: instanceWrapperId,
isDependencyTreeStatic: () => true,
isDependencyTreeDurable: () => false,
};
const { instance } = instanceWrapper;
const nativeMessageHandlers = this.metadataExplorer.explore(instance);
const isStatic = instanceWrapper.isDependencyTreeStatic();
const moduleRef = this.container.getModuleByKey(moduleKey);
const messageHandlers = nativeMessageHandlers.map(({ callback, isAckHandledManually, message, methodName }) => ({
message,
methodName,
callback: isStatic
? this.contextCreator.create(instance, callback, moduleKey, methodName, STATIC_CONTEXT)
: this.createRequestScopedHandler(instanceWrapper, moduleRef, moduleKey, methodName),
isAckHandledManually,
}));
this.inspectEntrypointDefinitions(instance, port, messageHandlers, instanceWrapperId);
if (this.appOptions.preview) {
return;
}
const observableServer = this.socketServerProvider.scanForSocketServer(options, port);
this.assignServerToProperties(instance, observableServer.server);
this.subscribeEvents(instanceWrapper, messageHandlers, observableServer, isStatic
? instance.handleConnection?.bind(instance)
: this.createRequestScopedEventHandler(instanceWrapper, moduleRef, moduleKey, 'handleConnection', observableServer.server), isStatic
? instance.handleDisconnect?.bind(instance)
: this.createRequestScopedEventHandler(instanceWrapper, moduleRef, moduleKey, 'handleDisconnect', observableServer.server));
}
subscribeEvents(instanceWrapper, subscribersMap, observableServer, connectionHandler, disconnectHandler) {
const { instance } = instanceWrapper;
const { init, disconnect, connection, server } = observableServer;
const adapter = this.config.getIoAdapter();
this.subscribeInitEvent(instance, init);
this.subscribeConnectionEvent(connectionHandler, connection);
this.subscribeDisconnectEvent(disconnectHandler, disconnect);
const handler = this.getConnectionHandler(this, instance, subscribersMap, disconnect, connection);
adapter.bindClientConnect(server, handler);
this.printSubscriptionLogs(instance, subscribersMap);
}
getConnectionHandler(context, instance, subscribersMap, disconnect, connection) {
const adapter = this.config.getIoAdapter();
return (...args) => {
const [client] = args;
connection.next(args);
context.subscribeMessages(subscribersMap, client, instance);
const disconnectHook = adapter.bindClientDisconnect;
disconnectHook &&
disconnectHook.call(adapter, client, (reason) => disconnect.next({ client, reason }));
};
}
subscribeInitEvent(instance, event) {
if (instance.afterInit) {
event.subscribe(instance.afterInit.bind(instance));
}
}
subscribeConnectionEvent(handlerOrGateway, event) {
const handler = typeof handlerOrGateway === 'function'
? handlerOrGateway
: handlerOrGateway?.handleConnection?.bind(handlerOrGateway);
if (handler) {
event
.pipe(distinctUntilChanged((prev, curr) => compareElementAt(prev, curr, 0)))
.subscribe((args) => handler(...args));
}
}
subscribeDisconnectEvent(handlerOrGateway, event) {
const handler = typeof handlerOrGateway === 'function'
? handlerOrGateway
: handlerOrGateway?.handleDisconnect?.bind(handlerOrGateway);
if (handler) {
event
.pipe(distinctUntilChanged((prev, curr) => {
const prevClient = prev?.client || prev;
const currClient = curr?.client || curr;
return prevClient === currClient;
}))
.subscribe((data) => {
if (data && typeof data === 'object' && 'client' in data) {
handler(data.client, data.reason);
}
else {
// Backward compatibility: if it's just the client
handler(data);
}
});
}
}
subscribeMessages(subscribersMap, client, instance) {
const adapter = this.config.getIoAdapter();
const handlers = subscribersMap.map(({ callback, message, isAckHandledManually }) => ({
message,
callback: callback.bind(instance, client),
isAckHandledManually,
}));
adapter.bindMessageHandlers(client, handlers, data => fromPromise(this.pickResult(data)).pipe(mergeAll()));
}
async pickResult(deferredResult) {
const result = await deferredResult;
if (isObservable(result)) {
return result;
}
if (result instanceof Promise) {
return fromPromise(result);
}
return of(result);
}
inspectEntrypointDefinitions(instance, port, messageHandlers, instanceWrapperId) {
messageHandlers.forEach(handler => {
this.graphInspector.insertEntrypointDefinition({
type: 'websocket',
methodName: handler.methodName,
className: instance.constructor?.name,
classNodeId: instanceWrapperId,
metadata: {
port,
key: handler.message,
message: handler.message,
},
}, instanceWrapperId);
});
}
createRequestScopedHandler(instanceWrapper, moduleRef, moduleKey, methodName) {
const { instance } = instanceWrapper;
const collection = moduleRef.providers;
const isTreeDurable = instanceWrapper.isDependencyTreeDurable();
return async (...args) => {
const [client] = args;
try {
const contextId = this.getContextId(client, isTreeDurable);
const contextInstance = await this.injector.loadPerContext(instance, moduleRef, collection, contextId);
return this.contextCreator.create(contextInstance, contextInstance[methodName], moduleKey, methodName, contextId, instanceWrapper.id)(...args);
}
catch (err) {
let exceptionFilter = this.exceptionFiltersCache.get(instance[methodName]);
if (!exceptionFilter) {
exceptionFilter = this.exceptionFiltersContext.create(instance, instance[methodName], moduleKey);
this.exceptionFiltersCache.set(instance[methodName], exceptionFilter);
}
const host = new ExecutionContextHost(args);
host.setType('ws');
exceptionFilter.handle(err, host);
}
};
}
createRequestScopedEventHandler(instanceWrapper, moduleRef, moduleKey, methodName, server) {
const { instance } = instanceWrapper;
const collection = moduleRef.providers;
const isTreeDurable = instanceWrapper.isDependencyTreeDurable();
const targetCallback = instance[methodName];
return async (...args) => {
const [client] = args;
let contextId;
try {
contextId = this.getContextId(client, isTreeDurable);
const contextInstance = await this.injector.loadPerContext(instance, moduleRef, collection, contextId);
this.assignServerToProperties(contextInstance, server);
const scopedMethod = contextInstance[methodName];
return scopedMethod?.apply(contextInstance, args);
}
catch (err) {
if (!targetCallback) {
throw err;
}
let exceptionFilter = this.exceptionFiltersCache.get(targetCallback);
if (!exceptionFilter) {
exceptionFilter = this.exceptionFiltersContext.create(instance, targetCallback, moduleKey);
this.exceptionFiltersCache.set(targetCallback, exceptionFilter);
}
const host = new ExecutionContextHost(args);
host.setType('ws');
exceptionFilter.handle(err, host);
}
finally {
if (methodName === 'handleDisconnect' && contextId) {
this.cleanupRequestScopedContext(instanceWrapper, contextId, client);
}
}
};
}
getContextId(request, isTreeDurable) {
const contextId = ContextIdFactory.getByRequest(request);
if (!request[REQUEST_CONTEXT_ID]) {
Object.defineProperty(request, REQUEST_CONTEXT_ID, {
value: contextId,
enumerable: false,
writable: false,
configurable: true,
});
const requestProviderValue = isTreeDurable
? contextId.payload
: Object.assign(request, contextId.payload);
this.container.registerRequestProvider(requestProviderValue, contextId);
}
return contextId;
}
cleanupRequestScopedContext(_instanceWrapper, _contextId, request) {
Reflect.deleteProperty(request, REQUEST_CONTEXT_ID);
}
assignServerToProperties(instance, server) {
for (const propertyKey of this.metadataExplorer.scanForServerHooks(instance)) {
Reflect.set(instance, propertyKey, server);
}
}
printSubscriptionLogs(instance, subscribersMap) {
const gatewayClassName = instance?.constructor?.name;
if (!gatewayClassName) {
return;
}
subscribersMap.forEach(({ message }) => this.logger.log(`${gatewayClassName} subscribed to the "${message}" message`));
}
}