diff --git a/src/http_jsc/websocket_client.rs b/src/http_jsc/websocket_client.rs index b5c43240de97..d2ace28b3ee5 100644 --- a/src/http_jsc/websocket_client.rs +++ b/src/http_jsc/websocket_client.rs @@ -2043,6 +2043,7 @@ pub enum ErrorCode { ProxyConnectionRefused = 35, ProxyTunnelFailed = 36, UnexpectedRsv1 = 37, + InvalidExtensionsHeader = 38, } // ────────────────────────────────────────────────────────────────────────── diff --git a/src/http_jsc/websocket_client/WebSocketUpgradeClient.rs b/src/http_jsc/websocket_client/WebSocketUpgradeClient.rs index e7c63e78ba66..4ea24e8bdafc 100644 --- a/src/http_jsc/websocket_client/WebSocketUpgradeClient.rs +++ b/src/http_jsc/websocket_client/WebSocketUpgradeClient.rs @@ -74,6 +74,109 @@ pub(crate) struct DeflateNegotiationResult { pub params: WebSocketDeflate::Params, } +/// Strips one layer of `"…"` quoting from an extension parameter value. +/// RFC 6455 §9.1 requires the unescaped content of a quoted-string value to +/// conform to the `token` ABNF, so no inner escapes need handling. +fn unquote(value: &[u8]) -> &[u8] { + if value.len() >= 2 && value[0] == b'"' && value[value.len() - 1] == b'"' { + &value[1..value.len() - 1] + } else { + value + } +} + +/// Parses a `…_max_window_bits` value from an extension negotiation response. +/// In a response the parameter must carry a decimal value between 8 and 15 +/// (RFC 7692 §8.1.2); a missing, malformed, or out-of-range value is `None`. +fn parse_window_bits(value: Option<&[u8]>) -> Option { + let value = unquote(value?); + // The value's grammar is `1*DIGIT` without leading zeroes; `parse_int` + // alone would also accept a `+` sign, `_` separators, and leading zeroes. + if value.is_empty() + || !value.iter().all(u8::is_ascii_digit) + || (value.len() > 1 && value[0] == b'0') + { + return None; + } + let bits = strings::parse_int::(value, 10).ok()?; + (WebSocketDeflate::Params::MIN_WINDOW_BITS..=WebSocketDeflate::Params::MAX_WINDOW_BITS) + .contains(&bits) + .then_some(bits) +} + +/// Validates one `Sec-WebSocket-Extensions` response header value against the +/// single `permessage-deflate; client_max_window_bits` offer we sent and +/// accumulates the accepted parameters into `deflate`. +/// +/// Returns `false` when the server lists an extension the client did not +/// offer (RFC 6455 §4.1), accepts `permessage-deflate` more than once, or +/// sends a `permessage-deflate` parameter that is unknown, repeated, or +/// carries an invalid value (RFC 7692 §8.1). The caller must then fail the +/// connection. The response may span several header lines; `deflate.enabled` +/// persists across calls so a repeated `permessage-deflate` element anywhere +/// in the response is rejected. +fn accept_extensions_response(value: &[u8], deflate: &mut DeflateNegotiationResult) -> bool { + let mut elements = HeaderValueIterator::init(value); + while let Some(element) = elements.next() { + let mut parts = element.split(|b| *b == b';'); + let name = strings::trim(parts.next().unwrap_or(b""), b" \t"); + if name != b"permessage-deflate" || deflate.enabled { + return false; + } + deflate.enabled = true; + + let mut seen_server_no_context_takeover = false; + let mut seen_client_no_context_takeover = false; + let mut seen_server_max_window_bits = false; + let mut seen_client_max_window_bits = false; + for part in parts { + let part = strings::trim(part, b" \t"); + let (key, value) = match strings::index_of_char_usize(part, b'=') { + Some(i) => ( + strings::trim(&part[..i], b" \t"), + Some(strings::trim(&part[i + 1..], b" \t")), + ), + None => (part, None), + }; + + match key { + // The `…_no_context_takeover` parameters carry no value + // (RFC 7692 §8.1.1). + b"server_no_context_takeover" + if value.is_none() && !seen_server_no_context_takeover => + { + seen_server_no_context_takeover = true; + deflate.params.server_no_context_takeover = 1; + } + b"client_no_context_takeover" + if value.is_none() && !seen_client_no_context_takeover => + { + seen_client_no_context_takeover = true; + deflate.params.client_no_context_takeover = 1; + } + b"server_max_window_bits" if !seen_server_max_window_bits => { + seen_server_max_window_bits = true; + let Some(bits) = parse_window_bits(value) else { + return false; + }; + deflate.params.server_max_window_bits = bits; + } + b"client_max_window_bits" if !seen_client_max_window_bits => { + seen_client_max_window_bits = true; + let Some(bits) = parse_window_bits(value) else { + return false; + }; + deflate.params.client_max_window_bits = bits; + } + // Anything else is a parameter that is not defined for use in + // a response, is repeated, or carries an unexpected value. + _ => return false, + } + } + } + true +} + #[derive(Clone, Copy, PartialEq, Eq)] enum State { Initializing, @@ -1386,90 +1489,17 @@ impl HTTPClient { header.name(), b"Sec-WebSocket-Extensions", ) { - // Per RFC 6455 §9.1, the server MUST NOT respond with an - // extension the client did not offer. Match upstream `ws` - // (lib/websocket.js: "Server sent a Sec-WebSocket-Extensions - // header but no extension was requested") and fail the - // handshake instead of silently accepting it. + // RFC 6455 §4.1 and RFC 7692 §8.1 require failing the + // connection on an extension we did not offer or an + // invalid `permessage-deflate` parameter, like `ws` does. // SAFETY: short-lived read. - if !unsafe { (*this).offered_permessage_deflate } { + if !unsafe { (*this).offered_permessage_deflate } + || !accept_extensions_response(header.value(), &mut deflate_result) + { // SAFETY: no `&mut Self` is live across this call. - unsafe { Self::terminate(this, ErrorCode::InvalidResponse) }; + unsafe { Self::terminate(this, ErrorCode::InvalidExtensionsHeader) }; return; } - // This is a simplified parser. A full parser would handle multiple extensions and quoted values. - for ext_str in header.value().split(|b| *b == b',') { - let mut ext_it = strings::trim(ext_str, b" \t").split(|b| *b == b';'); - let ext_name = strings::trim(ext_it.next().unwrap_or(b""), b" \t"); - if ext_name == b"permessage-deflate" { - deflate_result.enabled = true; - for param_str in ext_it { - let mut param_it = - strings::trim(param_str, b" \t").split(|b| *b == b'='); - let key = strings::trim(param_it.next().unwrap_or(b""), b" \t"); - let value = - strings::trim(param_it.next().unwrap_or(b""), b" \t"); - - if key == b"server_no_context_takeover" { - deflate_result.params.server_no_context_takeover = 1; - } else if key == b"client_no_context_takeover" { - deflate_result.params.client_no_context_takeover = 1; - } else if key == b"server_max_window_bits" { - if !value.is_empty() { - // Remove quotes if present - let trimmed_value = if value.len() >= 2 - && value[0] == b'"' - && value[value.len() - 1] == b'"' - { - &value[1..value.len() - 1] - } else { - value - }; - - if let Ok(bits) = - strings::parse_int::(trimmed_value, 10) - { - if bits >= WebSocketDeflate::Params::MIN_WINDOW_BITS - && bits - <= WebSocketDeflate::Params::MAX_WINDOW_BITS - { - deflate_result.params.server_max_window_bits = - bits; - } - } - } - } else if key == b"client_max_window_bits" { - if !value.is_empty() { - // Remove quotes if present - let trimmed_value = if value.len() >= 2 - && value[0] == b'"' - && value[value.len() - 1] == b'"' - { - &value[1..value.len() - 1] - } else { - value - }; - - if let Ok(bits) = - strings::parse_int::(trimmed_value, 10) - { - if bits >= WebSocketDeflate::Params::MIN_WINDOW_BITS - && bits - <= WebSocketDeflate::Params::MAX_WINDOW_BITS - { - deflate_result.params.client_max_window_bits = - bits; - } - } - } else { - // client_max_window_bits without value means use default (15) - deflate_result.params.client_max_window_bits = 15; - } - } - } - break; // Found and parsed permessage-deflate, stop. - } - } } } _ => {} @@ -1514,6 +1544,9 @@ impl HTTPClient { return; } + // The client requested one or more subprotocols but the server's 101 + // did not select any. Both RFC 6455 §4.1 and the WHATWG "establish a + // WebSocket connection" algorithm require failing the connection. // SAFETY: short-lived `&self` read. if !protocol_header_seen && !unsafe { (*this).subprotocols.is_empty() } { // SAFETY: no `&mut Self` is live across this call. diff --git a/src/jsc/bindings/webcore/WebSocket.cpp b/src/jsc/bindings/webcore/WebSocket.cpp index d3afe63e0b90..fb1c614e425e 100644 --- a/src/jsc/bindings/webcore/WebSocket.cpp +++ b/src/jsc/bindings/webcore/WebSocket.cpp @@ -1722,11 +1722,11 @@ void WebSocket::didFailWithErrorCode(Bun::WebSocketErrorCode code) break; } case Bun::WebSocketErrorCode::missing_client_protocol: { - didReceiveClose(CleanStatus::Clean, 1002, "Missing client protocol"_s); + didReceiveClose(CleanStatus::NotClean, 1002, "Server sent no subprotocol"_s, true); break; } case Bun::WebSocketErrorCode::mismatch_client_protocol: { - didReceiveClose(CleanStatus::Clean, 1002, "Mismatch client protocol"_s); + didReceiveClose(CleanStatus::NotClean, 1002, "Mismatch client protocol"_s, true); break; } case Bun::WebSocketErrorCode::timeout: { @@ -1832,6 +1832,10 @@ void WebSocket::didFailWithErrorCode(Bun::WebSocketErrorCode code) didReceiveClose(CleanStatus::NotClean, 1006, "Proxy tunnel failed"_s, true); break; } + case Bun::WebSocketErrorCode::invalid_extensions_header: { + didReceiveClose(CleanStatus::NotClean, 1002, "Invalid Sec-WebSocket-Extensions header"_s, true); + break; + } } // didReceiveClose already set m_state = CLOSED. The connect() ref diff --git a/src/jsc/bindings/webcore/WebSocketErrorCode.h b/src/jsc/bindings/webcore/WebSocketErrorCode.h index 02f478186a1e..6ca1df3d1d8b 100644 --- a/src/jsc/bindings/webcore/WebSocketErrorCode.h +++ b/src/jsc/bindings/webcore/WebSocketErrorCode.h @@ -43,6 +43,7 @@ enum class WebSocketErrorCode : int32_t { proxy_connection_refused = 35, proxy_tunnel_failed = 36, unexpected_rsv1 = 37, + invalid_extensions_header = 38, }; } diff --git a/test/js/web/websocket/websocket-permessage-deflate-edge-cases.test.ts b/test/js/web/websocket/websocket-permessage-deflate-edge-cases.test.ts index 02dd3a1aa50e..87dcf0b7704c 100644 --- a/test/js/web/websocket/websocket-permessage-deflate-edge-cases.test.ts +++ b/test/js/web/websocket/websocket-permessage-deflate-edge-cases.test.ts @@ -1,5 +1,5 @@ import { serve } from "bun"; -import { expect, setDefaultTimeout, test } from "bun:test"; +import { describe, expect, setDefaultTimeout, test } from "bun:test"; import crypto from "node:crypto"; import net from "node:net"; import { deflateRawSync, constants as zc } from "node:zlib"; @@ -570,3 +570,153 @@ test.each([false, true])( } }, ); + +// RFC 6455 section 4.1 requires the client to fail the connection when the +// server indicates an extension that was not offered, and RFC 7692 section +// 8.1 requires the same for a permessage-deflate response with an unknown +// parameter or an invalid parameter value. Use a raw TCP server so the +// Sec-WebSocket-Extensions response header is byte-exact. +describe("Sec-WebSocket-Extensions response validation", () => { + type Outcome = { opened: boolean; message: string | null; extensions: string; code: number; reason: string }; + + // Completes a correct handshake, optionally answering with the given + // Sec-WebSocket-Extensions value, then sends `frame`. Resolves once the + // client either receives the frame or observes the connection close. + async function negotiate(extensionsValue: string | null, frame: Buffer): Promise { + let request = Buffer.alloc(0); + using server = Bun.listen({ + hostname: "127.0.0.1", + port: 0, + socket: { + error() {}, + data(socket, data) { + request = Buffer.concat([request, data]); + const end = request.indexOf("\r\n\r\n"); + if (end === -1) return; + const key = /Sec-WebSocket-Key: *([A-Za-z0-9+/=]+)/i.exec(request.subarray(0, end).toString())?.[1]; + if (!key) { + socket.end(); + return; + } + const hasher = new Bun.CryptoHasher("sha1"); + hasher.update(key + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let response = + "HTTP/1.1 101 Switching Protocols\r\n" + + "Upgrade: websocket\r\n" + + "Connection: Upgrade\r\n" + + `Sec-WebSocket-Accept: ${hasher.digest("base64")}\r\n`; + if (extensionsValue !== null) { + response += `Sec-WebSocket-Extensions: ${extensionsValue}\r\n`; + } + socket.write(response + "\r\n"); + socket.write(frame); + socket.flush(); + }, + }, + }); + + const client = new WebSocket(`ws://127.0.0.1:${server.port}/`); + const { promise, resolve } = Promise.withResolvers(); + let opened = false; + client.onopen = () => { + opened = true; + }; + client.onerror = () => {}; + client.onmessage = event => { + resolve({ opened, message: event.data as string, extensions: client.extensions, code: 0, reason: "" }); + }; + client.onclose = event => { + resolve({ opened, message: null, extensions: client.extensions, code: event.code, reason: event.reason }); + }; + try { + return await promise; + } finally { + client.close(); + } + } + + // An unmasked, uncompressed text frame: FIN + text opcode, 5-byte payload. + const helloFrame = Buffer.from([0x81, 0x05, ...Buffer.from("hello")]); + + const invalid = [ + // The client only ever offers permessage-deflate. + "x-never-offered", + "permessage-deflate, x-never-offered", + "x-never-offered, permessage-deflate", + "permessage-deflate, permessage-deflate", + // Parameters that RFC 7692 does not define for use in a response. + "permessage-deflate; bogus_param=1", + "permessage-deflate; threshold=1024", + // window bits must be a decimal integer between 8 and 15, and a response + // (unlike an offer) must always carry the value. + "permessage-deflate; server_max_window_bits=20", + "permessage-deflate; server_max_window_bits=7", + "permessage-deflate; server_max_window_bits=potato", + "permessage-deflate; server_max_window_bits", + "permessage-deflate; server_max_window_bits=", + // The value's grammar is `1*DIGIT` without leading zeroes: no sign, + // no digit separators, no leading zero. + "permessage-deflate; server_max_window_bits=+10", + "permessage-deflate; server_max_window_bits=1_0", + "permessage-deflate; server_max_window_bits=08", + "permessage-deflate; client_max_window_bits=16", + "permessage-deflate; client_max_window_bits", + // The no_context_takeover parameters never carry a value. + "permessage-deflate; server_no_context_takeover=1", + "permessage-deflate; client_no_context_takeover=true", + // A parameter must not be repeated within an element. + "permessage-deflate; server_max_window_bits=10; server_max_window_bits=11", + "permessage-deflate; client_no_context_takeover; client_no_context_takeover", + ]; + + test.each(invalid)("fails the connection for %j", async value => { + expect(await negotiate(value, helloFrame)).toEqual({ + opened: false, + message: null, + extensions: "", + code: 1002, + reason: "Invalid Sec-WebSocket-Extensions header", + }); + }); + + const valid: [header: string | null, extensions: string][] = [ + [null, ""], + ["permessage-deflate", "permessage-deflate"], + [ + "permessage-deflate; server_no_context_takeover; client_no_context_takeover", + "permessage-deflate; server_no_context_takeover; client_no_context_takeover", + ], + [ + "permessage-deflate; server_max_window_bits=10; client_max_window_bits=12", + "permessage-deflate; server_max_window_bits=10; client_max_window_bits=12", + ], + // A quoted parameter value is legal (RFC 6455 section 9.1) and 15 is the default. + [ + 'permessage-deflate; server_max_window_bits="9"; client_max_window_bits=15', + "permessage-deflate; server_max_window_bits=9", + ], + ]; + + test.each(valid)("accepts %j", async (value, extensions) => { + expect(await negotiate(value, helloFrame)).toEqual({ + opened: true, + message: "hello", + extensions, + code: 0, + reason: "", + }); + }); + + test("still inflates frames after a validated negotiation", async () => { + // RSV1 + FIN + text opcode, with a raw-deflate payload per RFC 7692. + const compressed = deflateRawSync("hello"); + const frame = Buffer.concat([Buffer.from([0xc1, compressed.length]), compressed]); + expect(await negotiate("permessage-deflate; server_no_context_takeover", frame)).toEqual({ + opened: true, + message: "hello", + extensions: "permessage-deflate; server_no_context_takeover", + code: 0, + reason: "", + }); + }); +}); diff --git a/test/js/web/websocket/websocket-subprotocol-strict.test.ts b/test/js/web/websocket/websocket-subprotocol-strict.test.ts index 1d48911a424b..906dc3a41927 100644 --- a/test/js/web/websocket/websocket-subprotocol-strict.test.ts +++ b/test/js/web/websocket/websocket-subprotocol-strict.test.ts @@ -2,97 +2,122 @@ import { describe, expect, it, mock } from "bun:test"; import crypto from "node:crypto"; import net from "node:net"; -describe("WebSocket strict RFC 6455 subprotocol handling", () => { - async function createTestServer( - responseHeaders: string[], - ): Promise<{ port: number; [Symbol.asyncDispose]: () => Promise }> { - const server = net.createServer(); - let port: number; - - await new Promise(resolve => { - server.listen(0, () => { - port = (server.address() as any).port; - resolve(); - }); +// A byte-controlled RFC 6455 server: replies with a valid 101 plus whatever +// extra response headers the test crafts. Used to exercise the client-side +// handshake response validation in process_response(). +async function createTestServer( + responseHeaders: string[], +): Promise<{ port: number; [Symbol.asyncDispose]: () => Promise }> { + const server = net.createServer(); + let port: number; + + await new Promise(resolve => { + server.listen(0, "127.0.0.1", () => { + port = (server.address() as any).port; + resolve(); }); + }); - server.on("connection", socket => { - // Raw test server: tolerate client aborts, surface anything unexpected. - socket.on("error", (err: NodeJS.ErrnoException) => { - if (err.code !== "ECONNRESET" && err.code !== "EPIPE" && err.code !== "ECONNABORTED") throw err; - }); - let requestData = ""; + server.on("connection", socket => { + // Raw test server: tolerate client aborts, surface anything unexpected. + socket.on("error", (err: NodeJS.ErrnoException) => { + if (err.code !== "ECONNRESET" && err.code !== "EPIPE" && err.code !== "ECONNABORTED") throw err; + }); + let requestData = ""; - socket.on("data", data => { - requestData += data.toString(); + socket.on("data", data => { + requestData += data.toString(); - if (requestData.includes("\r\n\r\n")) { - const lines = requestData.split("\r\n"); - let websocketKey = ""; + if (requestData.includes("\r\n\r\n")) { + const lines = requestData.split("\r\n"); + let websocketKey = ""; - for (const line of lines) { - if (line.startsWith("Sec-WebSocket-Key:")) { - websocketKey = line.split(":")[1].trim(); - break; - } + for (const line of lines) { + if (line.startsWith("Sec-WebSocket-Key:")) { + websocketKey = line.split(":")[1].trim(); + break; } - - const acceptKey = crypto - .createHash("sha1") - .update(websocketKey + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11") - .digest("base64"); - - const response = [ - "HTTP/1.1 101 Switching Protocols", - "Upgrade: websocket", - "Connection: Upgrade", - `Sec-WebSocket-Accept: ${acceptKey}`, - ...responseHeaders, - "\r\n", - ].join("\r\n"); - - socket.write(response); } - }); + + const acceptKey = crypto + .createHash("sha1") + .update(websocketKey + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11") + .digest("base64"); + + const response = [ + "HTTP/1.1 101 Switching Protocols", + "Upgrade: websocket", + "Connection: Upgrade", + `Sec-WebSocket-Accept: ${acceptKey}`, + ...responseHeaders, + "\r\n", + ].join("\r\n"); + + socket.write(response); + } }); + }); - return { - port: port!, - [Symbol.asyncDispose]: async () => { - server.close(); - }, - }; + return { + port: port!, + [Symbol.asyncDispose]: async () => { + await new Promise((resolve, reject) => { + server.close(err => (err ? reject(err) : resolve())); + }); + }, + }; +} + +function connect(port: number, protocols: string[] | undefined) { + const url = `ws://127.0.0.1:${port}`; + // `undefined` means "no protocols argument at all", i.e. the default options. + return protocols === undefined ? new WebSocket(url) : new WebSocket(url, protocols); +} + +// A failed handshake must never reach `open`, must surface exactly one +// `error` event, and must close with `wasClean: false`. +async function expectConnectionFailure( + port: number, + protocols: string[] | undefined, + expectedCode = 1002, + expectedReason = "Mismatch client protocol", +) { + const { promise: closePromise, resolve: resolveClose, reject: rejectClose } = Promise.withResolvers(); + + const ws = connect(port, protocols); + const onerrorMock = mock(() => {}); + ws.onopen = () => rejectClose(new Error("handshake unexpectedly succeeded: open event fired")); + ws.onerror = onerrorMock; + ws.onclose = resolveClose; + + try { + const close = await closePromise; + expect(onerrorMock).toHaveBeenCalledTimes(1); + expect({ code: close.code, reason: close.reason, wasClean: close.wasClean }).toEqual({ + code: expectedCode, + reason: expectedReason, + wasClean: false, + }); + } finally { + ws.terminate(); } - - async function expectConnectionFailure(port: number, protocols: string[], expectedCode = 1002) { - const { promise: closePromise, resolve: resolveClose } = Promise.withResolvers(); - - const ws = new WebSocket(`ws://localhost:${port}`, protocols); - const onopenMock = mock(() => {}); - ws.onopen = onopenMock; - - ws.onclose = close => { - expect(close.code).toBe(expectedCode); - expect(close.reason).toBe("Mismatch client protocol"); - resolveClose(); - }; - - await closePromise; - expect(onopenMock).not.toHaveBeenCalled(); +} + +async function expectConnectionSuccess(port: number, protocols: string[] | undefined, expectedProtocol: string) { + const { promise: openPromise, resolve: resolveOpen, reject } = Promise.withResolvers(); + const ws = connect(port, protocols); + try { + ws.onopen = () => resolveOpen(); + ws.onerror = reject; + ws.onclose = e => reject(new Error(`closed: code=${e.code} reason=${e.reason}`)); + await openPromise; + expect(ws.protocol).toBe(expectedProtocol); + } finally { + ws.terminate(); } +} - async function expectConnectionSuccess(port: number, protocols: string[], expectedProtocol: string) { - const { promise: openPromise, resolve: resolveOpen, reject } = Promise.withResolvers(); - const ws = new WebSocket(`ws://localhost:${port}`, protocols); - try { - ws.onopen = () => resolveOpen(); - ws.onerror = reject; - await openPromise; - expect(ws.protocol).toBe(expectedProtocol); - } finally { - ws.terminate(); - } - } +describe("WebSocket strict RFC 6455 subprotocol handling", () => { // Multiple protocols in single header (comma-separated) - should fail it("should reject multiple comma-separated protocols", async () => { await using server = await createTestServer(["Sec-WebSocket-Protocol: chat, echo"]); @@ -156,7 +181,30 @@ describe("WebSocket strict RFC 6455 subprotocol handling", () => { await expectConnectionFailure(server.port, ["chat", "echo"]); }); + // RFC 6455 §4.1 / WHATWG "establish a WebSocket connection" step 4: if + // the client requested subprotocols, a 101 that selects none of them + // must fail the connection. + it("should reject a response with no Sec-WebSocket-Protocol when protocols were requested", async () => { + await using server = await createTestServer([]); + await expectConnectionFailure(server.port, ["chat", "echo"], 1002, "Server sent no subprotocol"); + }); + + it("should reject a response with no Sec-WebSocket-Protocol when a single protocol was requested", async () => { + await using server = await createTestServer([]); + await expectConnectionFailure(server.port, ["chat"], 1002, "Server sent no subprotocol"); + }); + // Valid cases - should succeed + it("should accept a response with no Sec-WebSocket-Protocol when none was requested", async () => { + await using server = await createTestServer([]); + await expectConnectionSuccess(server.port, undefined, ""); + }); + + it("should accept a response with no Sec-WebSocket-Protocol when the protocol list is empty", async () => { + await using server = await createTestServer([]); + await expectConnectionSuccess(server.port, [], ""); + }); + it("should accept single valid protocol (first in client list)", async () => { await using server = await createTestServer(["Sec-WebSocket-Protocol: chat"]); await expectConnectionSuccess(server.port, ["chat", "echo", "binary"], "chat"); @@ -197,18 +245,21 @@ describe("WebSocket strict RFC 6455 subprotocol handling", () => { await using server = await createTestServer([]); const { promise: closePromise, resolve: resolveClose } = Promise.withResolvers(); - const ws = new WebSocket(`ws://localhost:${server.port}`, ["chat", "echo"]); + const ws = new WebSocket(`ws://127.0.0.1:${server.port}`, ["chat", "echo"]); const onopenMock = mock(() => {}); ws.onopen = onopenMock; ws.onclose = close => resolveClose(close); const close = await closePromise; - expect(close.code).toBe(1002); - expect(close.reason).toBe("Missing client protocol"); + expect({ code: close.code, reason: close.reason, wasClean: close.wasClean }).toEqual({ + code: 1002, + reason: "Server sent no subprotocol", + wasClean: false, + }); expect(onopenMock).not.toHaveBeenCalled(); const { promise: openPromise, resolve: resolveOpen, reject } = Promise.withResolvers(); - const bare = new WebSocket(`ws://localhost:${server.port}`); + const bare = new WebSocket(`ws://127.0.0.1:${server.port}`); try { bare.onopen = () => resolveOpen(); bare.onerror = reject; @@ -220,3 +271,71 @@ describe("WebSocket strict RFC 6455 subprotocol handling", () => { } }); }); + +// RFC 6455 §4.1 step 4: a Sec-WebSocket-Extensions response that indicates an +// extension not present in the client's handshake must fail the connection. +// The default client offers only permessage-deflate. +describe("WebSocket strict RFC 6455 extension handling", () => { + it("should reject an extension the client never offered", async () => { + await using server = await createTestServer(["Sec-WebSocket-Extensions: x-bogus-ext"]); + await expectConnectionFailure(server.port, undefined, 1002, "Invalid Sec-WebSocket-Extensions header"); + }); + + it("should reject an unoffered extension listed after permessage-deflate", async () => { + await using server = await createTestServer(["Sec-WebSocket-Extensions: permessage-deflate, x-bogus-ext"]); + await expectConnectionFailure(server.port, undefined, 1002, "Invalid Sec-WebSocket-Extensions header"); + }); + + it("should reject an unoffered extension listed before permessage-deflate", async () => { + await using server = await createTestServer(["Sec-WebSocket-Extensions: x-bogus-ext, permessage-deflate"]); + await expectConnectionFailure(server.port, undefined, 1002, "Invalid Sec-WebSocket-Extensions header"); + }); + + it("should reject an unoffered extension carrying parameters", async () => { + await using server = await createTestServer(["Sec-WebSocket-Extensions: x-bogus-ext; foo=bar"]); + await expectConnectionFailure(server.port, undefined, 1002, "Invalid Sec-WebSocket-Extensions header"); + }); + + // RFC 7692 §5: the response must not list permessage-deflate more than once. + it("should reject a duplicate permessage-deflate in one header", async () => { + await using server = await createTestServer([ + "Sec-WebSocket-Extensions: permessage-deflate; server_max_window_bits=12, permessage-deflate; server_max_window_bits=10", + ]); + await expectConnectionFailure(server.port, undefined, 1002, "Invalid Sec-WebSocket-Extensions header"); + }); + + it("should reject permessage-deflate repeated across two Sec-WebSocket-Extensions headers", async () => { + await using server = await createTestServer([ + "Sec-WebSocket-Extensions: permessage-deflate", + "Sec-WebSocket-Extensions: permessage-deflate", + ]); + await expectConnectionFailure(server.port, undefined, 1002, "Invalid Sec-WebSocket-Extensions header"); + }); + + it("should still accept a plain permessage-deflate response", async () => { + await using server = await createTestServer(["Sec-WebSocket-Extensions: permessage-deflate"]); + const { promise: openPromise, resolve: resolveOpen, reject } = Promise.withResolvers(); + const ws = new WebSocket(`ws://127.0.0.1:${server.port}`); + try { + ws.onopen = () => resolveOpen(); + ws.onerror = reject; + ws.onclose = e => reject(new Error(`closed: code=${e.code} reason=${e.reason}`)); + await openPromise; + expect(ws.extensions).toContain("permessage-deflate"); + } finally { + ws.terminate(); + } + }); + + it("should still accept permessage-deflate with parameters", async () => { + await using server = await createTestServer([ + "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover", + ]); + await expectConnectionSuccess(server.port, undefined, ""); + }); + + it("should ignore empty list elements such as a trailing comma", async () => { + await using server = await createTestServer(["Sec-WebSocket-Extensions: permessage-deflate,"]); + await expectConnectionSuccess(server.port, undefined, ""); + }); +});