diff --git a/packages/client/connection/src/websocket-downlink.ts b/packages/client/connection/src/websocket-downlink.ts
index 09a0844ada..996edd2551 100644
--- a/packages/client/connection/src/websocket-downlink.ts
+++ b/packages/client/connection/src/websocket-downlink.ts
@@ -27,8 +27,8 @@ function send(socket: WebSocket, frame: RpcRequest): Promise {
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()
}
}
}
diff --git a/packages/client/connection/tests/client-apply.spec.ts b/packages/client/connection/tests/client-apply.spec.ts
index 322398d371..9bc645c847 100644
--- a/packages/client/connection/tests/client-apply.spec.ts
+++ b/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)
+ })
})
diff --git a/packages/client/connection/tests/websocket-downlink.spec.ts b/packages/client/connection/tests/websocket-downlink.spec.ts
index 254f761691..40bab41135 100644
--- a/packages/client/connection/tests/websocket-downlink.spec.ts
+++ b/packages/client/connection/tests/websocket-downlink.spec.ts
@@ -63,6 +63,16 @@ function read(socket: WebSocket): Promise {
return once(socket, 'message').then(([data]) => JSON.parse(String(data)) as ServerRequest)
}
+async function acceptedSocket(downlinks: WebSocketDownlinks): Promise {
+ const server = (downlinks as unknown as { server: { clients: Set } }).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(resolve => { release = resolve })
+ let finish!: () => void
+ const finished = new Promise(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(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')
+ })
})