From 91424230c77a01e4007592944840d53665f2bc2d Mon Sep 17 00:00:00 2001 From: Chris Wanstrath Date: Thu, 28 May 2026 14:51:30 -0700 Subject: [PATCH] Refactor WebSocket proxy to use Hono middleware --- src/server/index.tsx | 29 +++++---- src/server/proxy.ts | 146 +++++++++++++++++++------------------------ 2 files changed, 81 insertions(+), 94 deletions(-) diff --git a/src/server/index.tsx b/src/server/index.tsx index 92c8506..4ea0722 100644 --- a/src/server/index.tsx +++ b/src/server/index.tsx @@ -6,13 +6,26 @@ import syncRouter from './api/sync' import systemRouter from './api/system' import { Hype } from '@because/hype' import { cleanupStalePublishers } from './mdns' -import { extractSubdomain, proxySubdomain, proxyWebSocket, websocket } from './proxy' +import { upgradeWebSocket, websocket } from 'hono/bun' +import { extractSubdomain, proxySubdomain, wsProxyEvents } from './proxy' import { Shell } from './shell' -import type { Server } from 'bun' -import type { WsData } from './proxy' const app = new Hype({ layout: false, logging: !!process.env.DEBUG }) +// Subdomain proxy — runs before all Hono routes +app.use('*', async (c, next) => { + const subdomain = extractSubdomain(c.req.header('host') ?? '') + if (!subdomain) return next() + + if (c.req.header('upgrade')?.toLowerCase() === 'websocket') { + const events = wsProxyEvents(subdomain, c.req.raw) + if (!events) return c.text(`App "${subdomain}" not found or not running`, 502) + return upgradeWebSocket(c, events) + } + + return proxySubdomain(subdomain, c.req.raw) +}) + app.route('/api/apps', appsRouter) app.route('/api/events', eventsRouter) app.route('/api/sync', syncRouter) @@ -127,15 +140,5 @@ const defaults = app.defaults export default { ...defaults, maxRequestBodySize: 1024 * 1024 * 50, // 50MB - fetch(req: Request, server: Server) { - const subdomain = extractSubdomain(req.headers.get('host') ?? '') - if (subdomain) { - if (req.headers.get('upgrade')?.toLowerCase() === 'websocket') { - return proxyWebSocket(subdomain, req, server) - } - return proxySubdomain(subdomain, req) - } - return defaults.fetch.call(app, req, server) - }, websocket, } diff --git a/src/server/proxy.ts b/src/server/proxy.ts index c8d7e7a..a2f563b 100644 --- a/src/server/proxy.ts +++ b/src/server/proxy.ts @@ -1,20 +1,11 @@ -import type { Server, ServerWebSocket } from 'bun' +import type { WSContext, WSMessageReceive } from 'hono/ws' import { getAppBySubdomain } from '$apps' import { serveStatic } from '$static' export const perf = { timing: false } -export type { WsData } - -const pendingMessages = new Map, (string | ArrayBuffer | Uint8Array)[]>() -const upstreams = new Map, WebSocket>() - -interface WsData { - port: number - path: string - protocols: string[] - headers: Record -} +const upstreams = new WeakMap() +const pendingMessages = new WeakMap() export function extractSubdomain(host: string): string | null { // Strip port @@ -87,12 +78,10 @@ export async function proxySubdomain(subdomain: string, req: Request): Promise): Response | undefined { +export function wsProxyEvents(subdomain: string, req: Request) { const app = getAppBySubdomain(subdomain) - if (!app || app.state !== 'running' || !app.port) { - return new Response(`App "${subdomain}" not found or not running`, { status: 502 }) - } + if (!app || app.state !== 'running' || !app.port) return null const url = new URL(req.url) const path = url.pathname + url.search @@ -105,82 +94,77 @@ export function proxyWebSocket(subdomain: string, req: Request, server: Server = {} - if (protocolHeader) upgradeHeaders['sec-websocket-protocol'] = protocolHeader + const port = app.port - const ok = server.upgrade(req, { data: { port: app.port, path, protocols, headers: forwardHeaders } as WsData, headers: upgradeHeaders }) - if (ok) return undefined - return new Response('WebSocket upgrade failed', { status: 500 }) -} + return { + onOpen(_evt: Event, ws: WSContext) { + const upstream = new WebSocket(`ws://localhost:${port}${path}`, { + headers: { ...forwardHeaders, host: `localhost:${port}` }, + protocols, + }) -export const websocket = { - open(ws: ServerWebSocket) { - const { port, path } = ws.data - const upstream = new WebSocket(`ws://localhost:${port}${path}`, { - headers: { ...ws.data.headers, host: `localhost:${port}` }, - protocols: ws.data.protocols, - }) + upstream.binaryType = 'arraybuffer' + upstreams.set(ws, upstream) + pendingMessages.set(ws, []) - upstream.binaryType = 'arraybuffer' - upstreams.set(ws, upstream) - pendingMessages.set(ws, []) + const timeout = setTimeout(() => { + if (upstream.readyState !== WebSocket.OPEN) { + upstream.close() + ws.close() + } + }, 10_000) - const timeout = setTimeout(() => { - if (upstream.readyState !== WebSocket.OPEN) { - upstream.close() - ws.close() - } - }, 10_000) + upstream.addEventListener('open', () => { + clearTimeout(timeout) + const buffered = pendingMessages.get(ws) + if (buffered) { + for (const msg of buffered) upstream.send(msg) + pendingMessages.delete(ws) + } + }) - upstream.addEventListener('open', () => { - clearTimeout(timeout) - const buffered = pendingMessages.get(ws) - if (buffered) { - for (const msg of buffered) upstream.send(msg) + upstream.addEventListener('message', e => { + ws.send(e.data as string | ArrayBuffer) + }) + + upstream.addEventListener('close', () => { + clearTimeout(timeout) pendingMessages.delete(ws) + upstreams.delete(ws) + ws.close() + }) + + upstream.addEventListener('error', () => { + clearTimeout(timeout) + pendingMessages.delete(ws) + upstreams.delete(ws) + ws.close() + }) + }, + + onMessage(evt: MessageEvent, ws: WSContext) { + const upstream = upstreams.get(ws) + if (!upstream) return + if (upstream.readyState !== WebSocket.OPEN) { + const msg = typeof evt.data === 'string' ? evt.data : evt.data as ArrayBuffer + pendingMessages.get(ws)?.push(msg) + return } - }) + upstream.send(typeof evt.data === 'string' ? evt.data : evt.data as ArrayBuffer) + }, - upstream.addEventListener('message', e => { - // binaryType is 'arraybuffer' so data is always string | ArrayBuffer - ws.send(e.data as string | ArrayBuffer) - }) - - upstream.addEventListener('close', () => { - clearTimeout(timeout) + onClose(_evt: CloseEvent, ws: WSContext) { + const upstream = upstreams.get(ws) + if (upstream) { + upstream.close() + upstreams.delete(ws) + } pendingMessages.delete(ws) - upstreams.delete(ws) - ws.close() - }) - - upstream.addEventListener('error', () => { - clearTimeout(timeout) - pendingMessages.delete(ws) - upstreams.delete(ws) - ws.close() - }) - }, - - message(ws: ServerWebSocket, msg: string | ArrayBuffer | Uint8Array) { - const upstream = upstreams.get(ws) - if (!upstream) return - if (upstream.readyState !== WebSocket.OPEN) { - pendingMessages.get(ws)?.push(msg) - return - } - upstream.send(msg) - }, - - close(ws: ServerWebSocket) { - const upstream = upstreams.get(ws) - if (upstream) { - upstream.close() - upstreams.delete(ws) - } - pendingMessages.delete(ws) - }, + }, + } }