websocket-downlink.spec.ts 9.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308
  1. import { once } from 'node:events'
  2. import { createServer } from 'node:http'
  3. import type { AddressInfo } from 'node:net'
  4. import { afterEach, describe, expect, it, vi } from 'vitest'
  5. import WebSocket from 'ws'
  6. import type {
  7. ApiProxy, HostFrame, MuxFrame, RpcRequest, ServerRequest,
  8. } from '@deepseek-ai/dsh-host-apiproxy/api'
  9. import { RpcId } from '@deepseek-ai/dsh-host-apiproxy/api'
  10. import { HOST_EVENTS_PATH, MUX_EVENTS_PATH } from '../src/api-path.ts'
  11. import { WebSocketDownlinks } from '../src/websocket-downlink.ts'
  12. type MuxSource = (signal: AbortSignal) => AsyncIterable<RpcRequest<MuxFrame>>
  13. type HostSource = (signal: AbortSignal) => AsyncIterable<RpcRequest<HostFrame>>
  14. const running: (() => Promise<void>)[] = []
  15. afterEach(async () => {
  16. await Promise.all(running.splice(0).map(close => close()))
  17. })
  18. function untilAbort(signal: AbortSignal): Promise<void> {
  19. if (signal.aborted) return Promise.resolve()
  20. return new Promise((resolve) => {
  21. signal.addEventListener('abort', () => { resolve() }, { once: true })
  22. })
  23. }
  24. async function * idle<F>(signal: AbortSignal): AsyncGenerator<RpcRequest<F>> {
  25. await untilAbort(signal)
  26. }
  27. function api(mux: MuxSource, host: HostSource): ApiProxy {
  28. return {
  29. events: {
  30. mux: (_request, signal) => mux(signal),
  31. host: (_request, signal) => host(signal),
  32. },
  33. } as ApiProxy
  34. }
  35. async function serve(downlinks: WebSocketDownlinks): Promise<{
  36. origin: string
  37. close: () => Promise<void>
  38. }> {
  39. const server = createServer()
  40. server.on('upgrade', (request, socket, head) => {
  41. const pathname = new URL(request.url ?? '/', 'http://dsh.internal').pathname
  42. if (pathname === MUX_EVENTS_PATH) downlinks.handleMux(request, socket, head)
  43. else if (pathname === HOST_EVENTS_PATH) downlinks.handleHost(request, socket, head)
  44. else socket.destroy()
  45. })
  46. await new Promise<void>(resolve => server.listen(0, '127.0.0.1', resolve))
  47. const port = (server.address() as AddressInfo).port
  48. return {
  49. origin: `ws://127.0.0.1:${String(port)}`,
  50. close: async () => {
  51. await downlinks.close()
  52. await new Promise<void>(resolve => server.close(() => { resolve() }))
  53. },
  54. }
  55. }
  56. function read(socket: WebSocket): Promise<ServerRequest> {
  57. return once(socket, 'message').then(([data]) => JSON.parse(String(data)) as ServerRequest)
  58. }
  59. async function acceptedSocket(downlinks: WebSocketDownlinks): Promise<WebSocket> {
  60. const server = (downlinks as unknown as { server: { clients: Set<WebSocket> } }).server
  61. let accepted: WebSocket | undefined
  62. await vi.waitFor(() => {
  63. accepted = server.clients.values().next().value
  64. expect(accepted).toBeDefined()
  65. })
  66. return accepted as WebSocket
  67. }
  68. describe('WebSocket downlinks', () => {
  69. it('carries mux and host over independent downstream sockets and cancels each source on close', async () => {
  70. let muxAborted = false
  71. let hostAborted = false
  72. const downlinks = new WebSocketDownlinks(api(
  73. async function * (signal) {
  74. try {
  75. yield {
  76. rpcId: RpcId('mux-1'),
  77. payload: { type: 'session/subscribed', sessionId: 'session-1' as never, lastSeq: 4 },
  78. }
  79. await untilAbort(signal)
  80. } finally {
  81. muxAborted = true
  82. }
  83. },
  84. async function * (signal) {
  85. try {
  86. yield { rpcId: RpcId('host-1'), payload: { type: 'host/commands-changed' } }
  87. await untilAbort(signal)
  88. } finally {
  89. hostAborted = true
  90. }
  91. },
  92. ))
  93. const host = await serve(downlinks)
  94. running.push(host.close)
  95. const mux = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
  96. const hostSocket = new WebSocket(`${host.origin}${HOST_EVENTS_PATH}`)
  97. const muxFrame = read(mux)
  98. const hostFrame = read(hostSocket)
  99. expect(await muxFrame).toEqual({
  100. type: 'server-request',
  101. rpcId: 'mux-1',
  102. method: 'session/subscribed',
  103. payload: { type: 'session/subscribed', sessionId: 'session-1', lastSeq: 4 },
  104. })
  105. expect(await hostFrame).toEqual({
  106. type: 'server-request',
  107. rpcId: 'host-1',
  108. method: 'host/commands-changed',
  109. payload: { type: 'host/commands-changed' },
  110. })
  111. const muxClosed = once(mux, 'close')
  112. const hostClosed = once(hostSocket, 'close')
  113. mux.close()
  114. hostSocket.close()
  115. await Promise.all([muxClosed, hostClosed])
  116. await vi.waitFor(() => {
  117. expect(muxAborted).toBe(true)
  118. expect(hostAborted).toBe(true)
  119. })
  120. })
  121. it('rejects client messages because upstream remains HTTP', async () => {
  122. let aborted = false
  123. const downlinks = new WebSocketDownlinks(api(
  124. async function * (signal) {
  125. try {
  126. await untilAbort(signal)
  127. } finally {
  128. aborted = true
  129. }
  130. },
  131. idle,
  132. ))
  133. const host = await serve(downlinks)
  134. running.push(host.close)
  135. const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
  136. await once(socket, 'open')
  137. const closed = once(socket, 'close')
  138. socket.send('upstream payload')
  139. const [code, reason] = await closed as [number, Buffer]
  140. expect(code).toBe(1008)
  141. expect(String(reason)).toBe('downlink only')
  142. await vi.waitFor(() => { expect(aborted).toBe(true) })
  143. })
  144. it('sends stream/error before closing when a source fails', async () => {
  145. const downlinks = new WebSocketDownlinks(api(
  146. async function * () {
  147. throw new Error('mux source failed')
  148. },
  149. idle,
  150. ))
  151. const host = await serve(downlinks)
  152. running.push(host.close)
  153. const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
  154. const failure = read(socket)
  155. const closed = once(socket, 'close')
  156. expect((await failure).payload).toEqual({
  157. type: 'stream/error',
  158. error: { code: 'internal', message: 'Error: mux source failed', details: {} },
  159. })
  160. await closed
  161. })
  162. it('aborts the source when an accepted socket reports a transport error', async () => {
  163. let aborted = false
  164. const downlinks = new WebSocketDownlinks(api(
  165. async function * (signal) {
  166. try {
  167. await untilAbort(signal)
  168. } finally {
  169. aborted = true
  170. }
  171. },
  172. idle,
  173. ))
  174. const host = await serve(downlinks)
  175. running.push(host.close)
  176. const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
  177. await once(socket, 'open')
  178. const accepted = await acceptedSocket(downlinks)
  179. const closed = once(socket, 'close')
  180. accepted.emit('error', new Error('transport failed'))
  181. await closed
  182. expect(aborted).toBe(true)
  183. })
  184. it('drops a source frame that races after the client has closed', async () => {
  185. let release!: () => void
  186. const gate = new Promise<void>((resolve) => { release = resolve })
  187. let finish!: () => void
  188. const finished = new Promise<void>((resolve) => { finish = resolve })
  189. let sourceSignal: AbortSignal | undefined
  190. const downlinks = new WebSocketDownlinks(api(
  191. async function * (signal) {
  192. sourceSignal = signal
  193. try {
  194. await gate
  195. yield {
  196. rpcId: RpcId('late'),
  197. payload: { type: 'session/subscribed', sessionId: 'session-late' as never, lastSeq: 0 },
  198. }
  199. } finally {
  200. finish()
  201. }
  202. },
  203. idle,
  204. ))
  205. const host = await serve(downlinks)
  206. running.push(host.close)
  207. const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
  208. await once(socket, 'open')
  209. const closed = once(socket, 'close')
  210. socket.close()
  211. await closed
  212. await vi.waitFor(() => { expect(sourceSignal?.aborted).toBe(true) })
  213. release()
  214. await finished
  215. })
  216. it('contains socket send callback failures and closes the downlink', async () => {
  217. let release!: () => void
  218. const gate = new Promise<void>((resolve) => { release = resolve })
  219. const downlinks = new WebSocketDownlinks(api(
  220. async function * () {
  221. await gate
  222. yield {
  223. rpcId: RpcId('send-failure'),
  224. payload: { type: 'session/subscribed', sessionId: 'session-send' as never, lastSeq: 0 },
  225. }
  226. },
  227. idle,
  228. ))
  229. const host = await serve(downlinks)
  230. running.push(host.close)
  231. const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
  232. await once(socket, 'open')
  233. const accepted = await acceptedSocket(downlinks)
  234. const send = vi.spyOn(accepted, 'send').mockImplementation(((
  235. _data: unknown,
  236. optionsOrCallback?: unknown,
  237. callback?: (error?: Error) => void,
  238. ) => {
  239. const done = typeof optionsOrCallback === 'function'
  240. ? optionsOrCallback as (error?: Error) => void
  241. : callback
  242. done?.(new Error('socket send failed'))
  243. }) as WebSocket['send'])
  244. const closed = once(socket, 'close')
  245. release()
  246. await closed
  247. expect(send).toHaveBeenCalledTimes(2)
  248. send.mockRestore()
  249. })
  250. it('rejects when its acceptor has already closed', async () => {
  251. const downlinks = new WebSocketDownlinks(api(idle, idle))
  252. await downlinks.close()
  253. await expect(downlinks.close()).rejects.toThrow('The server is not running')
  254. })
  255. it('waits for source cleanup before teardown resolves', async () => {
  256. let cleanupStarted!: () => void
  257. const started = new Promise<void>((resolve) => { cleanupStarted = resolve })
  258. let releaseCleanup!: () => void
  259. const cleanupGate = new Promise<void>((resolve) => { releaseCleanup = resolve })
  260. let cleaned = false
  261. const downlinks = new WebSocketDownlinks(api(
  262. async function * (signal) {
  263. try {
  264. await untilAbort(signal)
  265. } finally {
  266. cleanupStarted()
  267. await cleanupGate
  268. cleaned = true
  269. }
  270. },
  271. idle,
  272. ))
  273. const host = await serve(downlinks)
  274. const socket = new WebSocket(`${host.origin}${MUX_EVENTS_PATH}`)
  275. await once(socket, 'open')
  276. let closed = false
  277. const closing = host.close().then(() => { closed = true })
  278. try {
  279. await started
  280. expect(closed).toBe(false)
  281. releaseCleanup()
  282. await closing
  283. expect(cleaned).toBe(true)
  284. } finally {
  285. releaseCleanup()
  286. await closing
  287. }
  288. })
  289. })