export type SocketState = "connecting" | "open" | "closing" | "closed"; export interface WebSocketClient { readonly state: SocketState; connect(): Promise; send(message: string): void; close(code?: number, reason?: string): void; onOpen(listener: () => void): () => void; onMessage(listener: (message: string) => void): () => void; onClose(listener: (code: number, reason: string) => void): () => void; onError(listener: (error: Event | Error) => void): () => void; } type ListenerMap = { open: Set<() => void>; message: Set<(message: string) => void>; close: Set<(code: number, reason: string) => void>; error: Set<(error: Event | Error) => void>; }; export type WebSocketConstructor = new ( url: string, protocols?: string | string[], ) => WebSocket; export class NestjWebSocketClient implements WebSocketClient { private socket?: WebSocket; private listeners: ListenerMap = { open: new Set(), message: new Set(), close: new Set(), error: new Set(), }; private currentState: SocketState = "closed"; get state(): SocketState { return this.currentState; } constructor( private readonly url: string, private readonly protocols?: string | string[], private readonly socketConstructor: WebSocketConstructor = WebSocket, ) {} async connect(): Promise { if ( this.socket && (this.currentState === "open" || this.currentState === "connecting") ) return; this.currentState = "connecting"; await new Promise((resolve, reject) => { const socket = new this.socketConstructor(this.url, this.protocols); this.socket = socket; socket.addEventListener("open", () => { this.currentState = "open"; this.listeners.open.forEach((f) => f()); resolve(); }, { once: true }); socket.addEventListener( "message", (event) => this.listeners.message.forEach((f) => f(String(event.data)) ), ); socket.addEventListener("close", (event) => { this.currentState = "closed"; this.listeners.close.forEach((f) => f(event.code, event.reason) ); }); socket.addEventListener("error", (event) => { this.listeners.error.forEach((f) => f(event)); reject(new Error("WebSocket connection failed")); }, { once: true }); }); } send(message: string): void { if (this.currentState !== "open") { throw new Error("WebSocket is not open"); } this.socket!.send(message); } close(code = 1000, reason = ""): void { this.currentState = "closing"; this.socket?.close(code, reason); } onOpen(listener: () => void) { return this.subscribe(this.listeners.open, listener); } onMessage(listener: (message: string) => void) { return this.subscribe(this.listeners.message, listener); } onClose(listener: (code: number, reason: string) => void) { return this.subscribe(this.listeners.close, listener); } onError(listener: (error: Event | Error) => void) { return this.subscribe(this.listeners.error, listener); } private subscribe(listeners: Set, listener: T): () => void { listeners.add(listener); return () => listeners.delete(listener); } } export function createWebSocketClient( url: string, protocols?: string | string[], ): WebSocketClient { return new NestjWebSocketClient(url, protocols); }