Skip to content
Closed
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
2 changes: 1 addition & 1 deletion Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -285,7 +285,7 @@ rand = "0.9"
rayon = "1.8"
rayon-core = "1.11.0"
regex = "1"
reqwest = { version = "0.12", features = ["stream", "json", "gzip", "brotli"] }
reqwest = { version = "0.12", features = ["stream", "json", "gzip", "brotli", "native-tls-alpn"] }
rolldown = { git = "https://github.com/rolldown/rolldown.git", tag = "v1.0.0-rc.3" }
rolldown_common = { git = "https://github.com/rolldown/rolldown.git", tag = "v1.0.0-rc.3" }
rolldown_error = { git = "https://github.com/rolldown/rolldown.git", tag = "v1.0.0-rc.3" }
Expand Down
11 changes: 9 additions & 2 deletions crates/bindings-typescript/src/react/SpacetimeDBProvider.ts
Original file line number Diff line number Diff line change
Expand Up @@ -64,9 +64,16 @@ export function SpacetimeDBProvider<
[key]
);

const reconnect = React.useCallback(
(builder: DbConnectionBuilder<any>) => {
ConnectionManager.rebuild(key, builder);
},
[key]
);

const contextValue = React.useMemo<ConnectionState>(
() => ({ ...state, getConnection }),
[state, getConnection]
() => ({ ...state, getConnection, reconnect }),
[state, getConnection, reconnect]
);

