From fb12826801c13d69234d3a488ad5deb2c4956058 Mon Sep 17 00:00:00 2001 From: devlikepro Date: Tue, 24 Jun 2025 16:36:13 +0700 Subject: [PATCH] [core] Refactor websocket authorization - use DI --- src/core/api/websocket.gateway.core.ts | 32 +++++++++++++++++----- src/core/app.module.core.ts | 2 ++ src/core/auth/WebSocketAuth.ts | 37 ++++++++++++++++---------- 3 files changed, 51 insertions(+), 20 deletions(-) diff --git a/src/core/api/websocket.gateway.core.ts b/src/core/api/websocket.gateway.core.ts index 27f01d62..16411c50 100644 --- a/src/core/api/websocket.gateway.core.ts +++ b/src/core/api/websocket.gateway.core.ts @@ -22,10 +22,18 @@ import { IncomingMessage } from 'http'; import * as url from 'url'; import { Server } from 'ws'; +export enum WebSocketCloseCode { + NORMAL = 1000, + GOING_AWAY = 1001, + PROTOCOL_ERROR = 1002, + UNSUPPORTED_DATA = 1003, + POLICY_VIOLATION = 1008, + INTERNAL_ERROR = 1011, +} + @WebSocketGateway({ path: '/ws', cors: true, - verifyClient: new WebSocketAuth().verifyClient, }) export class WebsocketGatewayCore implements @@ -43,7 +51,10 @@ export class WebsocketGatewayCore private heartbeat: WebsocketHeartbeatJob; private eventUnmask = new EventWildUnmask(WAHAEvents, WAHAEventsWild); - constructor(private manager: SessionManager) { + constructor( + private manager: SessionManager, + private auth: WebSocketAuth, + ) { this.logger = new Logger('WebsocketGateway'); this.heartbeat = new WebsocketHeartbeatJob( this.logger, @@ -53,14 +64,23 @@ export class WebsocketGatewayCore handleConnection(socket: WebSocket, request: IncomingMessage, ...args): any { // wsc - websocket client - const id = generatePrefixedId('wsc'); - socket.id = id; - this.logger.debug(`New client connected: ${request.url}`); + socket.id = generatePrefixedId('wsc'); + + if (!this.auth.validateRequest(request)) { + // Not authorized - close connection + socket.close(WebSocketCloseCode.POLICY_VIOLATION, 'Unauthorized'); + this.logger.debug( + `Unauthorized websocket connection attempt: ${request.url} - ${socket.id}`, + ); + return; + } + + this.logger.debug(`New client connected: ${request.url} - ${socket.id}`); const params = this.getParams(request); const session: string = params.session; const events: WAHAEvents[] = params.events; this.logger.debug( - `Client connected to session: '${session}', events: ${events}, ${id}`, + `Client connected to session: '${session}', events: ${events}, ${socket.id}`, ); const sub = this.manager diff --git a/src/core/app.module.core.ts b/src/core/app.module.core.ts index df51bb9d..b3e30ca4 100644 --- a/src/core/app.module.core.ts +++ b/src/core/app.module.core.ts @@ -16,6 +16,7 @@ import { import { WebsocketGatewayCore } from '@waha/core/api/websocket.gateway.core'; import { AuthMiddleware } from '@waha/core/auth/auth.middleware'; import { BasicAuthFunction } from '@waha/core/auth/basicAuth'; +import { WebSocketAuth } from '@waha/core/auth/WebSocketAuth'; import { GowsEngineConfigService } from '@waha/core/config/GowsEngineConfigService'; import { WebJSEngineConfigService } from '@waha/core/config/WebJSEngineConfigService'; import { MediaLocalStorageModule } from '@waha/core/media/local/media.local.storage.module'; @@ -173,6 +174,7 @@ const PROVIDERS = [ EngineConfigService, WebsocketGatewayCore, MediaLocalStorageConfig, + WebSocketAuth, ]; @Module({ diff --git a/src/core/auth/WebSocketAuth.ts b/src/core/auth/WebSocketAuth.ts index 36be62e7..37133f75 100644 --- a/src/core/auth/WebSocketAuth.ts +++ b/src/core/auth/WebSocketAuth.ts @@ -1,30 +1,39 @@ +import { Injectable } from '@nestjs/common'; +import { WhatsappConfigService } from '@waha/config.service'; +import { IncomingMessage } from 'http'; import * as url from 'url'; import { validateApiKey } from './apiKey.strategy'; +@Injectable() export class WebSocketAuth { - private key: string; + private readonly key: string; - constructor() { - this.key = process.env.WHATSAPP_API_KEY || process.env.WAHA_API_KEY || ''; + constructor(private config: WhatsappConfigService) { + this.key = this.config.getApiKey(); } - verifyClient = (info: any, callback: any) => { + validateRequest(request: IncomingMessage) { if (!this.key) { - callback(true); - return; + return true; } - // Do something with the info - let query = url.parse(info.req.url, true).query; - // case insensitive + const provided = this.getKeyFromQueryParams(request); + return validateApiKey(provided, this.key); + } + + private getKeyFromQueryParams(request: IncomingMessage) { + let query = url.parse(request.url, true).query; + // case-insensitive query params query = Object.keys(query).reduce((acc, key) => { acc[key.toLowerCase()] = query[key]; return acc; }, {}); - const apiKey = query['x-api-key']; - // @ts-ignore - const isValid = validateApiKey(apiKey, this.key); - callback(isValid); - }; + const provided = query['x-api-key']; + // Check if it's array - return first + if (Array.isArray(provided)) { + return provided[0]; + } + return provided; + } }