diff --git a/src/webrtc-signaling.ts b/src/webrtc-signaling.ts index c29f788..a2ede3f 100644 --- a/src/webrtc-signaling.ts +++ b/src/webrtc-signaling.ts @@ -6,7 +6,7 @@ import { prisma } from "./db"; import { IncomingMessage } from "http"; import { Socket } from "node:net"; import { Device } from "@prisma/client"; -import { Server, ServerResponse } from "node:http"; +import { STATUS_CODES, Server, ServerResponse } from "node:http"; import { cookieSessionMiddleware } from "."; import { effectiveSku, normalizeSku } from "./skus"; @@ -23,6 +23,23 @@ export interface DeviceConnection { export const activeConnections: Map = new Map(); export const inFlight: Set = new Set(); +/** + * Refuses an upgrade with a real HTTP response before closing the socket. + * + * Destroying the socket without a response makes the edge in front of the + * API report the request as a 504, so every device holding a revoked token + * showed up as a gateway failure on each of its 5-second retries. + */ +export function rejectUpgrade(socket: Socket, status: number) { + if (!socket.writable) { + return socket.destroy(); + } + socket.once("finish", socket.destroy); + socket.end( + `HTTP/1.1 ${status} ${STATUS_CODES[status]}\r\nConnection: close\r\nContent-Length: 0\r\n\r\n`, + ); +} + function toICEServers(str: string) { return str.split(",").filter(url => url.startsWith("stun:")); } @@ -59,7 +76,7 @@ export function registerWebSocketRouter( await handleClientSocketRequest(req, socket, head); } else { console.log(`[Webrtc] Unrecognized path: ${path}`); - return socket.destroy(); + return rejectUpgrade(socket, 404); } }); } @@ -78,7 +95,7 @@ async function handleDeviceSocketRequest( // Authenticate device const device = await authenticateDeviceRequest(req); if (!device) { - return socket.destroy(); + return rejectUpgrade(socket, 401); } // Inflight means that the device has connected, a client has connected to that device via HTTP, and they're now doing the signaling dance @@ -86,7 +103,7 @@ async function handleDeviceSocketRequest( console.log( `[Device] Device ${device.id} already has an inflight client connection.`, ); - return socket.destroy(); + return rejectUpgrade(socket, 409); } // Handle existing connections for this device @@ -116,11 +133,12 @@ async function handleDeviceSocketRequest( }); } catch (error) { console.error("Error handling device socket request:", error); - socket.destroy(); + rejectUpgrade(socket, 500); } } -// Authenticate the device connection +// Authenticate the device connection. Returns null only when the token can +// never authenticate: missing, unknown, or bound to a different device id. async function authenticateDeviceRequest(req: IncomingMessage) { const authHeader = req.headers["authorization"]; const secretToken = authHeader?.split(" ")?.[1]; @@ -130,24 +148,22 @@ async function authenticateDeviceRequest(req: IncomingMessage) { return null; } - try { - const device = await prisma.device.findFirst({ where: { secretToken } }); - if (!device) { - console.log("[Device] Invalid secret token provided."); - return null; - } - - const id = req.headers["x-device-id"] as string; - if (!id || id !== device.id) { - console.log("[Device] Invalid device ID or ID/token mismatch."); - return null; - } + // A failed lookup (database down, pool exhausted) must not read as a bad + // token: the caller answers 500 for it, which the device treats as + // transient, while 401 means the token itself will never work. + const device = await prisma.device.findFirst({ where: { secretToken } }); + if (!device) { + console.log("[Device] Invalid secret token provided."); + return null; + } - return device; - } catch (error) { - console.error("[Device] Error authenticating device:", error); + const id = req.headers["x-device-id"] as string; + if (!id || id !== device.id) { + console.log("[Device] Invalid device ID or ID/token mismatch."); return null; } + + return device; } // Setup the device WebSocket after authentication @@ -227,16 +243,16 @@ async function handleClientSocketRequest( cookieSessionMiddleware(req as any, {} as any, async () => { try { // Authenticate client and get device ID - const { deviceId, token } = await authenticateClientRequest(req as any); - if (!deviceId) { - return socket.destroy(); + const auth = await authenticateClientRequest(req as any); + if (auth.deviceId === null) { + return rejectUpgrade(socket, auth.status); } + const { deviceId, token } = auth; // Check if device is connected if (!activeConnections.has(deviceId)) { console.log(`[Client] Device ${deviceId} not connected.`); - socket.write("HTTP/1.1 404 Not Found\r\n\r\n"); - return socket.destroy(); + return rejectUpgrade(socket, 404); } // Complete the WebSocket upgrade @@ -245,23 +261,29 @@ async function handleClientSocketRequest( }); } catch (error) { console.error("Error in client WebSocket setup:", error); - socket.destroy(); + rejectUpgrade(socket, 500); } }); } catch (error) { console.error("Error handling client socket request:", error); - socket.destroy(); + rejectUpgrade(socket, 500); } } +type ClientAuth = + | { deviceId: string; token: string } + | { deviceId: null; status: 401 | 404 }; + // Authenticate the client connection -async function authenticateClientRequest(req: Request & { session: any }) { +async function authenticateClientRequest( + req: Request & { session: any }, +): Promise { const session = req.session; const token = session?.id_token; if (!token) { console.log("[Client] No authentication token."); - return { deviceId: null }; + return { deviceId: null, status: 401 }; } try { @@ -271,7 +293,7 @@ async function authenticateClientRequest(req: Request & { session: any }) { if (!deviceId) { console.log("[Client] No device ID provided."); - return { deviceId: null }; + return { deviceId: null, status: 404 }; } // Check if device exists and user has access @@ -282,13 +304,13 @@ async function authenticateClientRequest(req: Request & { session: any }) { if (!device) { console.log("[Client] Device not found or user doesn't have access."); - return { deviceId: null }; + return { deviceId: null, status: 404 }; } return { deviceId, token }; } catch (error) { console.error("[Client] Authentication error:", error); - return { deviceId: null }; + return { deviceId: null, status: 401 }; } } diff --git a/test/webrtc-signaling.test.ts b/test/webrtc-signaling.test.ts new file mode 100644 index 0000000..0e3caf2 --- /dev/null +++ b/test/webrtc-signaling.test.ts @@ -0,0 +1,193 @@ +import { afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; +import http from "node:http"; +import type { AddressInfo, Socket } from "node:net"; + +import { testPrisma } from "./setup"; + +// Prisma model delegates are proxies, so vi.spyOn cannot replace their +// methods. Wrap the real client instead and let a test fail the next +// device lookup, to simulate a database that is down or out of connections. +const db = vi.hoisted(() => ({ nextDeviceLookupError: null as Error | null })); +vi.mock("../src/db", async importOriginal => { + const { prisma } = await importOriginal(); + const device = new Proxy(prisma.device, { + get(target, prop, receiver) { + if (prop === "findFirst" && db.nextDeviceLookupError) { + const error = db.nextDeviceLookupError; + db.nextDeviceLookupError = null; + return () => Promise.reject(error); + } + return Reflect.get(target, prop, receiver); + }, + }); + return { + prisma: new Proxy(prisma, { + get: (target, prop, receiver) => + prop === "device" ? device : Reflect.get(target, prop, receiver), + }), + }; +}); + +// The signaling module imports the cookie-session middleware from src/index.ts, +// which starts the listening server. Replace it with one that reads the session +// from a test header, so client upgrades can carry a session without cookies. +vi.mock("../src/index", () => ({ + cookieSessionMiddleware: (req: any, _res: any, next: () => void) => { + const raw = req.headers["x-test-session"]; + req.session = raw ? JSON.parse(raw) : {}; + next(); + }, +})); + +const { activeConnections, registerWebSocketRouter } = await import( + "../src/webrtc-signaling" +); + +const GOOGLE_ID = "signaling-test-google-id"; +const DEVICE_ID = "signaling-test-device"; +const SECRET_TOKEN = "signaling-test-secret-token"; + +let server: http.Server; +let port: number; + +function unsignedJwt(payload: Record): string { + const encode = (value: unknown) => + Buffer.from(JSON.stringify(value)).toString("base64url"); + return `${encode({ alg: "none" })}.${encode(payload)}.`; +} + +function sessionHeader(sub: string): Record { + return { + "x-test-session": JSON.stringify({ + id_token: unsignedJwt({ iss: "https://accounts.google.com", sub }), + }), + }; +} + +interface UpgradeResult { + status: number; + /** Present only when the server completed the upgrade (101). */ + socket?: Socket; +} + +/** + * Sends a WebSocket upgrade request and resolves with the status the server + * answered with. Rejects if the server closes the socket without a response, + * which is what the handlers did before they wrote a status line. + */ +function upgrade(path: string, headers: Record = {}): Promise { + return new Promise((resolve, reject) => { + const req = http.request({ + host: "127.0.0.1", + port, + path, + headers: { + Connection: "Upgrade", + Upgrade: "websocket", + "Sec-WebSocket-Version": "13", + "Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ==", + ...headers, + }, + }); + req.on("upgrade", (res, socket) => resolve({ status: res.statusCode!, socket })); + req.on("response", res => { + res.resume(); + resolve({ status: res.statusCode! }); + }); + req.on("error", reject); + req.end(); + }); +} + +const deviceHeaders = (token: string, id = DEVICE_ID) => ({ + Authorization: `Bearer ${token}`, + "X-Device-ID": id, + "X-App-Version": "0.5.9", + "X-Device-SKU": "jetkvm-v2", +}); + +beforeAll(async () => { + server = http.createServer((_req, res) => res.writeHead(200).end()); + registerWebSocketRouter(server); + await new Promise(resolve => server.listen(0, "127.0.0.1", resolve)); + port = (server.address() as AddressInfo).port; +}); + +afterAll(async () => { + for (const conn of activeConnections.values()) conn.ws.terminate(); + await new Promise(resolve => server.close(() => resolve())); +}); + +describe("upgrade rejections", () => { + beforeEach(async () => { + activeConnections.clear(); + const user = await testPrisma.user.upsert({ + where: { googleId: GOOGLE_ID }, + update: {}, + create: { googleId: GOOGLE_ID }, + }); + await testPrisma.device.create({ + data: { id: DEVICE_ID, userId: user.id, secretToken: SECRET_TOKEN }, + }); + }); + + afterEach(async () => { + await testPrisma.device.deleteMany({ where: { id: DEVICE_ID } }); + await testPrisma.user.deleteMany({ where: { googleId: GOOGLE_ID } }); + }); + + it("answers 404 for an unknown path", async () => { + expect((await upgrade("/nope")).status).toBe(404); + }); + + it("answers 401 for a device upgrade without a token", async () => { + expect((await upgrade("/", { "X-Device-ID": DEVICE_ID })).status).toBe(401); + }); + + it("answers 401 for a device upgrade with a revoked token", async () => { + expect((await upgrade("/", deviceHeaders("not-the-token"))).status).toBe(401); + expect(activeConnections.has(DEVICE_ID)).toBe(false); + }); + + it("answers 401 when the device id does not match the token", async () => { + expect((await upgrade("/", deviceHeaders(SECRET_TOKEN, "other-device"))).status).toBe(401); + }); + + it("answers 500, not 401, when the token lookup itself fails", async () => { + db.nextDeviceLookupError = new Error("connection pool exhausted"); + expect((await upgrade("/", deviceHeaders(SECRET_TOKEN))).status).toBe(500); + expect(db.nextDeviceLookupError).toBeNull(); + }); + + it("completes the upgrade for a valid device token", async () => { + const result = await upgrade("/", deviceHeaders(SECRET_TOKEN)); + expect(result.status).toBe(101); + expect(activeConnections.get(DEVICE_ID)).toMatchObject({ + version: "0.5.9", + sku: "jetkvm-v2", + }); + + result.socket!.destroy(); + await vi.waitFor(() => expect(activeConnections.has(DEVICE_ID)).toBe(false)); + }); + + it("answers 401 for a client upgrade without a session", async () => { + expect((await upgrade(`/webrtc/signaling/client?id=${DEVICE_ID}`)).status).toBe(401); + }); + + it("answers 404 for a client upgrade to a device the user does not own", async () => { + const result = await upgrade( + `/webrtc/signaling/client?id=${DEVICE_ID}`, + sessionHeader("someone-else"), + ); + expect(result.status).toBe(404); + }); + + it("answers 404 for a client upgrade to a device that is offline", async () => { + const result = await upgrade( + `/webrtc/signaling/client?id=${DEVICE_ID}`, + sessionHeader(GOOGLE_ID), + ); + expect(result.status).toBe(404); + }); +});