React.useEffect(() => {
Expand Down
12 changes: 11 additions & 1 deletion crates/bindings-typescript/src/react/connection_state.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,16 @@
import type { DbConnectionImpl } from '../sdk/db_connection_impl';
import type {
DbConnectionBuilder,
DbConnectionImpl,
} from '../sdk/db_connection_impl';
import type { ConnectionState as ManagerConnectionState } from '../sdk/connection_manager';

export type ConnectionState = ManagerConnectionState & {
getConnection(): DbConnectionImpl<any> | null;
/**
* Tear down the current connection and reconnect using a fresh builder, for
* example to switch identity after sign-in or sign-out. The builder should
* carry the new token and the same uri + database name. Hooks re-bind to the
* new connection automatically.
*/
reconnect(builder: DbConnectionBuilder<any>): void;
};
6 changes: 5 additions & 1 deletion crates/bindings-typescript/src/sdk/connection_manager.ts
Original file line number Diff line number Diff line change
Expand Up @@ -315,14 +315,18 @@ class ConnectionManagerImpl {
*
* Pass `resumeSession: false` when the caller is deliberately changing
* identity — see {@link rebuild} — so the builder's own token wins.
*
* A builder with a token provider is never overwritten: the provider is
* called for every attempt and already decides the identity, and resuming
* its last token would reuse a short-lived credential after it expired.
*/
#buildManagedConnection<T extends DbConnectionImpl<any>>(
managed: ManagedConnection,
builder: DbConnectionBuilder<T>,
{ resumeSession = true }: { resumeSession?: boolean } = {}
): T {
managed.builder = builder;
if (resumeSession && managed.state.token) {
if (resumeSession && managed.state.token && !builder.hasTokenProvider()) {
builder.withToken(managed.state.token);
}
const connection = builder.build();
Expand Down
25 changes: 22 additions & 3 deletions crates/bindings-typescript/src/sdk/db_connection_builder.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,8 @@
import { DbConnectionImpl, type ConnectionEvent } from './db_connection_impl';
import {
DbConnectionImpl,
type ConnectionEvent,
type TokenProvider,
} from './db_connection_impl';
import { EventEmitter } from './event_emitter';
import type {
DbConnectionConfig,
Expand All @@ -22,7 +26,7 @@ export class DbConnectionBuilder<DbConnection extends DbConnectionImpl<any>> {
#uri?: URL;
#nameOrAddress?: string;
#identity?: Identity;
#token?: string;
#token?: string | TokenProvider;
#emitter: EventEmitter<ConnectionEvent> = new EventEmitter();
#compression: 'gzip' | 'brotli' | 'none' = 'gzip';
#lightMode: boolean = false;
Expand Down Expand Up @@ -76,13 +80,28 @@ export class DbConnectionBuilder<DbConnection extends DbConnectionImpl<any>> {
* is optional. You can store the token returned by the `onConnect` callback
* to use in future connections.
*
* Pass a function instead of a string for short-lived tokens, such as OIDC
* session JWTs. It is called, and may return a promise, on every connection
* attempt, including automatic reconnects by the framework providers, so
* each attempt authenticates with a fresh token. Returning `undefined`
* connects anonymously.
*
* @returns The `DbConnectionBuilder` instance.
*/
withToken(token?: string): this {
withToken(token?: string | TokenProvider): this {
this.#token = token;
return this;
}

/**
* Whether the token was set as a {@link TokenProvider}. The
* `ConnectionManager` does not resume a session's token over a provider,
* which already decides the identity for each attempt.
*/
hasTokenProvider(): boolean {
return typeof this.#token === 'function';
}

withWSFn(createWSFn: WebSocketFactory): this {
this.#createWSFn = createWSFn;
return this;
Expand Down
50 changes: 38 additions & 12 deletions crates/bindings-typescript/src/sdk/db_connection_impl.ts
Original file line number Diff line number Diff line change
Expand Up @@ -98,11 +98,20 @@ export type {

export type ConnectionEvent = 'connect' | 'disconnect' | 'connectError';

/**
* Supplies the auth token for one connection attempt. See
* {@link DbConnectionBuilder.withToken}.
*/
export type TokenProvider = () =>
| string
| undefined
| Promise<string | undefined>;

export type DbConnectionConfig<RemoteModule extends UntypedRemoteModule> = {
uri: URL;
nameOrAddress: string;
identity?: Identity;
token?: string;
token?: string | TokenProvider;
emitter: EventEmitter<ConnectionEvent>;
createWSFn: WebSocketFactory;
compression: 'gzip' | 'brotli' | 'none';
Expand Down Expand Up @@ -327,7 +336,6 @@ export class DbConnectionImpl<RemoteModule extends UntypedRemoteModule>
}

this.identity = identity;
this.token = token;

this.#remoteModule = remoteModule;
this.#emitter = emitter;
Expand Down Expand Up @@ -388,15 +396,23 @@ export class DbConnectionImpl<RemoteModule extends UntypedRemoteModule>
this.reducers = this.#makeReducers(remoteModule);
this.procedures = this.#makeProcedures(remoteModule);

this.wsPromise = createWSFn({
url,
nameOrAddress,
wsProtocol: [...PREFERRED_WS_PROTOCOLS],
authToken: token,
compression: compression,
lightMode: lightMode,
confirmedReads: confirmedReads,
})
// A token provider is asked for a fresh token on every attempt, so
// short-lived credentials such as OIDC session JWTs survive reconnects.
// Resolving it inside the chain routes a throw or rejection to
// `onConnectError` rather than out of `build()`. A plain string token is
// assigned synchronously, as before.
this.wsPromise = (async () => {
this.token = typeof token === 'function' ? await token() : token;
return createWSFn({
url,
nameOrAddress,
wsProtocol: [...PREFERRED_WS_PROTOCOLS],
authToken: this.token,
compression: compression,
lightMode: lightMode,
confirmedReads: confirmedReads,
});
})()
.then(v => {
this.ws = v;

Expand All @@ -423,7 +439,7 @@ export class DbConnectionImpl<RemoteModule extends UntypedRemoteModule>
})
.catch(e => {
stdbLogger('error', 'Error connecting to SpacetimeDB WS');
this.#emitter.emit('connectError', this, e);
this.#emitter.emit('connectError', this, toError(e));

return undefined;
});
Expand Down Expand Up @@ -1160,6 +1176,16 @@ export class DbConnectionImpl<RemoteModule extends UntypedRemoteModule>
this.#processMessage(data);
}
}
} catch (e) {
// A message that fails to apply (undecodable, or a callback threw)
// leaves the client cache inconsistent, and rethrowing from a WebSocket
// listener crashes Node hosts. Fail the connection instead: drop what is
// queued and close, so `onDisconnect` receives the error.
stdbLogger('error', 'Failed to process a server message', e);
this.#connectionError = toError(e);
this.#inboundQueueOffset = this.#inboundQueue.length;
this.isActive = false;
this.ws?.close();
} finally {
if (this.#inboundQueueOffset >= this.#inboundQueue.length) {
this.#inboundQueue.length = 0;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,10 @@ class MockBuilder {
return this;
}

hasTokenProvider(): boolean {
return false;
}

build(): MockConnection {
const connection = new MockConnection();
connection.token = this.token;
Expand Down
Original file line number Diff line number Diff line change
@@ -1,10 +1,14 @@
import { afterEach, beforeEach, describe, expect, test, vi } from 'vitest';
import { ConnectionId } from '../src';
import { ConnectionId, type TokenProvider } from '../src';
import { ServerMessage } from '../src/sdk/client_api/types.ts';
import {
CONNECTION_MANAGER_RECONNECT_MAX_DELAY_MS,
connectionManagerReconnectDelayMs,
ConnectionManager,
} from '../src/sdk/connection_manager.ts';
import WebsocketTestAdapter from '../src/sdk/websocket_test_adapter.ts';
import { DbConnection } from '../test-app/src/module_bindings/index.ts';
import { anIdentity } from './utils.ts';

type ErrorContextInterface = {
isActive: boolean;
Expand Down Expand Up @@ -144,6 +148,10 @@ class MockBuilder {
return this;
}

hasTokenProvider(): boolean {
return false;
}

build(): MockConnection {
const connection = new MockConnection(this.token);
this.buildCount += 1;
Expand Down Expand Up @@ -809,3 +817,83 @@ describe('ConnectionManager session continuity across rebuilds', () => {
ConnectionManager.release(key);
});
});

// These use a real `DbConnection` because the provider is resolved inside it,
// per connection attempt.
describe('ConnectionManager with a token provider', () => {
beforeEach(() => {
vi.useFakeTimers();
});

afterEach(() => {
vi.runOnlyPendingTimers();
vi.useRealTimers();
});

function providerBuilder(provider: TokenProvider) {
const sockets: WebsocketTestAdapter[] = [];
const authTokens: (string | undefined)[] = [];
const builder = DbConnection.builder()
.withUri('ws://127.0.0.1:1234')
.withDatabaseName('db')
.withToken(provider)
.withWSFn(args => {
authTokens.push(args.authToken);
const socket = new WebsocketTestAdapter();
sockets.push(socket);
return socket.openWebSocket(args);
});
return { builder, sockets, authTokens };
}

test('asks the provider again on reconnect instead of resuming its last token', async () => {
const key = nextKey();
let issued = 0;
const { builder, sockets, authTokens } = providerBuilder(
async () => `jwt-${++issued}`
);

ConnectionManager.retain(key, builder);
await vi.advanceTimersByTimeAsync(0);
sockets[0].acceptConnection();
sockets[0].sendToClient(
ServerMessage.InitialConnection({
identity: anIdentity,
token: 'jwt-1',
connectionId: ConnectionId.random(),
})
);
expect(ConnectionManager.getSnapshot(key)?.token).toBe('jwt-1');

// By now the short-lived token may have expired; the rebuild must not
// resume it.
sockets[0].close();
await vi.advanceTimersByTimeAsync(connectionManagerReconnectDelayMs(0));

expect(authTokens).toEqual(['jwt-1', 'jwt-2']);

ConnectionManager.release(key);
});

test('retries after the provider rejects', async () => {
const key = nextKey();
const provider = vi
.fn<TokenProvider>()
.mockRejectedValueOnce(new Error('token fetch failed'))
.mockResolvedValue('jwt');
const { builder, authTokens } = providerBuilder(provider);

ConnectionManager.retain(key, builder);
await vi.advanceTimersByTimeAsync(0);
expect(ConnectionManager.getSnapshot(key)?.connectionError?.message).toBe(
'token fetch failed'
);
expect(authTokens).toEqual([]);

await vi.advanceTimersByTimeAsync(connectionManagerReconnectDelayMs(0));
expect(provider).toHaveBeenCalledTimes(2);
expect(authTokens).toEqual(['jwt']);

ConnectionManager.release(key);
});
});
Loading
Loading