diff --git a/__tests__/properties/session.property.test.ts b/__tests__/properties/session.property.test.ts index b0a8df7a..5e3a8496 100644 --- a/__tests__/properties/session.property.test.ts +++ b/__tests__/properties/session.property.test.ts @@ -15,6 +15,9 @@ import { } from '../../router/handshake'; import { createClient } from '../../router/client'; import { createServer } from '../../router/server'; +import type { ClientTransport } from '../../transport/client'; +import type { Connection } from '../../transport/connection'; +import type { ServerTransport } from '../../transport/server'; import { closeAllConnections, numberOfConnections } from '../../testUtil'; import { createMockTransportNetwork } from '../../testUtil/fixtures/mockTransport'; import type { TestTransportOptions } from '../../testUtil/fixtures/transports'; @@ -152,12 +155,8 @@ const multiplexedSchedules: gs.Generator = gs.composite( function setup(opts?: TestTransportOptions): { network: ReturnType; - clientTransport: ReturnType< - ReturnType['getClientTransport'] - >; - serverTransport: ReturnType< - ReturnType['getServerTransport'] - >; + clientTransport: ClientTransport; + serverTransport: ServerTransport; client: ReturnType>; violations: Array; } { @@ -536,11 +535,8 @@ describe('re-handshake under faults', () => { throw new Error(`timed out waiting for ${what}`); } - const isConnected = ( - transport: ReturnType< - ReturnType['getClientTransport'] - >, - ) => numberOfConnections(transport) === 1; + const isConnected = (transport: ClientTransport) => + numberOfConnections(transport) === 1; const handshakeSchema = Type.Object({ token: Type.String() }); diff --git a/__tests__/protobuf.test.ts b/__tests__/protobuf.test.ts index 6ca580a6..12e31c38 100644 --- a/__tests__/protobuf.test.ts +++ b/__tests__/protobuf.test.ts @@ -1,5 +1,6 @@ /* eslint-disable @typescript-eslint/no-unsafe-member-access, @typescript-eslint/no-unsafe-call, @typescript-eslint/no-unsafe-assignment, @typescript-eslint/no-unsafe-return, @typescript-eslint/no-unsafe-argument */ import { beforeEach, describe, expect, test, vi } from 'vitest'; +import { Type } from 'typebox'; import { BinaryCodec, NaiveJsonCodec } from '../codec'; import { type ClientError, @@ -298,25 +299,29 @@ describe.each(protobufRouterMatrix)( }); test('protobuf handshake metadata is decoded through the router helpers', async () => { + const rejectionCodeSchema = Type.Union([Type.Literal('TOKEN_EXPIRED')]); const clientHandshakeOptions = createClientHandshakeOptions( AuthHandshakeSchema, () => ({ token: 'let-me-in' }), + undefined, + rejectionCodeSchema, ); const serverHandshakeOptions = createServerHandshakeOptions( AuthHandshakeSchema, (metadata) => ({ token: metadata.token, }), + undefined, + rejectionCodeSchema, ); - const clientTransport = getClientTransport( - 'client', - clientHandshakeOptions, - ); - const serverTransport = getServerTransport( - 'SERVER', - serverHandshakeOptions, - ); + const clientTransport = + getClientTransport('client'); + const serverTransport = getServerTransport< + (typeof serverHandshakeOptions)['schema'], + { token: string }, + typeof rejectionCodeSchema + >('SERVER', undefined); const TypedProtoService = createProtoService(); const testSvc = TypedProtoService.define(TestService, { echo: (request, ctx) => @@ -335,6 +340,7 @@ describe.each(protobufRouterMatrix)( TestService, clientTransport, serverTransport.clientId, + { handshakeOptions: clientHandshakeOptions }, ); await expect(client.echo({ text: 'hello' })).resolves.toMatchObject({ diff --git a/protobuf/client.ts b/protobuf/client.ts index 9984759e..00947531 100644 --- a/protobuf/client.ts +++ b/protobuf/client.ts @@ -6,6 +6,7 @@ import type { MessageInitShape, MessageShape, } from '@bufbuild/protobuf'; +import type { TSchema } from 'typebox'; import { Value } from 'typebox/value'; import { ClientTransport } from '../transport/client'; import { Connection } from '../transport/connection'; @@ -13,6 +14,7 @@ import { EventMap } from '../transport/events'; import { ControlFlags, ControlMessageCloseSchema, + type CustomHandshakeErrorCodeSchema, OpaqueTransportMessage, TransportClientId, cancelMessage, @@ -75,13 +77,16 @@ interface StartedMethodCall< /** * Creates a protobuf client for a single protobuf service descriptor. */ -export function createClient( +export function createClient< + Service extends DescService, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, +>( service: Service, - transport: ClientTransport, + transport: ClientTransport, serverId: TransportClientId, providedClientOptions: Partial< ClientOptions & { - handshakeOptions: ClientHandshakeOptions; + handshakeOptions: ClientHandshakeOptions; } > = {}, ): ProtobufClient { @@ -111,10 +116,13 @@ export function createClient( return client as ProtobufClient; } -function createMethodCaller( +function createMethodCaller< + Method extends DescMethod, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +>( service: DescService, method: Method, - transport: ClientTransport, + transport: ClientTransport, serverId: TransportClientId, clientOptions: ClientOptions, ): ClientMethod { @@ -235,9 +243,11 @@ function createMethodCaller( } } -function connectOnInvokeIfNeeded( +function connectOnInvokeIfNeeded< + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +>( clientOptions: ClientOptions, - transport: ClientTransport, + transport: ClientTransport, serverId: TransportClientId, ) { if (clientOptions.connectOnInvoke && !transport.sessions.has(serverId)) { @@ -245,10 +255,13 @@ function connectOnInvokeIfNeeded( } } -function startMethodCall( +function startMethodCall< + Method extends DescMethod, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +>( service: DescService, method: Method, - transport: ClientTransport, + transport: ClientTransport, serverId: TransportClientId, initialPayload: Uint8Array, procClosesWithInit: boolean, diff --git a/protobuf/handshake.ts b/protobuf/handshake.ts index f441ee4c..c87d2d39 100644 --- a/protobuf/handshake.ts +++ b/protobuf/handshake.ts @@ -11,6 +11,8 @@ import { type ServerHandshakeOptions, } from '../router/handshake'; import { + type CustomHandshakeErrorCode, + type CustomHandshakeErrorCodeSchema, HandshakeErrorCustomHandlerFatalResponseCodes, type TransportClientId, } from '../transport/message'; @@ -27,23 +29,36 @@ type ConstructHandshake = () => | MessageInitShape | Promise>; -type ValidateHandshake = ( +type ValidateHandshake< + Schema extends DescMessage, + ParsedMetadata, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +> = ( metadata: MessageShape, previousParsedMetadata?: ParsedMetadata, from?: TransportClientId, ) => | ParsedMetadata | ProtobufHandshakeFailureCode - | Promise; + | CustomHandshakeErrorCode + | Promise< + | ParsedMetadata + | ProtobufHandshakeFailureCode + | CustomHandshakeErrorCode + >; /** * Create client-side handshake options backed by a protobuf message type. */ -export function createClientHandshakeOptions( +export function createClientHandshakeOptions< + Schema extends DescMessage, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, +>( schema: Schema, construct: ConstructHandshake, eager?: boolean, -): ClientHandshakeOptions { + rejectionCodeSchema?: RejectionCodeSchema, +): ClientHandshakeOptions { return createTransportClientHandshakeOptions( HandshakeBytesSchema, async () => { @@ -52,6 +67,7 @@ export function createClientHandshakeOptions( return encodeMessageBytes(schema, metadata); }, eager, + rejectionCodeSchema, ); } @@ -61,11 +77,21 @@ export function createClientHandshakeOptions( export function createServerHandshakeOptions< Schema extends DescMessage, ParsedMetadata extends object = object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, >( schema: Schema, - validate: ValidateHandshake, + validate: ValidateHandshake< + Schema, + ParsedMetadata, + NoInfer + >, expiry?: (parsedMetadata: ParsedMetadata) => Date | undefined, -): ServerHandshakeOptions { + rejectionCodeSchema?: RejectionCodeSchema, +): ServerHandshakeOptions< + typeof HandshakeBytesSchema, + ParsedMetadata, + RejectionCodeSchema +> { return createTransportServerHandshakeOptions( HandshakeBytesSchema, async (metadata, previousParsedMetadata, from) => { @@ -79,5 +105,6 @@ export function createServerHandshakeOptions< return await validate(decoded, previousParsedMetadata, from); }, expiry, + rejectionCodeSchema, ); } diff --git a/protobuf/server.ts b/protobuf/server.ts index 164e65f1..f848aa13 100644 --- a/protobuf/server.ts +++ b/protobuf/server.ts @@ -16,6 +16,7 @@ import { EventMap } from '../transport/events'; import { ControlFlags, ControlMessageCloseSchema, + type CustomHandshakeErrorCodeSchema, OpaqueTransportMessage, TransportClientId, cancelMessage, @@ -133,11 +134,13 @@ export type Middleware = ( export interface ServerOptions< MetadataSchema extends TSchema, ParsedMetadata extends object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, > { readonly extendedContext?: object; readonly handshakeOptions?: ServerHandshakeOptions< MetadataSchema, - ParsedMetadata + ParsedMetadata, + RejectionCodeSchema >; readonly middlewares?: Array>; readonly maxCancelledStreamTombstonesPerSession?: number; @@ -146,6 +149,7 @@ export interface ServerOptions< class ProtobufServer< MetadataSchema extends TSchema, ParsedMetadata extends object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, > implements Server { readonly streams: Map; @@ -153,7 +157,8 @@ class ProtobufServer< private readonly transport: ServerTransport< Connection, MetadataSchema, - ParsedMetadata + ParsedMetadata, + RejectionCodeSchema >; private readonly methods: Map; @@ -171,9 +176,18 @@ class ProtobufServer< private unregisterTransportListeners: () => void; constructor( - transport: ServerTransport, + transport: ServerTransport< + Connection, + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >, services: ReadonlyArray, - options: ServerOptions = {}, + options: ServerOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + > = {}, ) { this.transport = transport; this.log = transport.log; @@ -979,10 +993,16 @@ class LRUSet { export function createServer< MetadataSchema extends TSchema, ParsedMetadata extends object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, >( - transport: ServerTransport, + transport: ServerTransport< + Connection, + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >, services: ReadonlyArray, - options?: ServerOptions, + options?: ServerOptions, ): Server { return new ProtobufServer(transport, services, options); } diff --git a/router/client.ts b/router/client.ts index 81d2bb7b..a14d4454 100644 --- a/router/client.ts +++ b/router/client.ts @@ -17,8 +17,9 @@ import { isStreamCancel, closeStreamMessage, cancelMessage, + type CustomHandshakeErrorCodeSchema, } from '../transport/message'; -import type { Static } from 'typebox'; +import type { Static, TSchema } from 'typebox'; import { Err, Result, AnyResultSchema } from './result'; import { EventMap } from '../transport/events'; import { Connection } from '../transport/connection'; @@ -241,12 +242,16 @@ const defaultClientOptions: ClientOptions = { // We are using any here because the ServiceContext is a server-side implementation // detail that doesn't affect the client interface // eslint-disable-next-line @typescript-eslint/no-explicit-any -export function createClient>( - transport: ClientTransport, +export function createClient< + // eslint-disable-next-line @typescript-eslint/no-explicit-any + ServiceSchemaMap extends AnyServiceSchemaMap, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, +>( + transport: ClientTransport, serverId: TransportClientId, providedClientOptions: Partial< ClientOptions & { - handshakeOptions: ClientHandshakeOptions; + handshakeOptions: ClientHandshakeOptions; } > = {}, ): Client { @@ -317,9 +322,9 @@ type AnyProcReturn = | ReturnType> | ReturnType>; -function handleProc( +function handleProc( procType: ValidProcType, - transport: ClientTransport, + transport: ClientTransport, serverId: TransportClientId, init: Static, serviceName: string, diff --git a/router/handshake.ts b/router/handshake.ts index 427b86b6..1b41a0e0 100644 --- a/router/handshake.ts +++ b/router/handshake.ts @@ -1,5 +1,7 @@ import type { Static, TSchema } from 'typebox'; import { + type CustomHandshakeErrorCode, + type CustomHandshakeErrorCodeSchema, HandshakeErrorCustomHandlerFatalResponseCodes, type TransportClientId, } from '../transport/message'; @@ -8,20 +10,27 @@ type ConstructHandshake = () => | Static | Promise>; -type ValidateHandshake = ( +type ValidateHandshake< + T extends TSchema, + ParsedMetadata, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +> = ( metadata: Static, previousParsedMetadata?: ParsedMetadata, from?: TransportClientId, ) => | Static + | CustomHandshakeErrorCode | ParsedMetadata | Promise< | Static + | CustomHandshakeErrorCode | ParsedMetadata >; export interface ClientHandshakeOptions< MetadataSchema extends TSchema = TSchema, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, > { /** * Schema for the metadata that the client sends to the server @@ -29,6 +38,15 @@ export interface ClientHandshakeOptions< */ schema: MetadataSchema; + /** + * Custom rejection codes the server may answer the handshake + * with, sent in the response's `code` field. Must match the server's + * {@link ServerHandshakeOptions.rejectionCodeSchema}: an unconfigured code + * is rejected as a malformed handshake response. Pass a TypeBox union of + * literals. + */ + rejectionCodeSchema?: RejectionCodeSchema; + /** * Gets the {@link HandshakeRequestMetadata} to send to the server. */ @@ -48,6 +66,7 @@ export interface ClientHandshakeOptions< export interface ServerHandshakeOptions< MetadataSchema extends TSchema = TSchema, ParsedMetadata extends object = object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, > { /** * Schema for the metadata that the server receives from the client @@ -55,6 +74,15 @@ export interface ServerHandshakeOptions< */ schema: MetadataSchema; + /** + * Custom rejection codes that {@link validate} may return. + * They travel in the handshake response's `code` field and are fatal like + * the built-in custom-handler codes. Clients must register the same codes + * in {@link ClientHandshakeOptions.rejectionCodeSchema} or they reject the + * response as malformed. Pass a TypeBox union of literals. + */ + rejectionCodeSchema?: RejectionCodeSchema; + /** * Parses the metadata sent by the client during the handshake into the * server-side {@link ParsedMetadata}, or returns a handshake failure code to @@ -67,7 +95,11 @@ export interface ServerHandshakeOptions< * confirm the presented id is the one the metadata authorizes before * returning parsed metadata. */ - validate: ValidateHandshake; + validate: ValidateHandshake< + MetadataSchema, + ParsedMetadata, + NoInfer + >; /** * When the credential expires (or undefined if it never does). The server @@ -84,21 +116,29 @@ export interface ServerHandshakeOptions< export function createClientHandshakeOptions< MetadataSchema extends TSchema = TSchema, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, >( schema: MetadataSchema, construct: ConstructHandshake, eager?: boolean, -): ClientHandshakeOptions { - return { schema, construct, eager }; + rejectionCodeSchema?: RejectionCodeSchema, +): ClientHandshakeOptions { + return { schema, construct, eager, rejectionCodeSchema }; } export function createServerHandshakeOptions< MetadataSchema extends TSchema = TSchema, ParsedMetadata extends object = object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, >( schema: MetadataSchema, - validate: ValidateHandshake, + validate: ValidateHandshake< + MetadataSchema, + ParsedMetadata, + NoInfer + >, expiry?: (parsedMetadata: ParsedMetadata) => Date | undefined, -): ServerHandshakeOptions { - return { schema, validate, expiry }; + rejectionCodeSchema?: RejectionCodeSchema, +): ServerHandshakeOptions { + return { schema, validate, expiry, rejectionCodeSchema }; } diff --git a/router/server.ts b/router/server.ts index e4548720..b8ba4616 100644 --- a/router/server.ts +++ b/router/server.ts @@ -29,6 +29,7 @@ import { cancelMessage, ProtocolVersion, TransportClientId, + type CustomHandshakeErrorCodeSchema, } from '../transport/message'; import { ProcedureHandlerContext } from './context'; import { Logger } from '../logging/log'; @@ -112,12 +113,14 @@ class RiverServer< MetadataSchema extends TSchema, ParsedMetadata extends object, Services extends AnyServiceSchemaMap, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, > implements Server { private transport: ServerTransport< Connection, MetadataSchema, - ParsedMetadata + ParsedMetadata, + RejectionCodeSchema >; private contextMap: Map; @@ -145,9 +148,18 @@ class RiverServer< private unregisterTransportListeners: () => void; constructor( - transport: ServerTransport, + transport: ServerTransport< + Connection, + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >, services: Services, - handshakeOptions?: ServerHandshakeOptions, + handshakeOptions?: ServerHandshakeOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >, extendedContext?: Context, maxCancelledStreamTombstonesPerSession = 200, middlewares: Array = [], @@ -1167,11 +1179,21 @@ export function createServer< // eslint-disable-next-line @typescript-eslint/no-explicit-any Services extends AnyServiceSchemaMap, Context extends MaybeDisposable = MaybeDisposable, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, >( - transport: ServerTransport, + transport: ServerTransport< + Connection, + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >, services: Services, providedServerOptions?: Partial<{ - handshakeOptions?: ServerHandshakeOptions; + handshakeOptions?: ServerHandshakeOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >; extendedContext?: Context; /** * Maximum number of cancelled streams to keep track of to avoid diff --git a/testUtil/fixtures/cleanup.ts b/testUtil/fixtures/cleanup.ts index 423d5122..32d8abb5 100644 --- a/testUtil/fixtures/cleanup.ts +++ b/testUtil/fixtures/cleanup.ts @@ -8,8 +8,12 @@ import { import { Server } from '../../router'; import { AnyServiceSchemaMap, MaybeDisposable } from '../../router/services'; import { numberOfConnections, testingSessionOptions } from '..'; +import type { TSchema } from 'typebox'; import { Value } from 'typebox/value'; -import { ControlMessageAckSchema } from '../../transport/message'; +import { + type CustomHandshakeErrorCodeSchema, + ControlMessageAckSchema, +} from '../../transport/message'; const waitUntilOptions = { timeout: 500, // account for possibility of conn backoff @@ -36,7 +40,9 @@ export async function advanceFakeTimersByConnectionBackoff() { await vi.advanceTimersByTimeAsync(500); } -export async function ensureTransportIsClean(t: Transport) { +export async function ensureTransportIsClean< + HandshakeFailureCode extends string, +>(t: Transport) { await advanceFakeTimersBySessionGrace(); await waitFor(() => expect( @@ -56,9 +62,9 @@ export function waitFor(cb: () => T | Promise) { return vi.waitFor(cb, waitUntilOptions); } -export async function ensureTransportBuffersAreEventuallyEmpty( - t: Transport, -) { +export async function ensureTransportBuffersAreEventuallyEmpty< + HandshakeFailureCode extends string, +>(t: Transport) { // wait for send buffers to be flushed // ignore heartbeat messages await waitFor(() => @@ -97,9 +103,10 @@ export async function ensureServerIsClean( ); } -export async function cleanupTransports( - transports: Array>, -) { +export async function cleanupTransports< + ConnType extends Connection, + HandshakeFailureCode extends string, +>(transports: Array>) { for (const t of transports) { if (t.getStatus() !== 'closed') { t.log?.info('*** end of test cleanup ***', { clientId: t.clientId }); @@ -108,16 +115,22 @@ export async function cleanupTransports( } } -export async function testFinishesCleanly({ +export async function testFinishesCleanly< + MetadataSchema extends TSchema, + ParsedMetadata extends object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +>({ clientTransports, serverTransport, server, }: Partial<{ - clientTransports: Array>; - // MetadataSchema and ParsedMetadata are not used in this test, - // so we can safely use any here - // eslint-disable-next-line @typescript-eslint/no-explicit-any - serverTransport: ServerTransport; + clientTransports: Array>; + serverTransport: ServerTransport< + Connection, + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >; server: Server; }>) { // pre-close invariants diff --git a/testUtil/fixtures/mockTransport.ts b/testUtil/fixtures/mockTransport.ts index a4d25177..fd7369f9 100644 --- a/testUtil/fixtures/mockTransport.ts +++ b/testUtil/fixtures/mockTransport.ts @@ -1,4 +1,5 @@ import { Transport, TransportClientId } from '../../transport'; +import type { CustomHandshakeErrorCodeSchema } from '../../transport/message'; import { ClientTransport } from '../../transport/client'; import { Connection } from '../../transport/connection'; import { ServerTransport } from '../../transport/server'; @@ -9,7 +10,10 @@ import { Duplex } from 'node:stream'; import { duplexPair } from '../duplex/duplexPair'; import { nanoid } from 'nanoid'; import type { TSchema } from 'typebox'; -import { ServerHandshakeOptions } from '../../router/handshake'; +import { + ClientHandshakeOptions, + ServerHandshakeOptions, +} from '../../router/handshake'; export class InMemoryConnection extends Connection { conn: Duplex; @@ -71,8 +75,10 @@ export function createMockTransportNetwork( // conn id -> [client->server, server->client] const connections = new Observable>({}); - const transports: Array> = []; - class MockClientTransport extends ClientTransport { + const transports: Array> = []; + class MockClientTransport< + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, + > extends ClientTransport { async createNewOutgoingConnection( to: TransportClientId, ): Promise { @@ -97,12 +103,14 @@ export function createMockTransportNetwork( } class MockServerTransport< - MetadataSchema extends TSchema = TSchema, - ParsedMetadata extends object = object, + MetadataSchema extends TSchema, + ParsedMetadata extends object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, > extends ServerTransport< InMemoryConnection, MetadataSchema, - ParsedMetadata + ParsedMetadata, + RejectionCodeSchema > { subscribeCleanup: () => void; @@ -135,8 +143,16 @@ export function createMockTransportNetwork( } return { - getClientTransport: (id, handshakeOptions) => { - const clientTransport = new MockClientTransport(id, opts?.client); + getClientTransport: < + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, + >( + id: TransportClientId, + handshakeOptions?: ClientHandshakeOptions, + ) => { + const clientTransport = new MockClientTransport( + id, + opts?.client, + ); if (handshakeOptions) { clientTransport.extendHandshake(handshakeOptions); } @@ -146,17 +162,23 @@ export function createMockTransportNetwork( return clientTransport; }, getServerTransport: < - MetadataSchema extends TSchema = TSchema, - ParsedMetadata extends object = object, + MetadataSchema extends TSchema, + ParsedMetadata extends object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, >( id = 'SERVER', handshakeOptions: - | ServerHandshakeOptions + | ServerHandshakeOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + > | undefined, ) => { const serverTransport = new MockServerTransport< MetadataSchema, - ParsedMetadata + ParsedMetadata, + RejectionCodeSchema >(id, opts?.server); if (handshakeOptions) { serverTransport.extendHandshake(handshakeOptions); diff --git a/testUtil/fixtures/transports.ts b/testUtil/fixtures/transports.ts index 75822f68..e4a4b3a5 100644 --- a/testUtil/fixtures/transports.ts +++ b/testUtil/fixtures/transports.ts @@ -16,7 +16,10 @@ import { ProvidedClientTransportOptions, ProvidedServerTransportOptions, } from '../../transport/options'; -import { TransportClientId } from '../../transport/message'; +import { + type CustomHandshakeErrorCodeSchema, + TransportClientId, +} from '../../transport/message'; import { ClientTransport } from '../../transport/client'; import { Connection } from '../../transport/connection'; import { ServerTransport } from '../../transport/server'; @@ -30,17 +33,29 @@ export interface TestTransportOptions { } export interface TestSetupHelpers { - getClientTransport: ( + getClientTransport: < + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, + >( id: TransportClientId, - handshakeOptions?: ClientHandshakeOptions, - ) => ClientTransport; + handshakeOptions?: ClientHandshakeOptions, + ) => ClientTransport; getServerTransport: < MetadataSchema extends TSchema = TSchema, ParsedMetadata extends object = object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, >( id?: TransportClientId, - handshakeOptions?: ServerHandshakeOptions, - ) => ServerTransport; + handshakeOptions?: ServerHandshakeOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >, + ) => ServerTransport< + Connection, + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >; simulatePhantomDisconnect: () => void; restartServer: () => Promise; cleanup: () => Promise | void; @@ -59,10 +74,11 @@ export const transports: Array = [ const port = await onWsServerReady(server); let wss = createWebSocketServer(server); + /* eslint-disable @typescript-eslint/no-explicit-any */ const transports: Array< - // eslint-disable-next-line @typescript-eslint/no-explicit-any - WebSocketClientTransport | WebSocketServerTransport + WebSocketClientTransport | WebSocketServerTransport > = []; + /* eslint-enable @typescript-eslint/no-explicit-any */ return { simulatePhantomDisconnect() { @@ -72,12 +88,21 @@ export const transports: Array = [ } } }, - getClientTransport: (id, handshakeOptions) => { - const clientTransport = new WebSocketClientTransport( - () => Promise.resolve(createLocalWebSocketClient(port)), - id, - opts?.client, - ); + getClientTransport: < + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, + >( + id: TransportClientId, + handshakeOptions?: ClientHandshakeOptions< + TSchema, + RejectionCodeSchema + >, + ) => { + const clientTransport = + new WebSocketClientTransport( + () => Promise.resolve(createLocalWebSocketClient(port)), + id, + opts?.client, + ); if (handshakeOptions) { clientTransport.extendHandshake(handshakeOptions); @@ -99,15 +124,21 @@ export const transports: Array = [ getServerTransport: < MetadataSchema extends TSchema, ParsedMetadata extends object, + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, >( id = 'SERVER', handshakeOptions: - | ServerHandshakeOptions + | ServerHandshakeOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + > | undefined, ) => { const serverTransport = new WebSocketServerTransport< MetadataSchema, - ParsedMetadata + ParsedMetadata, + RejectionCodeSchema >(wss, id, opts?.server); serverTransport.bindLogger((msg, ctx, level) => { @@ -128,7 +159,8 @@ export const transports: Array = [ return serverTransport as ServerTransport< Connection, MetadataSchema, - ParsedMetadata + ParsedMetadata, + RejectionCodeSchema >; }, async restartServer() { diff --git a/testUtil/index.ts b/testUtil/index.ts index da0c883c..a113fd7b 100644 --- a/testUtil/index.ts +++ b/testUtil/index.ts @@ -2,6 +2,7 @@ import NodeWs, { WebSocketServer } from 'ws'; import http from 'node:http'; import type { Static } from 'typebox'; import { + type CustomHandshakeErrorCodeSchema, OpaqueTransportMessage, PartialTransportMessage, currentProtocolVersion, @@ -18,7 +19,6 @@ import { SessionState } from '../transport/sessionStateMachine/common'; import { SessionStateGraph } from '../transport/sessionStateMachine/transitions'; import { BaseErrorSchemaType } from '../router/errors'; import { ClientTransport } from '../transport/client'; -import { ServerTransport } from '../transport/server'; import { getTracer } from '../tracing'; export { @@ -194,9 +194,11 @@ export function dummySession() { ); } -export function getClientSendFn( - clientTransport: ClientTransport, - serverTransport: ServerTransport, +export function getClientSendFn< + ClientRejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +>( + clientTransport: ClientTransport, + serverTransport: { clientId: string }, ) { const session = clientTransport.sessions.get(serverTransport.clientId) ?? @@ -208,9 +210,9 @@ export function getClientSendFn( ); } -export function getServerSendFn( - serverTransport: ServerTransport, - clientTransport: ClientTransport, +export function getServerSendFn( + serverTransport: Transport, + clientTransport: { clientId: string }, ) { const session = serverTransport.sessions.get(clientTransport.clientId); if (!session) { @@ -223,9 +225,10 @@ export function getServerSendFn( ); } -export function getTransportConnections( - transport: Transport, -): Array { +export function getTransportConnections< + ConnType extends Connection, + HandshakeFailureCode extends string, +>(transport: Transport): Array { const connections = []; for (const session of transport.sessions.values()) { if (session.state === SessionState.Connected) { @@ -236,15 +239,17 @@ export function getTransportConnections( return connections; } -export function numberOfConnections( - transport: Transport, -): number { +export function numberOfConnections< + ConnType extends Connection, + HandshakeFailureCode extends string, +>(transport: Transport): number { return getTransportConnections(transport).length; } -export function closeAllConnections( - transport: Transport, -) { +export function closeAllConnections< + ConnType extends Connection, + HandshakeFailureCode extends string, +>(transport: Transport) { for (const conn of getTransportConnections(transport)) { conn.close(); } diff --git a/transport/client.ts b/transport/client.ts index 4f9e440a..83e1a9b4 100644 --- a/transport/client.ts +++ b/transport/client.ts @@ -3,7 +3,10 @@ import { ClientHandshakeOptions } from '../router/handshake'; import { validationErrorToRiverErrors } from '../router/errors'; import { ControlMessageHandshakeResponseSchema, + ControlMessageHandshakeResponseSchemaWithCodes, ControlMessageRehandshakeRequestSchema, + type CustomHandshakeErrorCodeSchema, + type HandshakeErrorCode, HandshakeErrorRetriableResponseCodes, OpaqueTransportMessage, TransportClientId, @@ -11,6 +14,7 @@ import { handshakeRequestMessage, rehandshakeResponseMessage, } from './message'; +import type { TSchema } from 'typebox'; import { ClientTransportOptions, ProvidedClientTransportOptions, @@ -50,7 +54,8 @@ type ConstructedHandshakeMetadata = export abstract class ClientTransport< ConnType extends Connection, -> extends Transport { + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, +> extends Transport> { /** * The options for this transport. */ @@ -69,7 +74,16 @@ export abstract class ClientTransport< /** * Optional handshake options for this client. */ - handshakeExtensions?: ClientHandshakeOptions; + handshakeExtensions?: ClientHandshakeOptions; + + /** + * Handshake response schema extended with the custom + * rejection codes, when any are registered. + */ + protected handshakeResponseSchema: + | typeof ControlMessageHandshakeResponseSchema + | ReturnType = + ControlMessageHandshakeResponseSchema; /** * Handshake-metadata constructions prefetched when a connection attempt begins @@ -98,9 +112,17 @@ export abstract class ClientTransport< this.retryBudget = new LeakyBucketRateLimit(this.options); } - extendHandshake(options: ClientHandshakeOptions) { + extendHandshake = ( + options: ClientHandshakeOptions, + ) => { this.handshakeExtensions = options; - } + if (options.rejectionCodeSchema) { + this.handshakeResponseSchema = + ControlMessageHandshakeResponseSchemaWithCodes( + options.rejectionCodeSchema, + ); + } + }; protected handleRehandshakeMessage(message: OpaqueTransportMessage): void { if (!Value.Check(ControlMessageRehandshakeRequestSchema, message.payload)) { @@ -355,13 +377,13 @@ export abstract class ClientTransport< msg: OpaqueTransportMessage, ) { // invariant: msg is a handshake response - if (!Value.Check(ControlMessageHandshakeResponseSchema, msg.payload)) { + if (!Value.Check(this.handshakeResponseSchema, msg.payload)) { const reason = `received invalid handshake response`; this.rejectHandshakeResponse(session, reason, { ...session.loggingMetadata, transportMessage: msg, validationErrors: Value.Errors( - ControlMessageHandshakeResponseSchema, + this.handshakeResponseSchema, msg.payload, ).flatMap(validationErrorToRiverErrors), }); @@ -388,7 +410,8 @@ export abstract class ClientTransport< } else { this.protocolError({ type: ProtocolError.HandshakeFailed, - code: msg.payload.status.code, + code: msg.payload.status + .code as HandshakeErrorCode, message: reason, }); } diff --git a/transport/events.ts b/transport/events.ts index 5da1f81e..b8b70fa7 100644 --- a/transport/events.ts +++ b/transport/events.ts @@ -1,6 +1,8 @@ -import type { Static } from 'typebox'; import { Connection } from './connection'; -import { OpaqueTransportMessage, HandshakeErrorResponseCodes } from './message'; +import type { + BuiltInHandshakeErrorCode, + OpaqueTransportMessage, +} from './message'; import { Session, SessionState } from './sessionStateMachine'; import { SessionId } from './sessionStateMachine/common'; import { TransportStatus } from './transport'; @@ -16,7 +18,13 @@ export const ProtocolError = { export type ProtocolErrorType = (typeof ProtocolError)[keyof typeof ProtocolError]; -export interface EventMap { +/** + * Transport events. `HandshakeFailureCode` is the full set of codes observable + * on handshake-failed protocol errors, including built-in and custom codes. + */ +export interface EventMap< + HandshakeFailureCode extends string = BuiltInHandshakeErrorCode, +> { message: OpaqueTransportMessage; sessionStatus: | { @@ -36,7 +44,7 @@ export interface EventMap { protocolError: | { type: (typeof ProtocolError)['HandshakeFailed']; - code: Static; + code: HandshakeFailureCode; message: string; } | { @@ -52,12 +60,18 @@ export interface EventMap { } export type EventTypes = keyof EventMap; -export type EventHandler = ( - event: EventMap[K], -) => unknown; +export type EventHandler< + K extends EventTypes, + HandshakeFailureCode extends string = BuiltInHandshakeErrorCode, +> = (event: EventMap[K]) => unknown; -export class EventDispatcher { - private eventListeners: { [K in T]?: Set> } = {}; +export class EventDispatcher< + T extends EventTypes, + HandshakeFailureCode extends string, +> { + private eventListeners: { + [K in T]?: Set>; + } = {}; removeAllListeners() { this.eventListeners = {}; @@ -67,22 +81,33 @@ export class EventDispatcher { return this.eventListeners[eventType]?.size ?? 0; } - addEventListener(eventType: K, handler: EventHandler) { - if (!this.eventListeners[eventType]) { - this.eventListeners[eventType] = new Set(); + addEventListener( + eventType: K, + handler: EventHandler, + ) { + let listeners = this.eventListeners[eventType]; + if (!listeners) { + listeners = new Set(); + this.eventListeners[eventType] = listeners; } - this.eventListeners[eventType]?.add(handler); + listeners.add(handler); } - removeEventListener(eventType: K, handler: EventHandler) { + removeEventListener( + eventType: K, + handler: EventHandler, + ) { const handlers = this.eventListeners[eventType]; if (handlers) { this.eventListeners[eventType]?.delete(handler); } } - dispatchEvent(eventType: K, event: EventMap[K]) { + dispatchEvent( + eventType: K, + event: EventMap[K], + ) { const handlers = this.eventListeners[eventType]; if (handlers) { // copying ensures that adding more listeners in a handler doesn't diff --git a/transport/impls/ws/client.ts b/transport/impls/ws/client.ts index dea22b13..27dc6bc0 100644 --- a/transport/impls/ws/client.ts +++ b/transport/impls/ws/client.ts @@ -1,5 +1,8 @@ import { ClientTransport } from '../../client'; -import { TransportClientId } from '../../message'; +import { + type CustomHandshakeErrorCodeSchema, + TransportClientId, +} from '../../message'; import { ProvidedClientTransportOptions } from '../../options'; import { WebSocketConnection } from './connection'; import { WsLike } from './wslike'; @@ -9,7 +12,9 @@ import { WsLike } from './wslike'; * @class * @extends Transport */ -export class WebSocketClientTransport extends ClientTransport { +export class WebSocketClientTransport< + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, +> extends ClientTransport { /** * A function that returns a Promise that resolves to a websocket URL. */ diff --git a/transport/impls/ws/server.ts b/transport/impls/ws/server.ts index 9bcaddec..e1df56eb 100644 --- a/transport/impls/ws/server.ts +++ b/transport/impls/ws/server.ts @@ -1,4 +1,7 @@ -import { TransportClientId } from '../../message'; +import { + type CustomHandshakeErrorCodeSchema, + TransportClientId, +} from '../../message'; import { WebSocketServer } from 'ws'; import { WebSocketConnection } from './connection'; import { WsLike } from './wslike'; @@ -25,7 +28,13 @@ function cleanHeaders( export class WebSocketServerTransport< MetadataSchema extends TSchema = TSchema, ParsedMetadata extends object = object, -> extends ServerTransport { + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, +> extends ServerTransport< + WebSocketConnection, + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema +> { wss: WebSocketServer; constructor( diff --git a/transport/impls/ws/ws.test.ts b/transport/impls/ws/ws.test.ts index 285e66e7..f1f4c3f2 100644 --- a/transport/impls/ws/ws.test.ts +++ b/transport/impls/ws/ws.test.ts @@ -1,5 +1,6 @@ import http from 'node:http'; import { describe, test, expect, beforeEach } from 'vitest'; +import { Type } from 'typebox'; import { createWebSocketServer, onWsServerReady, @@ -23,6 +24,7 @@ import { import { PartialTransportMessage } from '../../message'; import type NodeWs from 'ws'; import { createPostTestCleanups } from '../../../testUtil/fixtures/cleanup'; +import { createClientHandshakeOptions } from '../../../router/handshake'; describe('sending and receiving across websockets works', async () => { let server: http.Server; @@ -42,6 +44,33 @@ describe('sending and receiving across websockets works', async () => { }; }); + test('custom handshake codes require an explicit transport type', () => { + const rejectionCodeSchema = Type.Union([Type.Literal('TOKEN_EXPIRED')]); + const handshakeOptions = createClientHandshakeOptions( + Type.Object({}), + () => ({}), + undefined, + rejectionCodeSchema, + ); + const getWs = () => { + throw new Error('not called'); + }; + const defaultTransport = new WebSocketClientTransport(getWs, 'client'); + + // @ts-expect-error a default transport cannot be widened to custom codes + const widenedTransport: WebSocketClientTransport< + typeof rejectionCodeSchema + > = defaultTransport; + + const typedTransport = new WebSocketClientTransport< + typeof rejectionCodeSchema + >(getWs, 'client'); + typedTransport.extendHandshake(handshakeOptions); + + expect(typedTransport.handshakeExtensions).toBe(handshakeOptions); + expect(widenedTransport).toBe(defaultTransport); + }); + test('basic send/receive', async () => { const clientTransport = new WebSocketClientTransport( () => Promise.resolve(createLocalWebSocketClient(port)), diff --git a/transport/message.test.ts b/transport/message.test.ts index cc2426da..61f20add 100644 --- a/transport/message.test.ts +++ b/transport/message.test.ts @@ -1,6 +1,8 @@ import { TransportMessage } from '.'; import { ControlFlags, + ControlMessageHandshakeResponseSchema, + ControlMessageHandshakeResponseSchemaWithCodes, handshakeRequestMessage, handshakeResponseMessage, isAck, @@ -8,6 +10,8 @@ import { isStreamOpen, } from './message'; import { describe, test, expect } from 'vitest'; +import { Type } from 'typebox'; +import { Value } from 'typebox/value'; const msg = ( to: string, @@ -105,6 +109,51 @@ describe('message helpers', () => { expect(mFail.payload.status.ok).toBe(false); }); + test('handshake response schema with custom error codes', () => { + const extended = ControlMessageHandshakeResponseSchemaWithCodes( + Type.Union([ + Type.Literal('REPL_NOT_FOUND'), + Type.Literal('TOKEN_EXPIRED'), + ]), + ); + const rejection = { + type: 'HANDSHAKE_RESP', + status: { + ok: false, + reason: 'rejected by handshake handler', + code: 'REPL_NOT_FOUND', + }, + }; + const protocolFailure = { + type: 'HANDSHAKE_RESP', + status: { + ok: false, + reason: 'bad', + code: 'SESSION_STATE_MISMATCH', + }, + }; + const unknownCode = { + type: 'HANDSHAKE_RESP', + status: { + ok: false, + reason: 'bad', + code: 'NOT_A_REGISTERED_CODE', + }, + }; + + expect(Value.Check(extended, rejection)).toBe(true); + expect(Value.Check(extended, protocolFailure)).toBe(true); + expect(Value.Check(extended, unknownCode)).toBe(false); + + // the base schema only knows the protocol-level codes + expect(Value.Check(ControlMessageHandshakeResponseSchema, rejection)).toBe( + false, + ); + expect( + Value.Check(ControlMessageHandshakeResponseSchema, protocolFailure), + ).toBe(true); + }); + test('default message has no control flags set', () => { const m = msg('a', 'b', 'stream', { test: 1 }, 'svc', 'proc'); diff --git a/transport/message.ts b/transport/message.ts index 5dbd0cd2..e06e4aac 100644 --- a/transport/message.ts +++ b/transport/message.ts @@ -1,4 +1,10 @@ -import { Type, type TSchema, type Static } from 'typebox'; +import { + Type, + type TLiteral, + type TSchema, + type TUnion, + type Static, +} from 'typebox'; import { PropagationContext } from '../tracing'; import { generateId } from './id'; // type-only: a value import closes a transport <-> router require cycle @@ -123,20 +129,71 @@ export const HandshakeErrorResponseCodes = Type.Union([ HandshakeErrorFatalResponseCodes, ]); -export const ControlMessageHandshakeResponseSchema = Type.Object({ - type: Type.Literal('HANDSHAKE_RESP'), - status: Type.Union([ - Type.Object({ - ok: Type.Literal(true), - sessionId: Type.String(), - }), - Type.Object({ - ok: Type.Literal(false), - reason: Type.String(), - code: HandshakeErrorResponseCodes, - }), - ]), -}); +/** + * A TypeBox union of literals declaring the custom handshake rejection codes. + */ +export type CustomHandshakeErrorCodeSchema = TUnion>>; + +/** + * The union of codes declared by a rejection code schema. Widened literal + * schemas (`TLiteral` rather than a specific literal) contribute no + * codes. + */ +export type CustomHandshakeErrorCode< + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +> = string extends Static + ? never + : Static; + +/** + * The protocol-level handshake error codes River can emit without any custom + * rejection codes. + */ +export type BuiltInHandshakeErrorCode = Static< + typeof HandshakeErrorResponseCodes +>; + +/** + * The protocol-level handshake error codes plus any custom rejection codes + * registered in the handshake options. + */ +export type HandshakeErrorCode< + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +> = BuiltInHandshakeErrorCode | CustomHandshakeErrorCode; + +const handshakeResponseSchema = (code: Code) => + Type.Object({ + type: Type.Literal('HANDSHAKE_RESP'), + status: Type.Union([ + Type.Object({ + ok: Type.Literal(true), + sessionId: Type.String(), + }), + Type.Object({ + ok: Type.Literal(false), + reason: Type.String(), + code, + }), + ]), + }); + +export const ControlMessageHandshakeResponseSchema = handshakeResponseSchema( + HandshakeErrorResponseCodes, +); + +/** + * A handshake response schema that additionally accepts custom + * rejection codes. Both peers must be configured with the same codes: an + * unconfigured peer rejects a custom code as a malformed response. + */ +export const ControlMessageHandshakeResponseSchemaWithCodes = < + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +>( + rejectionCodeSchema: RejectionCodeSchema, +) => + handshakeResponseSchema( + Type.Union([HandshakeErrorResponseCodes, rejectionCodeSchema]), + ); /** * Reserved stream id for the follow-up handshake (re-handshake) control @@ -257,14 +314,24 @@ export function handshakeRequestMessage({ */ export const SESSION_STATE_MISMATCH = 'session state mismatch'; -export function handshakeResponseMessage({ +export function handshakeResponseMessage< + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema, +>({ from, to, status, }: { from: TransportClientId; to: TransportClientId; - status: Static['status']; + // the code may be a custom rejection code, which is only + // known to peers that registered it in their handshake options + status: + | { ok: true; sessionId: string } + | { + ok: false; + reason: string; + code: HandshakeErrorCode; + }; }): TransportMessage> { return { id: generateId(), @@ -276,8 +343,10 @@ export function handshakeResponseMessage({ controlFlags: 0, payload: { type: 'HANDSHAKE_RESP', - status, - } satisfies Static, + status: status as Static< + typeof ControlMessageHandshakeResponseSchema + >['status'], + }, }; } diff --git a/transport/server.ts b/transport/server.ts index c41e43ce..6f69fdc2 100644 --- a/transport/server.ts +++ b/transport/server.ts @@ -5,7 +5,9 @@ import { ControlMessageHandshakeRequestSchema, ControlMessageRehandshakeResponseSchema, HandshakeErrorCustomHandlerFatalResponseCodes, - HandshakeErrorResponseCodes, + type CustomHandshakeErrorCode, + type CustomHandshakeErrorCodeSchema, + type HandshakeErrorCode, OpaqueTransportMessage, acceptedProtocolVersions, TransportClientId, @@ -36,7 +38,8 @@ export abstract class ServerTransport< ConnType extends Connection, MetadataSchema extends TSchema = TSchema, ParsedMetadata extends object = object, -> extends Transport { + RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never, +> extends Transport> { /** * The options for this transport. */ @@ -45,7 +48,11 @@ export abstract class ServerTransport< /** * Optional handshake options for the server. */ - handshakeExtensions?: ServerHandshakeOptions; + handshakeExtensions?: ServerHandshakeOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >; /** * A map of session handshake data for each session. @@ -71,10 +78,27 @@ export abstract class ServerTransport< }); } - extendHandshake( - options: ServerHandshakeOptions, - ) { + extendHandshake = ( + options: ServerHandshakeOptions< + MetadataSchema, + ParsedMetadata, + RejectionCodeSchema + >, + ) => { this.handshakeExtensions = options; + }; + + private isRejectionCode( + value: unknown, + ): value is + | Static + | CustomHandshakeErrorCode { + return ( + Value.Check(HandshakeErrorCustomHandlerFatalResponseCodes, value) || + (this.handshakeExtensions?.rejectionCodeSchema + ? Value.Check(this.handshakeExtensions.rejectionCodeSchema, value) + : false) + ); } protected deletePendingSession( @@ -204,15 +228,11 @@ export abstract class ServerTransport< return; } - if ( - Value.Check( - HandshakeErrorCustomHandlerFatalResponseCodes, - parsedMetadataOrFailureCode, - ) - ) { + if (this.isRejectionCode(parsedMetadataOrFailureCode)) { this.teardownForFailedRehandshake( session, 're-handshake metadata rejected by handshake handler', + parsedMetadataOrFailureCode, ); return; @@ -246,6 +266,7 @@ export abstract class ServerTransport< private teardownForFailedRehandshake( session: ServerSession, reason: string, + code: HandshakeErrorCode = 'REJECTED_BY_CUSTOM_HANDLER', ) { if (session._isConsumed) { return; @@ -263,7 +284,7 @@ export abstract class ServerTransport< this.protocolError({ type: ProtocolError.HandshakeFailed, - code: 'REJECTED_BY_CUSTOM_HANDLER', + code, message: reason, }); this.deleteSession(session, { unhealthy: true }); @@ -354,7 +375,7 @@ export abstract class ServerTransport< session: SessionWaitingForHandshake, to: TransportClientId, reason: string, - code: Static, + code: HandshakeErrorCode, metadata: MessageMetadata, ) { session.conn.telemetry?.span.setStatus({ @@ -497,12 +518,7 @@ export abstract class ServerTransport< } // handler rejected the connection - if ( - Value.Check( - HandshakeErrorCustomHandlerFatalResponseCodes, - parsedMetadataOrFailureCode, - ) - ) { + if (this.isRejectionCode(parsedMetadataOrFailureCode)) { this.rejectHandshakeRequest( session, msg.from, diff --git a/transport/transport.test.ts b/transport/transport.test.ts index ba1bc0d8..dcecc26e 100644 --- a/transport/transport.test.ts +++ b/transport/transport.test.ts @@ -21,10 +21,11 @@ import { } from '../testUtil/fixtures/cleanup'; import { testMatrix } from '../testUtil/fixtures/matrix'; import { PartialTransportMessage } from './message'; -import { Type } from 'typebox'; +import { Type, type Static } from 'typebox'; import { TestSetupHelpers } from '../testUtil/fixtures/transports'; import { createPostTestCleanups } from '../testUtil/fixtures/cleanup'; import { SessionState } from './sessionStateMachine'; +import { createServerHandshakeOptions } from '../router/handshake'; import { ProvidedClientTransportOptions, ProvidedTransportOptions, @@ -1945,5 +1946,166 @@ describe.each(testMatrix())( serverTransport, }); }); + + test('custom handler can reject with a registered handshake error code', async () => { + const schema = Type.Object({ + foo: Type.String(), + }); + + const rejectionCodeSchema = Type.Union([ + Type.Literal('REPL_NOT_FOUND'), + Type.Literal('TOKEN_EXPIRED'), + ]); + + type CustomHandshakeErrorCode = Static; + + interface ParsedMetadata { + foo: string; + } + + // compile-time: undeclared codes cannot be returned by validate + expect( + createServerHandshakeOptions< + typeof schema, + ParsedMetadata, + typeof rejectionCodeSchema + >( + schema, + // @ts-expect-error only declared rejection codes may be returned + async () => 'SOME_OTHER_CODE', + undefined, + rejectionCodeSchema, + ), + ).toBeDefined(); + + const broadCode = String('REPL_NOT_FOUND'); + const broadCodeSchema = Type.Union([Type.Literal(broadCode)]); + + expect( + createServerHandshakeOptions< + typeof schema, + ParsedMetadata, + typeof broadCodeSchema + >( + schema, + // @ts-expect-error widened literal schemas contribute no codes + async () => broadCode, + undefined, + broadCodeSchema, + ), + ).toBeDefined(); + + const parse = vi.fn(async (): Promise => { + return 'REPL_NOT_FOUND'; + }); + const serverTransport = getServerTransport< + typeof schema, + ParsedMetadata, + typeof rejectionCodeSchema + >('SERVER', { + schema, + validate: parse, + rejectionCodeSchema, + }); + + const clientTransport = getClientTransport('client', { + schema, + construct: async () => ({ foo: 'foo' }), + rejectionCodeSchema, + }); + + const clientHandshakeFailed = vi.fn(); + clientTransport.addEventListener('protocolError', clientHandshakeFailed); + const serverRejectedConnection = vi.fn(); + serverTransport.addEventListener( + 'protocolError', + serverRejectedConnection, + ); + clientTransport.connect(serverTransport.clientId); + + addPostTestCleanup(async () => { + clientTransport.removeEventListener( + 'protocolError', + clientHandshakeFailed, + ); + serverTransport.removeEventListener( + 'protocolError', + serverRejectedConnection, + ); + await cleanupTransports([clientTransport, serverTransport]); + }); + + await waitFor(() => { + expect(clientHandshakeFailed).toHaveBeenCalledTimes(1); + expect(clientHandshakeFailed).toHaveBeenCalledWith({ + type: ProtocolError.HandshakeFailed, + code: 'REPL_NOT_FOUND', + message: 'handshake failed: rejected by handshake handler', + }); + expect(serverRejectedConnection).toHaveBeenCalledWith({ + type: ProtocolError.HandshakeFailed, + code: 'REPL_NOT_FOUND', + message: 'rejected by handshake handler', + }); + }); + + await testFinishesCleanly({ + clientTransports: [clientTransport], + serverTransport, + }); + }); + + test('an unregistered handshake error code is rejected by an unconfigured client', async () => { + const schema = Type.Object({ + foo: Type.String(), + }); + + const rejectionCodeSchema = Type.Union([Type.Literal('REPL_NOT_FOUND')]); + + type CustomHandshakeErrorCode = Static; + + interface ParsedMetadata { + foo: string; + } + + const serverTransport = getServerTransport< + typeof schema, + ParsedMetadata, + typeof rejectionCodeSchema + >('SERVER', { + schema, + validate: async (): Promise => + 'REPL_NOT_FOUND', + rejectionCodeSchema, + }); + + // the client did not register the custom code: it must treat the + // response as malformed rather than accept an unknown code + const clientTransport = getClientTransport('client', { + schema, + construct: async () => ({ foo: 'foo' }), + }); + + const clientHandshakeFailed = vi.fn(); + clientTransport.addEventListener('protocolError', clientHandshakeFailed); + clientTransport.connect(serverTransport.clientId); + + addPostTestCleanup(async () => { + clientTransport.removeEventListener( + 'protocolError', + clientHandshakeFailed, + ); + await cleanupTransports([clientTransport]); + await cleanupTransports([serverTransport]); + }); + + await waitFor(() => { + expect(clientTransport.sessions.size).toBe(0); + }); + expect(clientHandshakeFailed).not.toHaveBeenCalled(); + + await testFinishesCleanly({ clientTransports: [clientTransport] }); + await testFinishesCleanly({ serverTransport }); + }); }, ); diff --git a/transport/transport.ts b/transport/transport.ts index 092ca8ff..390f5d47 100644 --- a/transport/transport.ts +++ b/transport/transport.ts @@ -1,4 +1,5 @@ import { + type BuiltInHandshakeErrorCode, OpaqueTransportMessage, PartialTransportMessage, TransportClientId, @@ -79,7 +80,10 @@ export interface SessionBackpressure { * ``` * @abstract */ -export abstract class Transport { +export abstract class Transport< + ConnType extends Connection, + HandshakeFailureCode extends string = BuiltInHandshakeErrorCode, +> { /** * The status of the transport. */ @@ -93,7 +97,7 @@ export abstract class Transport { /** * The event dispatcher for handling events of type EventTypes. */ - eventDispatcher: EventDispatcher; + eventDispatcher: EventDispatcher; /** * The options for this transport. @@ -114,7 +118,10 @@ export abstract class Transport { providedOptions?: ProvidedTransportOptions, ) { this.options = { ...defaultTransportOptions, ...providedOptions }; - this.eventDispatcher = new EventDispatcher(); + this.eventDispatcher = new EventDispatcher< + EventTypes, + HandshakeFailureCode + >(); this.clientId = clientId; this.status = 'open'; this.sessions = new Map(); @@ -149,10 +156,10 @@ export abstract class Transport { * @param the type of event to listen for * @param handler The message handler to add. */ - addEventListener>( - type: K, - handler: T, - ): void { + addEventListener< + K extends EventTypes, + T extends EventHandler, + >(type: K, handler: T): void { this.eventDispatcher.addEventListener(type, handler); } @@ -161,14 +168,16 @@ export abstract class Transport { * @param the type of event to un-listen on * @param handler The message handler to remove. */ - removeEventListener>( - type: K, - handler: T, - ): void { + removeEventListener< + K extends EventTypes, + T extends EventHandler, + >(type: K, handler: T): void { this.eventDispatcher.removeEventListener(type, handler); } - protected protocolError(message: EventMap['protocolError']) { + protected protocolError( + message: EventMap['protocolError'], + ) { this.eventDispatcher.dispatchEvent('protocolError', message); }