From f3c390b7ee8941a77516e79e687cdeab7a65784f Mon Sep 17 00:00:00 2001 From: Jonas Templestein <242550+jonastemplestein@users.noreply.github.com> Date: Thu, 11 Jun 2026 14:22:28 +0100 Subject: [PATCH 01/13] add websocket tunnel proof --- src/cli/bin.ts | 154 +++++++++++--- src/hosted/worker.ts | 10 +- src/index.ts | 198 +++++++++++++++++- src/server/worker.ts | 147 +++++++++++++- test/e2e.test.ts | 234 +++++++++++++++++++++- test/fixtures/bun-websocket-server.js | 24 +++ test/fixtures/capnweb-websocket-target.ts | 18 ++ 7 files changed, 736 insertions(+), 49 deletions(-) create mode 100644 test/fixtures/bun-websocket-server.js create mode 100644 test/fixtures/capnweb-websocket-target.ts diff --git a/src/cli/bin.ts b/src/cli/bin.ts index d7c3833..c61120e 100755 --- a/src/cli/bin.ts +++ b/src/cli/bin.ts @@ -15,7 +15,11 @@ import { color } from "./ansi.js"; import { CliFriendlyError } from "./cli-error.js"; import { CaptunTunnelConnectError, + createWebSocketForwardResponse, createCaptunTunnel, + type Fetcher, + type WebSocketFetcher, + type WebSocketForwardResult, HOSTED_CAPTUN_GATEWAY, randomConnectToken, } from "../index.js"; @@ -522,7 +526,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 +541,125 @@ 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 (request.headers.get("upgrade") === "websocket") { + 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`, - { status: 502 }, - ); - logRequest(tunnel.requestLogs, { - method: request.method, - path: `${url.pathname}${url.search}`, - rayId, - status: response.status, - startedAt: requestStartedAt, - }); - return response; - } + 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) => connectTargetWebSocket(tunnel, request), }; } +async function connectTargetWebSocket( + tunnel: ResolvedTunnel, + request: Request, +): Promise { + const url = new URL(request.url); + const requestStartedAt = performance.now(); + const rayId = request.headers.get("cf-ray") || "-"; + const targetUrl = new URL(`${tunnel.target}${url.pathname}${url.search}`); + targetUrl.protocol = targetUrl.protocol === "https:" ? "wss:" : "ws:"; + + const targetSocket = new WebSocket( + targetUrl, + request.headers + .get("sec-websocket-protocol") + ?.split(",") + .map((protocol) => protocol.trim()) + .filter(Boolean), + ); + + try { + await waitForWebSocketOpen(targetSocket); + logRequest(tunnel.requestLogs, { + method: request.method, + path: `${url.pathname}${url.search}`, + rayId, + status: 101, + startedAt: requestStartedAt, + }); + + return { + accepted: true, + headers: targetSocket.protocol ? [["sec-websocket-protocol", targetSocket.protocol]] : [], + response: createWebSocketForwardResponse(targetSocket, request.body), + }; + } catch { + targetSocket.close(); + const 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 { accepted: false, response }; + } +} + +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..eff23c9 100644 --- a/src/hosted/worker.ts +++ b/src/hosted/worker.ts @@ -102,10 +102,14 @@ export default { customHostname, tunnelName, }); - const response = await shard.forward( + const tunnelRequest = createTunnelForwardRequest(forwarded, { tunnelName, - createTunnelForwardRequest(forwarded, tunnelUrl), - ); + tunnelUrl, + }); + const response = + request.headers.get("upgrade") === "websocket" + ? await shard.fetch(tunnelRequest) + : await shard.forward(tunnelName, tunnelRequest); return stripSetCookieHeadersOutsideTunnel(response, new URL(tunnelUrl).hostname); }, } satisfies ExportedHandler; diff --git a/src/index.ts b/src/index.ts index da43598..69b9ff6 100644 --- a/src/index.ts +++ b/src/index.ts @@ -18,12 +18,16 @@ export interface Fetcher { fetch(request: Request): Response | Promise; } +export interface WebSocketFetcher { + connectWebSocket(request: Request): WebSocketForwardResult | 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 +35,16 @@ export interface RemoteFetcherCapability extends FetcherStub { onRpcBroken(callback: () => void): void; } +export type WebSocketMessage = string | Uint8Array; + +export type WebSocketFrame = + | { type: "message"; data: WebSocketMessage } + | { type: "close"; code?: number; reason?: string }; + +export type WebSocketForwardResult = + | { accepted: true; headers?: Array<[string, string]>; response: Response } + | { accepted: false; response: Response }; + export function fetcherStubFromRemoteCapability( remote: RemoteFetcherCapability, options: { onDisconnect?: () => void }, @@ -39,6 +53,7 @@ export function fetcherStubFromRemoteCapability( return { fetch: (request) => remote.fetch(request), + connectWebSocket: async (request) => await remote.connectWebSocket(request), ready: (tunnel) => remote.ready(tunnel), [Symbol.dispose]: () => remote[Symbol.dispose](), }; @@ -48,7 +63,7 @@ export function acceptFetcherCapabilityFromSocket( socket: WebSocket, options: { onDisconnect?: () => void } = {}, ): FetcherStub { - const remote = newWebSocketRpcSession(socket); + const remote = newWebSocketRpcSession(socket) as unknown as RemoteFetcherCapability; return fetcherStubFromRemoteCapability(remote, options); } @@ -134,9 +149,10 @@ 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 & { accept(): void; @@ -157,6 +173,7 @@ const WEBSOCKET_REJECTION_PROBE_TIMEOUT_MS = 500; /** Creates a public tunnel by exposing a local fetch implementation to a Tunnel Gateway. */ export async function createCaptunTunnel( options: Fetcher & { + connectWebSocket?: WebSocketFetcher["connectWebSocket"]; gateway?: string | URL; name?: string; token?: string; @@ -167,6 +184,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); @@ -216,12 +234,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; } @@ -232,6 +255,20 @@ class TunnelTargetFetcher extends RpcTarget implements TunnelClientCapability { ready(tunnel: TunnelReady) { this.onReady(tunnel); } + + async connectWebSocket(request: Request) { + if (this.fetcher.connectWebSocket) return this.fetcher.connectWebSocket(request); + + const response = await this.fetcher.fetch(createWebSocketUpgradeRequest(request)); + const webSocket = responseWebSocket(response); + if (!webSocket) return { accepted: false as const, response }; + + return { + accepted: true as const, + headers: [...response.headers], + response: createWebSocketForwardResponse(webSocket, request.body), + }; + } } function createWebSocket(url: string | URL, protocols: string[]) { @@ -358,3 +395,150 @@ export function acceptFetcherCapability( response: new Response(null, responseInit), }; } + +export function createWebSocketForwardResponse( + socket: WebSocket, + incoming: ReadableStream | null, +): Response { + acceptIfNeeded(socket); + if (incoming) { + void pipeFramesToWebSocket(incoming, socket); + } + + return new Response(webSocketFramesFromSocket(socket), { + headers: { "content-type": "application/x-captun-websocket-frames" }, + }); +} + +export async function pipeFramesToWebSocket(stream: ReadableStream, socket: WebSocket) { + for await (const frame of decodeWebSocketFrames(stream)) { + if (frame.type === "close") { + closeWebSocket(socket, frame.code, frame.reason); + return; + } + sendWebSocketMessage(socket, frame.data); + } +} + +function closeWebSocket(socket: WebSocket, code?: number, reason?: string) { + if (code === undefined || code === 1000 || (code >= 3000 && code <= 4999)) { + socket.close(code, reason); + return; + } + socket.close(); +} + +export function webSocketFramesFromSocket(socket: WebSocket): ReadableStream { + return new ReadableStream({ + start(controller) { + socket.addEventListener("message", (event) => { + void webSocketFrameFromMessage(event.data).then((frame) => { + controller.enqueue(encodeWebSocketFrame(frame)); + }); + }); + socket.addEventListener("close", (event) => { + controller.enqueue( + encodeWebSocketFrame({ type: "close", code: event.code, reason: event.reason }), + ); + controller.close(); + }); + socket.addEventListener("error", (event) => { + controller.error(event instanceof ErrorEvent ? event.error : new Error("WebSocket error")); + }); + }, + cancel() { + socket.close(1000, "stream canceled"); + }, + }); +} + +export async function* decodeWebSocketFrames( + stream: ReadableStream, +): AsyncIterable { + const decoder = new TextDecoder(); + let buffer = ""; + for await (const chunk of stream) { + buffer += decoder.decode(chunk, { stream: true }); + let newline = buffer.indexOf("\n"); + while (newline !== -1) { + const line = buffer.slice(0, newline); + buffer = buffer.slice(newline + 1); + if (line) yield decodeWebSocketFrame(line); + newline = buffer.indexOf("\n"); + } + } + buffer += decoder.decode(); + if (buffer) yield decodeWebSocketFrame(buffer); +} + +export async function webSocketFrameFromMessage(data: unknown): Promise { + if (typeof data === "string" || data instanceof Uint8Array) { + return { type: "message", data }; + } + if (data instanceof ArrayBuffer) return { type: "message", data: new Uint8Array(data) }; + if (data instanceof Blob) + return { type: "message", data: new Uint8Array(await data.arrayBuffer()) }; + return { type: "message", data: String(data) }; +} + +export function sendWebSocketMessage(socket: WebSocket, message: WebSocketMessage) { + if (typeof message === "string") { + socket.send(message); + return; + } + const bytes = new Uint8Array(message); + socket.send(bytes.buffer); +} + +export function encodeWebSocketFrame(frame: WebSocketFrame): Uint8Array { + const encoded = + frame.type === "message" + ? { + type: "message", + text: typeof frame.data === "string" ? frame.data : undefined, + bytes: typeof frame.data === "string" ? undefined : base64Encode(frame.data), + } + : frame; + return new TextEncoder().encode(`${JSON.stringify(encoded)}\n`); +} + +function decodeWebSocketFrame(line: string): WebSocketFrame { + const frame = JSON.parse(line) as { + type: "message" | "close"; + text?: string; + bytes?: string; + code?: number; + reason?: string; + }; + if (frame.type === "close") return { type: "close", code: frame.code, reason: frame.reason }; + return { + type: "message", + data: frame.bytes ? base64Decode(frame.bytes) : frame.text || "", + }; +} + +function base64Encode(bytes: Uint8Array) { + let binary = ""; + for (const byte of bytes) binary += String.fromCharCode(byte); + return btoa(binary); +} + +function base64Decode(value: string) { + const binary = atob(value); + return Uint8Array.from(binary, (char) => char.charCodeAt(0)); +} + +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..d334a79 100644 --- a/src/server/worker.ts +++ b/src/server/worker.ts @@ -2,9 +2,13 @@ import { DurableObject } from "cloudflare:workers"; import { acceptFetcherCapability, connectTokenFromRequest, + decodeWebSocketFrames, + encodeWebSocketFrame, GATEWAY_CONNECT_QUERY_PARAM, + sendWebSocketMessage, TUNNEL_CONNECT_DIAGNOSTIC_HEADER, TUNNEL_NAME_QUERY_PARAM, + webSocketFrameFromMessage, type FetcherStub, } from "../index.js"; import { @@ -27,6 +31,19 @@ export type CaptunEnv = { const TUNNEL_NAME_HEADER = "x-captun-tunnel-name"; const CUSTOM_HOSTNAME_RESERVED_TUNNEL_NAMES = ["captun", "gateway"]; +type WorkerWebSocket = WebSocket & { + accept(): void; +}; + +type WorkerWebSocketPairConstructor = new () => { + 0: WorkerWebSocket; + 1: WorkerWebSocket; +}; + +type WebSocketResponseInit = ResponseInit & { + webSocket: WebSocket; +}; + type CaptunShardBindingEnv = { CaptunServerShard: DurableObjectNamespace>; SHARD_COUNT?: string; @@ -91,7 +108,14 @@ 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 (!isGatewayConnectRequest(request)) { + if (request.headers.get("upgrade") !== "websocket") { + 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 (request.headers.get("upgrade") !== "websocket") { return this.diagnoseConnect(tunnelName, request); } @@ -136,6 +160,9 @@ export class CaptunServerShard< const tunnel = this.tunnels.get(tunnelName)?.fetcher; if (!tunnel) return new Response("No tunnel client connected\n", { status: 503 }); try { + if (request.headers.get("upgrade") === "websocket") { + return await forwardWebSocket(tunnel, request); + } return await tunnel.fetch(request); } catch { return new Response("Tunnel fetch failed\n", { status: 502 }); @@ -179,7 +206,13 @@ export default { customHostname: env.CUSTOM_HOSTNAME, tunnelName, }); - return shard.forward(tunnelName, createTunnelForwardRequest(forwarded, tunnelUrl)); + const tunnelRequest = createTunnelForwardRequest(forwarded, { + tunnelName, + tunnelUrl, + }); + + if (request.headers.get("upgrade") === "websocket") return shard.fetch(tunnelRequest); + return shard.forward(tunnelName, tunnelRequest); }, } satisfies ExportedHandler; @@ -238,12 +271,118 @@ 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 } | string, +): Request { const headers = new Headers(request.headers); - headers.set(TUNNEL_URL_HEADER, tunnelUrl); + if (typeof input === "string") { + headers.set(TUNNEL_URL_HEADER, input); + } else { + if (input.tunnelName) headers.set(TUNNEL_NAME_HEADER, input.tunnelName); + headers.set(TUNNEL_URL_HEADER, input.tunnelUrl); + } return new Request(request, { headers }); } +async function forwardWebSocket(tunnel: FetcherStub, request: Request) { + if (!tunnel.connectWebSocket) { + return new Response("Tunnel client cannot forward WebSockets\n", { status: 501 }); + } + + const WorkerWebSocketPair = ( + globalThis as typeof globalThis & { WebSocketPair: WorkerWebSocketPairConstructor } + ).WebSocketPair; + const pair = new WorkerWebSocketPair(); + const clientSocket = pair[0]; + const serverSocket = pair[1]; + serverSocket.accept(); + const incoming = new TransformStream(); + const writer = incoming.writable.getWriter(); + + try { + const result = await tunnel.connectWebSocket( + createWebSocketConnectRequest(request, incoming.readable), + ); + if (!result.accepted) { + serverSocket.close(1000, "WebSocket not accepted"); + return result.response; + } + + pipeWebSocketToRequestBody(serverSocket, writer); + pipeResponseBodyToWebSocket(result.response, serverSocket); + return new Response(null, { + status: 101, + webSocket: clientSocket, + headers: result.headers, + } as WebSocketResponseInit); + } catch (error) { + writer.releaseLock(); + serverSocket.close(1011, "WebSocket tunnel failed"); + throw error; + } +} + +function createWebSocketConnectRequest(request: Request, body: ReadableStream) { + const headers = new Headers(request.headers); + headers.delete("upgrade"); + headers.delete("connection"); + return new Request(request.url, { + method: "POST", + headers, + body, + // Required by Node-compatible Request implementations; harmless in Workers. + duplex: "half", + } as RequestInit & { duplex: "half" }); +} + +function pipeWebSocketToRequestBody( + socket: WebSocket, + writer: WritableStreamDefaultWriter, +) { + socket.addEventListener("message", (event) => { + void webSocketFrameFromMessage(event.data).then((frame) => + writer.write(encodeWebSocketFrame(frame)), + ); + }); + socket.addEventListener("close", (event) => { + void writer + .write(encodeWebSocketFrame({ type: "close", code: event.code, reason: event.reason })) + .finally(() => { + writer.close(); + writer.releaseLock(); + }); + }); + socket.addEventListener("error", () => { + void writer.abort(new Error("WebSocket error")); + }); +} + +function pipeResponseBodyToWebSocket(response: Response, socket: WebSocket) { + void (async () => { + if (!response.body) return; + try { + for await (const frame of decodeWebSocketFrames(response.body)) { + if (frame.type === "close") { + closeWebSocket(socket, frame.code, frame.reason); + return; + } + sendWebSocketMessage(socket, frame.data); + } + } catch { + socket.close(1011, "WebSocket tunnel failed"); + } + })(); +} + +function closeWebSocket(socket: WebSocket, code?: number, reason?: string) { + if (code === undefined || code === 1000 || (code >= 3000 && code <= 4999)) { + socket.close(code, reason); + return; + } + socket.close(); +} + 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..822da91 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -1,10 +1,19 @@ +import { spawn, type ChildProcessByStdio } from "node:child_process"; import { createHash } from "node:crypto"; +import net from "node:net"; +import { dirname, resolve } from "node:path"; +import type { Readable } from "node:stream"; +import { fileURLToPath } from "node:url"; + +import { createRouterClient } from "@orpc/server"; +import { newWebSocketRpcSession } from "capnweb"; import { expect, test, vi } from "vitest"; +import { createCaptunCliRouter } from "../src/cli/bin.js"; import { createCaptunTunnel } from "../src/index.js"; -import { createCaptunWorkerFixture } from "./miniflare.js"; +import { createCaptunWorkerFixture, createMiniflareWorkerFixture } from "./miniflare.js"; -vi.setConfig({ testTimeout: 15_000 }); +vi.setConfig({ testTimeout: 25_000 }); test.concurrent("forwards HTTP", async ({ task }) => { await using tunnel = await createTunnelFixture(task.name, async (request) => @@ -171,6 +180,59 @@ 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); +}); + +test("CLI tunnels WebSocket traffic to a local Bun server", async ({ task }) => { + await using app = await createBunWebSocketFixture(); + 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: String(app.port), + 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")); + try { + await waitForWebSocket(socket); + socket.send("hello-from-cli"); + await expect(readWebSocketMessage(socket)).resolves.toBe("echo:hello-from-cli"); + } finally { + socket.close(); + await delay(50); + shutdown.resolve(); + await runTunnel; + } +}); + async function createTunnelFixture( testName: string, fetch: (request: Request) => Response | Promise, @@ -238,3 +300,171 @@ function makeBytes(size: number) { function sha256(bytes: Uint8Array) { return createHash("sha256").update(bytes).digest("hex"); } + +async function createBunWebSocketFixture() { + const port = await getAvailablePort(); + const server = spawn("bun", ["run", "test/fixtures/bun-websocket-server.js"], { + cwd: resolve(dirname(fileURLToPath(import.meta.url)), ".."), + env: { ...process.env, PORT: String(port) }, + stdio: ["ignore", "pipe", "pipe"], + }); + const output = captureOutput(server); + + try { + await waitForTcp(port, server, output); + return { + port, + async [Symbol.asyncDispose]() { + await stopProcess(server); + }, + }; + } catch (error) { + await stopProcess(server); + throw new Error( + formatFixtureFailure(error instanceof Error ? error.message : String(error), output.logs()), + ); + } +} + +type ServerProcess = ChildProcessByStdio; + +async function getAvailablePort(): Promise { + return new Promise((resolvePort, reject) => { + const server = net.createServer(); + server.listen(0, "127.0.0.1", () => { + const address = server.address(); + if (!address || typeof address === "string") { + reject(new Error(`Failed to allocate a local port: ${String(address)}`)); + return; + } + + server.close((error) => { + if (error) reject(error); + else resolvePort(address.port); + }); + }); + server.on("error", reject); + }); +} + +async function waitForTcp(port: number, server: ServerProcess, output: CapturedProcessOutput) { + const startedAt = Date.now(); + while (Date.now() - startedAt < 15_000) { + const error = output.error(); + if (error) throw error; + if (server.exitCode !== null || server.signalCode) { + throw new Error( + `Bun server exited before port ${port} accepted connections\n\n${output.logs().trim() || "(none)"}`, + ); + } + + if (await canConnect(port)) return; + + await delay(100); + } + + throw new Error(`Timed out waiting for Bun server to accept connections on port ${port}`); +} + +function canConnect(port: number) { + return new Promise((resolveConnect) => { + const socket = net.connect(port, "127.0.0.1"); + socket.once("connect", () => { + socket.destroy(); + resolveConnect(true); + }); + socket.once("error", () => { + socket.destroy(); + resolveConnect(false); + }); + }); +} + +function captureOutput(child: ServerProcess) { + const chunks: string[] = []; + let processError: Error | undefined; + const capture = (chunk: string | Buffer) => { + chunks.push(String(chunk)); + if (chunks.length > 200) chunks.shift(); + }; + child.stdout.on("data", capture); + child.stderr.on("data", capture); + child.on("error", (error) => { + processError = error; + chunks.push(error.stack || error.message); + }); + + return { + logs: () => chunks.join(""), + error: () => processError, + }; +} + +interface CapturedProcessOutput { + logs(): string; + error(): Error | undefined; +} + +function formatFixtureFailure(message: string, serverLogs: string) { + return [message, "", "Server logs:", serverLogs.trim() || "(none)"].join("\n"); +} + +async function stopProcess(child: ServerProcess): Promise { + if (child.exitCode !== null || child.killed) return; + + child.kill("SIGINT"); + const exited = await Promise.race([ + new Promise((resolveExit) => child.once("exit", () => resolveExit(true))), + delay(5_000).then(() => false), + ]); + + if (!exited && child.exitCode === null && !child.killed) { + child.kill("SIGKILL"); + await new Promise((resolveExit) => child.once("exit", () => resolveExit())); + } +} + +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 readWebSocketMessage(socket: WebSocket) { + return new Promise((resolveMessage, rejectMessage) => { + socket.addEventListener( + "message", + async (event) => { + resolveMessage(await webSocketMessageText(event.data)); + }, + { once: true }, + ); + socket.addEventListener("error", () => rejectMessage(new Error("WebSocket error")), { + once: true, + }); + socket.addEventListener("close", () => rejectMessage(new Error("WebSocket closed")), { + once: true, + }); + }); +} + +async function webSocketMessageText(data: unknown) { + if (typeof data === "string") return data; + if (data instanceof Blob) return data.text(); + if (data instanceof ArrayBuffer) return new TextDecoder().decode(data); + if (data instanceof Uint8Array) return new TextDecoder().decode(data); + return String(data); +} + +function delay(ms: number): Promise { + return new Promise((resolveDelay) => { + setTimeout(resolveDelay, ms); + }); +} diff --git a/test/fixtures/bun-websocket-server.js b/test/fixtures/bun-websocket-server.js new file mode 100644 index 0000000..b7b23ab --- /dev/null +++ b/test/fixtures/bun-websocket-server.js @@ -0,0 +1,24 @@ +const server = Bun.serve({ + hostname: "127.0.0.1", + port: Number(process.env.PORT), + fetch(request, server) { + const url = new URL(request.url); + + if (url.pathname === "/ws") { + if (server.upgrade(request)) return; + return new Response("WebSocket upgrade failed\n", { status: 500 }); + } + + return new Response("ok\n"); + }, + websocket: { + message(socket, message) { + socket.send(`echo:${message}`); + }, + }, +}); + +process.on("SIGINT", () => { + server.stop(true); + process.exit(0); +}); diff --git a/test/fixtures/capnweb-websocket-target.ts b/test/fixtures/capnweb-websocket-target.ts new file mode 100644 index 0000000..37dc5f6 --- /dev/null +++ b/test/fixtures/capnweb-websocket-target.ts @@ -0,0 +1,18 @@ +import { newWorkersWebSocketRpcResponse, RpcTarget } from "capnweb"; + +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()); + } + + return new Response("Not found\n", { status: 404 }); + }, +}; From 0ba1b534aca97b4c41859007569afcc63f24e5a6 Mon Sep 17 00:00:00 2001 From: Jonas Templestein <242550+jonastemplestein@users.noreply.github.com> Date: Thu, 11 Jun 2026 14:58:57 +0100 Subject: [PATCH 02/13] Replace the WebSocket frame codec with Cap'n Web capability passing The tunnel already runs over Cap'n Web, which passes capabilities in both directions and serializes Uint8Array natively, so tunneled WebSockets don't need a bespoke wire format. connectWebSocket now receives a WebSocketHandle for the public socket and returns one for the local socket; each message is an ordered RPC call. This deletes the newline-JSON/base64 frame codec, the streamed-body POST bridge, and the listener-vs-close races that could drop or reorder frames and crash the CLI on an unhandled rejection. Also: - Pump conversions through a promise chain so Blob messages (Node delivers binary as Blob) cannot reorder or race the close event. - Forward only the negotiated subprotocol onto the public 101 instead of every local response header. - Skip set-cookie stripping for 101 responses, which cannot be reconstructed. - Compare the Upgrade header case-insensitively. - Drop the unused string overload of createTunnelForwardRequest and the unreachable 501 guard. - Replace the spawned Bun fixture with a miniflare echo worker and cover binary frames and close-code propagation end to end. --- src/cli/bin.ts | 70 +++---- src/hosted/worker.ts | 15 +- src/index.ts | 238 ++++++++++------------ src/server/worker.ts | 126 +++--------- test/e2e.test.ts | 214 +++++++------------ test/fixtures/bun-websocket-server.js | 24 --- test/fixtures/capnweb-websocket-target.ts | 36 ++++ 7 files changed, 289 insertions(+), 434 deletions(-) delete mode 100644 test/fixtures/bun-websocket-server.js diff --git a/src/cli/bin.ts b/src/cli/bin.ts index c61120e..42640d7 100755 --- a/src/cli/bin.ts +++ b/src/cli/bin.ts @@ -15,13 +15,16 @@ import { color } from "./ansi.js"; import { CliFriendlyError } from "./cli-error.js"; import { CaptunTunnelConnectError, - createWebSocketForwardResponse, createCaptunTunnel, - type Fetcher, - type WebSocketFetcher, - type WebSocketForwardResult, 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"; @@ -544,7 +547,7 @@ async function connectTunnelWithRetry( function makeTunnelFetcher(tunnel: ResolvedTunnel): Fetcher & WebSocketFetcher { return { fetch: async (request: Request) => { - if (request.headers.get("upgrade") === "websocket") { + if (isWebSocketUpgradeRequest(request)) { return new Response("Use connectWebSocket for WebSocket tunnel requests\n", { status: 400, }); @@ -583,20 +586,28 @@ function makeTunnelFetcher(tunnel: ResolvedTunnel): Fetcher & WebSocketFetcher { } }, - connectWebSocket: (request) => connectTargetWebSocket(tunnel, request), + connectWebSocket: (request, remote) => connectTargetWebSocket(tunnel, request, remote), }; } async function connectTargetWebSocket( tunnel: ResolvedTunnel, request: Request, -): Promise { + remote: WebSocketHandle, +): Promise { const url = new URL(request.url); const requestStartedAt = performance.now(); - const rayId = request.headers.get("cf-ray") || "-"; + 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, request.headers @@ -608,34 +619,25 @@ async function connectTargetWebSocket( try { await waitForWebSocketOpen(targetSocket); - logRequest(tunnel.requestLogs, { - method: request.method, - path: `${url.pathname}${url.search}`, - rayId, - status: 101, - startedAt: requestStartedAt, - }); - - return { - accepted: true, - headers: targetSocket.protocol ? [["sec-websocket-protocol", targetSocket.protocol]] : [], - response: createWebSocketForwardResponse(targetSocket, request.body), - }; } catch { targetSocket.close(); - const 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 { accepted: false, response }; + log(502); + return { + accepted: false, + response: new Response( + `Request reached the captun cli, but ${targetUrl.origin} did not accept the WebSocket\n`, + { status: 502 }, + ), + }; } + + log(101); + pipeWebSocketToHandle(targetSocket, remote); + return { + accepted: true, + protocol: targetSocket.protocol || undefined, + socket: webSocketHandleFromSocket(targetSocket), + }; } async function waitForWebSocketOpen(socket: WebSocket) { diff --git a/src/hosted/worker.ts b/src/hosted/worker.ts index eff23c9..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"; @@ -106,17 +107,16 @@ export default { tunnelName, tunnelUrl, }); - const response = - request.headers.get("upgrade") === "websocket" - ? await shard.fetch(tunnelRequest) - : await shard.forward(tunnelName, tunnelRequest); + 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 }); } @@ -162,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"; } @@ -181,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 69b9ff6..05400b5 100644 --- a/src/index.ts +++ b/src/index.ts @@ -19,7 +19,10 @@ export interface Fetcher { } export interface WebSocketFetcher { - connectWebSocket(request: Request): WebSocketForwardResult | Promise; + connectWebSocket( + request: Request, + remote: WebSocketHandle, + ): WebSocketConnectResult | Promise; } export type TunnelReady = { @@ -37,12 +40,19 @@ export interface RemoteFetcherCapability extends FetcherStub { export type WebSocketMessage = string | Uint8Array; -export type WebSocketFrame = - | { type: "message"; data: WebSocketMessage } - | { type: "close"; code?: number; reason?: string }; +/** + * 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 WebSocketForwardResult = - | { accepted: true; headers?: Array<[string, string]>; response: Response } +export type WebSocketConnectResult = + | { accepted: true; protocol?: string; socket: WebSocketHandle } | { accepted: false; response: Response }; export function fetcherStubFromRemoteCapability( @@ -53,7 +63,7 @@ export function fetcherStubFromRemoteCapability( return { fetch: (request) => remote.fetch(request), - connectWebSocket: async (request) => await remote.connectWebSocket(request), + connectWebSocket: (request, handle) => remote.connectWebSocket(request, handle), ready: (tunnel) => remote.ready(tunnel), [Symbol.dispose]: () => remote[Symbol.dispose](), }; @@ -256,17 +266,19 @@ class TunnelTargetFetcher extends RpcTarget implements TunnelClientCapability { this.onReady(tunnel); } - async connectWebSocket(request: Request) { - if (this.fetcher.connectWebSocket) return this.fetcher.connectWebSocket(request); + 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 webSocket = responseWebSocket(response); - if (!webSocket) return { accepted: false as const, response }; + const socket = responseWebSocket(response); + if (!socket) return { accepted: false as const, response }; + acceptIfNeeded(socket); + pipeWebSocketToHandle(socket, remote); return { accepted: true as const, - headers: [...response.headers], - response: createWebSocketForwardResponse(webSocket, request.body), + protocol: response.headers.get("sec-websocket-protocol") || undefined, + socket: webSocketHandleFromSocket(socket), }; } } @@ -396,136 +408,112 @@ export function acceptFetcherCapability( }; } -export function createWebSocketForwardResponse( - socket: WebSocket, - incoming: ReadableStream | null, -): Response { - acceptIfNeeded(socket); - if (incoming) { - void pipeFramesToWebSocket(incoming, socket); - } +export function isWebSocketUpgradeRequest(request: Request): boolean { + return request.headers.get("upgrade")?.toLowerCase() === "websocket"; +} - return new Response(webSocketFramesFromSocket(socket), { - headers: { "content-type": "application/x-captun-websocket-frames" }, - }); +/** 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); } -export async function pipeFramesToWebSocket(stream: ReadableStream, socket: WebSocket) { - for await (const frame of decodeWebSocketFrames(stream)) { - if (frame.type === "close") { - closeWebSocket(socket, frame.code, frame.reason); - return; - } - sendWebSocketMessage(socket, frame.data); +class SocketHandle extends RpcTarget implements WebSocketHandle { + constructor(private socket: WebSocket) { + super(); } -} -function closeWebSocket(socket: WebSocket, code?: number, reason?: string) { - if (code === undefined || code === 1000 || (code >= 3000 && code <= 4999)) { - socket.close(code, reason); - return; + send(message: WebSocketMessage) { + this.socket.send(message); + } + + close(code?: number, reason?: string) { + closeWebSocket(this.socket, code, reason); } - socket.close(); } -export function webSocketFramesFromSocket(socket: WebSocket): ReadableStream { - return new ReadableStream({ - start(controller) { - socket.addEventListener("message", (event) => { - void webSocketFrameFromMessage(event.data).then((frame) => { - controller.enqueue(encodeWebSocketFrame(frame)); - }); - }); - socket.addEventListener("close", (event) => { - controller.enqueue( - encodeWebSocketFrame({ type: "close", code: event.code, reason: event.reason }), - ); - controller.close(); - }); - socket.addEventListener("error", (event) => { - controller.error(event instanceof ErrorEvent ? event.error : new Error("WebSocket error")); - }); - }, - cancel() { - socket.close(1000, "stream canceled"); - }, +/** + * 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(() => { + forward(remote.close(code, reason)); + disposeStub(remote); + }); + }; + + socket.addEventListener("message", (event) => { + if (finished) return; + enqueue(async () => { + forward(remote.send(await webSocketMessage(event.data))); + }); }); + socket.addEventListener("close", (event) => finish(event.code, event.reason)); + socket.addEventListener("error", () => finish(1011, "WebSocket error")); } -export async function* decodeWebSocketFrames( - stream: ReadableStream, -): AsyncIterable { - const decoder = new TextDecoder(); - let buffer = ""; - for await (const chunk of stream) { - buffer += decoder.decode(chunk, { stream: true }); - let newline = buffer.indexOf("\n"); - while (newline !== -1) { - const line = buffer.slice(0, newline); - buffer = buffer.slice(newline + 1); - if (line) yield decodeWebSocketFrame(line); - newline = buffer.indexOf("\n"); - } - } - buffer += decoder.decode(); - if (buffer) yield decodeWebSocketFrame(buffer); -} +type StubLike = { dup?(): unknown; [Symbol.dispose]?(): void }; -export async function webSocketFrameFromMessage(data: unknown): Promise { - if (typeof data === "string" || data instanceof Uint8Array) { - return { type: "message", data }; - } - if (data instanceof ArrayBuffer) return { type: "message", data: new Uint8Array(data) }; - if (data instanceof Blob) - return { type: "message", data: new Uint8Array(await data.arrayBuffer()) }; - return { type: "message", data: String(data) }; +function dupStub(handle: WebSocketHandle): WebSocketHandle { + const dup = (handle as WebSocketHandle & StubLike).dup; + return typeof dup === "function" ? (dup.call(handle) as WebSocketHandle) : handle; } -export function sendWebSocketMessage(socket: WebSocket, message: WebSocketMessage) { - if (typeof message === "string") { - socket.send(message); - return; - } - const bytes = new Uint8Array(message); - socket.send(bytes.buffer); -} - -export function encodeWebSocketFrame(frame: WebSocketFrame): Uint8Array { - const encoded = - frame.type === "message" - ? { - type: "message", - text: typeof frame.data === "string" ? frame.data : undefined, - bytes: typeof frame.data === "string" ? undefined : base64Encode(frame.data), - } - : frame; - return new TextEncoder().encode(`${JSON.stringify(encoded)}\n`); -} - -function decodeWebSocketFrame(line: string): WebSocketFrame { - const frame = JSON.parse(line) as { - type: "message" | "close"; - text?: string; - bytes?: string; - code?: number; - reason?: string; - }; - if (frame.type === "close") return { type: "close", code: frame.code, reason: frame.reason }; - return { - type: "message", - data: frame.bytes ? base64Decode(frame.bytes) : frame.text || "", - }; +function disposeStub(handle: WebSocketHandle) { + (handle as WebSocketHandle & StubLike)[Symbol.dispose]?.(); } -function base64Encode(bytes: Uint8Array) { - let binary = ""; - for (const byte of bytes) binary += String.fromCharCode(byte); - return btoa(binary); +/** + * 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). + */ +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()); + } + return String(data); } -function base64Decode(value: string) { - const binary = atob(value); - return Uint8Array.from(binary, (char) => char.charCodeAt(0)); +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) { diff --git a/src/server/worker.ts b/src/server/worker.ts index d334a79..c88e61b 100644 --- a/src/server/worker.ts +++ b/src/server/worker.ts @@ -2,14 +2,14 @@ import { DurableObject } from "cloudflare:workers"; import { acceptFetcherCapability, connectTokenFromRequest, - decodeWebSocketFrames, - encodeWebSocketFrame, GATEWAY_CONNECT_QUERY_PARAM, - sendWebSocketMessage, + isWebSocketUpgradeRequest, + pipeWebSocketToHandle, TUNNEL_CONNECT_DIAGNOSTIC_HEADER, TUNNEL_NAME_QUERY_PARAM, - webSocketFrameFromMessage, + webSocketHandleFromSocket, type FetcherStub, + type WebSocketConnectResult, } from "../index.js"; import { captunShardName, @@ -109,14 +109,14 @@ export class CaptunServerShard< if (!tunnelUrl) return new Response("Missing tunnel URL\n", { status: 404 }); if (!isGatewayConnectRequest(request)) { - if (request.headers.get("upgrade") !== "websocket") { + 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 (request.headers.get("upgrade") !== "websocket") { + if (!isWebSocketUpgradeRequest(request)) { return this.diagnoseConnect(tunnelName, request); } @@ -160,7 +160,7 @@ export class CaptunServerShard< const tunnel = this.tunnels.get(tunnelName)?.fetcher; if (!tunnel) return new Response("No tunnel client connected\n", { status: 503 }); try { - if (request.headers.get("upgrade") === "websocket") { + if (isWebSocketUpgradeRequest(request)) { return await forwardWebSocket(tunnel, request); } return await tunnel.fetch(request); @@ -211,14 +211,14 @@ export default { tunnelUrl, }); - if (request.headers.get("upgrade") === "websocket") return shard.fetch(tunnelRequest); + 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 }); } @@ -244,7 +244,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"; } @@ -273,114 +273,42 @@ export function createTunnelConnectRequest(input: { export function createTunnelForwardRequest( request: Request, - input: { tunnelName?: string; tunnelUrl: string } | string, + input: { tunnelName: string; tunnelUrl: string }, ): Request { const headers = new Headers(request.headers); - if (typeof input === "string") { - headers.set(TUNNEL_URL_HEADER, input); - } else { - if (input.tunnelName) headers.set(TUNNEL_NAME_HEADER, input.tunnelName); - headers.set(TUNNEL_URL_HEADER, input.tunnelUrl); - } + headers.set(TUNNEL_NAME_HEADER, input.tunnelName); + headers.set(TUNNEL_URL_HEADER, input.tunnelUrl); return new Request(request, { headers }); } async function forwardWebSocket(tunnel: FetcherStub, request: Request) { - if (!tunnel.connectWebSocket) { - return new Response("Tunnel client cannot forward WebSockets\n", { status: 501 }); - } - const WorkerWebSocketPair = ( globalThis as typeof globalThis & { WebSocketPair: WorkerWebSocketPairConstructor } ).WebSocketPair; const pair = new WorkerWebSocketPair(); - const clientSocket = pair[0]; const serverSocket = pair[1]; serverSocket.accept(); - const incoming = new TransformStream(); - const writer = incoming.writable.getWriter(); + let result: WebSocketConnectResult; try { - const result = await tunnel.connectWebSocket( - createWebSocketConnectRequest(request, incoming.readable), - ); - if (!result.accepted) { - serverSocket.close(1000, "WebSocket not accepted"); - return result.response; - } - - pipeWebSocketToRequestBody(serverSocket, writer); - pipeResponseBodyToWebSocket(result.response, serverSocket); - return new Response(null, { - status: 101, - webSocket: clientSocket, - headers: result.headers, - } as WebSocketResponseInit); + result = await tunnel.connectWebSocket(request, webSocketHandleFromSocket(serverSocket)); } catch (error) { - writer.releaseLock(); serverSocket.close(1011, "WebSocket tunnel failed"); throw error; } -} - -function createWebSocketConnectRequest(request: Request, body: ReadableStream) { - const headers = new Headers(request.headers); - headers.delete("upgrade"); - headers.delete("connection"); - return new Request(request.url, { - method: "POST", - headers, - body, - // Required by Node-compatible Request implementations; harmless in Workers. - duplex: "half", - } as RequestInit & { duplex: "half" }); -} - -function pipeWebSocketToRequestBody( - socket: WebSocket, - writer: WritableStreamDefaultWriter, -) { - socket.addEventListener("message", (event) => { - void webSocketFrameFromMessage(event.data).then((frame) => - writer.write(encodeWebSocketFrame(frame)), - ); - }); - socket.addEventListener("close", (event) => { - void writer - .write(encodeWebSocketFrame({ type: "close", code: event.code, reason: event.reason })) - .finally(() => { - writer.close(); - writer.releaseLock(); - }); - }); - socket.addEventListener("error", () => { - void writer.abort(new Error("WebSocket error")); - }); -} - -function pipeResponseBodyToWebSocket(response: Response, socket: WebSocket) { - void (async () => { - if (!response.body) return; - try { - for await (const frame of decodeWebSocketFrames(response.body)) { - if (frame.type === "close") { - closeWebSocket(socket, frame.code, frame.reason); - return; - } - sendWebSocketMessage(socket, frame.data); - } - } catch { - socket.close(1011, "WebSocket tunnel failed"); - } - })(); -} - -function closeWebSocket(socket: WebSocket, code?: number, reason?: string) { - if (code === undefined || code === 1000 || (code >= 3000 && code <= 4999)) { - socket.close(code, reason); - return; + if (!result.accepted) { + serverSocket.close(1000, "WebSocket not accepted"); + return result.response; } - socket.close(); + + 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]?.(); + return new Response(null, { + status: 101, + webSocket: pair[0], + headers: result.protocol ? { "sec-websocket-protocol": result.protocol } : undefined, + } as WebSocketResponseInit); } function constantTimeEqual(actual: Uint8Array, expected: Uint8Array) { diff --git a/test/e2e.test.ts b/test/e2e.test.ts index 822da91..9bad2cc 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -1,9 +1,4 @@ -import { spawn, type ChildProcessByStdio } from "node:child_process"; import { createHash } from "node:crypto"; -import net from "node:net"; -import { dirname, resolve } from "node:path"; -import type { Readable } from "node:stream"; -import { fileURLToPath } from "node:url"; import { createRouterClient } from "@orpc/server"; import { newWebSocketRpcSession } from "capnweb"; @@ -13,7 +8,7 @@ import { createCaptunCliRouter } from "../src/cli/bin.js"; import { createCaptunTunnel } from "../src/index.js"; import { createCaptunWorkerFixture, createMiniflareWorkerFixture } from "./miniflare.js"; -vi.setConfig({ testTimeout: 25_000 }); +vi.setConfig({ testTimeout: 15_000 }); test.concurrent("forwards HTTP", async ({ task }) => { await using tunnel = await createTunnelFixture(task.name, async (request) => @@ -199,8 +194,35 @@ test.concurrent("forwards WebSocket Cap'n Web RPC to a local fetch handler", asy await delay(50); }); -test("CLI tunnels WebSocket traffic to a local Bun server", async ({ task }) => { - await using app = await createBunWebSocketFixture(); +// 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 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")); + await waitForWebSocket(socket); + 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("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(); @@ -212,7 +234,7 @@ test("CLI tunnels WebSocket traffic to a local Bun server", async ({ task }) => const client = createRouterClient(router); const runTunnel = client.tunnel({ - target: String(app.port), + target: app.origin, gateway: server.gateway, name: tunnelName(task.name), token: server.token, @@ -224,7 +246,17 @@ test("CLI tunnels WebSocket traffic to a local Bun server", async ({ task }) => try { await waitForWebSocket(socket); socket.send("hello-from-cli"); - await expect(readWebSocketMessage(socket)).resolves.toBe("echo: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); @@ -301,129 +333,6 @@ function sha256(bytes: Uint8Array) { return createHash("sha256").update(bytes).digest("hex"); } -async function createBunWebSocketFixture() { - const port = await getAvailablePort(); - const server = spawn("bun", ["run", "test/fixtures/bun-websocket-server.js"], { - cwd: resolve(dirname(fileURLToPath(import.meta.url)), ".."), - env: { ...process.env, PORT: String(port) }, - stdio: ["ignore", "pipe", "pipe"], - }); - const output = captureOutput(server); - - try { - await waitForTcp(port, server, output); - return { - port, - async [Symbol.asyncDispose]() { - await stopProcess(server); - }, - }; - } catch (error) { - await stopProcess(server); - throw new Error( - formatFixtureFailure(error instanceof Error ? error.message : String(error), output.logs()), - ); - } -} - -type ServerProcess = ChildProcessByStdio; - -async function getAvailablePort(): Promise { - return new Promise((resolvePort, reject) => { - const server = net.createServer(); - server.listen(0, "127.0.0.1", () => { - const address = server.address(); - if (!address || typeof address === "string") { - reject(new Error(`Failed to allocate a local port: ${String(address)}`)); - return; - } - - server.close((error) => { - if (error) reject(error); - else resolvePort(address.port); - }); - }); - server.on("error", reject); - }); -} - -async function waitForTcp(port: number, server: ServerProcess, output: CapturedProcessOutput) { - const startedAt = Date.now(); - while (Date.now() - startedAt < 15_000) { - const error = output.error(); - if (error) throw error; - if (server.exitCode !== null || server.signalCode) { - throw new Error( - `Bun server exited before port ${port} accepted connections\n\n${output.logs().trim() || "(none)"}`, - ); - } - - if (await canConnect(port)) return; - - await delay(100); - } - - throw new Error(`Timed out waiting for Bun server to accept connections on port ${port}`); -} - -function canConnect(port: number) { - return new Promise((resolveConnect) => { - const socket = net.connect(port, "127.0.0.1"); - socket.once("connect", () => { - socket.destroy(); - resolveConnect(true); - }); - socket.once("error", () => { - socket.destroy(); - resolveConnect(false); - }); - }); -} - -function captureOutput(child: ServerProcess) { - const chunks: string[] = []; - let processError: Error | undefined; - const capture = (chunk: string | Buffer) => { - chunks.push(String(chunk)); - if (chunks.length > 200) chunks.shift(); - }; - child.stdout.on("data", capture); - child.stderr.on("data", capture); - child.on("error", (error) => { - processError = error; - chunks.push(error.stack || error.message); - }); - - return { - logs: () => chunks.join(""), - error: () => processError, - }; -} - -interface CapturedProcessOutput { - logs(): string; - error(): Error | undefined; -} - -function formatFixtureFailure(message: string, serverLogs: string) { - return [message, "", "Server logs:", serverLogs.trim() || "(none)"].join("\n"); -} - -async function stopProcess(child: ServerProcess): Promise { - if (child.exitCode !== null || child.killed) return; - - child.kill("SIGINT"); - const exited = await Promise.race([ - new Promise((resolveExit) => child.once("exit", () => resolveExit(true))), - delay(5_000).then(() => false), - ]); - - if (!exited && child.exitCode === null && !child.killed) { - child.kill("SIGKILL"); - await new Promise((resolveExit) => child.once("exit", () => resolveExit())); - } -} - function waitForWebSocket(socket: WebSocket) { if (socket.readyState === WebSocket.OPEN) return Promise.resolve(); return new Promise((resolveOpen, rejectOpen) => { @@ -437,15 +346,9 @@ function waitForWebSocket(socket: WebSocket) { }); } -function readWebSocketMessage(socket: WebSocket) { - return new Promise((resolveMessage, rejectMessage) => { - socket.addEventListener( - "message", - async (event) => { - resolveMessage(await webSocketMessageText(event.data)); - }, - { 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, }); @@ -455,12 +358,31 @@ function readWebSocketMessage(socket: WebSocket) { }); } +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; - if (data instanceof Blob) return data.text(); - if (data instanceof ArrayBuffer) return new TextDecoder().decode(data); - if (data instanceof Uint8Array) return new TextDecoder().decode(data); - return String(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 { diff --git a/test/fixtures/bun-websocket-server.js b/test/fixtures/bun-websocket-server.js deleted file mode 100644 index b7b23ab..0000000 --- a/test/fixtures/bun-websocket-server.js +++ /dev/null @@ -1,24 +0,0 @@ -const server = Bun.serve({ - hostname: "127.0.0.1", - port: Number(process.env.PORT), - fetch(request, server) { - const url = new URL(request.url); - - if (url.pathname === "/ws") { - if (server.upgrade(request)) return; - return new Response("WebSocket upgrade failed\n", { status: 500 }); - } - - return new Response("ok\n"); - }, - websocket: { - message(socket, message) { - socket.send(`echo:${message}`); - }, - }, -}); - -process.on("SIGINT", () => { - server.stop(true); - process.exit(0); -}); diff --git a/test/fixtures/capnweb-websocket-target.ts b/test/fixtures/capnweb-websocket-target.ts index 37dc5f6..590bb5f 100644 --- a/test/fixtures/capnweb-websocket-target.ts +++ b/test/fixtures/capnweb-websocket-target.ts @@ -6,13 +6,49 @@ class DummyCapability extends RpcTarget { } } +type WorkerWebSocketPairConstructor = new () => { + 0: WebSocket; + 1: WebSocket & { accept(): void }; +}; + 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(); return new Response("Not found\n", { status: 404 }); }, }; + +/** Echoes text as `echo:`, binary untouched, and `close-with: ` as a close. */ +function webSocketEchoResponse() { + 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; + } + socket.send(`echo:${event.data}`); + })(); + }); + return new Response(null, { + status: 101, + webSocket: pair[0], + } as ResponseInit & { webSocket: WebSocket }); +} From c49bf3e3d9c45e121f00f45fca6f3414ee5b11b2 Mon Sep 17 00:00:00 2001 From: Jonas Templestein <242550+jonastemplestein@users.noreply.github.com> Date: Thu, 11 Jun 2026 15:12:55 +0100 Subject: [PATCH 03/13] Close orphaned public WebSockets and prove nested tunnels end to end A tunnel client disconnect now closes that tunnel's forwarded public WebSockets with 1001 instead of leaving idle sockets open until their next send fails. The shard tracks forwarded sockets per active tunnel and closes them on disconnect and on same-name reconnect. New e2e coverage for the ways a tunneled WebSocket server should just work: - Captun through Captun: the CLI exposes an inner gateway through an outer tunnel; RPC and echo traffic cross two gateways and two tunnel clients. - The hosted gateway forwards WebSockets (was only covered on the core worker). - Concurrent WebSockets over one tunnel do not cross talk. - Subprotocol negotiation reaches the public 101 on both the CLI path and the local-fetch fallback path. - A local server that rejects the upgrade fails the public handshake instead of hanging. - Disconnecting the tunnel client closes public sockets with 1001. Also dispose the piped handle stub even when the final close call throws. --- src/index.ts | 7 +- src/server/worker.ts | 44 +++++- test/e2e.test.ts | 155 +++++++++++++++++++++- test/fixtures/capnweb-websocket-target.ts | 7 +- 4 files changed, 198 insertions(+), 15 deletions(-) diff --git a/src/index.ts b/src/index.ts index 05400b5..5aa63af 100644 --- a/src/index.ts +++ b/src/index.ts @@ -455,8 +455,11 @@ export function pipeWebSocketToHandle(socket: WebSocket, handle: WebSocketHandle if (finished) return; finished = true; enqueue(() => { - forward(remote.close(code, reason)); - disposeStub(remote); + try { + forward(remote.close(code, reason)); + } finally { + disposeStub(remote); + } }); }; diff --git a/src/server/worker.ts b/src/server/worker.ts index c88e61b..4c351fe 100644 --- a/src/server/worker.ts +++ b/src/server/worker.ts @@ -53,6 +53,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 = @@ -128,14 +130,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 }); @@ -157,13 +170,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 { if (isWebSocketUpgradeRequest(request)) { return await forwardWebSocket(tunnel, request); } - return await tunnel.fetch(request); + return await tunnel.fetcher.fetch(request); } catch { return new Response("Tunnel fetch failed\n", { status: 502 }); } @@ -281,7 +294,7 @@ export function createTunnelForwardRequest( return new Request(request, { headers }); } -async function forwardWebSocket(tunnel: FetcherStub, request: Request) { +async function forwardWebSocket(tunnel: ActiveTunnel, request: Request) { const WorkerWebSocketPair = ( globalThis as typeof globalThis & { WebSocketPair: WorkerWebSocketPairConstructor } ).WebSocketPair; @@ -291,7 +304,10 @@ async function forwardWebSocket(tunnel: FetcherStub, request: Request) { let result: WebSocketConnectResult; try { - result = await tunnel.connectWebSocket(request, webSocketHandleFromSocket(serverSocket)); + result = await tunnel.fetcher.connectWebSocket( + request, + webSocketHandleFromSocket(serverSocket), + ); } catch (error) { serverSocket.close(1011, "WebSocket tunnel failed"); throw error; @@ -304,6 +320,8 @@ async function forwardWebSocket(tunnel: FetcherStub, request: Request) { 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], @@ -311,6 +329,18 @@ async function forwardWebSocket(tunnel: FetcherStub, request: Request) { } 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 9bad2cc..a43a7ce 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -6,7 +6,11 @@ import { expect, test, vi } from "vitest"; import { createCaptunCliRouter } from "../src/cli/bin.js"; import { createCaptunTunnel } from "../src/index.js"; -import { createCaptunWorkerFixture, createMiniflareWorkerFixture } from "./miniflare.js"; +import { + createCaptunWorkerFixture, + createHostedCaptunWorkerFixture, + createMiniflareWorkerFixture, +} from "./miniflare.js"; vi.setConfig({ testTimeout: 15_000 }); @@ -196,7 +200,7 @@ test.concurrent("forwards WebSocket Cap'n Web RPC to a local fetch handler", asy // 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 and close codes", async ({ task }) => { +test.concurrent("forwards WebSocket messages, subprotocols, and close codes", async ({ task }) => { await using target = await createMiniflareWorkerFixture({ entryPoint: "test/fixtures/capnweb-websocket-target.ts", durableObjects: {}, @@ -206,8 +210,9 @@ test.concurrent("forwards WebSocket messages and close codes", async ({ task }) target.worker.fetch(request.url, request), ); - const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + 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"); @@ -217,6 +222,91 @@ test.concurrent("forwards WebSocket messages and close codes", async ({ task }) 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); +}); + +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", @@ -242,9 +332,10 @@ test("CLI tunnels WebSocket traffic to a local target", async ({ task }) => { }); const tunnel = await ready.promise; - const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws")); + const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws"), ["alpha", "beta"]); try { await waitForWebSocket(socket); + expect(socket).toMatchObject({ protocol: "alpha" }); socket.send("hello-from-cli"); await expect(nextWebSocketMessage(socket).then(webSocketMessageText)).resolves.toBe( "echo:hello-from-cli", @@ -265,6 +356,62 @@ test("CLI tunnels WebSocket traffic to a local target", async ({ task }) => { } }); +// 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, diff --git a/test/fixtures/capnweb-websocket-target.ts b/test/fixtures/capnweb-websocket-target.ts index 590bb5f..22215f7 100644 --- a/test/fixtures/capnweb-websocket-target.ts +++ b/test/fixtures/capnweb-websocket-target.ts @@ -17,14 +17,14 @@ export default { if (url.pathname === "/rpc") { return newWorkersWebSocketRpcResponse(request, new DummyCapability()); } - if (url.pathname === "/ws") return webSocketEchoResponse(); + 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() { +function webSocketEchoResponse(request: Request) { const WorkerWebSocketPair = ( globalThis as typeof globalThis & { WebSocketPair: WorkerWebSocketPairConstructor } ).WebSocketPair; @@ -47,8 +47,11 @@ function webSocketEchoResponse() { 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 ResponseInit & { webSocket: WebSocket }); } From 3034d56e7e357c8e3d0e6c44776c65f90b8dff40 Mon Sep 17 00:00:00 2001 From: Jonas Templestein <242550+jonastemplestein@users.noreply.github.com> Date: Thu, 11 Jun 2026 15:42:26 +0100 Subject: [PATCH 04/13] Forward handshake headers to the local WebSocket server from the CLI The CLI's WebSocket path dropped Cookie, Authorization, Origin, and other handshake headers that tunneled HTTP requests forward, so local servers doing handshake auth worked over HTTP but not over WebSockets. Node's WebSocket (undici) accepts a headers option; forward everything except hop-by-hop headers and the Sec-WebSocket-* family, which the local client negotiates itself. The e2e CLI test now asserts the local server sees the public client's cookie and authorization headers. Found by Cursor Bugbot on the PR. --- src/cli/bin.ts | 25 +++++++++++++++++++---- test/e2e.test.ts | 15 +++++++++++++- test/fixtures/capnweb-websocket-target.ts | 9 ++++++++ 3 files changed, 44 insertions(+), 5 deletions(-) diff --git a/src/cli/bin.ts b/src/cli/bin.ts index 42640d7..8c5290d 100755 --- a/src/cli/bin.ts +++ b/src/cli/bin.ts @@ -608,14 +608,15 @@ async function connectTargetWebSocket( const targetUrl = new URL(`${tunnel.target}${url.pathname}${url.search}`); targetUrl.protocol = targetUrl.protocol === "https:" ? "wss:" : "ws:"; - const targetSocket = new WebSocket( - targetUrl, - request.headers + 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); @@ -640,6 +641,22 @@ async function connectTargetWebSocket( }; } +/** + * 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"); diff --git a/test/e2e.test.ts b/test/e2e.test.ts index a43a7ce..85a5996 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -332,10 +332,23 @@ test("CLI tunnels WebSocket traffic to a local target", async ({ task }) => { }); const tunnel = await ready.promise; - const socket = new WebSocket(`${tunnel.url}/ws`.replace(/^http/, "ws"), ["alpha", "beta"]); + 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", diff --git a/test/fixtures/capnweb-websocket-target.ts b/test/fixtures/capnweb-websocket-target.ts index 22215f7..89a3210 100644 --- a/test/fixtures/capnweb-websocket-target.ts +++ b/test/fixtures/capnweb-websocket-target.ts @@ -44,6 +44,15 @@ function webSocketEchoResponse(request: Request) { 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}`); })(); }); From b9ee62d9e17d22cdda28673cdf0d08783cae9565 Mon Sep 17 00:00:00 2001 From: Jonas Templestein <242550+jonastemplestein@users.noreply.github.com> Date: Thu, 11 Jun 2026 15:51:19 +0100 Subject: [PATCH 05/13] Make Workers-style WebSocket handling work on every runtime MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A plain fetch handler can now answer tunneled WebSockets the Workers way on any runtime. Two primitives were missing outside workerd and are now exported: WebSocketPair (the native pair where the runtime has one, otherwise an in-memory pair with Workers accept()/queuing semantics) and createWebSocketResponse (a real 101 upgrade response in Workers; elsewhere a Response that carries the socket to the tunnel bridge, since other runtimes reject status 101). connectWebSocket remains the low-level hook for dialing out to a separate local WebSocket server, which is what the CLI does. The new e2e test runs a Workers-style handler in plain Node — welcome message sent before the handler returns, echo, subprotocol selection, and close codes in both directions — with no workerd on the target side. --- src/index.ts | 86 ++++++++++++++++++++++++++++++++++++++++++++++++ test/e2e.test.ts | 54 +++++++++++++++++++++++++++++- 2 files changed, 139 insertions(+), 1 deletion(-) diff --git a/src/index.ts b/src/index.ts index 2be5acf..acd16bf 100644 --- a/src/index.ts +++ b/src/index.ts @@ -434,6 +434,92 @@ 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); diff --git a/test/e2e.test.ts b/test/e2e.test.ts index 85a5996..ec07f8f 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -5,7 +5,12 @@ import { newWebSocketRpcSession } from "capnweb"; import { expect, test, vi } from "vitest"; import { createCaptunCliRouter } from "../src/cli/bin.js"; -import { createCaptunTunnel } from "../src/index.js"; +import { + createCaptunTunnel, + createWebSocketResponse, + isWebSocketUpgradeRequest, + WebSocketPair, +} from "../src/index.js"; import { createCaptunWorkerFixture, createHostedCaptunWorkerFixture, @@ -252,6 +257,53 @@ test.concurrent("forwards concurrent WebSockets over one tunnel without cross-ta 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 }); +}); + test.concurrent("fails the public WebSocket when the local server rejects it", async ({ task }) => { await using tunnel = await createTunnelFixture( task.name, From ef7de9728afd64b79eb6be827d0e6b8568e69278 Mon Sep 17 00:00:00 2001 From: Jonas Templestein <242550+jonastemplestein@users.noreply.github.com> Date: Thu, 11 Jun 2026 16:27:55 +0100 Subject: [PATCH 06/13] Cap tunneled WebSocket messages at 16MiB and document WebSocket support MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Robustness testing against a deployed gateway showed a single ~24MiB message kills the entire tunnel: base64 over the Cap'n Web leg overflows its ~32MiB frame limit and every connection on the tunnel gets closed with 1001. The pipe now closes just the offending socket with 4009 before forwarding; covered by an e2e test that checks the tunnel survives. Also: a short WebSockets section in the README, and pnpm test now runs the core test/ suite only — example tests depend on the Bun/Deno versions CI pins, so they move to pnpm test:examples, which CI runs as its own step. --- .github/workflows/ci.yml | 2 ++ README.md | 31 +++++++++++++++++++++++++++++++ package.json | 3 ++- src/index.ts | 12 +++++++++++- test/e2e.test.ts | 27 +++++++++++++++++++++++++++ 5 files changed, 73 insertions(+), 2 deletions(-) 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..125fc54 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", diff --git a/src/index.ts b/src/index.ts index acd16bf..fe54e30 100644 --- a/src/index.ts +++ b/src/index.ts @@ -574,13 +574,23 @@ export function pipeWebSocketToHandle(socket: WebSocket, handle: WebSocketHandle socket.addEventListener("message", (event) => { if (finished) return; enqueue(async () => { - forward(remote.send(await webSocketMessage(event.data))); + 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 ((typeof message === "string" ? message.length : message.byteLength) > 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; + type StubLike = { dup?(): unknown; [Symbol.dispose]?(): void }; function dupStub(handle: WebSocketHandle): WebSocketHandle { diff --git a/test/e2e.test.ts b/test/e2e.test.ts index ec07f8f..bb08fd3 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -304,6 +304,33 @@ test.concurrent("answers WebSockets Workers-style from a Node fetch handler", as await expect(serverClosed.promise).resolves.toMatchObject({ code: 4002 }); }); +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 }); + + // 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, From 665c18dc2c984ea04ad70c61970321c32b02d3bd Mon Sep 17 00:00:00 2001 From: Jonas Templestein <242550+jonastemplestein@users.noreply.github.com> Date: Thu, 11 Jun 2026 16:37:58 +0100 Subject: [PATCH 07/13] Reject the public upgrade when the local socket dies right after opening If the local server closes immediately after the WebSocket handshake, the CLI now returns the 502 rejection instead of accepting an already-dead socket. A close that lands after acceptance was already propagated by the pipe; this narrows the remaining race where readyState flips before the close event dispatches. Found by Cursor Bugbot on the PR. --- src/cli/bin.ts | 3 +++ 1 file changed, 3 insertions(+) diff --git a/src/cli/bin.ts b/src/cli/bin.ts index 8c5290d..f46ab76 100755 --- a/src/cli/bin.ts +++ b/src/cli/bin.ts @@ -620,6 +620,9 @@ async function connectTargetWebSocket( 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); From e7eebf2ae865f2b80e4d2df02e30be4ba2e0d8ca Mon Sep 17 00:00:00 2001 From: Misha Kaletsky <15040698+mmkal@users.noreply.github.com> Date: Thu, 11 Jun 2026 18:01:55 +0100 Subject: [PATCH 08/13] Expose silent corruption of untunnelable WebSocket payloads MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The in-memory WebSocketPair accepts any value in send(), and the tunnel pipe's message normalization falls back to String(data) for types it doesn't recognize — so a handler bug like send({object}) reaches the public client as the text frame "[object Object]" instead of failing. This test asserts the sane behavior (the socket closes instead of delivering corrupted text) and fails against the current fallback; the fix follows in the next commit. Co-Authored-By: Claude Fable 5 --- test/e2e.test.ts | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/test/e2e.test.ts b/test/e2e.test.ts index bb08fd3..438b4b6 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -304,6 +304,33 @@ test.concurrent("answers WebSockets Workers-style from a Node fetch handler", as 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", From ca8414a137f2c0181ebbe597b6a1b8e693a97960 Mon Sep 17 00:00:00 2001 From: Misha Kaletsky <15040698+mmkal@users.noreply.github.com> Date: Thu, 11 Jun 2026 18:02:33 +0100 Subject: [PATCH 09/13] Fail the socket instead of stringifying untunnelable payloads webSocketMessage now throws on payload types it doesn't recognize instead of falling back to String(data). The pipe's existing error handling turns the throw into a 1011 close of just that socket, so a handler bug surfaces as a failed connection rather than the public client silently receiving "[object Object]". Co-Authored-By: Claude Fable 5 --- src/index.ts | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/src/index.ts b/src/index.ts index fe54e30..2b18c14 100644 --- a/src/index.ts +++ b/src/index.ts @@ -607,7 +607,9 @@ function disposeStub(handle: WebSocketHandle) { * 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). + * (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; @@ -620,7 +622,9 @@ async function webSocketMessage(data: unknown): Promise { if (typeof (data as Blob | null)?.arrayBuffer === "function") { return new Uint8Array(await (data as Blob).arrayBuffer()); } - return String(data); + throw new TypeError( + `Cannot tunnel WebSocket message ${Object.prototype.toString.call(data)}; expected string or binary`, + ); } function closeWebSocket(socket: WebSocket, code?: number, reason?: string) { From d1923bac365f6a2ea692828cd15e4a8b81a40e36 Mon Sep 17 00:00:00 2001 From: Misha Kaletsky <15040698+mmkal@users.noreply.github.com> Date: Thu, 11 Jun 2026 18:03:33 +0100 Subject: [PATCH 10/13] Measure the oversized-message cap in wire bytes, not string length MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A string's .length counts UTF-16 code units; multi-byte characters serialize to up to 3x that in UTF-8 and JSON escapes up to 6x, so a string under 16M units could still overflow the ~32MiB Cap'n Web frame limit and kill every connection on the tunnel — the exact failure the cap exists to prevent. Strings now measure their JSON-escaped UTF-8 size, with a fast path that skips encoding when even the worst case fits. Co-Authored-By: Claude Fable 5 --- src/index.ts | 15 ++++++++++++++- test/e2e.test.ts | 8 ++++++++ 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/src/index.ts b/src/index.ts index 2b18c14..f32c3ec 100644 --- a/src/index.ts +++ b/src/index.ts @@ -578,7 +578,7 @@ export function pipeWebSocketToHandle(socket: WebSocket, handle: WebSocketHandle // 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 ((typeof message === "string" ? message.length : message.byteLength) > MAX_MESSAGE_BYTES) { + if (messageByteLength(message) > MAX_MESSAGE_BYTES) { closeWebSocket(socket, 4009, "Message too large to tunnel"); return; } @@ -591,6 +591,19 @@ export function pipeWebSocketToHandle(socket: WebSocket, handle: WebSocketHandle 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 { diff --git a/test/e2e.test.ts b/test/e2e.test.ts index 438b4b6..6be9c31 100644 --- a/test/e2e.test.ts +++ b/test/e2e.test.ts @@ -347,6 +347,14 @@ test.concurrent("closes only the socket that sends an oversized message", async 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); From 03486b928edc41b3b707caaba39ac2866a8cca7f Mon Sep 17 00:00:00 2001 From: Misha Kaletsky <15040698+mmkal@users.noreply.github.com> Date: Thu, 11 Jun 2026 18:04:01 +0100 Subject: [PATCH 11/13] Declare node >=22.4 in engines The CLI's WebSocket forwarding relies on undici's non-standard WebSocket options ({ protocols, headers }); older Nodes either lack the global WebSocket or silently drop the forwarded handshake headers rather than erroring. Fail loudly at install time instead. Co-Authored-By: Claude Fable 5 --- package.json | 3 +++ 1 file changed, 3 insertions(+) diff --git a/package.json b/package.json index 125fc54..2903bc0 100644 --- a/package.json +++ b/package.json @@ -127,5 +127,8 @@ "optional": true } }, + "engines": { + "node": ">=22.4" + }, "packageManager": "pnpm@10.11.1" } From adc502dd590dd5411824270fff4408927f67f5e7 Mon Sep 17 00:00:00 2001 From: Misha Kaletsky <15040698+mmkal@users.noreply.github.com> Date: Thu, 11 Jun 2026 18:05:31 +0100 Subject: [PATCH 12/13] Export the Worker WebSocket types instead of declaring them thrice WorkerWebSocket, WorkerWebSocketPairConstructor, and WebSocketResponseInit were duplicated across src/index.ts, src/server/worker.ts, and the e2e fixture. The public WebSocketPair export was already typed with the private constructor type, so exporting it also fixes a leaked-private-type hole. Co-Authored-By: Claude Fable 5 --- src/index.ts | 6 +++--- src/server/worker.ts | 15 ++------------- test/fixtures/capnweb-websocket-target.ts | 9 +++------ 3 files changed, 8 insertions(+), 22 deletions(-) diff --git a/src/index.ts b/src/index.ts index f32c3ec..830098d 100644 --- a/src/index.ts +++ b/src/index.ts @@ -164,16 +164,16 @@ type TunnelClientCapability = Fetcher & 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; }; diff --git a/src/server/worker.ts b/src/server/worker.ts index 4c351fe..6115bc7 100644 --- a/src/server/worker.ts +++ b/src/server/worker.ts @@ -10,6 +10,8 @@ import { webSocketHandleFromSocket, type FetcherStub, type WebSocketConnectResult, + type WebSocketResponseInit, + type WorkerWebSocketPairConstructor, } from "../index.js"; import { captunShardName, @@ -31,19 +33,6 @@ export type CaptunEnv = { const TUNNEL_NAME_HEADER = "x-captun-tunnel-name"; const CUSTOM_HOSTNAME_RESERVED_TUNNEL_NAMES = ["captun", "gateway"]; -type WorkerWebSocket = WebSocket & { - accept(): void; -}; - -type WorkerWebSocketPairConstructor = new () => { - 0: WorkerWebSocket; - 1: WorkerWebSocket; -}; - -type WebSocketResponseInit = ResponseInit & { - webSocket: WebSocket; -}; - type CaptunShardBindingEnv = { CaptunServerShard: DurableObjectNamespace>; SHARD_COUNT?: string; diff --git a/test/fixtures/capnweb-websocket-target.ts b/test/fixtures/capnweb-websocket-target.ts index 89a3210..c2189c0 100644 --- a/test/fixtures/capnweb-websocket-target.ts +++ b/test/fixtures/capnweb-websocket-target.ts @@ -1,16 +1,13 @@ import { newWorkersWebSocketRpcResponse, RpcTarget } from "capnweb"; +import type { WebSocketResponseInit, WorkerWebSocketPairConstructor } from "../../src/index.js"; + class DummyCapability extends RpcTarget { ping(value: string) { return `pong:${value}`; } } -type WorkerWebSocketPairConstructor = new () => { - 0: WebSocket; - 1: WebSocket & { accept(): void }; -}; - export default { fetch(request: Request) { const url = new URL(request.url); @@ -62,5 +59,5 @@ function webSocketEchoResponse(request: Request) { status: 101, webSocket: pair[0], headers: protocol ? { "sec-websocket-protocol": protocol } : undefined, - } as ResponseInit & { webSocket: WebSocket }); + } as WebSocketResponseInit); } From 89d2dd5b077b723836d27b3f2f1103e18560a35d Mon Sep 17 00:00:00 2001 From: Misha Kaletsky <15040698+mmkal@users.noreply.github.com> Date: Thu, 11 Jun 2026 18:06:56 +0100 Subject: [PATCH 13/13] Type the RPC session as RemoteFetcherCapability instead of double-casting MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit newWebSocketRpcSession's type parameter describes the remote interface, and capnweb's stub provides onRpcBroken itself — so the session can be declared as RemoteFetcherCapability directly. The as-unknown-as cast was papering over a wrong type argument. Co-Authored-By: Claude Fable 5 --- src/index.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/index.ts b/src/index.ts index 830098d..de7e6e0 100644 --- a/src/index.ts +++ b/src/index.ts @@ -73,7 +73,7 @@ export function acceptFetcherCapabilityFromSocket( socket: WebSocket, options: { onDisconnect?: () => void } = {}, ): FetcherStub { - const remote = newWebSocketRpcSession(socket) as unknown as RemoteFetcherCapability; + const remote = newWebSocketRpcSession(socket); return fetcherStubFromRemoteCapability(remote, options); }