Selaa lähdekoodia

test(web): cover WebSocket downlink races

imccyu 1 kuukausi sitten
vanhempi
sitoutus
c6d0cbd8de

+ 2 - 3
packages/client/connection/src/websocket-downlink.ts

@@ -27,8 +27,8 @@ function send(socket: WebSocket, frame: RpcRequest<Frame>): Promise<void> {
       return
     }
     socket.send(JSON.stringify(serverRequest(frame)), (error) => {
-      if (error === undefined) resolve()
-      else reject(error)
+      if (error) reject(error)
+      else resolve()
     })
   })
 }
@@ -129,7 +129,6 @@ export class WebSocketDownlinks {
     } finally {
       abort.abort()
       if (socket.readyState === WebSocket.OPEN) socket.close()
-      else if (socket.readyState === WebSocket.CONNECTING) socket.terminate()
     }
   }
 }

+ 14 - 0
packages/client/connection/tests/client-apply.spec.ts

@@ -189,4 +189,18 @@ describe('connection client apply', () => {
     abort.abort()
     await expect(pending).resolves.toMatchObject({ done: true })
   })
+
+  it('closes a WebSocket immediately when its signal was already aborted', async () => {
+    ;(globalThis as Win).location = {
+      hostname: 'localhost', search: '', origin: 'http://localhost:3080',
+    }
+    ;(globalThis as WebSocketGlobal).WebSocket = FakeWebSocket as unknown as typeof WebSocket
+    const client = (await mount()).api
+    const abort = new AbortController()
+    abort.abort()
+    const iterator = client.events.mux({}, abort.signal)[Symbol.asyncIterator]()
+    await expect(iterator.next()).resolves.toMatchObject({ done: true })
+    expect(sockets).toHaveLength(1)
+    expect(sockets[0]?.readyState).toBe(FakeWebSocket.CLOSED)
+  })
 })

+ 101 - 0
packages/client/connection/tests/websocket-downlink.spec.ts

@@ -63,6 +63,16 @@ function read(socket: WebSocket): Promise<ServerRequest> {
   return once(socket, 'message').then(([data]) => JSON.parse(String(data)) as ServerRequest)
 }
 
+async function acceptedSocket(downlinks: WebSocketDownlinks): Promise<WebSocket> {
+  const server = (downlinks as unknown as { server: { clients: Set<WebSocket> } }).server
+  let accepted: WebSocket | undefined
+  await vi.waitFor(() => {
+    accepted = server.clients.values().next().value
+    expect(accepted).toBeDefined()
+  })
+  return accepted as WebSocket
+}
+
 describe('WebSocket downlinks', () => {
   it('carries mux and host over independent downstream sockets and cancels each source on close', async () => {
     let muxAborted = false
@@ -161,4 +171,95 @@ describe('WebSocket downlinks', () => {
     })
     await closed
   })
+
+  it('aborts the source when an accepted socket reports a transport error', async () => {
+    let aborted = false
+    const downlinks = new WebSocketDownlinks(api(
+      async function * (signal) {
+        try {
+          await untilAbort(signal)
+        } finally {
+          aborted = true
+        }
+      },
+      idle,
+    ))
+    const host = await serve(downlinks)
+    running.push(host.close)
+    const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
+    await once(socket, 'open')
+    const accepted = await acceptedSocket(downlinks)
+    const closed = once(socket, 'close')
+    accepted.emit('error', new Error('transport failed'))
+    await closed
+    expect(aborted).toBe(true)
+  })
+
+  it('drops a source frame that races after the client has closed', async () => {
+    let release!: () => void
+    const gate = new Promise<void>(resolve => { release = resolve })
+    let finish!: () => void
+    const finished = new Promise<void>(resolve => { finish = resolve })
+    let sourceSignal: AbortSignal | undefined
+    const downlinks = new WebSocketDownlinks(api(
+      async function * (signal) {
+        sourceSignal = signal
+        try {
+          await gate
+          yield { rpcId: RpcId('late'), payload: { type: 'host/commands-changed' } }
+        } finally {
+          finish()
+        }
+      },
+      idle,
+    ))
+    const host = await serve(downlinks)
+    running.push(host.close)
+    const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
+    await once(socket, 'open')
+    const closed = once(socket, 'close')
+    socket.close()
+    await closed
+    await vi.waitFor(() => { expect(sourceSignal?.aborted).toBe(true) })
+    release()
+    await finished
+  })
+
+  it('contains socket send callback failures and closes the downlink', async () => {
+    let release!: () => void
+    const gate = new Promise<void>(resolve => { release = resolve })
+    const downlinks = new WebSocketDownlinks(api(
+      async function * () {
+        await gate
+        yield { rpcId: RpcId('send-failure'), payload: { type: 'host/commands-changed' } }
+      },
+      idle,
+    ))
+    const host = await serve(downlinks)
+    running.push(host.close)
+    const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
+    await once(socket, 'open')
+    const accepted = await acceptedSocket(downlinks)
+    const send = vi.spyOn(accepted, 'send').mockImplementation(((
+      _data: unknown,
+      optionsOrCallback?: unknown,
+      callback?: (error?: Error) => void,
+    ) => {
+      const done = typeof optionsOrCallback === 'function'
+        ? optionsOrCallback as (error?: Error) => void
+        : callback
+      done?.(new Error('socket send failed'))
+    }) as WebSocket['send'])
+    const closed = once(socket, 'close')
+    release()
+    await closed
+    expect(send).toHaveBeenCalledTimes(2)
+    send.mockRestore()
+  })
+
+  it('rejects when its acceptor has already closed', async () => {
+    const downlinks = new WebSocketDownlinks(api(idle, idle))
+    await downlinks.close()
+    await expect(downlinks.close()).rejects.toThrow('The server is not running')
+  })
 })