From b4162de3339bee8caafc48aec397a06ec08a41e5 Mon Sep 17 00:00:00 2001 From: Marco Fernstaedt Date: Mon, 17 Aug 2026 17:40:29 +0000 Subject: [PATCH] fix(desktop): bind headers to scoped WebSocket URL --- apps/desktop/electron/main.ts | 61 +++----- .../electron/remote-ws-headers.test.ts | 133 ++++++++++++++++++ apps/desktop/electron/remote-ws-headers.ts | 85 +++++++++++ 3 files changed, 235 insertions(+), 44 deletions(-) create mode 100644 apps/desktop/electron/remote-ws-headers.test.ts create mode 100644 apps/desktop/electron/remote-ws-headers.ts diff --git a/apps/desktop/electron/main.ts b/apps/desktop/electron/main.ts index 5f7c598679..06b2eae673 100644 --- a/apps/desktop/electron/main.ts +++ b/apps/desktop/electron/main.ts @@ -292,6 +292,11 @@ import { revalidatePooledRemoteBackends, revalidateRemoteConnection } from './remote-liveness' +import { + applyRemoteRequestHeaders, + createRegistryGatewayWsUrlHandler, + createRemoteWsHeaderStore +} from './remote-ws-headers' import { missingRendererAssets } from './renderer-bundle' import { attachRendererConsoleCapture, formatRendererBoundaryReport } from './renderer-log' import { @@ -1405,7 +1410,7 @@ let connectionConfigCacheMtime = null let connectionRegistryCache = null let connectionRegistryCacheMtime = null let remoteHeaderRulesInstalled = false -const remoteWsHeadersByUrl = new Map>() +const remoteWsHeaderStore = createRemoteWsHeaderStore() const hermesLog = [] const previewWatchers = new Map() let previewShortcutActive = false @@ -8686,25 +8691,11 @@ function encryptIncomingRemoteHeaders(raw, existing, options: { allowPlainText?: } function rememberRemoteWsHeaders(wsUrl, headers = {}) { - if (!wsUrl || Object.keys(headers).length === 0) { - return - } - - remoteWsHeadersByUrl.set(String(wsUrl), headers as Record) - - while (remoteWsHeadersByUrl.size > 100) { - const oldest = remoteWsHeadersByUrl.keys().next().value - - if (!oldest) { - break - } - - remoteWsHeadersByUrl.delete(oldest) - } + remoteWsHeaderStore.remember(wsUrl, headers) } function headersForRemoteRequest(requestUrl) { - const exactWsHeaders = remoteWsHeadersByUrl.get(String(requestUrl)) + const exactWsHeaders = remoteWsHeaderStore.headersFor(requestUrl) if (exactWsHeaders && Object.keys(exactWsHeaders).length > 0) { return exactWsHeaders @@ -8730,15 +8721,7 @@ function installRemoteHeaderRules() { remoteHeaderRulesInstalled = true session.defaultSession.webRequest.onBeforeSendHeaders((details, callback) => { - const headers = headersForRemoteRequest(details.url) - - if (Object.keys(headers).length === 0) { - callback({}) - - return - } - - callback({ requestHeaders: { ...details.requestHeaders, ...headers } }) + applyRemoteRequestHeaders(details, callback, headersForRemoteRequest) }) } @@ -13850,25 +13833,15 @@ ipcMain.handle('hermes:agents:roster', async () => { // Registry-scoped fresh WS URL: the (connectionId, profile) analogue of // hermes:gateway:ws-url. Same single-use-ticket discipline for OAuth sources. +const registryGatewayWsUrlHandler = createRegistryGatewayWsUrlHandler({ + ensureBackend: ensureRegistryBackend, + mintTicket: mintGatewayWsTicket, + buildTicketUrl: buildGatewayWsUrlWithTicket, + rememberHeaders: rememberRemoteWsHeaders +}) + ipcMain.handle('hermes:gateway:ws-url-for', async (_event, payload) => { - const { connectionId, profile } = payload && typeof payload === 'object' ? (payload as any) : ({} as any) - - return gatewayWsUrlIpcResult(async () => { - const connection: any = await ensureRegistryBackend(connectionId, profile) - - if (connection.authMode === 'oauth') { - const ticket = await mintGatewayWsTicket(connection.baseUrl, connection.headers) - const wsUrl = buildGatewayWsUrlWithTicket(connection.baseUrl, ticket) - - rememberRemoteWsHeaders(wsUrl, connection.headers) - - return registryGatewayWsUrl(connection, wsUrl) - } - - rememberRemoteWsHeaders(connection.wsUrl, connection.headers) - - return registryGatewayWsUrl(connection, connection.wsUrl) - }) + return gatewayWsUrlIpcResult(() => registryGatewayWsUrlHandler(payload)) }) // Fan out `hermes update` to every eligible registered connection at once. diff --git a/apps/desktop/electron/remote-ws-headers.test.ts b/apps/desktop/electron/remote-ws-headers.test.ts new file mode 100644 index 0000000000..19806ad901 --- /dev/null +++ b/apps/desktop/electron/remote-ws-headers.test.ts @@ -0,0 +1,133 @@ +import { describe, expect, it, vi } from 'vitest' + +import { + applyRemoteRequestHeaders, + createRegistryGatewayWsUrlHandler, + createRemoteWsHeaderStore, + type RegistryGatewayWsConnection +} from './remote-ws-headers' + +const accessHeaders = { + 'CF-Access-Client-Id': 'client-id', + 'CF-Access-Client-Secret': 'client-secret' +} + +function createHarness(connection: RegistryGatewayWsConnection) { + const store = createRemoteWsHeaderStore() + const ensureBackend = vi.fn(async () => connection) + const mintTicket = vi.fn(async () => 'fresh-ticket') + + const handler = createRegistryGatewayWsUrlHandler({ + ensureBackend, + mintTicket, + buildTicketUrl: baseUrl => `${baseUrl.replace(/^https:/, 'wss:')}/api/ws?region=us&ticket=fresh-ticket&profile=old`, + rememberHeaders: store.remember + }) + + return { ensureBackend, handler, mintTicket, store } +} + +function expectRequestHeaders( + store: ReturnType, + url: string, + expected: Record | undefined +) { + const callback = vi.fn() + + applyRemoteRequestHeaders({ url, requestHeaders: { Origin: 'app://hermes' } }, callback, store.headersFor) + + expect(callback).toHaveBeenCalledOnce() + expect(callback).toHaveBeenCalledWith(expected ? { requestHeaders: { Origin: 'app://hermes', ...expected } } : {}) +} + +function expectNoHeadersForNearbyUrls(store: ReturnType, exactUrl: string) { + const exact = new URL(exactUrl) + const unscoped = new URL(exact) + unscoped.searchParams.delete('profile') + const sibling = new URL(exact) + sibling.pathname = '/api/ws/sibling' + const otherProfile = new URL(exact) + otherProfile.searchParams.set('profile', 'analysis') + const otherCredential = new URL(exact) + + if (otherCredential.searchParams.has('ticket')) { + otherCredential.searchParams.set('ticket', 'other-ticket') + } else { + otherCredential.searchParams.set('token', 'other-token') + } + + const reordered = new URL(exact) + const entries = [...reordered.searchParams.entries()].reverse() + reordered.search = '' + + for (const [name, value] of entries) { + reordered.searchParams.append(name, value) + } + + for (const url of [unscoped, sibling, otherProfile, otherCredential, reordered]) { + expect(store.headersFor(url.toString())).toEqual({}) + expectRequestHeaders(store, url.toString(), undefined) + } +} + +describe('registry gateway WebSocket headers', () => { + it('token path binds headers to the exact profile scoped URL', async () => { + const { ensureBackend, handler, mintTicket, store } = createHarness({ + authMode: 'token', + baseUrl: 'https://gateway.example', + wsUrl: 'wss://gateway.example/api/ws?token=secret&trace=one&profile=old', + headers: accessHeaders, + profile: 'research', + sharedRemote: true + }) + + const result = await handler({ connectionId: 'remote-one', profile: 'research' }) + const expectedUrl = 'wss://gateway.example/api/ws?token=secret&trace=one&profile=research' + + expect(result).toBe(expectedUrl) + expect(ensureBackend).toHaveBeenCalledWith('remote-one', 'research') + expect(mintTicket).not.toHaveBeenCalled() + expect(store.headersFor(result)).toEqual(accessHeaders) + expectRequestHeaders(store, result, accessHeaders) + expectNoHeadersForNearbyUrls(store, result) + }) + + it('OAuth path binds headers to the exact fresh profile scoped URL', async () => { + const { handler, mintTicket, store } = createHarness({ + authMode: 'oauth', + baseUrl: 'https://gateway.example', + wsUrl: 'wss://gateway.example/api/ws?ticket=stale', + headers: accessHeaders, + profile: 'research', + sharedRemote: true + }) + + const result = await handler({ connectionId: 'cloud-one', profile: 'research' }) + const expectedUrl = 'wss://gateway.example/api/ws?region=us&ticket=fresh-ticket&profile=research' + + expect(result).toBe(expectedUrl) + expect(mintTicket).toHaveBeenCalledOnce() + expect(mintTicket).toHaveBeenCalledWith('https://gateway.example', accessHeaders) + expect(store.headersFor(result)).toEqual(accessHeaders) + expectRequestHeaders(store, result, accessHeaders) + expectNoHeadersForNearbyUrls(store, result) + }) + + it('sharedRemote false preserves the original URL and exact header behavior', async () => { + const { handler, store } = createHarness({ + authMode: 'token', + baseUrl: 'https://gateway.example', + wsUrl: 'wss://gateway.example/api/ws?trace=one&token=secret', + headers: accessHeaders, + profile: 'research', + sharedRemote: false + }) + + const result = await handler({ connectionId: 'remote-one', profile: 'research' }) + + expect(result).toBe('wss://gateway.example/api/ws?trace=one&token=secret') + expect(store.headersFor(result)).toEqual(accessHeaders) + expectRequestHeaders(store, result, accessHeaders) + expect(store.headersFor('wss://gateway.example/api/ws?token=secret&trace=one')).toEqual({}) + }) +}) diff --git a/apps/desktop/electron/remote-ws-headers.ts b/apps/desktop/electron/remote-ws-headers.ts new file mode 100644 index 0000000000..77fe935dd6 --- /dev/null +++ b/apps/desktop/electron/remote-ws-headers.ts @@ -0,0 +1,85 @@ +import { registryGatewayWsUrl } from './plugin-profile-routes' + +export interface RegistryGatewayWsConnection { + authMode: string + baseUrl: string + wsUrl: string + headers?: Record + profile?: null | string + sharedRemote?: boolean +} + +interface RegistryGatewayWsUrlDependencies { + ensureBackend: (connectionId: unknown, profile: unknown) => Promise + mintTicket: (baseUrl: string, headers?: Record) => Promise + buildTicketUrl: (baseUrl: string, ticket: string) => string + rememberHeaders: (wsUrl: string, headers?: Record) => void +} + +interface RemoteRequestDetails { + url: string + requestHeaders?: Record +} + +type RemoteRequestCallback = (result: { requestHeaders?: Record }) => void + +export function createRemoteWsHeaderStore(limit = 100) { + const headersByUrl = new Map>() + + const remember = (wsUrl: string, headers: Record = {}) => { + if (!wsUrl || Object.keys(headers).length === 0) { + return + } + + headersByUrl.set(String(wsUrl), headers) + + while (headersByUrl.size > limit) { + const oldest = headersByUrl.keys().next().value + + if (!oldest) { + break + } + + headersByUrl.delete(oldest) + } + } + + const headersFor = (requestUrl: string): Record => headersByUrl.get(String(requestUrl)) ?? {} + + return { headersFor, remember } +} + +export function applyRemoteRequestHeaders( + details: RemoteRequestDetails, + callback: RemoteRequestCallback, + headersForRequest: (requestUrl: string) => Record +) { + const headers = headersForRequest(details.url) + + if (Object.keys(headers).length === 0) { + callback({}) + + return + } + + callback({ requestHeaders: { ...details.requestHeaders, ...headers } }) +} + +export function createRegistryGatewayWsUrlHandler(dependencies: RegistryGatewayWsUrlDependencies) { + return async (payload: unknown): Promise => { + const { connectionId, profile } = payload && typeof payload === 'object' ? (payload as any) : ({} as any) + const connection = await dependencies.ensureBackend(connectionId, profile) + let wsUrl = connection.wsUrl + + if (connection.authMode === 'oauth') { + const ticket = await dependencies.mintTicket(connection.baseUrl, connection.headers) + wsUrl = dependencies.buildTicketUrl(connection.baseUrl, ticket) + } + + const finalWsUrl = registryGatewayWsUrl(connection, wsUrl) + + dependencies.rememberHeaders(finalWsUrl, connection.headers) + + return finalWsUrl + } +}