Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,8 @@ Pass an `AbortSignal` to integrate shutdown with your process lifecycle.

Enable the built-in `ws` transport with `websocket: true`. It defaults to a
1 MiB maximum message payload with compression disabled; pass
`websocket: { maxPayload, perMessageDeflate }` to override those settings.
`websocket: { maxPayload, maxRejectionBodyBytes, perMessageDeflate }` to override
those settings. Rejected upgrade bodies are capped at 64 KiB by default.

```ts
router.ws("/echo", (socket) => {
Expand Down
1 change: 1 addition & 0 deletions src/contracts.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ import type { PerMessageDeflateOptions } from "ws";

export interface NodeWebSocketOptions {
readonly maxPayload?: number;
readonly maxRejectionBodyBytes?: number;
readonly perMessageDeflate?: boolean | PerMessageDeflateOptions;
readonly allowedOrigins?: readonly string[];
}
Expand Down
60 changes: 54 additions & 6 deletions src/websocket.ts
Original file line number Diff line number Diff line change
Expand Up @@ -36,11 +36,51 @@ function socketLike(socket: WebSocket): WebSocketLike {
};
}

async function rejectUpgrade(socket: Duplex, response: Response): Promise<void> {
const body = Buffer.from(await response.arrayBuffer());
async function readRejectionBody(
response: Response,
maxBytes: number,
): Promise<Buffer | undefined> {
const lengthHeader = response.headers.get("content-length");
const declaredLength = lengthHeader === null ? undefined : Number(lengthHeader);
if (
declaredLength !== undefined &&
Number.isFinite(declaredLength) &&
declaredLength > maxBytes
) {
await response.body?.cancel("WebSocket rejection body exceeded maxRejectionBodyBytes");
return undefined;
}
if (!response.body) return Buffer.alloc(0);
const reader = response.body.getReader();
const chunks: Buffer[] = [];
let length = 0;
while (true) {
const part = await reader.read();
if (part.done) return Buffer.concat(chunks, length);
length += part.value.byteLength;
if (length > maxBytes) {
await reader.cancel("WebSocket rejection body exceeded maxRejectionBodyBytes");
return undefined;
}
chunks.push(Buffer.from(part.value));
}
}

async function rejectUpgrade(
socket: Duplex,
response: Response,
maxBodyBytes: number,
): Promise<void> {
const buffered = await readRejectionBody(response, maxBodyBytes);
if (!buffered) {
response = new Response("WebSocket rejection body exceeded configured limit", { status: 500 });
}
const body = buffered ?? Buffer.from(await response.arrayBuffer());
const lines = [`HTTP/1.1 ${response.status} ${response.statusText || "Rejected"}`];
response.headers.forEach((value, name) => lines.push(`${name}: ${value}`));
if (!response.headers.has("content-length")) lines.push(`content-length: ${body.byteLength}`);
response.headers.forEach((value, name) => {
if (name !== "content-length" && name !== "transfer-encoding") lines.push(`${name}: ${value}`);
});
lines.push(`content-length: ${body.byteLength}`);
lines.push("connection: close", "", "");
socket.end(Buffer.concat([Buffer.from(lines.join("\r\n")), body]));
}
Expand All @@ -51,6 +91,10 @@ export function installWebSockets(
options: NodeWebSocketOptions = {},
handlerOptions: NodeHandlerOptions = {},
): { close(): void } {
const maxRejectionBodyBytes = options.maxRejectionBodyBytes ?? 65_536;
if (!Number.isInteger(maxRejectionBodyBytes) || maxRejectionBodyBytes <= 0) {
throw new TypeError("WebSocket maxRejectionBodyBytes must be a positive integer.");
}
const webSockets = new WebSocketServer({
noServer: true,
maxPayload: options.maxPayload ?? 1_048_576,
Expand All @@ -74,7 +118,11 @@ export function installWebSockets(
normalizedOrigin = undefined;
}
if (!normalizedOrigin || !allowedOrigins.includes(normalizedOrigin)) {
await rejectUpgrade(socket, new Response("Forbidden", { status: 403 }));
await rejectUpgrade(
socket,
new Response("Forbidden", { status: 403 }),
maxRejectionBodyBytes,
);
return;
}
let marker: Response | undefined;
Expand All @@ -88,7 +136,7 @@ export function installWebSockets(
},
});
if (!accepted || response !== accepted.response) {
await rejectUpgrade(socket, response);
await rejectUpgrade(socket, response, maxRejectionBodyBytes);
return;
}
webSockets.handleUpgrade(request, socket, head, (webSocket) => {
Expand Down
35 changes: 35 additions & 0 deletions tests/node.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,41 @@ describe("Node adapter", () => {
await new Promise<void>((resolve) => server.close(() => resolve()));
});

it("should stop buffering oversized WebSocket rejection bodies", async () => {
let cancelled = false;
const body = new ReadableStream({
pull(controller) {
controller.enqueue(new TextEncoder().encode("123456"));
},
cancel() {
cancelled = true;
},
});
const server = await listen(
{ fetch: async () => new Response(body, { status: 401 }) },
{
host: "127.0.0.1",
websocket: { maxRejectionBodyBytes: 8 },
},
);
const address = server.address();
if (!address || typeof address === "string") throw new Error("Expected TCP address");
const socket = new WebSocket(`ws://127.0.0.1:${address.port}/rejected`, {
origin: `http://127.0.0.1:${address.port}`,
});
socket.on("error", () => undefined);
const [, response] = await once(socket, "unexpected-response");
const chunks: Buffer[] = [];
response.on("data", (chunk) => chunks.push(Buffer.from(chunk)));
await once(response, "end");
expect(response.statusCode).toBe(500);
expect(Buffer.concat(chunks).toString()).toBe(
"WebSocket rejection body exceeded configured limit",
);
expect(cancelled).toBe(true);
await new Promise<void>((resolve) => server.close(() => resolve()));
});

it("should reject untrusted Host and absolute-form request targets", async () => {
const server = await listen(
createServerApp({ routes: [{ path: "/", handler: (ctx) => ctx.ok() }] }),
Expand Down