@assistant-ui/react-mcp 0.1.16 → 0.1.17

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -7,6 +7,11 @@ import type {
7
7
  } from "@modelcontextprotocol/client";
8
8
  import type { MCPStorage } from "../resources/storage/types";
9
9
  import type { MCPAuthConfig } from "../mcp-scope";
10
+ import type { MCPPersistedAuthState } from "./types";
11
+ import {
12
+ isAuthStateForServerUrl,
13
+ normalizeMcpServerUrl,
14
+ } from "../utils/serverUrl";
10
15
 
11
16
  const STATE_PREFIX = "aui-mcp:";
12
17
 
@@ -49,6 +54,7 @@ export function decodeServerIdFromState(state: string): string | null {
49
54
 
50
55
  export type CreateOAuthProviderOptions = {
51
56
  serverId: string;
57
+ serverUrl: string;
52
58
  /** Must be `auth.type === "oauth"`. */
53
59
  config: Extract<MCPAuthConfig, { type: "oauth" }>;
54
60
  storage: MCPStorage;
@@ -58,16 +64,64 @@ export type CreateOAuthProviderOptions = {
58
64
  };
59
65
 
60
66
  type OAuthProviderCache = {
67
+ token?: string | undefined;
61
68
  tokens?: OAuthTokens | undefined;
69
+ tokensClientId?: string | undefined;
62
70
  clientInformation?: OAuthClientInformationFull | undefined;
71
+ clientInformationSource?: MCPPersistedAuthState["clientInformationSource"];
63
72
  codeVerifier?: string | undefined;
64
73
  state?: string | undefined;
65
74
  discoveryState?: OAuthDiscoveryState | undefined;
66
75
  };
67
76
 
68
- type OAuthProviderPersistence = {
77
+ type OAuthConfig = Extract<MCPAuthConfig, { type: "oauth" }>;
78
+
79
+ type OAuthCredentialState = {
80
+ tokens?: OAuthTokens | undefined;
81
+ tokensClientId?: string | undefined;
82
+ clientInformation?: OAuthClientInformationFull | undefined;
83
+ clientInformationSource?: "registered" | undefined;
84
+ };
85
+
86
+ const registeredClientId = (
87
+ state: OAuthCredentialState | null | undefined,
88
+ ): string | undefined =>
89
+ state?.clientInformationSource === "registered"
90
+ ? state.clientInformation?.client_id
91
+ : undefined;
92
+
93
+ export const hasUsableOAuthTokens = (
94
+ state: OAuthCredentialState | null | undefined,
95
+ config: OAuthConfig,
96
+ ): boolean => {
97
+ const clientId = config.clientId ?? registeredClientId(state);
98
+ return (
99
+ clientId !== undefined &&
100
+ state?.tokens !== undefined &&
101
+ state.tokensClientId === clientId
102
+ );
103
+ };
104
+
105
+ const hasUsableRegisteredClientInformation = (
106
+ state: OAuthCredentialState | null | undefined,
107
+ config: OAuthConfig,
108
+ ): boolean => {
109
+ const clientId = registeredClientId(state);
110
+ return (
111
+ clientId !== undefined &&
112
+ (config.clientId === undefined || config.clientId === clientId)
113
+ );
114
+ };
115
+
116
+ type OAuthProviderEndpointCache = {
117
+ serverUrl: string;
69
118
  cached: OAuthProviderCache | null;
70
119
  cachePromise: Promise<OAuthProviderCache> | null;
120
+ invalidated: boolean;
121
+ };
122
+
123
+ type OAuthProviderPersistence = {
124
+ endpoint: OAuthProviderEndpointCache | null;
71
125
  queue: Promise<void>;
72
126
  invalidated: boolean;
73
127
  };
@@ -111,7 +165,11 @@ const persistenceByIdentity = new WeakMap<
111
165
  const getPersistence = (
112
166
  storage: MCPStorage,
113
167
  serverId: string,
114
- ): OAuthProviderPersistence => {
168
+ serverUrl: string,
169
+ ): {
170
+ persistence: OAuthProviderPersistence;
171
+ endpoint: OAuthProviderEndpointCache;
172
+ } => {
115
173
  const identity = getStorageIdentity(storage);
116
174
  let byServerId = persistenceByIdentity.get(identity);
117
175
  if (!byServerId) {
@@ -122,14 +180,25 @@ const getPersistence = (
122
180
  let persistence = byServerId.get(serverId);
123
181
  if (!persistence) {
124
182
  persistence = {
125
- cached: null,
126
- cachePromise: null,
183
+ endpoint: null,
127
184
  queue: Promise.resolve(),
128
185
  invalidated: false,
129
186
  };
130
187
  byServerId.set(serverId, persistence);
131
188
  }
132
- return persistence;
189
+
190
+ let endpoint = persistence.endpoint;
191
+ if (endpoint?.serverUrl !== serverUrl) {
192
+ if (endpoint) endpoint.invalidated = true;
193
+ endpoint = {
194
+ serverUrl,
195
+ cached: null,
196
+ cachePromise: null,
197
+ invalidated: false,
198
+ };
199
+ persistence.endpoint = endpoint;
200
+ }
201
+ return { persistence, endpoint };
133
202
  };
134
203
 
135
204
  /**
@@ -155,7 +224,10 @@ export const clearOAuthProviderAuthState = async (
155
224
  byServerId.delete(serverId);
156
225
  if (byServerId.size === 0) persistenceByIdentity.delete(identity);
157
226
 
158
- await Promise.allSettled([persistence.cachePromise, persistence.queue]);
227
+ if (persistence.endpoint) persistence.endpoint.invalidated = true;
228
+ const cachePromise = persistence.endpoint?.cachePromise;
229
+ if (cachePromise) await Promise.allSettled([cachePromise]);
230
+ await persistence.queue;
159
231
  await storage.clearAuthState(serverId);
160
232
  };
161
233
 
@@ -167,15 +239,28 @@ export const clearOAuthProviderAuthState = async (
167
239
  export function createOAuthProvider(
168
240
  opts: CreateOAuthProviderOptions,
169
241
  ): OAuthClientProvider {
170
- const { serverId, config, storage, redirectUri, onAuthorizationUrl } = opts;
171
- const persistence = getPersistence(storage, serverId);
242
+ const {
243
+ serverId,
244
+ serverUrl,
245
+ config,
246
+ storage,
247
+ redirectUri,
248
+ onAuthorizationUrl,
249
+ } = opts;
250
+ const normalizedServerUrl = normalizeMcpServerUrl(serverUrl);
251
+ const { persistence, endpoint } = getPersistence(
252
+ storage,
253
+ serverId,
254
+ normalizedServerUrl,
255
+ );
172
256
  let pendingState: string | undefined;
173
257
 
174
- // The cache is shared with every other provider for this (storage, serverId),
175
- // so a statically configured client stays a read-time overlay owned by this
176
- // provider. Writing it into the cache would leak this provider's registration
177
- // to a replacement built for a different, or absent, clientId.
178
- const staticClientInformation = (():
258
+ // The cache is shared with every other provider for this storage, server id,
259
+ // and server URL, so a statically configured client stays a read-time overlay
260
+ // owned by this provider. Writing it into the cache would leak this provider's
261
+ // registration to a replacement built for a different, or absent, clientId.
262
+ // The SDK's write-backs, its issuer stamp included, replace the overlay.
263
+ const configuredClientInformation = ():
179
264
  | OAuthClientInformationFull
180
265
  | undefined => {
181
266
  if (!config.clientId) return undefined;
@@ -185,50 +270,92 @@ export function createOAuthProvider(
185
270
  };
186
271
  if (config.clientSecret) ci.client_secret = config.clientSecret;
187
272
  return ci;
188
- })();
273
+ };
274
+ let clientInformationOverlay = configuredClientInformation();
275
+
276
+ const activeClientId = (cache: OAuthProviderCache): string | undefined =>
277
+ clientInformationOverlay?.client_id ?? registeredClientId(cache);
189
278
 
190
279
  const loadCache = (): Promise<OAuthProviderCache> => {
191
- if (persistence.cached) return Promise.resolve(persistence.cached);
192
- if (persistence.cachePromise) return persistence.cachePromise;
193
-
194
- persistence.cachePromise = storage.loadAuthState(serverId).then(
195
- (persisted) => {
196
- const initial: OAuthProviderCache = {};
197
- if (persisted?.tokens) initial.tokens = persisted.tokens;
198
- if (persisted?.clientInformation)
199
- initial.clientInformation = persisted.clientInformation;
200
- if (persisted?.codeVerifier)
201
- initial.codeVerifier = persisted.codeVerifier;
202
- if (persisted?.state) initial.state = persisted.state;
203
- if (persisted?.discoveryState)
204
- initial.discoveryState = persisted.discoveryState;
205
- persistence.cached = initial;
206
- return initial;
207
- },
208
- (error) => {
209
- persistence.cachePromise = null;
210
- throw error;
211
- },
212
- );
213
- return persistence.cachePromise;
280
+ if (endpoint.invalidated) return Promise.resolve({});
281
+ if (endpoint.cached) return Promise.resolve(endpoint.cached);
282
+ if (endpoint.cachePromise) return endpoint.cachePromise;
283
+
284
+ endpoint.cachePromise = persistence.queue
285
+ .then(() => storage.loadAuthState(serverId))
286
+ .then(
287
+ async (persisted) => {
288
+ const initial: OAuthProviderCache = {};
289
+ let needsMigration = false;
290
+ if (endpoint.invalidated) return initial;
291
+ if (
292
+ persisted &&
293
+ isAuthStateForServerUrl(persisted, normalizedServerUrl)
294
+ ) {
295
+ if (
296
+ hasUsableRegisteredClientInformation(persisted, config) &&
297
+ persisted.clientInformation
298
+ ) {
299
+ initial.clientInformation = persisted.clientInformation;
300
+ initial.clientInformationSource = "registered";
301
+ } else if (
302
+ persisted?.clientInformation ||
303
+ persisted?.clientInformationSource !== undefined
304
+ ) {
305
+ needsMigration = true;
306
+ }
307
+ if (hasUsableOAuthTokens(persisted, config)) {
308
+ initial.tokens = persisted.tokens;
309
+ initial.tokensClientId = persisted.tokensClientId;
310
+ } else if (
311
+ persisted?.tokens ||
312
+ persisted?.tokensClientId !== undefined
313
+ ) {
314
+ needsMigration = true;
315
+ }
316
+ if (persisted?.token) initial.token = persisted.token;
317
+ if (persisted?.codeVerifier)
318
+ initial.codeVerifier = persisted.codeVerifier;
319
+ if (persisted?.state) initial.state = persisted.state;
320
+ if (persisted?.discoveryState)
321
+ initial.discoveryState = persisted.discoveryState;
322
+ }
323
+ endpoint.cached = initial;
324
+ if (needsMigration) await persist().catch(() => {});
325
+ return initial;
326
+ },
327
+ (error) => {
328
+ endpoint.cachePromise = null;
329
+ throw error;
330
+ },
331
+ );
332
+ return endpoint.cachePromise;
214
333
  };
215
334
 
216
- const persist = () => {
335
+ function persist() {
217
336
  const task = persistence.queue.then(async () => {
218
- if (persistence.invalidated) return;
219
- const c = persistence.cached;
337
+ if (persistence.invalidated || endpoint.invalidated) return;
338
+ const c = endpoint.cached;
220
339
  if (!c) return;
221
340
  const next: Parameters<typeof storage.saveAuthState>[1] = {};
222
- if (c.tokens) next.tokens = c.tokens;
223
- if (c.clientInformation) next.clientInformation = c.clientInformation;
341
+ if (hasUsableOAuthTokens(c, config) && c.tokens && c.tokensClientId) {
342
+ next.tokens = c.tokens;
343
+ next.tokensClientId = c.tokensClientId;
344
+ }
345
+ if (c.clientInformation && c.clientInformationSource === "registered") {
346
+ next.clientInformation = c.clientInformation;
347
+ next.clientInformationSource = "registered";
348
+ }
349
+ if (c.token) next.token = c.token;
224
350
  if (c.codeVerifier) next.codeVerifier = c.codeVerifier;
225
351
  if (c.state) next.state = c.state;
226
352
  if (c.discoveryState) next.discoveryState = c.discoveryState;
353
+ next.serverUrl = normalizedServerUrl;
227
354
  await storage.saveAuthState(serverId, next);
228
355
  });
229
356
  persistence.queue = task.catch(() => {});
230
357
  return task;
231
- };
358
+ }
232
359
 
233
360
  const clientMetadata: OAuthClientMetadata = {
234
361
  client_name: "assistant-ui",
@@ -258,20 +385,37 @@ export function createOAuthProvider(
258
385
  },
259
386
  async clientInformation() {
260
387
  const c = await loadCache();
261
- return staticClientInformation ?? c.clientInformation;
388
+ if (clientInformationOverlay) return clientInformationOverlay;
389
+ if (c.clientInformationSource !== "registered") return undefined;
390
+ return c.clientInformation;
262
391
  },
263
392
  async saveClientInformation(info) {
393
+ if (clientInformationOverlay) {
394
+ clientInformationOverlay = info as OAuthClientInformationFull;
395
+ return;
396
+ }
264
397
  const c = await loadCache();
265
398
  c.clientInformation = info as OAuthClientInformationFull;
399
+ c.clientInformationSource = "registered";
400
+ if (c.tokensClientId !== c.clientInformation.client_id) {
401
+ delete c.tokens;
402
+ delete c.tokensClientId;
403
+ }
266
404
  await persist();
267
405
  },
268
406
  async tokens() {
269
407
  const c = await loadCache();
408
+ const clientId = activeClientId(c);
409
+ if (clientId === undefined || c.tokensClientId !== clientId)
410
+ return undefined;
270
411
  return c.tokens;
271
412
  },
272
413
  async saveTokens(tokens) {
273
414
  const c = await loadCache();
274
415
  c.tokens = tokens;
416
+ const clientId = activeClientId(c);
417
+ if (clientId) c.tokensClientId = clientId;
418
+ else delete c.tokensClientId;
275
419
  delete c.state;
276
420
  await persist();
277
421
  },
@@ -305,8 +449,19 @@ export function createOAuthProvider(
305
449
  },
306
450
  async invalidateCredentials(scope) {
307
451
  const c = await loadCache();
308
- if (scope === "all" || scope === "tokens") delete c.tokens;
309
- if (scope === "all" || scope === "client") delete c.clientInformation;
452
+ if (scope === "all" || scope === "tokens") {
453
+ delete c.tokens;
454
+ delete c.tokensClientId;
455
+ }
456
+ if (scope === "all" || scope === "client") {
457
+ delete c.clientInformation;
458
+ delete c.clientInformationSource;
459
+ if (!config.clientId) {
460
+ delete c.tokens;
461
+ delete c.tokensClientId;
462
+ }
463
+ clientInformationOverlay = configuredClientInformation();
464
+ }
310
465
  if (scope === "all" || scope === "verifier") {
311
466
  delete c.codeVerifier;
312
467
  delete c.state;
package/src/auth/types.ts CHANGED
@@ -5,11 +5,15 @@ import type {
5
5
  } from "@modelcontextprotocol/client";
6
6
 
7
7
  export type MCPPersistedAuthState = {
8
+ /** MCP server URL this authentication state belongs to. Required with credentials. */
9
+ serverUrl?: string;
8
10
  tokens?: OAuthTokens;
11
+ tokensClientId?: string;
9
12
  clientInformation?: OAuthClientInformationFull;
13
+ clientInformationSource?: "registered";
10
14
  codeVerifier?: string;
11
15
  state?: string;
12
16
  discoveryState?: OAuthDiscoveryState;
13
- /** Bearer token (entered at add-form time). */
17
+ /** Host-persisted bearer token. Must be paired with serverUrl. */
14
18
  token?: string;
15
19
  };
@@ -120,6 +120,9 @@ const useMcpManagerResource = (
120
120
 
121
121
  useEffect(() => {
122
122
  const signal = { cancelled: false };
123
+ // Hydration reads persisted records asynchronously; there is no earlier
124
+ // point than mount at which to start it.
125
+ // eslint-disable-next-line react-hooks/set-state-in-effect
123
126
  void hydrate(signal);
124
127
  return () => {
125
128
  signal.cancelled = true;