Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 7 additions & 11 deletions __tests__/properties/session.property.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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';
Expand Down Expand Up @@ -152,12 +155,8 @@ const multiplexedSchedules: gs.Generator<MultiplexedSchedule> = gs.composite(

function setup(opts?: TestTransportOptions): {
network: ReturnType<typeof createMockTransportNetwork>;
clientTransport: ReturnType<
ReturnType<typeof createMockTransportNetwork>['getClientTransport']
>;
serverTransport: ReturnType<
ReturnType<typeof createMockTransportNetwork>['getServerTransport']
>;
clientTransport: ClientTransport<Connection>;
serverTransport: ServerTransport<Connection>;
client: ReturnType<typeof createClient<typeof services>>;
violations: Array<string>;
} {
Expand Down Expand Up @@ -536,11 +535,8 @@ describe('re-handshake under faults', () => {
throw new Error(`timed out waiting for ${what}`);
}

const isConnected = (
transport: ReturnType<
ReturnType<typeof createMockTransportNetwork>['getClientTransport']
>,
) => numberOfConnections(transport) === 1;
const isConnected = (transport: ClientTransport<Connection>) =>
numberOfConnections(transport) === 1;

const handshakeSchema = Type.Object({ token: Type.String() });

Expand Down
22 changes: 14 additions & 8 deletions __tests__/protobuf.test.ts
Original file line number Diff line number Diff line change
@@ -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,
Expand Down Expand Up @@ -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<typeof rejectionCodeSchema>('client');
const serverTransport = getServerTransport<
(typeof serverHandshakeOptions)['schema'],
{ token: string },
typeof rejectionCodeSchema
>('SERVER', undefined);
const TypedProtoService = createProtoService<object, { token: string }>();
const testSvc = TypedProtoService.define(TestService, {
echo: (request, ctx) =>
Expand All @@ -335,6 +340,7 @@ describe.each(protobufRouterMatrix)(
TestService,
clientTransport,
serverTransport.clientId,
{ handshakeOptions: clientHandshakeOptions },
);

await expect(client.echo({ text: 'hello' })).resolves.toMatchObject({
Expand Down
31 changes: 22 additions & 9 deletions protobuf/client.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,13 +6,15 @@ 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';
import { EventMap } from '../transport/events';
import {
ControlFlags,
ControlMessageCloseSchema,
type CustomHandshakeErrorCodeSchema,
OpaqueTransportMessage,
TransportClientId,
cancelMessage,
Expand Down Expand Up @@ -75,13 +77,16 @@ interface StartedMethodCall<
/**
* Creates a protobuf client for a single protobuf service descriptor.
*/
export function createClient<Service extends DescService>(
export function createClient<
Service extends DescService,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never,
>(
service: Service,
transport: ClientTransport<Connection>,
transport: ClientTransport<Connection, RejectionCodeSchema>,
serverId: TransportClientId,
providedClientOptions: Partial<
ClientOptions & {
handshakeOptions: ClientHandshakeOptions;
handshakeOptions: ClientHandshakeOptions<TSchema, RejectionCodeSchema>;
}
> = {},
): ProtobufClient<Service> {
Expand Down Expand Up @@ -111,10 +116,13 @@ export function createClient<Service extends DescService>(
return client as ProtobufClient<Service>;
}

function createMethodCaller<Method extends DescMethod>(
function createMethodCaller<
Method extends DescMethod,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema,
>(
service: DescService,
method: Method,
transport: ClientTransport<Connection>,
transport: ClientTransport<Connection, RejectionCodeSchema>,
serverId: TransportClientId,
clientOptions: ClientOptions,
): ClientMethod<Method> {
Expand Down Expand Up @@ -235,20 +243,25 @@ function createMethodCaller<Method extends DescMethod>(
}
}

function connectOnInvokeIfNeeded(
function connectOnInvokeIfNeeded<
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema,
>(
clientOptions: ClientOptions,
transport: ClientTransport<Connection>,
transport: ClientTransport<Connection, RejectionCodeSchema>,
serverId: TransportClientId,
) {
if (clientOptions.connectOnInvoke && !transport.sessions.has(serverId)) {
transport.connect(serverId);
}
}

function startMethodCall<Method extends DescMethod>(
function startMethodCall<
Method extends DescMethod,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema,
>(
service: DescService,
method: Method,
transport: ClientTransport<Connection>,
transport: ClientTransport<Connection, RejectionCodeSchema>,
serverId: TransportClientId,
initialPayload: Uint8Array,
procClosesWithInit: boolean,
Expand Down
39 changes: 33 additions & 6 deletions protobuf/handshake.ts
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@ import {
type ServerHandshakeOptions,
} from '../router/handshake';
import {
type CustomHandshakeErrorCode,
type CustomHandshakeErrorCodeSchema,
HandshakeErrorCustomHandlerFatalResponseCodes,
type TransportClientId,
} from '../transport/message';
Expand All @@ -27,23 +29,36 @@ type ConstructHandshake<Schema extends DescMessage> = () =>
| MessageInitShape<Schema>
| Promise<MessageInitShape<Schema>>;

type ValidateHandshake<Schema extends DescMessage, ParsedMetadata> = (
type ValidateHandshake<
Schema extends DescMessage,
ParsedMetadata,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema,
> = (
metadata: MessageShape<Schema>,
previousParsedMetadata?: ParsedMetadata,
from?: TransportClientId,
) =>
| ParsedMetadata
| ProtobufHandshakeFailureCode
| Promise<ParsedMetadata | ProtobufHandshakeFailureCode>;
| CustomHandshakeErrorCode<RejectionCodeSchema>
| Promise<
| ParsedMetadata
| ProtobufHandshakeFailureCode
| CustomHandshakeErrorCode<RejectionCodeSchema>
>;

/**
* Create client-side handshake options backed by a protobuf message type.
*/
export function createClientHandshakeOptions<Schema extends DescMessage>(
export function createClientHandshakeOptions<
Schema extends DescMessage,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never,
>(
schema: Schema,
construct: ConstructHandshake<Schema>,
eager?: boolean,
): ClientHandshakeOptions<typeof HandshakeBytesSchema> {
rejectionCodeSchema?: RejectionCodeSchema,
): ClientHandshakeOptions<typeof HandshakeBytesSchema, RejectionCodeSchema> {
return createTransportClientHandshakeOptions(
HandshakeBytesSchema,
async () => {
Expand All @@ -52,6 +67,7 @@ export function createClientHandshakeOptions<Schema extends DescMessage>(
return encodeMessageBytes(schema, metadata);
},
eager,
rejectionCodeSchema,
);
}

Expand All @@ -61,11 +77,21 @@ export function createClientHandshakeOptions<Schema extends DescMessage>(
export function createServerHandshakeOptions<
Schema extends DescMessage,
ParsedMetadata extends object = object,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never,
>(
schema: Schema,
validate: ValidateHandshake<Schema, ParsedMetadata>,
validate: ValidateHandshake<
Schema,
ParsedMetadata,
NoInfer<RejectionCodeSchema>
>,
expiry?: (parsedMetadata: ParsedMetadata) => Date | undefined,
): ServerHandshakeOptions<typeof HandshakeBytesSchema, ParsedMetadata> {
rejectionCodeSchema?: RejectionCodeSchema,
): ServerHandshakeOptions<
typeof HandshakeBytesSchema,
ParsedMetadata,
RejectionCodeSchema
> {
return createTransportServerHandshakeOptions(
HandshakeBytesSchema,
async (metadata, previousParsedMetadata, from) => {
Expand All @@ -79,5 +105,6 @@ export function createServerHandshakeOptions<
return await validate(decoded, previousParsedMetadata, from);
},
expiry,
rejectionCodeSchema,
);
}
32 changes: 26 additions & 6 deletions protobuf/server.ts
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import { EventMap } from '../transport/events';
import {
ControlFlags,
ControlMessageCloseSchema,
type CustomHandshakeErrorCodeSchema,
OpaqueTransportMessage,
TransportClientId,
cancelMessage,
Expand Down Expand Up @@ -133,11 +134,13 @@ export type Middleware<ParsedMetadata extends object = object> = (
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<Middleware<ParsedMetadata>>;
readonly maxCancelledStreamTombstonesPerSession?: number;
Expand All @@ -146,14 +149,16 @@ export interface ServerOptions<
class ProtobufServer<
MetadataSchema extends TSchema,
ParsedMetadata extends object,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema,
> implements Server
{
readonly streams: Map<StreamId, ProcStream>;

private readonly transport: ServerTransport<
Connection,
MetadataSchema,
ParsedMetadata
ParsedMetadata,
RejectionCodeSchema
>;

private readonly methods: Map<string, RegisteredMethod>;
Expand All @@ -171,9 +176,18 @@ class ProtobufServer<
private unregisterTransportListeners: () => void;

constructor(
transport: ServerTransport<Connection, MetadataSchema, ParsedMetadata>,
transport: ServerTransport<
Connection,
MetadataSchema,
ParsedMetadata,
RejectionCodeSchema
>,
services: ReadonlyArray<AnyProtoService>,
options: ServerOptions<MetadataSchema, ParsedMetadata> = {},
options: ServerOptions<
MetadataSchema,
ParsedMetadata,
RejectionCodeSchema
> = {},
) {
this.transport = transport;
this.log = transport.log;
Expand Down Expand Up @@ -979,10 +993,16 @@ class LRUSet<T> {
export function createServer<
MetadataSchema extends TSchema,
ParsedMetadata extends object,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never,
>(
transport: ServerTransport<Connection, MetadataSchema, ParsedMetadata>,
transport: ServerTransport<
Connection,
MetadataSchema,
ParsedMetadata,
RejectionCodeSchema
>,
services: ReadonlyArray<AnyProtoService>,
options?: ServerOptions<MetadataSchema, ParsedMetadata>,
options?: ServerOptions<MetadataSchema, ParsedMetadata, RejectionCodeSchema>,
): Server {
return new ProtobufServer(transport, services, options);
}
17 changes: 11 additions & 6 deletions router/client.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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';
Expand Down Expand Up @@ -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<ServiceSchemaMap extends AnyServiceSchemaMap<any>>(
transport: ClientTransport<Connection>,
export function createClient<
// eslint-disable-next-line @typescript-eslint/no-explicit-any
ServiceSchemaMap extends AnyServiceSchemaMap<any>,
RejectionCodeSchema extends CustomHandshakeErrorCodeSchema = never,
>(
transport: ClientTransport<Connection, RejectionCodeSchema>,
serverId: TransportClientId,
providedClientOptions: Partial<
ClientOptions & {
handshakeOptions: ClientHandshakeOptions;
handshakeOptions: ClientHandshakeOptions<TSchema, RejectionCodeSchema>;
}
> = {},
): Client<ServiceSchemaMap> {
Expand Down Expand Up @@ -317,9 +322,9 @@ type AnyProcReturn =
| ReturnType<StreamFn<AnyService, string>>
| ReturnType<SubscriptionFn<AnyService, string>>;

function handleProc(
function handleProc<RejectionCodeSchema extends CustomHandshakeErrorCodeSchema>(
procType: ValidProcType,
transport: ClientTransport<Connection>,
transport: ClientTransport<Connection, RejectionCodeSchema>,
serverId: TransportClientId,
init: Static<PayloadType>,
serviceName: string,
Expand Down
Loading
Loading