fix(desktop): bind headers to scoped WebSocket URL
This commit is contained in:
committed by
Teknium
parent
4ba2608524
commit
b4162de333
@@ -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<string, Record<string, string>>()
|
||||
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<string, string>)
|
||||
|
||||
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.
|
||||
|
||||
@@ -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<typeof createRemoteWsHeaderStore>,
|
||||
url: string,
|
||||
expected: Record<string, string> | 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<typeof createRemoteWsHeaderStore>, 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({})
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,85 @@
|
||||
import { registryGatewayWsUrl } from './plugin-profile-routes'
|
||||
|
||||
export interface RegistryGatewayWsConnection {
|
||||
authMode: string
|
||||
baseUrl: string
|
||||
wsUrl: string
|
||||
headers?: Record<string, string>
|
||||
profile?: null | string
|
||||
sharedRemote?: boolean
|
||||
}
|
||||
|
||||
interface RegistryGatewayWsUrlDependencies {
|
||||
ensureBackend: (connectionId: unknown, profile: unknown) => Promise<RegistryGatewayWsConnection>
|
||||
mintTicket: (baseUrl: string, headers?: Record<string, string>) => Promise<string>
|
||||
buildTicketUrl: (baseUrl: string, ticket: string) => string
|
||||
rememberHeaders: (wsUrl: string, headers?: Record<string, string>) => void
|
||||
}
|
||||
|
||||
interface RemoteRequestDetails {
|
||||
url: string
|
||||
requestHeaders?: Record<string, string>
|
||||
}
|
||||
|
||||
type RemoteRequestCallback = (result: { requestHeaders?: Record<string, string> }) => void
|
||||
|
||||
export function createRemoteWsHeaderStore(limit = 100) {
|
||||
const headersByUrl = new Map<string, Record<string, string>>()
|
||||
|
||||
const remember = (wsUrl: string, headers: Record<string, string> = {}) => {
|
||||
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<string, string> => headersByUrl.get(String(requestUrl)) ?? {}
|
||||
|
||||
return { headersFor, remember }
|
||||
}
|
||||
|
||||
export function applyRemoteRequestHeaders(
|
||||
details: RemoteRequestDetails,
|
||||
callback: RemoteRequestCallback,
|
||||
headersForRequest: (requestUrl: string) => Record<string, string>
|
||||
) {
|
||||
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<string> => {
|
||||
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
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user