File size: 11,951 Bytes
00a912e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
/**
 * WebSocket server wrapper — provides a socket.io-compatible API on top of
 * the lightweight `ws` library.
 *
 * Protocol: JSON messages of the form `{ "event": "<name>", "data": <payload> }`.
 *
 * Features replicated from socket.io:
 *   - Typed event emitter (on/emit)
 *   - Room management (join/leave/to/except)
 *   - Middleware (connection handshake guard)
 *   - socket.data, socket.id, socket.handshake
 *   - Graceful close
 */

import { createServer, type IncomingMessage, type Server as HttpServer } from 'node:http'
import { TIMING } from '@music-together/shared'
import { nanoid } from 'nanoid'
import { WebSocket, WebSocketServer, type RawData, type AddressInfo } from 'ws'
import { logger } from './utils/logger.js'

// ---------------------------------------------------------------------------
// Types
// ---------------------------------------------------------------------------

export interface Handshake {
  headers: Record<string, string | string[] | undefined>
}

type EventHandler = (...args: any[]) => void

export class TypedSocket<
  ClientToServerEvents extends Record<string, any> = Record<string, any>,
  ServerToClientEvents extends Record<string, any> = Record<string, any>,
  SocketData extends Record<string, any> = Record<string, any>,
> {
  readonly id: string
  readonly handshake: Handshake
  data: SocketData = {} as SocketData

  /** @internal */
  public readonly ws: WebSocket
  private handlers = new Map<string, Set<EventHandler>>()
  private rooms = new Set<string>()
  /** Reference to the server, for room/broadcast operations */
  private server: TypedServer<ClientToServerEvents, ServerToClientEvents, SocketData>

  constructor(
    ws: WebSocket,
    req: IncomingMessage,
    server: TypedServer<ClientToServerEvents, ServerToClientEvents, SocketData>,
  ) {
    this.id = nanoid(12)
    this.ws = ws
    this.server = server
    this.handshake = { headers: req.headers as Record<string, string | string[] | undefined> }

    this.ws.on('message', (raw: RawData) => {
      try {
        const msg = JSON.parse(raw.toString())
        if (msg && typeof msg.event === 'string') {
          this.dispatch(msg.event, msg.data)
        }
      } catch (err) {
        logger.warn('Failed to parse WebSocket message', { socketId: this.id })
      }
    })

    this.ws.on('close', () => {
      this.dispatch('disconnect', 'transport close')
      this.server.removeSocket(this)
    })

    this.ws.on('error', (err) => {
      logger.warn('WebSocket error', { socketId: this.id, error: err.message })
    })
  }

  // -- Event emitter --------------------------------------------------------

  on<E extends keyof ClientToServerEvents & string>(event: E, handler: EventHandler): this
  on(event: 'disconnect', handler: (reason: string) => void): this
  on(event: string, handler: EventHandler): this {
    let set = this.handlers.get(event)
    if (!set) {
      set = new Set()
      this.handlers.set(event, set)
    }
    set.add(handler)
    return this
  }

  off<E extends keyof ClientToServerEvents & string>(event: E, handler: EventHandler): this
  off(event: 'disconnect', handler: (reason: string) => void): this
  off(event: string, handler: EventHandler): this {
    this.handlers.get(event)?.delete(handler)
    return this
  }

  emit<E extends keyof ServerToClientEvents & string>(event: E, ...args: any[]): boolean {
    if (this.ws.readyState !== WebSocket.OPEN) return false
    const data = args.length <= 1 ? args[0] : args
    this.ws.send(JSON.stringify({ event, data }))
    return true
  }

  private dispatch(event: string, ...args: any[]): void {
    const set = this.handlers.get(event)
    if (!set) return
    for (const handler of set) {
      try {
        handler.call(this, ...args)
      } catch (err) {
        logger.error('Event handler error', err, { socketId: this.id, event })
      }
    }
  }

  // -- Room management ------------------------------------------------------

  join(room: string): void {
    this.rooms.add(room)
    this.server.addSocketToRoom(room, this)
  }

  leave(room: string): void {
    this.rooms.delete(room)
    this.server.removeSocketFromRoom(room, this)
  }

  /** Returns a broadcaster that sends to all sockets in `room` except this one */
  to(room: string): Broadcaster<ServerToClientEvents> {
    return new Broadcaster(this.server, room, this.id)
  }

  // -- Connection control ---------------------------------------------------

  get connected(): boolean {
    return this.ws.readyState === WebSocket.OPEN
  }

  disconnect(close = true): void {
    if (close) {
      this.ws.close()
    }
  }

  // -- Internal (used by TypedServer) ---------------------------------------

  /** @internal */
  _sendToRoom(room: string, event: string, data: any, exceptSocketId?: string): void {
    if (this.ws.readyState !== WebSocket.OPEN) return
    // This socket sends to its own room members excluding itself
    // (used by socket.to(room).emit())
  }

  /** @internal */
  _rooms(): ReadonlySet<string> {
    return this.rooms
  }
}

// ---------------------------------------------------------------------------
// Broadcaster — returned by .to(room)
// ---------------------------------------------------------------------------

class Broadcaster<ServerToClientEvents extends Record<string, any>> {
  private server: TypedServer<any, ServerToClientEvents, any>
  private room: string
  private exceptId?: string

  constructor(server: TypedServer<any, ServerToClientEvents, any>, room: string, exceptId?: string) {
    this.server = server
    this.room = room
    this.exceptId = exceptId
  }

  except(socketId: string): this {
    this.exceptId = socketId
    return this
  }

