From 23b7310d2ce7cfd0a727fbe32349fb6b76fa620d Mon Sep 17 00:00:00 2001 From: Mayank Mehra Date: Fri, 21 Aug 2026 09:07:16 -0700 Subject: [PATCH] Guard consumed rehandshake sessions --- __tests__/e2e.test.ts | 80 +++++++++++++++++++++++++++++++++++++++++++ transport/server.ts | 14 +++++--- 2 files changed, 89 insertions(+), 5 deletions(-) diff --git a/__tests__/e2e.test.ts b/__tests__/e2e.test.ts index be360d73..2a59000f 100644 --- a/__tests__/e2e.test.ts +++ b/__tests__/e2e.test.ts @@ -44,6 +44,8 @@ import { createServerHandshakeOptions, } from '../router/handshake'; import { RehandshakeStreamId } from '../transport/message'; +import { createPromiseWithResolvers } from '../transport/promises'; +import { SessionState } from '../transport/sessionStateMachine'; import { TestSetupHelpers } from '../testUtil/fixtures/transports'; describe.each(testMatrix())( @@ -1544,6 +1546,84 @@ describe.each(testMatrix())( await advanceFakeTimersBySessionGrace(); }); + test('a rejected re-handshake after disconnect leaves cleanup to the replacement session', async () => { + const requestSchema = Type.Object({ token: Type.String() }); + + type ParsedMetadata = Static; + + let token = 'token-v1'; + const refreshValidation = createPromiseWithResolvers< + ParsedMetadata | 'REJECTED_BY_CUSTOM_HANDLER' + >(); + const refreshValidated = vi.fn(); + const construct = vi.fn(() => ({ token })); + const clientTransport = getClientTransport( + 'client', + createClientHandshakeOptions(requestSchema, construct), + ); + const validate = vi.fn( + async ( + metadata: ParsedMetadata, + ): Promise => { + if (metadata.token === 'token-v1') { + return { token: metadata.token }; + } + + const result = await refreshValidation.promise; + refreshValidated(); + + return result; + }, + ); + const serverTransport = getServerTransport< + typeof requestSchema, + ParsedMetadata + >( + 'SERVER', + createServerHandshakeOptions( + requestSchema, + validate, + ), + ); + addPostTestCleanup(async () => { + await cleanupTransports([clientTransport, serverTransport]); + }); + + const protocolError = vi.fn(); + serverTransport.addEventListener('protocolError', protocolError); + clientTransport.connect(serverTransport.clientId); + await waitFor(() => + expect(serverTransport.sessions.get('client')?.state).toBe( + SessionState.Connected, + ), + ); + + clientTransport.reconnectOnConnectionDrop = false; + token = 'token-v2'; + expect(serverTransport.requestRehandshake('client')).toBe(true); + await waitFor(() => expect(validate).toHaveBeenCalledTimes(2)); + + closeAllConnections(clientTransport); + await waitFor(() => + expect(serverTransport.sessions.get('client')?.state).toBe( + SessionState.NoConnection, + ), + ); + + refreshValidation.resolve('REJECTED_BY_CUSTOM_HANDLER'); + await waitFor(() => expect(refreshValidated).toHaveBeenCalledOnce()); + + expect(protocolError).not.toHaveBeenCalled(); + expect(serverTransport.sessions.get('client')?.state).toBe( + SessionState.NoConnection, + ); + + await advanceFakeTimersBySessionGrace(); + await waitFor(() => + expect(serverTransport.sessions.has('client')).toBe(false), + ); + }); + test('an in-flight handler observes refreshed metadata mid-stream', async () => { const requestSchema = Type.Object({ token: Type.String() }); diff --git a/transport/server.ts b/transport/server.ts index ef42d22d..c41e43ce 100644 --- a/transport/server.ts +++ b/transport/server.ts @@ -238,20 +238,24 @@ export abstract class ServerTransport< /** * Tears down a session whose re-handshake failed (rejected, malformed, timed - * out, or a thrown validator). No-ops if {@link session} is no longer the live - * session for its peer — a transparent reconnect keeps the same id, so callers - * reaching here after an async gap can't accidentally close the session that - * replaced it. + * out, or a thrown validator). No-ops if {@link session} has been consumed or + * is no longer the live session for its peer — a transparent reconnect keeps + * the same id, so callers reaching here after an async gap can't accidentally + * close the session that replaced it. */ private teardownForFailedRehandshake( session: ServerSession, reason: string, ) { - if (this.sessions.get(session.to) !== session) { + if (session._isConsumed) { return; } const to = session.to; + if (this.sessions.get(to) !== session) { + return; + } + this.log?.warn(`tearing down session to ${to}: ${reason}`, { ...session.loggingMetadata, connectedTo: to,