UNPKG

@nestjs/websockets

Version:

Nest - modern, fast, powerful node.js web framework (@websockets)

283 lines (282 loc) 13.3 kB
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`)); } }