import { createHash } from 'crypto' import { request as requestHttp, type IncomingMessage } from 'http' import { request as requestHttps } from 'https' import type { Duplex } from 'stream' export const WEBSOCKET_PROXY_PATH = '/__websocket-proxy' export const WEBSOCKET_PROXY_HANDSHAKE_TIMEOUT_MS = 10_000 const WEBSOCKET_PROTOCOL_TOKEN = /^[!#$%&'*+\-.^_`|~0-9A-Za-z]+$/ const WEBSOCKET_KEY = /^[+/0-9A-Za-z]{22}==$/ type UpgradeRejection = { status: number message: string headers?: Record } const STRIP_UPSTREAM_HEADERS = new Set([ 'host', 'connection', 'upgrade', 'transfer-encoding', 'content-length', 'sec-websocket-accept', 'sec-websocket-extensions', 'sec-websocket-key', 'sec-websocket-protocol', 'sec-websocket-version', ]) function rejectUpgrade( socket: Duplex, status: number, message: string, headers: Record = {}, ) { try { const extraHeaders = Object.entries(headers) .map(([name, value]) => `${name}: ${value}\r\n`) .join('') socket.write( `HTTP/1.1 ${status} ${message}\r\nConnection: close\r\nContent-Type: text/plain\r\nContent-Length: ${message.length}\r\n${extraHeaders}\r\n${message}`, ) } catch {} socket.destroy() } function hasHeaderToken(value: string | undefined, token: string): boolean { return (value || '').split(',').some((part) => part.trim().toLowerCase() === token) } function hasValidProtocolOffer(value: string | undefined): boolean { if (value === undefined) return true if (!WEBSOCKET_PROTOCOL_TOKEN.test(value[0] || '')) return false if (!WEBSOCKET_PROTOCOL_TOKEN.test(value[value.length - 1] || '')) return false const protocols = new Set() for (const part of value.split(',')) { const protocol = part.trim() if (!WEBSOCKET_PROTOCOL_TOKEN.test(protocol) || protocols.has(protocol)) { return false } protocols.add(protocol) } return true } function validateUpgradeRequest(req: IncomingMessage): UpgradeRejection | null { if (req.method !== 'GET') { return { status: 405, message: 'invalid websocket http method' } } if ( !hasHeaderToken(req.headers.connection, 'upgrade') || !hasHeaderToken(req.headers.upgrade, 'websocket') ) { return { status: 400, message: 'invalid websocket upgrade headers' } } const key = req.headers['sec-websocket-key'] if ( typeof key !== 'string' || !WEBSOCKET_KEY.test(key) || Buffer.from(key, 'base64').byteLength !== 16 || Buffer.from(key, 'base64').toString('base64') !== key ) { return { status: 400, message: 'invalid websocket key' } } if (req.headers['sec-websocket-version'] !== '13') { return { status: 400, message: 'invalid websocket version', headers: { 'Sec-WebSocket-Version': '13' }, } } if (!hasValidProtocolOffer(req.headers['sec-websocket-protocol'])) { return { status: 400, message: 'invalid websocket protocol offer' } } return null } function isSameOriginUpgrade(req: IncomingMessage): boolean { const origin = req.headers.origin const host = req.headers.host if (!origin || !host) return false try { return new URL(origin).host === host } catch { return false } } function decodeProxyHeaders(encoded: string | null): Record { if (!encoded) return {} const base64 = encoded.replace(/-/g, '+').replace(/_/g, '/') const padded = base64 + '='.repeat((4 - (base64.length % 4)) % 4) const parsed = JSON.parse(Buffer.from(padded, 'base64').toString('utf8')) if (!parsed || typeof parsed !== 'object' || Array.isArray(parsed)) return {} const headers: Record = {} for (const [key, value] of Object.entries(parsed)) { if (value == null) continue if (STRIP_UPSTREAM_HEADERS.has(key.toLowerCase())) continue headers[key] = Array.isArray(value) ? value.join(', ') : String(value) } return headers } function getDefaultWebSocketOrigin(targetUrl: URL): string { const origin = new URL(targetUrl.href) origin.protocol = targetUrl.protocol === 'wss:' ? 'https:' : 'http:' return origin.origin } function createUpstreamHeaders( req: IncomingMessage, targetUrl: URL, headers: Record, forwardOrigin: boolean, ): Record { const upstreamHeaders = { ...headers } const originKeys = Object.keys(upstreamHeaders).filter( (key) => key.toLowerCase() === 'origin', ) if (forwardOrigin && originKeys.length === 0) { upstreamHeaders.origin = getDefaultWebSocketOrigin(targetUrl) } else if (!forwardOrigin) { for (const key of originKeys) delete upstreamHeaders[key] } upstreamHeaders.connection = 'Upgrade' upstreamHeaders.upgrade = 'websocket' const key = req.headers['sec-websocket-key'] if (typeof key === 'string') upstreamHeaders['sec-websocket-key'] = key const version = req.headers['sec-websocket-version'] if (typeof version === 'string') upstreamHeaders['sec-websocket-version'] = version const protocol = req.headers['sec-websocket-protocol'] if (typeof protocol === 'string') upstreamHeaders['sec-websocket-protocol'] = protocol const extensions = req.headers['sec-websocket-extensions'] if (typeof extensions === 'string') { upstreamHeaders['sec-websocket-extensions'] = extensions } return upstreamHeaders } export function isWebSocketProxyRequestUrl(rawUrl: string | undefined): boolean { if (!rawUrl) return false try { return new URL(rawUrl, 'http://localhost').pathname === WEBSOCKET_PROXY_PATH } catch { return false } } export function handleWebSocketProxyUpgrade( req: IncomingMessage, socket: Duplex, head: Buffer, forwardOrigin = true, ): boolean { if (!isWebSocketProxyRequestUrl(req.url)) return false if (!isSameOriginUpgrade(req)) { rejectUpgrade(socket, 403, 'forbidden websocket proxy origin') return true } const upgradeRejection = validateUpgradeRequest(req) if (upgradeRejection) { rejectUpgrade( socket, upgradeRejection.status, upgradeRejection.message, upgradeRejection.headers, ) return true } let targetUrl: URL let headers: Record try { const requestUrl = new URL(req.url || '/', 'http://localhost') const target = requestUrl.searchParams.get('url') if (!target) { rejectUpgrade(socket, 400, 'missing websocket proxy url') return true } targetUrl = new URL(target) if (targetUrl.protocol !== 'ws:' && targetUrl.protocol !== 'wss:') { rejectUpgrade(socket, 400, 'invalid websocket proxy protocol') return true } headers = decodeProxyHeaders(requestUrl.searchParams.get('headers')) } catch { rejectUpgrade(socket, 400, 'invalid websocket proxy request') return true } const upstreamHeaders = createUpstreamHeaders(req, targetUrl, headers, forwardOrigin) let settled = false const request = targetUrl.protocol === 'wss:' ? requestHttps : requestHttp const upstreamRequestUrl = new URL(targetUrl) upstreamRequestUrl.protocol = targetUrl.protocol === 'wss:' ? 'https:' : 'http:' const upstreamRequest = request(upstreamRequestUrl, { method: 'GET', headers: upstreamHeaders, }) let destroyUpstreamSocket = () => {} upstreamRequest.once('socket', (upstreamSocket) => { destroyUpstreamSocket = () => upstreamSocket.destroy() }) const fail = (status: number, message: string) => { if (settled) return settled = true clearTimeout(timer) upstreamRequest.destroy() destroyUpstreamSocket() rejectUpgrade(socket, status, message) } const timer = setTimeout(() => { fail(504, 'upstream websocket handshake timeout') }, WEBSOCKET_PROXY_HANDSHAKE_TIMEOUT_MS) socket.once('close', () => { if (!settled) { settled = true clearTimeout(timer) } upstreamRequest.destroy() destroyUpstreamSocket() }) upstreamRequest.once('response', (response) => { response.resume() fail(502, 'upstream websocket rejected upgrade') }) upstreamRequest.once('error', () => { fail(502, 'upstream websocket error') }) upstreamRequest.once('upgrade', (response, upstreamSocket, upstreamHead) => { if (settled) { upstreamSocket.destroy() return } const accept = response.headers['sec-websocket-accept'] const upgrade = response.headers.upgrade const requestKey = req.headers['sec-websocket-key'] const expectedAccept = typeof requestKey === 'string' ? createHash('sha1') .update(`${requestKey}258EAFA5-E914-47DA-95CA-C5AB0DC85B11`) .digest('base64') : '' if ( typeof accept !== 'string' || accept !== expectedAccept || typeof upgrade !== 'string' || upgrade.toLowerCase() !== 'websocket' ) { upstreamSocket.destroy() fail(502, 'upstream websocket handshake is invalid') return } const selectedProtocol = response.headers['sec-websocket-protocol'] const requestedProtocols = (req.headers['sec-websocket-protocol'] || '') .split(',') .map((protocol) => protocol.trim()) if ( typeof selectedProtocol === 'string' && !requestedProtocols.includes(selectedProtocol) ) { upstreamSocket.destroy() fail(502, 'upstream websocket selected an unrequested protocol') return } settled = true clearTimeout(timer) const responseHeaders = [ 'HTTP/1.1 101 Switching Protocols', 'Upgrade: websocket', 'Connection: Upgrade', `Sec-WebSocket-Accept: ${accept}`, ] if (typeof selectedProtocol === 'string') { responseHeaders.push(`Sec-WebSocket-Protocol: ${selectedProtocol}`) } const selectedExtensions = response.headers['sec-websocket-extensions'] if (typeof selectedExtensions === 'string') { responseHeaders.push(`Sec-WebSocket-Extensions: ${selectedExtensions}`) } // the upstream and browser used the same websocket key and extension offer, // so their negotiated byte streams can be relayed without decoding frames. socket.write(`${responseHeaders.join('\r\n')}\r\n\r\n`) if (head.length > 0) upstreamSocket.write(head) if (upstreamHead.length > 0) socket.write(upstreamHead) socket.once('close', () => upstreamSocket.destroy()) socket.once('error', () => upstreamSocket.destroy()) upstreamSocket.once('error', () => socket.destroy()) socket.pipe(upstreamSocket) upstreamSocket.pipe(socket) }) upstreamRequest.end() return true }