diff --git a/packages/relay/src/cloudflare-adapter.test.ts b/packages/relay/src/cloudflare-adapter.test.ts index 49cce3273..eeca70e70 100644 --- a/packages/relay/src/cloudflare-adapter.test.ts +++ b/packages/relay/src/cloudflare-adapter.test.ts @@ -1,6 +1,9 @@ import { afterEach, describe, expect, it, vi } from "vitest"; import relayWorker, { RelayDurableObject } from "./cloudflare-adapter.js"; +type DurableObjectStateArg = ConstructorParameters[0]; +type RelayEnvArg = Parameters[1]; + type MockSocket = WebSocket & { send: ReturnType; close: ReturnType; @@ -68,17 +71,19 @@ async function withMockWebSocketPair( } } +const swallow = () => undefined; + describe("RelayDurableObject versioning", () => { it("accepts legacy v1 client sockets without connectionId", async () => { const { state } = createMockState(); await withMockWebSocketPair(async () => { - const relay = new RelayDurableObject(state as any); + const relay = new RelayDurableObject(state as unknown as DurableObjectStateArg); const req = new Request("https://relay.test/ws?role=client&serverId=srv_test&v=1", { headers: { Upgrade: "websocket", }, }); - await relay.fetch(req).catch(() => undefined); + await relay.fetch(req).catch(swallow); expect(state.acceptWebSocket).toHaveBeenCalled(); }); }); @@ -86,11 +91,11 @@ describe("RelayDurableObject versioning", () => { it("assigns a connectionId when v2 client connects without one", async () => { const { state } = createMockState(); await withMockWebSocketPair(async ({ serverWs }) => { - const relay = new RelayDurableObject(state as any); + const relay = new RelayDurableObject(state as unknown as DurableObjectStateArg); const req = new Request("https://relay.test/ws?role=client&serverId=srv_test&v=2", { headers: { Upgrade: "websocket" }, }); - await relay.fetch(req).catch(() => undefined); + await relay.fetch(req).catch(swallow); expect(state.acceptWebSocket).toHaveBeenCalled(); const attachment = serverWs.deserializeAttachment(); expect(attachment).toMatchObject({ @@ -117,8 +122,10 @@ describe("RelayDurableObject control nudge/reset behavior", () => { setTagSockets(`client:${clientId}`, []); setTagSockets(`server:${clientId}`, []); - const relay = new RelayDurableObject(state as any); - (relay as any).nudgeOrResetControlForConnection(clientId); + const relay = new RelayDurableObject(state as unknown as DurableObjectStateArg); + ( + relay as unknown as { nudgeOrResetControlForConnection(id: string): void } + ).nudgeOrResetControlForConnection(clientId); vi.advanceTimersByTime(15_000); @@ -143,8 +150,10 @@ describe("RelayDurableObject control nudge/reset behavior", () => { setTagSockets(`client:${clientId}`, [client]); setTagSockets(`server:${clientId}`, []); - const relay = new RelayDurableObject(state as any); - (relay as any).nudgeOrResetControlForConnection(clientId); + const relay = new RelayDurableObject(state as unknown as DurableObjectStateArg); + ( + relay as unknown as { nudgeOrResetControlForConnection(id: string): void } + ).nudgeOrResetControlForConnection(clientId); vi.advanceTimersByTime(10_000); expect(control.send).toHaveBeenCalledTimes(1); @@ -166,7 +175,7 @@ describe("RelayDurableObject control nudge/reset behavior", () => { setTagSockets("client", [existingClient]); await withMockWebSocketPair(async () => { - const relay = new RelayDurableObject(state as any); + const relay = new RelayDurableObject(state as unknown as DurableObjectStateArg); const req = new Request( "https://relay.test/ws?role=client&serverId=srv_test&connectionId=clt_same_session&v=2", { @@ -176,7 +185,7 @@ describe("RelayDurableObject control nudge/reset behavior", () => { }, ); - await relay.fetch(req).catch(() => undefined); + await relay.fetch(req).catch(swallow); expect(existingClient.close).not.toHaveBeenCalled(); }); }); @@ -206,7 +215,7 @@ describe("RelayDurableObject control nudge/reset behavior", () => { setTagSockets("client", [stillConnectedClient]); setTagSockets(`client:${clientId}`, [stillConnectedClient]); - const relay = new RelayDurableObject(state as any); + const relay = new RelayDurableObject(state as unknown as DurableObjectStateArg); relay.webSocketClose( disconnectedClient as unknown as WebSocket, 1001, @@ -231,7 +240,7 @@ describe("relay worker endpoint routing", () => { const response = await relayWorker.fetch( new Request("https://relay.test/ws?serverId=srv_test&role=server"), - { RELAY: { idFromName, get } } as any, + { RELAY: { idFromName, get } } as unknown as RelayEnvArg, ); expect(idFromName).toHaveBeenCalledWith("relay-v1:srv_test"); @@ -248,7 +257,7 @@ describe("relay worker endpoint routing", () => { const response = await relayWorker.fetch( new Request("https://relay.test/ws?serverId=srv_test&role=server&v=2"), - { RELAY: { idFromName, get } } as any, + { RELAY: { idFromName, get } } as unknown as RelayEnvArg, ); expect(idFromName).toHaveBeenCalledWith("relay-v2:srv_test"); @@ -263,7 +272,7 @@ describe("relay worker endpoint routing", () => { const response = await relayWorker.fetch( new Request("https://relay.test/ws?serverId=srv_test&role=server&v=nope"), - { RELAY: { idFromName, get } } as any, + { RELAY: { idFromName, get } } as unknown as RelayEnvArg, ); expect(response.status).toBe(400); diff --git a/packages/relay/src/cloudflare-adapter.ts b/packages/relay/src/cloudflare-adapter.ts index 656edd966..f19f6c00f 100644 --- a/packages/relay/src/cloudflare-adapter.ts +++ b/packages/relay/src/cloudflare-adapter.ts @@ -128,6 +128,36 @@ export class RelayDurableObject { } } + private closeExistingServerSockets(args: { + isServerControl: boolean; + isServerData: boolean; + resolvedConnectionId: string; + }): void { + if (args.isServerControl) { + for (const ws of this.state.getWebSockets("server-control")) { + ws.close(1008, "Replaced by new connection"); + } + } else if (args.isServerData) { + for (const ws of this.state.getWebSockets(`server:${args.resolvedConnectionId}`)) { + ws.close(1008, "Replaced by new connection"); + } + } + } + + private handleControlKeepalive(ws: WebSocket, message: string): void { + try { + const parsed = JSON.parse(message) as unknown as { type?: unknown }; + if (parsed?.type !== "ping") return; + try { + ws.send(JSON.stringify({ type: "pong", ts: Date.now() })); + } catch { + // ignore + } + } catch { + // ignore non-JSON control payloads + } + } + private nudgeOrResetControlForConnection(connectionId: string): void { // If the daemon's control WS becomes half-open, the DO can't reliably detect it via ws.send errors // (Cloudflare may accept writes even if the other side is no longer reading). @@ -270,15 +300,7 @@ export class RelayDurableObject { // - server-control: single per serverId // - server-data: single per connectionId // - client: many sockets per connectionId are allowed - if (isServerControl) { - for (const ws of this.state.getWebSockets("server-control")) { - ws.close(1008, "Replaced by new connection"); - } - } else if (isServerData) { - for (const ws of this.state.getWebSockets(`server:${resolvedConnectionId}`)) { - ws.close(1008, "Replaced by new connection"); - } - } + this.closeExistingServerSockets({ isServerControl, isServerData, resolvedConnectionId }); const [client, server] = this.createWebSocketPair(); @@ -302,9 +324,15 @@ export class RelayDurableObject { }; (server as WebSocketWithAttachment).serializeAttachment(attachment); - console.log( - `[Relay DO] v2:${role}${isServerControl ? "(control)" : ""}${isServerData ? `(data:${resolvedConnectionId})` : role === "client" ? `(${resolvedConnectionId})` : ""} connected to session ${serverId}`, - ); + let roleSuffix = ""; + if (isServerControl) { + roleSuffix = "(control)"; + } else if (isServerData) { + roleSuffix = `(data:${resolvedConnectionId})`; + } else if (role === "client") { + roleSuffix = `(${resolvedConnectionId})`; + } + console.log(`[Relay DO] v2:${role}${roleSuffix} connected to session ${serverId}`); if (role === "client") { this.notifyControls({ type: "connected", connectionId: resolvedConnectionId }); @@ -387,18 +415,7 @@ export class RelayDurableObject { if (!connectionId) { // Control channel: support simple app-level keepalive. if (typeof message === "string") { - try { - const parsed = JSON.parse(message) as unknown as { type?: unknown }; - if (parsed?.type === "ping") { - try { - ws.send(JSON.stringify({ type: "pong", ts: Date.now() })); - } catch { - // ignore - } - } - } catch { - // ignore non-JSON control payloads - } + this.handleControlKeepalive(ws, message); } return; } diff --git a/packages/relay/src/crypto.test.ts b/packages/relay/src/crypto.test.ts index 8163df069..127693720 100644 --- a/packages/relay/src/crypto.test.ts +++ b/packages/relay/src/crypto.test.ts @@ -105,7 +105,8 @@ describe("crypto", () => { const ciphertext = await encrypt(correctKey, "secret"); - expect(() => decrypt(wrongKey, ciphertext)).toThrow(); + const tryDecrypt = () => decrypt(wrongKey, ciphertext); + expect(tryDecrypt).toThrow(); }); it("produces different ciphertext for same plaintext (random IV)", async () => { diff --git a/packages/relay/src/e2e.test.ts b/packages/relay/src/e2e.test.ts index b9ef966e7..d432402be 100644 --- a/packages/relay/src/e2e.test.ts +++ b/packages/relay/src/e2e.test.ts @@ -38,6 +38,14 @@ async function sleep(ms: number): Promise { await new Promise((resolve) => setTimeout(resolve, ms)); } +function rawToText(raw: unknown): string { + if (typeof raw === "string") return raw; + if (raw && typeof (raw as { toString?: unknown }).toString === "function") { + return (raw as { toString(): string }).toString(); + } + return ""; +} + function spawnRelayDevServer(port: number): ChildProcess { return spawn( process.execPath, @@ -69,6 +77,16 @@ function assertRelayStillRunning(relayProcess: ChildProcess): void { } } +function tryConnect(port: number): Promise { + return new Promise((resolve, reject) => { + const socket = net.connect(port, "127.0.0.1", () => { + socket.end(); + resolve(); + }); + socket.on("error", reject); + }); +} + async function waitForServer( port: number, relayProcess: ChildProcess, @@ -78,13 +96,7 @@ async function waitForServer( while (Date.now() - start < timeout) { assertRelayStillRunning(relayProcess); try { - await new Promise((resolve, reject) => { - const socket = net.connect(port, "127.0.0.1", () => { - socket.end(); - resolve(); - }); - socket.on("error", reject); - }); + await tryConnect(port); return; } catch { await sleep(100); @@ -105,18 +117,24 @@ async function waitForRelayWebSocketReady( const probeUrl = `ws://127.0.0.1:${port}/ws?serverId=${serverId}&role=server&v=2`; const opened = await new Promise((resolve) => { const ws = new WebSocket(probeUrl); + let settled = false; + const settle = (value: boolean) => { + if (settled) return; + settled = true; + resolve(value); + }; const timer = setTimeout(() => { ws.terminate(); - resolve(false); + settle(false); }, 5000); ws.once("open", () => { clearTimeout(timer); ws.close(1000, "probe"); - resolve(true); + settle(true); }); ws.once("error", () => { clearTimeout(timer); - resolve(false); + settle(false); }); }); if (opened) { @@ -161,21 +179,16 @@ async function stopRelayProcess(relayProcess: ChildProcess): Promise { relayPort = await getAvailablePort(); relayProcess = spawnRelayDevServer(relayPort); + const hasContent = (line: string) => line.trim().length > 0; relayProcess.stdout?.on("data", (data: Buffer) => { - const lines = data - .toString() - .split("\n") - .filter((l) => l.trim()); + const lines = data.toString().split("\n").filter(hasContent); for (const line of lines) { // eslint-disable-next-line no-console console.log(`[relay] ${line}`); } }); relayProcess.stderr?.on("data", (data: Buffer) => { - const lines = data - .toString() - .split("\n") - .filter((l) => l.trim()); + const lines = data.toString().split("\n").filter(hasContent); for (const line of lines) { // eslint-disable-next-line no-console console.error(`[relay] ${line}`); @@ -242,12 +255,7 @@ async function stopRelayProcess(relayProcess: ChildProcess): Promise { ); const onMessage = (raw: unknown) => { try { - const text = - typeof raw === "string" - ? raw - : raw && typeof (raw as any).toString === "function" - ? (raw as any).toString() - : ""; + const text = rawToText(raw); const msg = JSON.parse(text); if (msg?.type === "connected" && msg.connectionId === connectionId) { clearTimeout(timeout); @@ -396,12 +404,7 @@ async function stopRelayProcess(relayProcess: ChildProcess): Promise { const timeout = setTimeout(() => reject(new Error("timed out waiting for connected")), 5000); const onMessage = (raw: unknown) => { try { - const text = - typeof raw === "string" - ? raw - : raw && typeof (raw as any).toString === "function" - ? (raw as any).toString() - : ""; + const text = rawToText(raw); const msg = JSON.parse(text); if (msg?.type === "connected" && msg.connectionId === connectionId) { clearTimeout(timeout); @@ -469,8 +472,6 @@ async function stopRelayProcess(relayProcess: ChildProcess): Promise { }); it("wrong key cannot decrypt", async () => { - const serverId = "wrong-key-test-" + Date.now(); - // Setup - daemon and client with correct keys const daemonKeyPair = await generateKeyPair(); const clientKeyPair = await generateKeyPair(); diff --git a/packages/relay/src/encrypted-channel.ts b/packages/relay/src/encrypted-channel.ts index 06273b3b1..b443128d2 100644 --- a/packages/relay/src/encrypted-channel.ts +++ b/packages/relay/src/encrypted-channel.ts @@ -60,8 +60,12 @@ function buildInvalidHelloError(rawText: string, parsed?: unknown): Error { const parsedRecord = parsed && typeof parsed === "object" ? (parsed as Record) : null; const rawType = parsedRecord?.type; - const receivedType = - typeof rawType === "string" ? rawType : rawType === undefined ? "undefined" : typeof rawType; + function describeType(value: unknown): string { + if (typeof value === "string") return value; + if (value === undefined) return "undefined"; + return typeof value; + } + const receivedType = describeType(rawType); const hasKey = typeof parsedRecord?.key === "string" && parsedRecord.key.trim().length > 0; const compact = rawText.replace(/\s+/g, " ").trim(); const preview = compact.length > 160 ? `${compact.slice(0, 157)}...` : compact; @@ -285,39 +289,7 @@ export class EncryptedChannel { if (parsed.type === "e2ee_hello" && typeof parsed.key === "string") { if (this.options.daemonKeyPair) { - try { - const clientPublicKey = importPublicKey(parsed.key); - const nextSharedKey = deriveSharedKey( - this.options.daemonKeyPair.secretKey, - clientPublicKey, - ); - - // If it's the same client key (handshake retry), re-send - // "ready" but do not re-key. Re-keying here would desync - // the channel and cause decrypt failures. - if (keysEqual(nextSharedKey, this.sharedKey)) { - this.transport.send( - JSON.stringify({ type: "e2ee_ready" } satisfies E2EEReadyMessage), - ); - return null; - } - - // Different key implies a new client connection (common with relays - // where the daemon's socket stays open while the client reconnects). - // Re-key and re-send "ready". Drop any queued sends to avoid leaking - // messages between logical client sessions. - this.state = "handshaking"; - this.sharedKey = nextSharedKey; - this.pendingSends = []; - this.transport.send( - JSON.stringify({ type: "e2ee_ready" } satisfies E2EEReadyMessage), - ); - this.state = "open"; - await this.flushPendingSends(); - return null; - } catch (error) { - throw error; - } + await this.handleDaemonRehello(parsed.key); } return null; } @@ -398,6 +370,31 @@ export class EncryptedChannel { } } + private async handleDaemonRehello(clientKeyB64: string): Promise { + if (!this.options.daemonKeyPair) return; + const clientPublicKey = importPublicKey(clientKeyB64); + const nextSharedKey = deriveSharedKey(this.options.daemonKeyPair.secretKey, clientPublicKey); + + // If it's the same client key (handshake retry), re-send + // "ready" but do not re-key. Re-keying here would desync + // the channel and cause decrypt failures. + if (keysEqual(nextSharedKey, this.sharedKey)) { + this.transport.send(JSON.stringify({ type: "e2ee_ready" } satisfies E2EEReadyMessage)); + return; + } + + // Different key implies a new client connection (common with relays + // where the daemon's socket stays open while the client reconnects). + // Re-key and re-send "ready". Drop any queued sends to avoid leaking + // messages between logical client sessions. + this.state = "handshaking"; + this.sharedKey = nextSharedKey; + this.pendingSends = []; + this.transport.send(JSON.stringify({ type: "e2ee_ready" } satisfies E2EEReadyMessage)); + this.state = "open"; + await this.flushPendingSends(); + } + close(code = 1000, reason = "Normal closure"): void { this.state = "closed"; this.transport.close(code, reason); diff --git a/packages/relay/src/live-relay.e2e.test.ts b/packages/relay/src/live-relay.e2e.test.ts index 3794cfd37..444006a0b 100644 --- a/packages/relay/src/live-relay.e2e.test.ts +++ b/packages/relay/src/live-relay.e2e.test.ts @@ -29,6 +29,62 @@ async function withRetry( throw lastError instanceof Error ? lastError : new Error(String(lastError)); } +function waitOpen(ws: WebSocket, label: string): Promise { + return new Promise((resolve, reject) => { + const timeout = setTimeout( + () => reject(new Error(`Timed out opening ${label} websocket`)), + 10_000, + ); + const onOpen = () => { + clearTimeout(timeout); + resolve(); + }; + const onError = (err: Error) => { + clearTimeout(timeout); + reject(err); + }; + ws.once("open", onOpen); + ws.once("error", onError); + }); +} + +function waitForConnected(ws: WebSocket, connectionId: string): Promise { + return new Promise((resolve, reject) => { + const timeout = setTimeout(() => reject(new Error("Timed out waiting for connected")), 10_000); + const onMessage = (raw: WebSocket.RawData) => { + try { + const msg = JSON.parse(raw.toString()); + if (msg && msg.type === "connected" && msg.connectionId === connectionId) { + clearTimeout(timeout); + resolve(); + } + } catch { + // ignore + } + }; + ws.on("message", onMessage); + }); +} + +function waitForOnceMessage( + ws: WebSocket, + mode: T, + timeoutError: string, +): Promise { + return new Promise((resolve, reject) => { + const timeout = setTimeout(() => reject(new Error(timeoutError)), 10_000); + const onMessage = (data: WebSocket.RawData) => { + clearTimeout(timeout); + resolve( + (mode === "string" ? data.toString() : (data as Buffer)) as T extends "string" + ? string + : Buffer, + ); + }; + ws.once("message", onMessage); + }); +} + describe("Live relay (relay.paseo.sh) E2E", () => { const liveIt = process.env.RUN_LIVE_RELAY_E2E === "1" ? it : it.skip; @@ -63,45 +119,13 @@ describe("Live relay (relay.paseo.sh) E2E", () => { const clientWs = new WebSocket(clientUrl); let daemonWs: WebSocket | null = null; - const waitOpen = (ws: WebSocket, label: string) => - new Promise((resolve, reject) => { - const timeout = setTimeout( - () => reject(new Error(`Timed out opening ${label} websocket`)), - 10_000, - ); - ws.once("open", () => { - clearTimeout(timeout); - resolve(); - }); - ws.once("error", (err) => { - clearTimeout(timeout); - reject(err); - }); - }); - try { await Promise.all([ waitOpen(daemonControlWs, "server-control"), waitOpen(clientWs, "client"), ]); - await new Promise((resolve, reject) => { - const timeout = setTimeout( - () => reject(new Error("Timed out waiting for connected")), - 10_000, - ); - daemonControlWs.on("message", (raw) => { - try { - const msg = JSON.parse(raw.toString()); - if (msg && msg.type === "connected" && msg.connectionId === connectionId) { - clearTimeout(timeout); - resolve(); - } - } catch { - // ignore - } - }); - }); + await waitForConnected(daemonControlWs, connectionId); daemonWs = new WebSocket(serverDataUrl); await waitOpen(daemonWs, "server-data"); @@ -110,16 +134,11 @@ describe("Live relay (relay.paseo.sh) E2E", () => { // Client sends hello with its public key (not encrypted). clientWs.send(JSON.stringify({ type: "hello", key: clientPubKeyB64 })); - const daemonReceivedHello = await new Promise((resolve, reject) => { - const timeout = setTimeout( - () => reject(new Error("Timed out waiting for hello")), - 10_000, - ); - daemonWs!.once("message", (data) => { - clearTimeout(timeout); - resolve(data.toString()); - }); - }); + const daemonReceivedHello = await waitForOnceMessage( + daemonWs, + "string", + "Timed out waiting for hello", + ); const hello = JSON.parse(daemonReceivedHello) as { type: string; @@ -139,16 +158,11 @@ describe("Live relay (relay.paseo.sh) E2E", () => { const ciphertextFromClient = await encrypt(clientSharedKey, plaintextFromClient); clientWs.send(Buffer.from(ciphertextFromClient)); - const daemonReceivedCiphertext = await new Promise((resolve, reject) => { - const timeout = setTimeout( - () => reject(new Error("Timed out waiting for encrypted message")), - 10_000, - ); - daemonWs!.once("message", (data) => { - clearTimeout(timeout); - resolve(data as Buffer); - }); - }); + const daemonReceivedCiphertext = await waitForOnceMessage( + daemonWs, + "buffer", + "Timed out waiting for encrypted message", + ); const decryptedOnDaemon = await decrypt( daemonSharedKey, @@ -163,16 +177,11 @@ describe("Live relay (relay.paseo.sh) E2E", () => { const ciphertextFromDaemon = await encrypt(daemonSharedKey, plaintextFromDaemon); daemonWs!.send(Buffer.from(ciphertextFromDaemon)); - const clientReceivedCiphertext = await new Promise((resolve, reject) => { - const timeout = setTimeout( - () => reject(new Error("Timed out waiting for encrypted response")), - 10_000, - ); - clientWs.once("message", (data) => { - clearTimeout(timeout); - resolve(data as Buffer); - }); - }); + const clientReceivedCiphertext = await waitForOnceMessage( + clientWs, + "buffer", + "Timed out waiting for encrypted response", + ); const decryptedOnClient = await decrypt( clientSharedKey,