  emit<E extends keyof ServerToClientEvents & string>(event: E, ...args: any[]): void {
    const data = args.length <= 1 ? args[0] : args
    const msg = JSON.stringify({ event, data })
    const sockets = this.server.getSocketsInRoom(this.room)
    for (const s of sockets) {
      if (this.exceptId && s.id === this.exceptId) continue
      if (s.ws.readyState === WebSocket.OPEN) {
        s.ws.send(msg)
      }
    }
  }
}

// ---------------------------------------------------------------------------
// TypedServer
// ---------------------------------------------------------------------------

type MiddlewareFn<SD extends Record<string, any>> = (
  socket: TypedSocket<any, any, SD>,
  next: (err?: Error) => void,
) => void

export class TypedServer<
  ClientToServerEvents extends Record<string, any> = Record<string, any>,
  ServerToClientEvents extends Record<string, any> = Record<string, any>,
  SocketData extends Record<string, any> = Record<string, any>,
> {
  private wss: WebSocketServer
  private sockets = new Set<TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>>()
  private roomMap = new Map<string, Set<TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>>>()
  private middlewares: MiddlewareFn<SocketData>[] = []
  private connectionHandlers: ((socket: TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>) => void)[] = []
  private heartbeatTimer: ReturnType<typeof setInterval>
  private aliveSockets = new Map<WebSocket, boolean>()

  constructor(httpServer: HttpServer) {
    this.wss = new WebSocketServer({ noServer: true })

    httpServer.on('upgrade', (request, socket, head) => {
      const pathname = new URL(request.url || '', `http://${request.headers.host}`).pathname

      if (pathname === '/ws') {
        this.wss.handleUpgrade(request, socket, head, (ws) => {
          this.wss.emit('connection', ws, request)
        })
      } else {
        socket.destroy()
      }
    })

    this.wss.on('connection', (ws: WebSocket, req: IncomingMessage) => {
      this.aliveSockets.set(ws, true)
      ws.on('pong', () => this.aliveSockets.set(ws, true))
      this.handleConnection(ws, req)
    })

    this.heartbeatTimer = setInterval(() => {
      for (const socket of this.sockets) {
        const ws = socket.ws
        if (this.aliveSockets.get(ws) === false) {
          logger.warn('WebSocket heartbeat timed out', { socketId: socket.id })
          ws.terminate()
          continue
        }
        this.aliveSockets.set(ws, false)
        ws.ping()
      }
    }, TIMING.WEBSOCKET_HEARTBEAT_INTERVAL_MS)
    this.heartbeatTimer.unref()
  }

  // -- Middleware ------------------------------------------------------------

  use(fn: MiddlewareFn<SocketData>): this {
    this.middlewares.push(fn)
    return this
  }

  // -- Connection event -----------------------------------------------------

  on(event: 'connection', handler: (socket: TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>) => void): this {
    this.connectionHandlers.push(handler)
    return this
  }

  // -- Room / broadcast -----------------------------------------------------

  /** Broadcast an event to every connected socket, regardless of room. */
  emit<E extends keyof ServerToClientEvents & string>(event: E, ...args: any[]): void {
    const data = args.length <= 1 ? args[0] : args
    const message = JSON.stringify({ event, data })
    for (const socket of this.sockets) {
      if (socket.ws.readyState === WebSocket.OPEN) socket.ws.send(message)
    }
  }

  to(room: string): Broadcaster<ServerToClientEvents> {
    return new Broadcaster(this, room)
  }

  /** @internal */
  addSocketToRoom(room: string, socket: TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>): void {
    let set = this.roomMap.get(room)
    if (!set) {
      set = new Set()
      this.roomMap.set(room, set)
    }
    set.add(socket)
  }

  /** @internal */
  removeSocketFromRoom(room: string, socket: TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>): void {
    const set = this.roomMap.get(room)
    if (!set) return
    set.delete(socket)
    if (set.size === 0) this.roomMap.delete(room)
  }

  /** @internal */
  removeSocket(socket: TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>): void {
    this.sockets.delete(socket)
    this.aliveSockets.delete(socket.ws)
    // Clean up all room memberships
    for (const [room, set] of this.roomMap) {
      set.delete(socket)
      if (set.size === 0) this.roomMap.delete(room)
    }
  }

  /** @internal */
  getSocketsInRoom(room: string): TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>[] {
    const set = this.roomMap.get(room)
    return set ? Array.from(set) : []
  }

  // -- Close ----------------------------------------------------------------

  close(cb?: () => void): void {
    clearInterval(this.heartbeatTimer)
    for (const s of this.sockets) {
      try {
        s.ws.close()
      } catch {
        // ignore
      }
    }
    this.wss.close(cb)
  }

  // -- Internal connection handler ------------------------------------------

  private handleConnection(ws: WebSocket, req: IncomingMessage): void {
    const socket = new TypedSocket<ClientToServerEvents, ServerToClientEvents, SocketData>(ws, req, this)

    // Run middlewares sequentially
    const runMiddleware = (index: number): void => {
      if (index >= this.middlewares.length) {
        // All middleware passed — register and notify
        this.sockets.add(socket)
        for (const handler of this.connectionHandlers) {
          handler(socket)
        }
        return
      }
      this.middlewares[index]!(socket, (err) => {
        if (err) {
          logger.warn('WebSocket middleware rejected connection', {
            socketId: socket.id,
            error: err.message,
          })
          // Send error and close
          try {
            ws.send(JSON.stringify({ event: 'connect_error', data: { message: err.message } }))
            ws.close()
          } catch {
            // ignore
          }
          return
        }
        runMiddleware(index + 1)
      })
    }

    runMiddleware(0)
  }
}