diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e980c7d..0d24f38 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -23,6 +23,8 @@ jobs: - run: pnpm install - run: pnpm build - run: pnpm test + # Example tests need the pinned Bun/Deno versions above; local runs skip them. + - run: pnpm test:examples - run: pnpm lint - name: arethetypeswrong run: | diff --git a/README.md b/README.md index 57eeadc..4c32238 100644 --- a/README.md +++ b/README.md @@ -65,6 +65,37 @@ await new Promise(() => {}); // stay alive until killed That's all you need! No local ports, just a fetch function. +### WebSockets + +Tunnels forward WebSockets too: `npx captun 3000` exposes any local WebSocket +server (socket.io, `ws`, Bun, Deno, ...) with handshake headers, subprotocols, +binary messages, and close codes passing through. In code, a fetch handler +accepts WebSockets Workers-style on any runtime: + +```ts +import { + createCaptunTunnel, + createWebSocketResponse, + isWebSocketUpgradeRequest, + WebSocketPair, +} from "captun"; + +await createCaptunTunnel({ + fetch(request) { + if (!isWebSocketUpgradeRequest(request)) return new Response("hello"); + const pair = new WebSocketPair(); + pair[1].accept(); + pair[1].addEventListener("message", (event) => pair[1].send(`echo:${event.data}`)); + return createWebSocketResponse(pair[0]); + }, +}); +``` + +Connections are relayed message by message over the tunnel, so ping/pong and +compression are per-hop, close codes outside 1000/3000–4999 degrade to a plain +close, and messages are capped at 16MiB so one oversized frame can't take down +the tunnel. + ### Vite plugin `captun/vite` serves your Vite dev server (and `vite preview`) through a public tunnel URL — handy for receiving webhooks against local code, sharing work in progress, or pointing remote devices and agents at your dev server. diff --git a/package.json b/package.json index ebf6aff..2903bc0 100644 --- a/package.json +++ b/package.json @@ -87,7 +87,8 @@ "deploy:worker": "pnpm run build:hosted-browser-module && wrangler deploy", "deploy:hosted": "pnpm run build:hosted-browser-module && wrangler deploy --config wrangler.hosted.jsonc", "dev": "wrangler dev", - "test": "vitest run", + "test": "vitest run test/", + "test:examples": "vitest run examples/", "test:unit": "vitest run test/worker.test.ts", "cli": "tsx src/cli/bin.ts", "smoke": "./scripts/smoke-test.sh", @@ -126,5 +127,8 @@ "optional": true } }, + "engines": { + "node": ">=22.4" + }, "packageManager": "pnpm@10.11.1" } diff --git a/src/cli/bin.ts b/src/cli/bin.ts index d7c3833..f46ab76 100755 --- a/src/cli/bin.ts +++ b/src/cli/bin.ts @@ -17,7 +17,14 @@ import { CaptunTunnelConnectError, createCaptunTunnel, HOSTED_CAPTUN_GATEWAY, + isWebSocketUpgradeRequest, + pipeWebSocketToHandle, randomConnectToken, + webSocketHandleFromSocket, + type Fetcher, + type WebSocketConnectResult, + type WebSocketFetcher, + type WebSocketHandle, } from "../index.js"; import { assertLocalTargetAcceptingConnections } from "./local-target.js"; import { withSpinner } from "./spinner.js"; @@ -522,7 +529,8 @@ async function connectTunnelWithRetry( gateway: tunnel.gateway, name: tunnel.name, token: tunnel.token, - fetch: fetcher, + fetch: fetcher.fetch, + connectWebSocket: fetcher.connectWebSocket, }), ); } catch (error) { @@ -536,42 +544,144 @@ async function connectTunnelWithRetry( throw new Error("unreachable"); } -function makeTunnelFetcher(tunnel: ResolvedTunnel) { - return async (request: Request) => { - if (isCaptunHealthRequest(request)) return captunHealthResponse(); +function makeTunnelFetcher(tunnel: ResolvedTunnel): Fetcher & WebSocketFetcher { + return { + fetch: async (request: Request) => { + if (isWebSocketUpgradeRequest(request)) { + return new Response("Use connectWebSocket for WebSocket tunnel requests\n", { + status: 400, + }); + } - const url = new URL(request.url); - const requestStartedAt = performance.now(); - const rayId = request.headers.get("cf-ray") || "-"; - try { - const response = await fetch( - new Request(`${tunnel.target}${url.pathname}${url.search}`, request), - ); - logRequest(tunnel.requestLogs, { - method: request.method, - path: `${url.pathname}${url.search}`, - rayId, - status: response.status, - startedAt: requestStartedAt, - }); - return response; - } catch { - const response = new Response( - `Request reached the captun cli, but ${tunnel.target} is not accepting connections\n`, + if (isCaptunHealthRequest(request)) return captunHealthResponse(); + + const url = new URL(request.url); + const requestStartedAt = performance.now(); + const rayId = request.headers.get("cf-ray") || "-"; + try { + const response = await fetch( + new Request(`${tunnel.target}${url.pathname}${url.search}`, request), + ); + logRequest(tunnel.requestLogs, { + method: request.method, + path: `${url.pathname}${url.search}`, + rayId, + status: response.status, + startedAt: requestStartedAt, + }); + return response; + } catch { + const response = new Response( + `Request reached the captun cli, but ${tunnel.target} is not accepting connections\n`, + { status: 502 }, + ); + logRequest(tunnel.requestLogs, { + method: request.method, + path: `${url.pathname}${url.search}`, + rayId, + status: response.status, + startedAt: requestStartedAt, + }); + return response; + } + }, + + connectWebSocket: (request, remote) => connectTargetWebSocket(tunnel, request, remote), + }; +} + +async function connectTargetWebSocket( + tunnel: ResolvedTunnel, + request: Request, + remote: WebSocketHandle, +): Promise { + const url = new URL(request.url); + const requestStartedAt = performance.now(); + const log = (status: number) => + logRequest(tunnel.requestLogs, { + method: request.method, + path: `${url.pathname}${url.search}`, + rayId: request.headers.get("cf-ray") || "-", + status, + startedAt: requestStartedAt, + }); + + const targetUrl = new URL(`${tunnel.target}${url.pathname}${url.search}`); + targetUrl.protocol = targetUrl.protocol === "https:" ? "wss:" : "ws:"; + const targetSocket = new WebSocket(targetUrl, { + protocols: request.headers + .get("sec-websocket-protocol") + ?.split(",") + .map((protocol) => protocol.trim()) + .filter(Boolean), + headers: forwardedHandshakeHeaders(request.headers), + // Node's WebSocket (undici) accepts { protocols, headers }; the DOM type doesn't. + } as unknown as string[]); + + try { + await waitForWebSocketOpen(targetSocket); + // The local server may have closed right after the handshake; reject the + // public upgrade cleanly instead of accepting an already-dead socket. + if (targetSocket.readyState !== WebSocket.OPEN) throw new Error("WebSocket closed after open"); + } catch { + targetSocket.close(); + log(502); + return { + accepted: false, + response: new Response( + `Request reached the captun cli, but ${targetUrl.origin} did not accept the WebSocket\n`, { status: 502 }, - ); - logRequest(tunnel.requestLogs, { - method: request.method, - path: `${url.pathname}${url.search}`, - rayId, - status: response.status, - startedAt: requestStartedAt, - }); - return response; - } + ), + }; + } + + log(101); + pipeWebSocketToHandle(targetSocket, remote); + return { + accepted: true, + protocol: targetSocket.protocol || undefined, + socket: webSocketHandleFromSocket(targetSocket), }; } +/** + * Handshake headers to forward to the local WebSocket server, so cookie or + * token auth behaves the same as tunneled HTTP requests. Hop-by-hop headers + * and the Sec-WebSocket-* family stay out: the local WebSocket client + * negotiates its own handshake (subprotocols are passed separately). + */ +function forwardedHandshakeHeaders(headers: Headers) { + const skip = new Set(["connection", "host", "keep-alive", "te", "trailer", "upgrade"]); + const forwarded: Record = {}; + for (const [name, value] of headers) { + if (skip.has(name) || name.startsWith("sec-websocket-") || name.startsWith("proxy-")) continue; + forwarded[name] = value; + } + return forwarded; +} + +async function waitForWebSocketOpen(socket: WebSocket) { + if (socket.readyState === WebSocket.OPEN) return; + if (socket.readyState !== WebSocket.CONNECTING) throw new Error("WebSocket closed before open"); + + const listeners = new AbortController(); + await new Promise((resolveOpen, rejectOpen) => { + const settle = (callback: () => void) => { + listeners.abort(); + callback(); + }; + socket.addEventListener("open", () => settle(resolveOpen), { signal: listeners.signal }); + socket.addEventListener("error", () => settle(() => rejectOpen(new Error("WebSocket error"))), { + signal: listeners.signal, + }); + socket.addEventListener( + "close", + () => settle(() => rejectOpen(new Error("WebSocket closed before open"))), + { signal: listeners.signal }, + ); + }); +} + function tunnelConnectError(tunnel: ResolvedTunnel, cause: unknown) { const hostname = new URL(tunnel.gateway).hostname; const message = cause instanceof Error ? cause.message : String(cause); diff --git a/src/hosted/worker.ts b/src/hosted/worker.ts index 17913fa..fb5ac6f 100644 --- a/src/hosted/worker.ts +++ b/src/hosted/worker.ts @@ -9,6 +9,7 @@ import { import { connectTokenFromRequest, GATEWAY_CONNECT_QUERY_PARAM, + isWebSocketUpgradeRequest, TUNNEL_CONNECT_DIAGNOSTIC_HEADER, TUNNEL_NAME_QUERY_PARAM, } from "../index.js"; @@ -102,17 +103,20 @@ export default { customHostname, tunnelName, }); - const response = await shard.forward( + const tunnelRequest = createTunnelForwardRequest(forwarded, { tunnelName, - createTunnelForwardRequest(forwarded, tunnelUrl), - ); + tunnelUrl, + }); + const response = isWebSocketUpgradeRequest(request) + ? await shard.fetch(tunnelRequest) + : await shard.forward(tunnelName, tunnelRequest); return stripSetCookieHeadersOutsideTunnel(response, new URL(tunnelUrl).hostname); }, } satisfies ExportedHandler; async function connectTunnel(request: Request, env: HostedCaptunEnv) { const diagnostic = isConnectDiagnostic(request); - if (!diagnostic && request.headers.get("upgrade") !== "websocket") { + if (!diagnostic && !isWebSocketUpgradeRequest(request)) { return new Response("Expected WebSocket upgrade\n", { status: 400 }); } @@ -158,7 +162,7 @@ function isGatewayConnectRequest(request: Request) { } function isConnectDiagnostic(request: Request) { - if (request.headers.get("upgrade") === "websocket") return false; + if (isWebSocketUpgradeRequest(request)) return false; return request.headers.get(TUNNEL_CONNECT_DIAGNOSTIC_HEADER) === "1"; } @@ -177,6 +181,9 @@ function createForwardedRequest(request: Request, customHostname: string | undef } function stripSetCookieHeadersOutsideTunnel(response: Response, tunnelHostname: string) { + // An upgraded response carries the webSocket and cannot be reconstructed. + if (response.status === 101) return response; + const setCookies = setCookieHeaders(response.headers); if (setCookies.length === 0) return response; diff --git a/src/index.ts b/src/index.ts index bf2bae1..de7e6e0 100644 --- a/src/index.ts +++ b/src/index.ts @@ -18,12 +18,19 @@ export interface Fetcher { fetch(request: Request): Response | Promise; } +export interface WebSocketFetcher { + connectWebSocket( + request: Request, + remote: WebSocketHandle, + ): WebSocketConnectResult | Promise; +} + export type TunnelReady = { url: string; token?: string; }; -export interface FetcherStub extends Fetcher, Disposable { +export interface FetcherStub extends Fetcher, WebSocketFetcher, Disposable { ready(tunnel: TunnelReady): void | Promise; } @@ -31,6 +38,23 @@ export interface RemoteFetcherCapability extends FetcherStub { onRpcBroken(callback: () => void): void; } +export type WebSocketMessage = string | Uint8Array; + +/** + * A tunneled WebSocket as a Cap'n Web capability: each side of a tunneled + * connection holds a handle to the socket on the other side and forwards + * messages by calling it. Cap'n Web delivers calls in order, so no extra + * framing is needed. + */ +export interface WebSocketHandle { + send(message: WebSocketMessage): unknown; + close(code?: number, reason?: string): unknown; +} + +export type WebSocketConnectResult = + | { accepted: true; protocol?: string; socket: WebSocketHandle } + | { accepted: false; response: Response }; + export function fetcherStubFromRemoteCapability( remote: RemoteFetcherCapability, options: { onDisconnect?: () => void }, @@ -39,6 +63,7 @@ export function fetcherStubFromRemoteCapability( return { fetch: (request) => remote.fetch(request), + connectWebSocket: (request, handle) => remote.connectWebSocket(request, handle), ready: (tunnel) => remote.ready(tunnel), [Symbol.dispose]: () => remote[Symbol.dispose](), }; @@ -48,7 +73,7 @@ export function acceptFetcherCapabilityFromSocket( socket: WebSocket, options: { onDisconnect?: () => void } = {}, ): FetcherStub { - const remote = newWebSocketRpcSession(socket); + const remote = newWebSocketRpcSession(socket); return fetcherStubFromRemoteCapability(remote, options); } @@ -134,20 +159,21 @@ export class CaptunTunnelConnectError extends Error { } } -type TunnelClientCapability = Fetcher & { - ready(tunnel: TunnelReady): void | Promise; -}; +type TunnelClientCapability = Fetcher & + WebSocketFetcher & { + ready(tunnel: TunnelReady): void | Promise; + }; -type WorkerWebSocket = WebSocket & { +export type WorkerWebSocket = WebSocket & { accept(): void; }; -type WorkerWebSocketPairConstructor = new () => { +export type WorkerWebSocketPairConstructor = new () => { 0: WorkerWebSocket; 1: WorkerWebSocket; }; -type WebSocketResponseInit = ResponseInit & { +export type WebSocketResponseInit = ResponseInit & { webSocket: WebSocket; }; @@ -156,6 +182,13 @@ const WEBSOCKET_REJECTION_PROBE_TIMEOUT_MS = 500; /** Options for {@link createCaptunTunnel}. */ export type CreateCaptunTunnelOptions = Fetcher & { + /** + * Low-level hook for forwarding public WebSockets, e.g. by dialing out to a + * separate local WebSocket server the way the CLI does. Without it, `fetch` + * handles WebSockets too: a returned Worker-style response with a + * `webSocket` is bridged automatically. + */ + connectWebSocket?: WebSocketFetcher["connectWebSocket"]; /** * Tunnel Gateway URL. Defaults to the hosted `https://captun.sh` service. * After `npx captun deploy`, pass your own gateway URL here. @@ -183,6 +216,7 @@ export async function createCaptunTunnel( const socket = createWebSocket(connect.url, connect.protocols); const fetcher = new TunnelTargetFetcher({ fetch: options.fetch, + connectWebSocket: options.connectWebSocket, ready: (tunnel) => ready.resolve(tunnel), }); const session = newWebSocketRpcSession(socket, fetcher); @@ -232,12 +266,17 @@ function randomTunnelName() { } class TunnelTargetFetcher extends RpcTarget implements TunnelClientCapability { - private fetcher: Fetcher; + private fetcher: Fetcher & Partial; private onReady: (tunnel: TunnelReady) => void; - constructor(options: { fetch: Fetcher["fetch"]; ready: (tunnel: TunnelReady) => void }) { + constructor(options: { + fetch: Fetcher["fetch"]; + connectWebSocket?: WebSocketFetcher["connectWebSocket"]; + ready: (tunnel: TunnelReady) => void; + }) { super(); this.fetcher = { fetch: options.fetch }; + if (options.connectWebSocket) this.fetcher.connectWebSocket = options.connectWebSocket; this.onReady = options.ready; } @@ -248,6 +287,22 @@ class TunnelTargetFetcher extends RpcTarget implements TunnelClientCapability { ready(tunnel: TunnelReady) { this.onReady(tunnel); } + + async connectWebSocket(request: Request, remote: WebSocketHandle) { + if (this.fetcher.connectWebSocket) return this.fetcher.connectWebSocket(request, remote); + + const response = await this.fetcher.fetch(createWebSocketUpgradeRequest(request)); + const socket = responseWebSocket(response); + if (!socket) return { accepted: false as const, response }; + + acceptIfNeeded(socket); + pipeWebSocketToHandle(socket, remote); + return { + accepted: true as const, + protocol: response.headers.get("sec-websocket-protocol") || undefined, + socket: webSocketHandleFromSocket(socket), + }; + } } function createWebSocket(url: string | URL, protocols: string[]) { @@ -374,3 +429,242 @@ export function acceptFetcherCapability( response: new Response(null, responseInit), }; } + +export function isWebSocketUpgradeRequest(request: Request): boolean { + return request.headers.get("upgrade")?.toLowerCase() === "websocket"; +} + +class PairedWebSocketPair { + 0: PairedWebSocket; + 1: PairedWebSocket; + + constructor() { + const first = new PairedWebSocket(); + const second = new PairedWebSocket(); + first.peer = second; + second.peer = first; + this[0] = first; + this[1] = second; + } +} + +class PairedWebSocket extends EventTarget { + peer!: PairedWebSocket; + private accepted = false; + private closed = false; + private queued: Event[] = []; + + accept() { + if (this.accepted) return; + this.accepted = true; + // Flush asynchronously so listeners attached right after accept() still + // receive events that arrived earlier (matching Workers queuing). + queueMicrotask(() => { + for (const event of this.queued.splice(0)) this.dispatchEvent(event); + }); + } + + send(message: unknown) { + if (this.closed) throw new TypeError("Can't call WebSocket send() after close()."); + this.peer.deliver(Object.assign(new Event("message"), { data: message })); + } + + close(code?: number, reason?: string) { + if (this.closed) return; + this.closed = true; + this.peer.closed = true; + const detail = { code: code ?? 1005, reason: reason ?? "" }; + this.peer.deliver(Object.assign(new Event("close"), detail)); + this.deliver(Object.assign(new Event("close"), detail)); + } + + private deliver(event: Event) { + if (!this.accepted) { + this.queued.push(event); + return; + } + // Microtasks dispatch in order, after any pending accept() flush. + queueMicrotask(() => this.dispatchEvent(event)); + } +} + +/** + * `WebSocketPair` on every runtime: the Workers-native pair where the runtime + * provides one, otherwise an in-memory pair, so a plain `fetch` handler can + * answer tunneled WebSockets Workers-style anywhere. Pair sockets implement + * the surface Workers code uses — `accept()`, `send()`, `close()`, and + * message/close events — not every WebSocket property. + */ +export const WebSocketPair: WorkerWebSocketPairConstructor = + (globalThis as { WebSocketPair?: WorkerWebSocketPairConstructor }).WebSocketPair ?? + (PairedWebSocketPair as unknown as WorkerWebSocketPairConstructor); + +/** + * Wraps one end of a WebSocketPair in the Response a tunneled `fetch` handler + * returns to accept the WebSocket. In Workers this is a real 101 upgrade + * response; on other runtimes the Response only carries the socket to the + * tunnel bridge (the public 101 is produced by the gateway), because their + * Response constructors reject status 101. + */ +export function createWebSocketResponse( + webSocket: WebSocket, + init?: { protocol?: string }, +): Response { + const headers = init?.protocol ? { "sec-websocket-protocol": init.protocol } : undefined; + try { + return new Response(null, { status: 101, webSocket, headers } as WebSocketResponseInit); + } catch { + const response = new Response(null, { headers }); + Object.defineProperty(response, "webSocket", { value: webSocket }); + return response; + } +} + +/** Exposes a local WebSocket as a WebSocketHandle the other side of a tunnel can call. */ +export function webSocketHandleFromSocket(socket: WebSocket): WebSocketHandle { + return new SocketHandle(socket); +} + +class SocketHandle extends RpcTarget implements WebSocketHandle { + constructor(private socket: WebSocket) { + super(); + } + + send(message: WebSocketMessage) { + this.socket.send(message); + } + + close(code?: number, reason?: string) { + closeWebSocket(this.socket, code, reason); + } +} + +/** + * Forwards every message and the final close from a local WebSocket to the + * remote side's handle. Conversions are chained so messages arrive in order + * even when a runtime delivers binary frames as Blobs (async to read). + */ +export function pipeWebSocketToHandle(socket: WebSocket, handle: WebSocketHandle): void { + // Cap'n Web disposes stubs received as call arguments when the call + // returns; dup() keeps the capability alive for the socket's lifetime. + const remote = dupStub(handle); + const closeFailed = () => closeWebSocket(socket, 1011, "WebSocket tunnel failed"); + let pending = Promise.resolve(); + const enqueue = (action: () => void | Promise) => { + pending = pending.then(action).catch(closeFailed); + }; + // A failed forward means the tunnel side is gone; close our side too. + const forward = (call: unknown) => { + void Promise.resolve(call).catch(closeFailed); + }; + + let finished = false; + const finish = (code?: number, reason?: string) => { + if (finished) return; + finished = true; + enqueue(() => { + try { + forward(remote.close(code, reason)); + } finally { + disposeStub(remote); + } + }); + }; + + socket.addEventListener("message", (event) => { + if (finished) return; + enqueue(async () => { + const message = await webSocketMessage(event.data); + // Cap'n Web frames the tunnel leg as base64 JSON with a ~32MiB frame + // limit; an oversized message would kill the whole tunnel (closing + // every connection on it) instead of just this socket. + if (messageByteLength(message) > MAX_MESSAGE_BYTES) { + closeWebSocket(socket, 4009, "Message too large to tunnel"); + return; + } + forward(remote.send(message)); + }); + }); + socket.addEventListener("close", (event) => finish(event.code, event.reason)); + socket.addEventListener("error", () => finish(1011, "WebSocket error")); +} + +const MAX_MESSAGE_BYTES = 16 * 1024 * 1024; + +/** + * What a message costs on the Cap'n Web wire: bytes for binary (base64 + * inflation fits in the 2x headroom below the ~32MiB frame limit), and + * JSON-escaped UTF-8 bytes for strings — `length` counts UTF-16 units, which + * undercounts multi-byte characters by up to 3x and escapes by up to 6x. + */ +function messageByteLength(message: WebSocketMessage) { + if (typeof message !== "string") return message.byteLength; + // Fast path: even at the 6-bytes-per-unit worst case it fits. + if (message.length * 6 <= MAX_MESSAGE_BYTES) return message.length; + return new TextEncoder().encode(JSON.stringify(message)).byteLength; +} + +type StubLike = { dup?(): unknown; [Symbol.dispose]?(): void }; + +function dupStub(handle: WebSocketHandle): WebSocketHandle { + const dup = (handle as WebSocketHandle & StubLike).dup; + return typeof dup === "function" ? (dup.call(handle) as WebSocketHandle) : handle; +} + +function disposeStub(handle: WebSocketHandle) { + (handle as WebSocketHandle & StubLike)[Symbol.dispose]?.(); +} + +/** + * Normalizes a message event payload to string | Uint8Array. Checks are + * realm-safe (no bare instanceof) because tunneled sockets can come from + * other contexts — e.g. miniflare delivers Blobs from its own realm — and + * the copy yields a true Uint8Array, which Cap'n Web's serializer requires + * (a Node Buffer's prototype is not Uint8Array.prototype). Anything else + * throws, which fails just that socket (the pipe closes it with 1011) + * rather than delivering a stringified payload. + */ +async function webSocketMessage(data: unknown): Promise { + if (typeof data === "string") return data; + if (ArrayBuffer.isView(data)) { + return new Uint8Array(data.buffer, data.byteOffset, data.byteLength).slice(); + } + if (Object.prototype.toString.call(data) === "[object ArrayBuffer]") { + return new Uint8Array(data as ArrayBuffer).slice(); + } + if (typeof (data as Blob | null)?.arrayBuffer === "function") { + return new Uint8Array(await (data as Blob).arrayBuffer()); + } + throw new TypeError( + `Cannot tunnel WebSocket message ${Object.prototype.toString.call(data)}; expected string or binary`, + ); +} + +function closeWebSocket(socket: WebSocket, code?: number, reason?: string) { + // WebSocket.close() only accepts 1000 or 3000-4999; other codes (1001, + // 1011, ...) are receive-only and must fall back to a bare close. + try { + if (code === undefined || code === 1000 || (code >= 3000 && code <= 4999)) { + socket.close(code, reason); + } else { + socket.close(); + } + } catch { + // Already closed or closing. + } +} + +function acceptIfNeeded(socket: WebSocket) { + const maybeWorkerSocket = socket as WebSocket & { accept?: () => void }; + if (typeof maybeWorkerSocket.accept === "function") maybeWorkerSocket.accept(); +} + +function responseWebSocket(response: Response): WebSocket | undefined { + return (response as Response & { webSocket?: WebSocket | null }).webSocket || undefined; +} + +function createWebSocketUpgradeRequest(request: Request) { + const headers = new Headers(request.headers); + headers.set("upgrade", "websocket"); + return new Request(request.url, { headers }); +} diff --git a/src/server/worker.ts b/src/server/worker.ts index f7a48ff..6115bc7 100644 --- a/src/server/worker.ts +++ b/src/server/worker.ts @@ -3,9 +3,15 @@ import { acceptFetcherCapability, connectTokenFromRequest, GATEWAY_CONNECT_QUERY_PARAM, + isWebSocketUpgradeRequest, + pipeWebSocketToHandle, TUNNEL_CONNECT_DIAGNOSTIC_HEADER, TUNNEL_NAME_QUERY_PARAM, + webSocketHandleFromSocket, type FetcherStub, + type WebSocketConnectResult, + type WebSocketResponseInit, + type WorkerWebSocketPairConstructor, } from "../index.js"; import { captunShardName, @@ -36,6 +42,8 @@ type ActiveTunnel = { url: string; token?: string; fetcher: FetcherStub; + /** Public WebSockets forwarded to this tunnel, closed when the tunnel client goes away. */ + sockets: Set; }; export type TunnelAdmission = @@ -91,8 +99,15 @@ export class CaptunServerShard< const tunnelUrl = request.headers.get(TUNNEL_URL_HEADER); if (!tunnelUrl) return new Response("Missing tunnel URL\n", { status: 404 }); - // Non-upgrade requests are diagnostic probes: run admission, skip the upgrade. - if (request.headers.get("upgrade") !== "websocket") { + if (!isGatewayConnectRequest(request)) { + if (!isWebSocketUpgradeRequest(request)) { + return new Response("Expected WebSocket upgrade\n", { status: 400 }); + } + return this.forward(tunnelName, request); + } + + // Non-upgrade connect requests are diagnostic probes: run admission, skip the upgrade. + if (!isWebSocketUpgradeRequest(request)) { return this.diagnoseConnect(tunnelName, request); } @@ -104,14 +119,25 @@ export class CaptunServerShard< }); if (!admission.ok) return admission.response; - activeTunnel?.fetcher[Symbol.dispose](); + if (activeTunnel) { + activeTunnel.fetcher[Symbol.dispose](); + closeTunnelSockets(activeTunnel); + } const { response, fetcher } = acceptFetcherCapability({ request, onDisconnect: () => { - if (this.tunnels.get(tunnelName)?.fetcher === fetcher) this.tunnels.delete(tunnelName); + const active = this.tunnels.get(tunnelName); + if (active?.fetcher !== fetcher) return; + this.tunnels.delete(tunnelName); + closeTunnelSockets(active); }, }); - const tunnel = { url: tunnelUrl, token: admission.token, fetcher }; + const tunnel: ActiveTunnel = { + url: tunnelUrl, + token: admission.token, + fetcher, + sockets: new Set(), + }; this.tunnels.set(tunnelName, tunnel); queueMicrotask(() => { void fetcher.ready({ url: tunnel.url, token: tunnel.token }); @@ -133,10 +159,13 @@ export class CaptunServerShard< } async forward(tunnelName: string, request: Request): Promise { - const tunnel = this.tunnels.get(tunnelName)?.fetcher; + const tunnel = this.tunnels.get(tunnelName); if (!tunnel) return new Response("No tunnel client connected\n", { status: 503 }); try { - return await tunnel.fetch(request); + if (isWebSocketUpgradeRequest(request)) { + return await forwardWebSocket(tunnel, request); + } + return await tunnel.fetcher.fetch(request); } catch { return new Response("Tunnel fetch failed\n", { status: 502 }); } @@ -179,13 +208,19 @@ export default { customHostname: env.CUSTOM_HOSTNAME, tunnelName, }); - return shard.forward(tunnelName, createTunnelForwardRequest(forwarded, tunnelUrl)); + const tunnelRequest = createTunnelForwardRequest(forwarded, { + tunnelName, + tunnelUrl, + }); + + if (isWebSocketUpgradeRequest(request)) return shard.fetch(tunnelRequest); + return shard.forward(tunnelName, tunnelRequest); }, } satisfies ExportedHandler; function connectTunnel(request: Request, env: CaptunEnv) { const diagnostic = isConnectDiagnostic(request); - if (!diagnostic && request.headers.get("upgrade") !== "websocket") { + if (!diagnostic && !isWebSocketUpgradeRequest(request)) { return new Response("Expected WebSocket upgrade\n", { status: 400 }); } @@ -211,7 +246,7 @@ function isGatewayConnectRequest(request: Request) { } function isConnectDiagnostic(request: Request) { - if (request.headers.get("upgrade") === "websocket") return false; + if (isWebSocketUpgradeRequest(request)) return false; return request.headers.get(TUNNEL_CONNECT_DIAGNOSTIC_HEADER) === "1"; } @@ -238,12 +273,63 @@ export function createTunnelConnectRequest(input: { return new Request(input.request, { headers }); } -export function createTunnelForwardRequest(request: Request, tunnelUrl: string): Request { +export function createTunnelForwardRequest( + request: Request, + input: { tunnelName: string; tunnelUrl: string }, +): Request { const headers = new Headers(request.headers); - headers.set(TUNNEL_URL_HEADER, tunnelUrl); + headers.set(TUNNEL_NAME_HEADER, input.tunnelName); + headers.set(TUNNEL_URL_HEADER, input.tunnelUrl); return new Request(request, { headers }); } +async function forwardWebSocket(tunnel: ActiveTunnel, request: Request) { + const WorkerWebSocketPair = ( + globalThis as typeof globalThis & { WebSocketPair: WorkerWebSocketPairConstructor } + ).WebSocketPair; + const pair = new WorkerWebSocketPair(); + const serverSocket = pair[1]; + serverSocket.accept(); + + let result: WebSocketConnectResult; + try { + result = await tunnel.fetcher.connectWebSocket( + request, + webSocketHandleFromSocket(serverSocket), + ); + } catch (error) { + serverSocket.close(1011, "WebSocket tunnel failed"); + throw error; + } + if (!result.accepted) { + serverSocket.close(1000, "WebSocket not accepted"); + return result.response; + } + + pipeWebSocketToHandle(serverSocket, result.socket); + // The pipe dup()ed its own reference to the tunnel client's handle; release ours. + (result.socket as Partial)[Symbol.dispose]?.(); + tunnel.sockets.add(serverSocket); + serverSocket.addEventListener("close", () => tunnel.sockets.delete(serverSocket)); + return new Response(null, { + status: 101, + webSocket: pair[0], + headers: result.protocol ? { "sec-websocket-protocol": result.protocol } : undefined, + } as WebSocketResponseInit); +} + +function closeTunnelSockets(tunnel: ActiveTunnel) { + for (const socket of tunnel.sockets) { + try { + // workerd allows close(1001) (going away) even though browser clients don't. + socket.close(1001, "Tunnel client disconnected"); + } catch { + // Already closed. + } + } + tunnel.sockets.clear(); +} + function constantTimeEqual(actual: Uint8Array, expected: Uint8Array) { if (actual.length !== expected.length) return false; let diff = 0; diff --git a/test/e2e.test.ts b/test/e2e.test.ts index aa572a0..6be9c31 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -1,8 +1,21 @@ import { createHash } from "node:crypto"; + +import { createRouterClient } from "@orpc/server"; +import { newWebSocketRpcSession } from "capnweb"; import { expect, test, vi } from "vitest"; -import { createCaptunTunnel } from "../src/index.js"; -import { createCaptunWorkerFixture } from "./miniflare.js"; +import { createCaptunCliRouter } from "../src/cli/bin.js"; +import { + createCaptunTunnel, + createWebSocketResponse, + isWebSocketUpgradeRequest, + WebSocketPair, +} from "../src/index.js"; +import { + createCaptunWorkerFixture, + createHostedCaptunWorkerFixture, + createMiniflareWorkerFixture, +} from "./miniflare.js"; vi.setConfig({ testTimeout: 15_000 }); @@ -171,6 +184,361 @@ test.concurrent("uploads multipart form data", async ({ task }) => { }); }); +test.concurrent("forwards WebSocket Cap'n Web RPC to a local fetch handler", async ({ task }) => { + await using target = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using tunnel = await createTunnelFixture(task.name, (request) => + target.worker.fetch(request.url, request), + ); + + const rpc = newWebSocketRpcSession<{ ping(value: string): Promise }>( + `${tunnel.url}/rpc`.replace(/^http/, "ws"), + ); + + await expect(rpc.ping("dummy-capability")).resolves.toBe("pong:dummy-capability"); + rpc[Symbol.dispose](); + await delay(50); +}); + +// Binary frames are covered by the CLI test below: miniflare's getWorker() +// fetch proxy used here corrupts binary WebSocket messages to "[object Blob]". +test.concurrent("forwards WebSocket messages, subprotocols, and close codes", async ({ task }) => { + await using target = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using tunnel = await createTunnelFixture(task.name, (request) => + target.worker.fetch(request.url, request), + ); + + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws"), ["alpha", "beta"]); + await waitForWebSocket(socket); + expect(socket).toMatchObject({ protocol: "alpha" }); + const closed = nextWebSocketClose(socket); + + socket.send("ping"); + await expect(nextWebSocketMessage(socket).then(webSocketMessageText)).resolves.toBe("echo:ping"); + + socket.send("close-with:4001 done"); + await expect(closed).resolves.toMatchObject({ code: 4001, reason: "done" }); +}); + +test.concurrent("forwards concurrent WebSockets over one tunnel without cross-talk", async ({ + task, +}) => { + await using target = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using tunnel = await createTunnelFixture(task.name, (request) => + target.worker.fetch(request.url, request), + ); + + const first = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + const second = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await Promise.all([waitForWebSocket(first), waitForWebSocket(second)]); + + first.send("from-first"); + second.send("from-second"); + await expect(nextWebSocketMessage(first).then(webSocketMessageText)).resolves.toBe( + "echo:from-first", + ); + await expect(nextWebSocketMessage(second).then(webSocketMessageText)).resolves.toBe( + "echo:from-second", + ); + + first.close(); + second.close(); + await delay(50); +}); + +// Workers-style WebSocket handling from a plain fetch handler running in +// Node, using the library's runtime-agnostic WebSocketPair — no workerd on +// the target side. +test.concurrent("answers WebSockets Workers-style from a Node fetch handler", async ({ task }) => { + const serverClosed = Promise.withResolvers<{ code: number; reason: string }>(); + + await using tunnel = await createTunnelFixture(task.name, (request) => { + if (!isWebSocketUpgradeRequest(request)) return new Response("http ok\n"); + + const pair = new WebSocketPair(); + const server = pair[1]; + server.accept(); + server.send("welcome"); + server.addEventListener("message", (event) => { + const data = (event as MessageEvent).data as string; + if (data === "close-me") { + server.close(4002, "client asked"); + return; + } + server.send(`node-echo:${data}`); + }); + server.addEventListener("close", (event) => { + const { code, reason } = event as { code?: number; reason?: string }; + serverClosed.resolve({ code: code ?? 0, reason: reason ?? "" }); + }); + return createWebSocketResponse(pair[0], { protocol: "node-pair" }); + }); + + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws"), ["node-pair"]); + await waitForWebSocket(socket); + expect(socket).toMatchObject({ protocol: "node-pair" }); + + // The welcome was sent before the handler even returned; it must not be lost. + await expect(nextWebSocketMessage(socket).then(webSocketMessageText)).resolves.toBe("welcome"); + + socket.send("hello"); + await expect(nextWebSocketMessage(socket).then(webSocketMessageText)).resolves.toBe( + "node-echo:hello", + ); + + // Server-initiated close reaches the public client with its code. + const closed = nextWebSocketClose(socket); + socket.send("close-me"); + await expect(closed).resolves.toMatchObject({ code: 4002, reason: "client asked" }); + await expect(serverClosed.promise).resolves.toMatchObject({ code: 4002 }); +}); + +// The in-memory WebSocketPair accepts any value in send(), but only strings +// and bytes can cross the tunnel. A bogus payload must fail that socket, not +// reach the public client stringified as "[object Object]". +test.concurrent("closes the socket instead of stringifying an untunnelable payload", async ({ + task, +}) => { + await using tunnel = await createTunnelFixture(task.name, (request) => { + if (!isWebSocketUpgradeRequest(request)) return new Response("http ok\n"); + + const pair = new WebSocketPair(); + const server = pair[1]; + server.accept(); + server.addEventListener("message", () => { + server.send({ not: "a websocket payload" } as unknown as string); + }); + return createWebSocketResponse(pair[0]); + }); + + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await waitForWebSocket(socket); + const closed = nextWebSocketClose(socket); + + socket.send("trigger"); + await expect(nextWebSocketMessage(socket)).rejects.toThrow("WebSocket closed"); + await closed; +}); + +test.concurrent("closes only the socket that sends an oversized message", async ({ task }) => { + await using target = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using tunnel = await createTunnelFixture(task.name, (request) => + target.worker.fetch(request.url, request), + ); + + const oversized = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await waitForWebSocket(oversized); + const closed = nextWebSocketClose(oversized); + oversized.send(new Uint8Array(17 * 1024 * 1024)); + await expect(closed).resolves.toMatchObject({ code: 4009 }); + + // Text is measured in wire bytes, not string length: ~10.5M UTF-16 units + // (under the cap) but ~20MiB of UTF-8. + const oversizedText = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await waitForWebSocket(oversizedText); + const textClosed = nextWebSocketClose(oversizedText); + oversizedText.send("💥".repeat(5 * 1024 * 1024)); + await expect(textClosed).resolves.toMatchObject({ code: 4009 }); + + // The tunnel itself survives; other connections keep working. + const second = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await waitForWebSocket(second); + second.send("still-alive"); + await expect(nextWebSocketMessage(second).then(webSocketMessageText)).resolves.toBe( + "echo:still-alive", + ); + second.close(); + await delay(50); +}); + +test.concurrent("fails the public WebSocket when the local server rejects it", async ({ task }) => { + await using tunnel = await createTunnelFixture( + task.name, + () => new Response("No WebSockets here\n", { status: 404 }), + ); + + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await expect(waitForWebSocket(socket)).rejects.toThrow(); +}); + +test.concurrent("closes public WebSockets when the tunnel client disconnects", async ({ task }) => { + await using target = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using server = await createServerFixture(); + const tunnel = await createCaptunTunnel({ + gateway: server.gateway, + name: tunnelName(task.name), + token: server.token, + fetch: (request) => target.worker.fetch(request.url, request), + }); + + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await waitForWebSocket(socket); + const closed = nextWebSocketClose(socket); + + tunnel[Symbol.dispose](); + await expect(closed).resolves.toMatchObject({ code: 1001 }); +}); + +test.concurrent("forwards WebSockets through the hosted gateway", async ({ task }) => { + await using target = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using gateway = await createHostedCaptunWorkerFixture(); + using tunnel = await createCaptunTunnel({ + gateway: gateway.origin, + name: tunnelName(task.name), + fetch: (request) => target.worker.fetch(request.url, request), + }); + + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + await waitForWebSocket(socket); + socket.send("hosted"); + await expect(nextWebSocketMessage(socket).then(webSocketMessageText)).resolves.toBe( + "echo:hosted", + ); + socket.close(); + await delay(50); +}); + +test("CLI tunnels WebSocket traffic to a local target", async ({ task }) => { + await using app = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using server = await createServerFixture(); + const ready = Promise.withResolvers<{ url: string }>(); + const shutdown = Promise.withResolvers(); + const router = createCaptunCliRouter({ + readConfig: async () => undefined, + waitForShutdown: () => shutdown.promise, + onTunnelReady: ({ url }) => ready.resolve({ url }), + }); + const client = createRouterClient(router); + + const runTunnel = client.tunnel({ + target: app.origin, + gateway: server.gateway, + name: tunnelName(task.name), + token: server.token, + requestLogs: false, + }); + + const tunnel = await ready.promise; + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws"), { + protocols: ["alpha", "beta"], + headers: { cookie: "session=tunnel-test", authorization: "Bearer tunnel-test" }, + // Node's WebSocket (undici) accepts { protocols, headers }; the DOM type doesn't. + } as unknown as string[]); + try { + await waitForWebSocket(socket); + expect(socket).toMatchObject({ protocol: "alpha" }); + + socket.send("handshake-headers"); + await expect( + nextWebSocketMessage(socket).then(webSocketMessageText).then(JSON.parse), + ).resolves.toMatchObject({ + cookie: "session=tunnel-test", + authorization: "Bearer tunnel-test", + }); + + socket.send("hello-from-cli"); + await expect(nextWebSocketMessage(socket).then(webSocketMessageText)).resolves.toBe( + "echo:hello-from-cli", + ); + + const bytes = new Uint8Array([0, 1, 2, 250, 255]); + socket.send(bytes); + await expect(nextWebSocketMessage(socket).then(webSocketMessageBytes)).resolves.toEqual(bytes); + + const closed = nextWebSocketClose(socket); + socket.send("close-with:4001 done"); + await expect(closed).resolves.toMatchObject({ code: 4001, reason: "done" }); + } finally { + socket.close(); + await delay(50); + shutdown.resolve(); + await runTunnel; + } +}); + +// Captun through Captun: the CLI exposes an entire inner gateway through an +// outer tunnel, so public traffic crosses two gateways and two tunnel clients +// before reaching the worker. (The inner tunnel still connects to its gateway +// directly: a Gateway Connect Request cannot ride through a tunnel because +// gateways claim any request carrying the connect query param before tunnel +// routing.) +test("tunnels Captun through Captun", async ({ task }) => { + await using target = await createMiniflareWorkerFixture({ + entryPoint: "test/fixtures/capnweb-websocket-target.ts", + durableObjects: {}, + bindings: {}, + }); + await using inner = await createTunnelFixture(`${task.name} inner`, (request) => + target.worker.fetch(request.url, request), + ); + + // The outer tunnel exposes the inner gateway like any local server. + await using outerServer = await createServerFixture(); + const ready = Promise.withResolvers<{ url: string }>(); + const shutdown = Promise.withResolvers(); + const router = createCaptunCliRouter({ + readConfig: async () => undefined, + waitForShutdown: () => shutdown.promise, + onTunnelReady: ({ url }) => ready.resolve({ url }), + }); + const runTunnel = createRouterClient(router).tunnel({ + target: new URL(inner.url).origin, + gateway: outerServer.gateway, + name: tunnelName(`${task.name} outer`), + token: outerServer.token, + requestLogs: false, + }); + + const outer = await ready.promise; + const nestedUrl = `${outer.url}${new URL(inner.url).pathname}`; + const rpc = newWebSocketRpcSession<{ ping(value: string): Promise }>( + `${nestedUrl}/rpc`.replace(/^http/, "ws"), + ); + const socket = new WebSocket(`${nestedUrl}/ws`.replace(/^http/, "ws")); + try { + await expect(rpc.ping("captun-through-captun")).resolves.toBe("pong:captun-through-captun"); + + await waitForWebSocket(socket); + socket.send("through-two-tunnels"); + await expect(nextWebSocketMessage(socket).then(webSocketMessageText)).resolves.toBe( + "echo:through-two-tunnels", + ); + } finally { + socket.close(); + rpc[Symbol.dispose](); + await delay(50); + shutdown.resolve(); + await runTunnel; + } +}); + async function createTunnelFixture( testName: string, fetch: (request: Request) => Response | Promise, @@ -238,3 +606,61 @@ function makeBytes(size: number) { function sha256(bytes: Uint8Array) { return createHash("sha256").update(bytes).digest("hex"); } + +function waitForWebSocket(socket: WebSocket) { + if (socket.readyState === WebSocket.OPEN) return Promise.resolve(); + return new Promise((resolveOpen, rejectOpen) => { + socket.addEventListener("open", () => resolveOpen(), { once: true }); + socket.addEventListener("error", () => rejectOpen(new Error("WebSocket error")), { + once: true, + }); + socket.addEventListener("close", () => rejectOpen(new Error("WebSocket closed")), { + once: true, + }); + }); +} + +function nextWebSocketMessage(socket: WebSocket) { + return new Promise((resolveMessage, rejectMessage) => { + socket.addEventListener("message", (event) => resolveMessage(event.data), { once: true }); + socket.addEventListener("error", () => rejectMessage(new Error("WebSocket error")), { + once: true, + }); + socket.addEventListener("close", () => rejectMessage(new Error("WebSocket closed")), { + once: true, + }); + }); +} + +function nextWebSocketClose(socket: WebSocket) { + return new Promise<{ code: number; reason: string }>((resolveClose) => { + socket.addEventListener( + "close", + (event) => { + const { code, reason } = event as { code?: number; reason?: string }; + resolveClose({ code: code ?? 0, reason: reason ?? "" }); + }, + { once: true }, + ); + }); +} + +async function webSocketMessageText(data: unknown) { + if (typeof data === "string") return data; + return new TextDecoder().decode(await webSocketMessageBytes(data)); +} + +async function webSocketMessageBytes(data: unknown): Promise { + if (data instanceof Blob) return new Uint8Array(await data.arrayBuffer()); + if (data instanceof ArrayBuffer) return new Uint8Array(data); + if (data instanceof Uint8Array) return new Uint8Array(data); + throw new Error( + `Expected a binary WebSocket message, got ${typeof data}: ${JSON.stringify(data)?.slice(0, 200)}`, + ); +} + +function delay(ms: number): Promise { + return new Promise((resolveDelay) => { + setTimeout(resolveDelay, ms); + }); +} diff --git a/test/fixtures/capnweb-websocket-target.ts b/test/fixtures/capnweb-websocket-target.ts new file mode 100644 index 0000000..c2189c0 --- /dev/null +++ b/test/fixtures/capnweb-websocket-target.ts @@ -0,0 +1,63 @@ +import { newWorkersWebSocketRpcResponse, RpcTarget } from "capnweb"; + +import type { WebSocketResponseInit, WorkerWebSocketPairConstructor } from "../../src/index.js"; + +class DummyCapability extends RpcTarget { + ping(value: string) { + return `pong:${value}`; + } +} + +export default { + fetch(request: Request) { + const url = new URL(request.url); + if (url.pathname === "/rpc") { + return newWorkersWebSocketRpcResponse(request, new DummyCapability()); + } + if (url.pathname === "/ws") return webSocketEchoResponse(request); + + return new Response("Not found\n", { status: 404 }); + }, +}; + +/** Echoes text as `echo:`, binary untouched, and `close-with: ` as a close. */ +function webSocketEchoResponse(request: Request) { + const WorkerWebSocketPair = ( + globalThis as typeof globalThis & { WebSocketPair: WorkerWebSocketPairConstructor } + ).WebSocketPair; + const pair = new WorkerWebSocketPair(); + const socket = pair[1]; + socket.accept(); + socket.addEventListener("message", (event) => { + void (async () => { + if (typeof event.data !== "string") { + // Binary may arrive as ArrayBuffer or Blob depending on the runtime. + const blob = event.data as Blob; + socket.send(typeof blob.arrayBuffer === "function" ? await blob.arrayBuffer() : event.data); + return; + } + if (event.data.startsWith("close-with:")) { + const [code, ...reason] = event.data.slice("close-with:".length).split(" "); + socket.close(Number(code), reason.join(" ")); + return; + } + if (event.data === "handshake-headers") { + socket.send( + JSON.stringify({ + cookie: request.headers.get("cookie"), + authorization: request.headers.get("authorization"), + }), + ); + return; + } + socket.send(`echo:${event.data}`); + })(); + }); + // Select the first offered subprotocol so tests can assert negotiation. + const protocol = request.headers.get("sec-websocket-protocol")?.split(",")[0]?.trim(); + return new Response(null, { + status: 101, + webSocket: pair[0], + headers: protocol ? { "sec-websocket-protocol": protocol } : undefined, + } as WebSocketResponseInit); +}