fix(desktop): guard stale refresh ownership

This commit is contained in:
Zeus-Deus
2026-08-25 22:44:58 +02:00
committed by Teknium
parent b3bdf0c816
commit 693dd5d042
6 changed files with 151 additions and 12 deletions

View File

@@ -234,7 +234,7 @@ function Harness({
}: {
beforeConnectionSwitch?: () => void
refreshHermesConfig?: (force?: boolean, shouldPublish?: () => boolean) => Promise<void>
refreshSessions?: () => Promise<void>
refreshSessions?: (shouldPublish?: () => boolean) => Promise<void>
} = {}) {
useGatewayBoot({
beforeConnectionSwitch,
@@ -721,6 +721,50 @@ describe('useGatewayBoot remote reconnect loop (real hook, fake socket)', () =>
expect(FakeWebSocket.instances).toHaveLength(socketCount)
})
it('passes switch ownership through a session refresh held across a newer switch', async () => {
const staleRefresh = deferred<void>()
const publications: string[] = []
let switchRefresh = 0
const refreshSessions = vi.fn(async (shouldPublish?: () => boolean) => {
// Initial boot remains a compatible zero-argument caller.
if (!shouldPublish) {
return
}
switchRefresh += 1
const label = switchRefresh === 1 ? 'settings-a' : 'settings-b'
if (switchRefresh === 1) {
await staleRefresh.promise
}
if (shouldPublish()) {
publications.push(label)
}
})
render(<Harness refreshSessions={refreshSessions} />)
await flushAsync()
act(() => connectionApplied?.())
await vi.waitFor(() => expect(switchRefresh).toBe(1))
act(() => connectionApplied?.())
await vi.waitFor(() => expect(switchRefresh).toBe(2))
await flushAsync()
expect(publications).toEqual(['settings-b'])
await act(async () => {
staleRefresh.resolve()
await vi.advanceTimersByTimeAsync(0)
})
expect(publications).toEqual(['settings-b'])
expect($gatewaySwitching.get()).toBe(false)
})
it('a superseded Settings switch cannot publish delayed cwd or config work after the winner', async () => {
const desktop = fakeDesktop() as ReturnType<typeof fakeDesktop> & Record<string, unknown>
const staleSanitize = deferred<{ cwd: string }>()

View File

@@ -135,7 +135,7 @@ interface GatewayBootOptions {
) => void
onGatewayReady: (gateway: HermesGateway | null) => void
refreshHermesConfig: (force?: boolean, shouldPublish?: () => boolean) => Promise<void>
refreshSessions: () => Promise<void>
refreshSessions: (shouldPublish?: () => boolean) => Promise<void>
}
export function useGatewayBoot({
@@ -556,7 +556,7 @@ export function useGatewayBoot({
seedDefaultCwd(ownsSwitch),
refreshActiveProfile().catch(() => undefined),
callbacksRef.current.refreshHermesConfig(false, ownsSwitch).catch(() => undefined),
callbacksRef.current.refreshSessions().catch(() => undefined)
callbacksRef.current.refreshSessions(ownsSwitch).catch(() => undefined)
])
if (!ownsSwitch()) {

View File

@@ -8,12 +8,17 @@ import {
$cronSessions,
$messagingPlatformTotals,
$messagingSessions,
$messagingTruncated,
$sessionProfilesTruncated,
$sessionProfilesUsage,
$sessions,
$sessionsLoading,
setCronSessions,
setMessagingPlatformTotals,
setMessagingSessions,
setMessagingTruncated,
setSessionProfilesTruncated,
setSessionProfilesUsage,
setSessions,
setSessionsLoading
} from '@/store/session'
@@ -104,6 +109,8 @@ beforeEach(() => {
setMessagingSessions([])
setMessagingPlatformTotals({})
setMessagingTruncated(false)
setSessionProfilesTruncated({})
setSessionProfilesUsage({})
setSessionsLoading(false)
})
@@ -114,6 +121,8 @@ afterEach(() => {
setMessagingSessions([])
setMessagingPlatformTotals({})
setMessagingTruncated(false)
setSessionProfilesTruncated({})
setSessionProfilesUsage({})
setSessionsLoading(false)
})
@@ -250,6 +259,49 @@ describe('refreshSessions identity + loading hygiene', () => {
expect(loadingStates).toEqual([false, true, false])
})
it('does not let a superseded owner publish or release a newer switch loading barrier', async () => {
const pending = deferred<SidebarSessionsResponse>()
let ownsRefresh = true
listSidebarSessions.mockReturnValue(pending.promise)
const { result } = renderHook(() => useSessionListActions({ profileScope: 'default' }))
const refresh = result.current.refreshSessions(() => ownsRefresh)
expect($sessionsLoading.get()).toBe(true)
ownsRefresh = false
setSessions([row('winner')])
setCronSessions([row('winner-cron', { source: 'cron' })])
setMessagingSessions([row('winner-message', { source: 'signal' })])
setMessagingTruncated(true)
setSessionProfilesTruncated({ winner: true })
setSessionProfilesUsage({ winner: { cost_usd: 2, tokens: 20 } })
setSessionsLoading(true)
await act(async () => {
pending.resolve({
recents: {
profiles_truncated: { stale: true },
profiles_usage: { stale: { cost_usd: 1, tokens: 10 } },
sessions: [row('stale')]
},
cron: { sessions: [row('stale-cron', { source: 'cron' })] },
messaging: { sessions: [row('stale-message', { source: 'telegram' })] }
})
await refresh
})
expect($sessions.get().map(session => session.id)).toEqual(['winner'])
expect($cronSessions.get().map(session => session.id)).toEqual(['winner-cron'])
expect($messagingSessions.get().map(session => session.id)).toEqual(['winner-message'])
expect($messagingTruncated.get()).toBe(true)
expect($sessionProfilesTruncated.get()).toEqual({ winner: true })
expect($sessionProfilesUsage.get()).toEqual({ winner: { cost_usd: 2, tokens: 20 } })
expect($sessionsLoading.get()).toBe(true)
expect(getCronJobs).not.toHaveBeenCalled()
})
it('clears initial loading after a failed source activation advances the gateway epoch', async () => {
const pending = deferred<SidebarSessionsResponse>()
listSidebarSessions.mockReturnValue(pending.promise)

View File

@@ -224,11 +224,11 @@ export function useSessionListActions({ profileScope }: UseSessionListActionsArg
}, [profileScope])
/** Refresh every sidebar session slice without committing an obsolete profile response. */
const refreshSessions = useCallback(async () => {
const refreshSessions = useCallback(async (shouldPublish: () => boolean = () => true) => {
const sessionProfile = sidebarProfileForScope(profileScope)
const activationEpoch = gatewayActivationEpoch()
if (sidebarProfileForScope(profileScopeRef.current) !== sessionProfile) {
if (!shouldPublish() || sidebarProfileForScope(profileScopeRef.current) !== sessionProfile) {
return
}
@@ -240,7 +240,7 @@ export function useSessionListActions({ profileScope }: UseSessionListActionsArg
// $sessionsLoading subscriber twice per turn for no visible change.
const showLoading = $sessions.get().length === 0
if (showLoading) {
if (showLoading && shouldPublish()) {
setSessionsLoading(true)
}
@@ -270,6 +270,7 @@ export function useSessionListActions({ profileScope }: UseSessionListActionsArg
})
if (
shouldPublish() &&
refreshSessionsRequestRef.current === requestId &&
sidebarProfileForScope(profileScopeRef.current) === sessionProfile &&
gatewayActivationEpoch() === activationEpoch
@@ -332,16 +333,16 @@ export function useSessionListActions({ profileScope }: UseSessionListActionsArg
setMessagingTruncated(result.messaging.sessions.length >= MESSAGING_SECTION_LIMIT)
}
} finally {
// The request id is enough here: a newer refresh owns its own loading
// state, while a failed source activation still needs the old request to
// clear the spinner even though it advanced the gateway epoch.
if (showLoading && refreshSessionsRequestRef.current === requestId) {
// Request identity preserves the zero-argument refresh contract across a
// failed activation epoch; an explicit owner predicate is stronger and
// must never release a newer switch's loading barrier.
if (showLoading && shouldPublish() && refreshSessionsRequestRef.current === requestId) {
setSessionsLoading(false)
}
}
// Cron *jobs* are a distinct API (getCronJobs), not a session slice.
if (sidebarProfileForScope(profileScopeRef.current) === sessionProfile) {
if (shouldPublish() && sidebarProfileForScope(profileScopeRef.current) === sessionProfile) {
void refreshCronJobs()
}
}, [profileScope, refreshCronJobs])

View File

@@ -132,6 +132,48 @@ describe('hermes-ws-recovery-v1 silent blackhole', () => {
vi.useRealTimers()
})
it('keeps a replacement socket open when a superseded connect times out', async () => {
vi.useFakeTimers()
vi.stubGlobal('WebSocket', { OPEN: LoopbackSocket.OPEN })
const backend = new RecoveryBackend()
const sockets: LoopbackSocket[] = []
const states: string[] = []
const client = new JsonRpcGatewayClient({
connectTimeoutMs: 50,
heartbeatIntervalMs: 0,
socketFactory: () => {
const socket = new LoopbackSocket(sockets.length + 1, backend)
sockets.push(socket)
return socket as unknown as WebSocket
}
})
client.onState(state => states.push(state))
const staleConnect = client.connect('ws://gateway.test/a')
const staleRejection = expect(staleConnect).rejects.toThrow('WebSocket connection failed')
client.close()
const currentConnect = client.connect('ws://gateway.test/b')
sockets[1].open()
await currentConnect
expect(client.connectionState).toBe('open')
await vi.advanceTimersByTimeAsync(50)
await staleRejection
expect(client.connectionState).toBe('open')
expect(sockets[1].readyState).toBe(LoopbackSocket.OPEN)
expect(states.slice(states.lastIndexOf('open'))).toEqual(['open'])
client.close()
})
it('recovers the persisted final on a replacement socket without duplicating the prompt', async () => {
vi.useFakeTimers()
vi.stubGlobal('WebSocket', { OPEN: LoopbackSocket.OPEN })

View File

@@ -272,9 +272,9 @@ export class JsonRpcGatewayClient {
}
this.socket = null
this.setState('error')
}
this.setState('error')
reject(new Error(this.options.connectErrorMessage))
}, this.options.connectTimeoutMs)
}