@assistant-ui/react-mcp 0.1.14 → 0.1.16

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.
Files changed (52) hide show
  1. package/dist/auth/createOAuthProvider.d.ts +7 -1
  2. package/dist/auth/createOAuthProvider.d.ts.map +1 -1
  3. package/dist/auth/createOAuthProvider.js +115 -31
  4. package/dist/auth/createOAuthProvider.js.map +1 -1
  5. package/dist/auth/types.d.ts +1 -0
  6. package/dist/auth/types.d.ts.map +1 -1
  7. package/dist/hooks/useMcpOAuthCallback.d.ts.map +1 -1
  8. package/dist/hooks/useMcpOAuthCallback.js +5 -6
  9. package/dist/hooks/useMcpOAuthCallback.js.map +1 -1
  10. package/dist/primitives/server/McpServerIcon.js.map +1 -1
  11. package/dist/primitives/server/McpServerOAuthLink.js.map +1 -1
  12. package/dist/resources/McpManagerResource.d.ts.map +1 -1
  13. package/dist/resources/McpManagerResource.js +4 -2
  14. package/dist/resources/McpManagerResource.js.map +1 -1
  15. package/dist/resources/McpServerResource.d.ts +2 -1
  16. package/dist/resources/McpServerResource.d.ts.map +1 -1
  17. package/dist/resources/McpServerResource.js +107 -21
  18. package/dist/resources/McpServerResource.js.map +1 -1
  19. package/dist/resources/storage/McpLocalStorage.d.ts +7 -0
  20. package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
  21. package/dist/resources/storage/McpLocalStorage.js +132 -37
  22. package/dist/resources/storage/McpLocalStorage.js.map +1 -1
  23. package/dist/resources/storage/McpMemoryStorage.d.ts.map +1 -1
  24. package/dist/resources/storage/McpMemoryStorage.js +9 -6
  25. package/dist/resources/storage/McpMemoryStorage.js.map +1 -1
  26. package/dist/resources/storage/types.d.ts +12 -0
  27. package/dist/resources/storage/types.d.ts.map +1 -1
  28. package/dist/utils/createMcpId.d.ts +9 -0
  29. package/dist/utils/createMcpId.d.ts.map +1 -0
  30. package/dist/utils/createMcpId.js +11 -0
  31. package/dist/utils/createMcpId.js.map +1 -0
  32. package/dist/utils/invokeMcpCallback.d.ts.map +1 -1
  33. package/dist/utils/invokeMcpCallback.js +2 -12
  34. package/dist/utils/invokeMcpCallback.js.map +1 -1
  35. package/package.json +8 -8
  36. package/src/auth/createOAuthProvider.test.ts +407 -2
  37. package/src/auth/createOAuthProvider.ts +171 -40
  38. package/src/auth/types.ts +1 -0
  39. package/src/hooks/useMcpOAuthCallback.test.ts +74 -1
  40. package/src/hooks/useMcpOAuthCallback.tsx +11 -8
  41. package/src/resources/McpManagerResource.test.ts +128 -0
  42. package/src/resources/McpManagerResource.ts +4 -5
  43. package/src/resources/McpServerResource.test.ts +420 -18
  44. package/src/resources/McpServerResource.ts +148 -27
  45. package/src/resources/storage/McpLocalStorage.test.ts +71 -1
  46. package/src/resources/storage/McpLocalStorage.ts +69 -47
  47. package/src/resources/storage/McpMemoryStorage.test.ts +99 -0
  48. package/src/resources/storage/McpMemoryStorage.ts +23 -17
  49. package/src/resources/storage/types.ts +12 -0
  50. package/src/utils/createMcpId.test.ts +25 -0
  51. package/src/utils/createMcpId.ts +10 -0
  52. package/src/utils/invokeMcpCallback.ts +3 -21
@@ -1,8 +1,11 @@
1
1
  import type { OAuthDiscoveryState } from "@modelcontextprotocol/client";
2
- import { describe, expect, it } from "vitest";
2
+ import { describe, expect, it, vi } from "vitest";
3
3
  import type { MCPStorage } from "../resources/storage/types";
4
4
  import type { MCPPersistedAuthState } from "./types";
5
- import { createOAuthProvider } from "./createOAuthProvider";
5
+ import {
6
+ clearOAuthProviderAuthState,
7
+ createOAuthProvider,
8
+ } from "./createOAuthProvider";
6
9
 
7
10
  const discoveryState: OAuthDiscoveryState = {
8
11
  authorizationServerUrl: "https://auth.example.com",
@@ -36,6 +39,23 @@ const createStorage = (initial: MCPPersistedAuthState | null = null) => {
36
39
  return { storage, getState: () => state };
37
40
  };
38
41
 
42
+ const createSharedStorages = (scopeId: string) => {
43
+ let state: MCPPersistedAuthState | null = null;
44
+ const create = (): MCPStorage => ({
45
+ scopeId,
46
+ loadCustomServers: async () => [],
47
+ saveCustomServers: async () => {},
48
+ loadAuthState: async () => state,
49
+ saveAuthState: async (_serverId, next) => {
50
+ state = next;
51
+ },
52
+ clearAuthState: async () => {
53
+ state = null;
54
+ },
55
+ });
56
+ return { create, getState: () => state };
57
+ };
58
+
39
59
  const createProvider = (storage: MCPStorage) =>
40
60
  createOAuthProvider({
41
61
  serverId: "docs",
@@ -45,6 +65,55 @@ const createProvider = (storage: MCPStorage) =>
45
65
  onAuthorizationUrl: () => {},
46
66
  });
47
67
 
68
+ describe("createOAuthProvider callback state", () => {
69
+ it("persists the generated state with the PKCE verifier", async () => {
70
+ const { storage, getState } = createStorage();
71
+ const provider = createProvider(storage);
72
+
73
+ const state = await provider.state?.();
74
+ await provider.saveCodeVerifier("pkce-verifier");
75
+
76
+ expect(state).toMatch(/^aui-mcp:ZG9jcw\./);
77
+ expect(getState()).toEqual({
78
+ codeVerifier: "pkce-verifier",
79
+ state,
80
+ });
81
+ });
82
+
83
+ it("consumes callback state when tokens are saved", async () => {
84
+ const { storage, getState } = createStorage({
85
+ codeVerifier: "pkce-verifier",
86
+ state: "aui-mcp:ZG9jcw.nonce",
87
+ });
88
+ const provider = createProvider(storage);
89
+
90
+ await provider.saveTokens({
91
+ access_token: "access-token",
92
+ token_type: "bearer",
93
+ });
94
+
95
+ expect(getState()).toEqual({
96
+ tokens: { access_token: "access-token", token_type: "bearer" },
97
+ codeVerifier: "pkce-verifier",
98
+ });
99
+ });
100
+
101
+ it.each(["verifier", "all"] as const)(
102
+ "clears callback state through the %s invalidation scope",
103
+ async (scope) => {
104
+ const { storage, getState } = createStorage({
105
+ codeVerifier: "pkce-verifier",
106
+ state: "aui-mcp:ZG9jcw.nonce",
107
+ });
108
+ const provider = createProvider(storage);
109
+
110
+ await provider.invalidateCredentials?.(scope);
111
+
112
+ expect(getState()).toEqual({});
113
+ },
114
+ );
115
+ });
116
+
48
117
  describe("createOAuthProvider discovery state", () => {
49
118
  it("persists discovery state alongside the PKCE verifier", async () => {
50
119
  const { storage, getState } = createStorage({
@@ -84,3 +153,339 @@ describe("createOAuthProvider discovery state", () => {
84
153
  },
85
154
  );
86
155
  });
156
+
157
+ describe("createOAuthProvider persistence", () => {
158
+ it("loads persisted auth state once for concurrent reads", async () => {
159
+ let resolveLoad!: (value: MCPPersistedAuthState | null) => void;
160
+ const loadAuthState = vi.fn(
161
+ () =>
162
+ new Promise<MCPPersistedAuthState | null>((resolve) => {
163
+ resolveLoad = resolve;
164
+ }),
165
+ );
166
+ const { storage } = createStorage();
167
+ storage.loadAuthState = loadAuthState;
168
+ const provider = createProvider(storage);
169
+
170
+ const tokens = provider.tokens();
171
+ const clientInformation = provider.clientInformation();
172
+
173
+ expect(loadAuthState).toHaveBeenCalledTimes(1);
174
+ resolveLoad(null);
175
+ await Promise.all([tokens, clientInformation]);
176
+ });
177
+
178
+ it("retries loading persisted auth state after a failure", async () => {
179
+ const failure = new Error("storage unavailable");
180
+ const loadAuthState = vi
181
+ .fn<() => Promise<MCPPersistedAuthState | null>>()
182
+ .mockRejectedValueOnce(failure)
183
+ .mockResolvedValueOnce({ codeVerifier: "pkce-verifier" });
184
+ const { storage } = createStorage();
185
+ storage.loadAuthState = loadAuthState;
186
+ const provider = createProvider(storage);
187
+
188
+ const tokens = provider.tokens();
189
+ const clientInformation = provider.clientInformation();
190
+
191
+ await expect(tokens).rejects.toBe(failure);
192
+ await expect(clientInformation).rejects.toBe(failure);
193
+ expect(loadAuthState).toHaveBeenCalledTimes(1);
194
+
195
+ await expect(provider.codeVerifier()).resolves.toBe("pkce-verifier");
196
+ expect(loadAuthState).toHaveBeenCalledTimes(2);
197
+ });
198
+
199
+ it("serializes writes so newer auth state is not overwritten", async () => {
200
+ const { storage } = createStorage();
201
+ const pendingWrites: Array<() => void> = [];
202
+ let persisted: MCPPersistedAuthState | null = null;
203
+ storage.saveAuthState = async (_serverId, next) => {
204
+ await new Promise<void>((resolve) => pendingWrites.push(resolve));
205
+ persisted = next;
206
+ };
207
+ const provider = createProvider(storage);
208
+ await provider.tokens();
209
+
210
+ const tokenSave = provider.saveTokens({
211
+ access_token: "access-token",
212
+ token_type: "bearer",
213
+ });
214
+ await vi.waitFor(() => expect(pendingWrites).toHaveLength(1));
215
+
216
+ const verifierSave = provider.saveCodeVerifier("pkce-verifier");
217
+ for (let i = 0; i < 10; i++) await Promise.resolve();
218
+ expect(pendingWrites).toHaveLength(1);
219
+
220
+ pendingWrites.shift()!();
221
+ await vi.waitFor(() => expect(pendingWrites).toHaveLength(1));
222
+ pendingWrites.shift()!();
223
+ await Promise.all([tokenSave, verifierSave]);
224
+
225
+ expect(persisted).toEqual({
226
+ tokens: { access_token: "access-token", token_type: "bearer" },
227
+ codeVerifier: "pkce-verifier",
228
+ });
229
+ });
230
+
231
+ it("continues persisting after a failed auth state write", async () => {
232
+ const { storage } = createStorage();
233
+ const failure = new Error("storage unavailable");
234
+ let saveCount = 0;
235
+ let persisted: MCPPersistedAuthState | null = null;
236
+ let rejectFirstSave!: (reason: unknown) => void;
237
+ storage.saveAuthState = async (_serverId, next) => {
238
+ saveCount += 1;
239
+ if (saveCount === 1) {
240
+ await new Promise<void>((_resolve, reject) => {
241
+ rejectFirstSave = reject;
242
+ });
243
+ }
244
+ persisted = next;
245
+ };
246
+ const provider = createProvider(storage);
247
+ await provider.tokens();
248
+
249
+ const tokenSave = provider.saveTokens({
250
+ access_token: "access-token",
251
+ token_type: "bearer",
252
+ });
253
+ const verifierSave = provider.saveCodeVerifier("pkce-verifier");
254
+
255
+ await vi.waitFor(() => expect(saveCount).toBe(1));
256
+ for (let i = 0; i < 10; i++) await Promise.resolve();
257
+ expect(saveCount).toBe(1);
258
+
259
+ const tokenSaveResult = expect(tokenSave).rejects.toBe(failure);
260
+ rejectFirstSave(failure);
261
+ await tokenSaveResult;
262
+ await expect(verifierSave).resolves.toBeUndefined();
263
+ expect(saveCount).toBe(2);
264
+ expect(persisted).toEqual({
265
+ tokens: { access_token: "access-token", token_type: "bearer" },
266
+ codeVerifier: "pkce-verifier",
267
+ });
268
+ });
269
+ });
270
+
271
+ describe("createOAuthProvider persistence across provider instances", () => {
272
+ it("shares one auth state load across provider instances", async () => {
273
+ let resolveLoad!: (value: MCPPersistedAuthState | null) => void;
274
+ const loadAuthState = vi.fn(
275
+ () =>
276
+ new Promise<MCPPersistedAuthState | null>((resolve) => {
277
+ resolveLoad = resolve;
278
+ }),
279
+ );
280
+ const { storage } = createStorage();
281
+ storage.loadAuthState = loadAuthState;
282
+ const provider = createProvider(storage);
283
+ const replacementProvider = createProvider(storage);
284
+
285
+ const tokens = provider.tokens();
286
+ const clientInformation = replacementProvider.clientInformation();
287
+
288
+ expect(loadAuthState).toHaveBeenCalledTimes(1);
289
+ resolveLoad(null);
290
+ await Promise.all([tokens, clientInformation]);
291
+ });
292
+
293
+ it("serializes writes across provider instances", async () => {
294
+ const { storage } = createStorage();
295
+ const pendingWrites: Array<() => void> = [];
296
+ let persisted: MCPPersistedAuthState | null = null;
297
+ storage.saveAuthState = async (_serverId, next) => {
298
+ await new Promise<void>((resolve) => pendingWrites.push(resolve));
299
+ persisted = next;
300
+ };
301
+ const provider = createProvider(storage);
302
+ const replacementProvider = createProvider(storage);
303
+ await Promise.all([provider.tokens(), replacementProvider.tokens()]);
304
+
305
+ const tokenSave = provider.saveTokens({
306
+ access_token: "access-token",
307
+ token_type: "bearer",
308
+ });
309
+ await vi.waitFor(() => expect(pendingWrites).toHaveLength(1));
310
+
311
+ const verifierSave = replacementProvider.saveCodeVerifier("pkce-verifier");
312
+ for (let i = 0; i < 10; i++) await Promise.resolve();
313
+ expect(pendingWrites).toHaveLength(1);
314
+
315
+ pendingWrites.shift()!();
316
+ await vi.waitFor(() => expect(pendingWrites).toHaveLength(1));
317
+ pendingWrites.shift()!();
318
+ await Promise.all([tokenSave, verifierSave]);
319
+
320
+ expect(persisted).toEqual({
321
+ tokens: { access_token: "access-token", token_type: "bearer" },
322
+ codeVerifier: "pkce-verifier",
323
+ });
324
+ });
325
+
326
+ it("clears after a pending write and fences the discarded provider", async () => {
327
+ const { storage, getState } = createStorage();
328
+ const pendingWrites: Array<() => void> = [];
329
+ const saveAuthState = storage.saveAuthState;
330
+ storage.saveAuthState = async (serverId, next) => {
331
+ await new Promise<void>((resolve) => pendingWrites.push(resolve));
332
+ await saveAuthState(serverId, next);
333
+ };
334
+ const clearAuthState = vi.spyOn(storage, "clearAuthState");
335
+ const provider = createProvider(storage);
336
+ await provider.tokens();
337
+
338
+ const save = provider.saveTokens({
339
+ access_token: "access-token",
340
+ token_type: "bearer",
341
+ });
342
+ await vi.waitFor(() => expect(pendingWrites).toHaveLength(1));
343
+
344
+ const clear = clearOAuthProviderAuthState(storage, "docs");
345
+ for (let i = 0; i < 10; i++) await Promise.resolve();
346
+ expect(clearAuthState).not.toHaveBeenCalled();
347
+
348
+ pendingWrites.shift()!();
349
+ await expect(save).resolves.toBeUndefined();
350
+ await clear;
351
+ expect(clearAuthState).toHaveBeenCalledTimes(1);
352
+ expect(getState()).toBeNull();
353
+
354
+ storage.saveAuthState = saveAuthState;
355
+ await provider.saveCodeVerifier("late-verifier");
356
+ expect(getState()).toBeNull();
357
+ });
358
+
359
+ it("fences a queued write when clearing through a same-scope storage", async () => {
360
+ const { create, getState } = createSharedStorages("same-scope-clear");
361
+ const storage = create();
362
+ const replacement = create();
363
+ const pendingWrites: Array<() => void> = [];
364
+ const saveAuthState = storage.saveAuthState;
365
+ storage.saveAuthState = async (serverId, next) => {
366
+ await new Promise<void>((resolve) => pendingWrites.push(resolve));
367
+ await saveAuthState(serverId, next);
368
+ };
369
+ const clearAuthState = vi.spyOn(replacement, "clearAuthState");
370
+ const provider = createProvider(storage);
371
+ await provider.tokens();
372
+
373
+ const save = provider.saveTokens({
374
+ access_token: "access-token",
375
+ token_type: "bearer",
376
+ });
377
+ await vi.waitFor(() => expect(pendingWrites).toHaveLength(1));
378
+
379
+ const clear = clearOAuthProviderAuthState(replacement, "docs");
380
+ for (let i = 0; i < 10; i++) await Promise.resolve();
381
+ expect(clearAuthState).not.toHaveBeenCalled();
382
+
383
+ pendingWrites.shift()!();
384
+ await expect(save).resolves.toBeUndefined();
385
+ await clear;
386
+ expect(clearAuthState).toHaveBeenCalledTimes(1);
387
+ expect(getState()).toBeNull();
388
+
389
+ storage.saveAuthState = saveAuthState;
390
+ await provider.saveCodeVerifier("late-verifier");
391
+ expect(getState()).toBeNull();
392
+ });
393
+
394
+ it("keeps differently-scoped storages on separate persistence", async () => {
395
+ const first = createSharedStorages("scope-a");
396
+ const second = createSharedStorages("scope-b");
397
+ const firstProvider = createProvider(first.create());
398
+ const secondProvider = createProvider(second.create());
399
+
400
+ await firstProvider.saveTokens({
401
+ access_token: "first-token",
402
+ token_type: "bearer",
403
+ });
404
+ await secondProvider.saveTokens({
405
+ access_token: "second-token",
406
+ token_type: "bearer",
407
+ });
408
+ await clearOAuthProviderAuthState(second.create(), "docs");
409
+ await firstProvider.saveCodeVerifier("first-verifier");
410
+
411
+ expect(first.getState()).toEqual({
412
+ tokens: { access_token: "first-token", token_type: "bearer" },
413
+ codeVerifier: "first-verifier",
414
+ });
415
+ expect(second.getState()).toBeNull();
416
+ });
417
+
418
+ it("re-derives static client information for a replacement provider", async () => {
419
+ const { storage } = createStorage();
420
+ const provider = createOAuthProvider({
421
+ serverId: "docs",
422
+ config: { type: "oauth", clientId: "client-a" },
423
+ storage,
424
+ redirectUri: "http://localhost/callback",
425
+ onAuthorizationUrl: () => {},
426
+ });
427
+ await expect(provider.clientInformation()).resolves.toEqual({
428
+ client_id: "client-a",
429
+ redirect_uris: ["http://localhost/callback"],
430
+ });
431
+
432
+ const replacementProvider = createOAuthProvider({
433
+ serverId: "docs",
434
+ config: { type: "oauth", clientId: "client-b", clientSecret: "secret-b" },
435
+ storage,
436
+ redirectUri: "http://localhost/callback-2",
437
+ onAuthorizationUrl: () => {},
438
+ });
439
+ await expect(replacementProvider.clientInformation()).resolves.toEqual({
440
+ client_id: "client-b",
441
+ client_secret: "secret-b",
442
+ redirect_uris: ["http://localhost/callback-2"],
443
+ });
444
+ });
445
+
446
+ it("does not leak static client information to a dynamic provider", async () => {
447
+ const { storage } = createStorage();
448
+ const staticProvider = createOAuthProvider({
449
+ serverId: "docs",
450
+ config: { type: "oauth", clientId: "client-a" },
451
+ storage,
452
+ redirectUri: "http://localhost/callback",
453
+ onAuthorizationUrl: () => {},
454
+ });
455
+ await expect(staticProvider.clientInformation()).resolves.toEqual({
456
+ client_id: "client-a",
457
+ redirect_uris: ["http://localhost/callback"],
458
+ });
459
+
460
+ const dynamicProvider = createProvider(storage);
461
+ await expect(dynamicProvider.clientInformation()).resolves.toBeUndefined();
462
+ });
463
+
464
+ it("keeps a provider built while the clear is in flight usable", async () => {
465
+ const { storage, getState } = createStorage();
466
+ let releaseClear: (() => void) | undefined;
467
+ const clearAuthState = storage.clearAuthState;
468
+ storage.clearAuthState = async (serverId) => {
469
+ await new Promise<void>((resolve) => {
470
+ releaseClear = resolve;
471
+ });
472
+ await clearAuthState(serverId);
473
+ };
474
+ const provider = createProvider(storage);
475
+ await provider.saveTokens({
476
+ access_token: "access-token",
477
+ token_type: "bearer",
478
+ });
479
+
480
+ const clear = clearOAuthProviderAuthState(storage, "docs");
481
+ await vi.waitFor(() => expect(releaseClear).toBeTypeOf("function"));
482
+
483
+ const replacementProvider = createProvider(storage);
484
+ releaseClear!();
485
+ await clear;
486
+ expect(getState()).toBeNull();
487
+
488
+ await replacementProvider.saveCodeVerifier("new-verifier");
489
+ expect(getState()).toEqual({ codeVerifier: "new-verifier" });
490
+ });
491
+ });
@@ -57,6 +57,108 @@ export type CreateOAuthProviderOptions = {
57
57
  onAuthorizationUrl: (url: URL) => void;
58
58
  };
59
59
 
60
+ type OAuthProviderCache = {
61
+ tokens?: OAuthTokens | undefined;
62
+ clientInformation?: OAuthClientInformationFull | undefined;
63
+ codeVerifier?: string | undefined;
64
+ state?: string | undefined;
65
+ discoveryState?: OAuthDiscoveryState | undefined;
66
+ };
67
+
68
+ type OAuthProviderPersistence = {
69
+ cached: OAuthProviderCache | null;
70
+ cachePromise: Promise<OAuthProviderCache> | null;
71
+ queue: Promise<void>;
72
+ invalidated: boolean;
73
+ };
74
+
75
+ // scopeId, not object identity, is what addresses the same persisted data, so
76
+ // storages sharing one share an anchor and an unscoped storage is its own
77
+ // identity. Every storage declaring a scope holds that scope's anchor, so the
78
+ // coordination state below is collected once the last of them is gone.
79
+ const anchorByStorage = new WeakMap<MCPStorage, object>();
80
+ const anchorByScope = new Map<string, WeakRef<object>>();
81
+ const anchorRegistry = new FinalizationRegistry<string>((scopeId) => {
82
+ if (!anchorByScope.get(scopeId)?.deref()) anchorByScope.delete(scopeId);
83
+ });
84
+
85
+ const getStorageIdentity = (storage: MCPStorage): object => {
86
+ const existing = anchorByStorage.get(storage);
87
+ if (existing) return existing;
88
+
89
+ const { scopeId } = storage;
90
+ if (scopeId === undefined) return storage;
91
+
92
+ let anchor = anchorByScope.get(scopeId)?.deref();
93
+ if (!anchor) {
94
+ anchor = {};
95
+ anchorByScope.set(scopeId, new WeakRef(anchor));
96
+ anchorRegistry.register(anchor, scopeId);
97
+ }
98
+ anchorByStorage.set(storage, anchor);
99
+ return anchor;
100
+ };
101
+
102
+ // McpServerResource builds a fresh provider for every transport, so the cache,
103
+ // the in-flight load, and the write queue have to outlive any one provider.
104
+ // saveAuthState replaces the whole record, so two providers writing their own
105
+ // snapshots concurrently would drop whichever field the loser had added.
106
+ const persistenceByIdentity = new WeakMap<
107
+ object,
108
+ Map<string, OAuthProviderPersistence>
109
+ >();
110
+
111
+ const getPersistence = (
112
+ storage: MCPStorage,
113
+ serverId: string,
114
+ ): OAuthProviderPersistence => {
115
+ const identity = getStorageIdentity(storage);
116
+ let byServerId = persistenceByIdentity.get(identity);
117
+ if (!byServerId) {
118
+ byServerId = new Map();
119
+ persistenceByIdentity.set(identity, byServerId);
120
+ }
121
+
122
+ let persistence = byServerId.get(serverId);
123
+ if (!persistence) {
124
+ persistence = {
125
+ cached: null,
126
+ cachePromise: null,
127
+ queue: Promise.resolve(),
128
+ invalidated: false,
129
+ };
130
+ byServerId.set(serverId, persistence);
131
+ }
132
+ return persistence;
133
+ };
134
+
135
+ /**
136
+ * Clears persisted OAuth state after the in-flight load and every queued write
137
+ * for that server have settled, so a discarded provider cannot recreate the
138
+ * record it was mid-save on.
139
+ */
140
+ export const clearOAuthProviderAuthState = async (
141
+ storage: MCPStorage,
142
+ serverId: string,
143
+ ): Promise<void> => {
144
+ const identity = getStorageIdentity(storage);
145
+ const byServerId = persistenceByIdentity.get(identity);
146
+ const persistence = byServerId?.get(serverId);
147
+ if (!byServerId || !persistence) {
148
+ await storage.clearAuthState(serverId);
149
+ return;
150
+ }
151
+
152
+ // Detaching the entry before awaiting keeps a provider built during the clear
153
+ // on a fresh generation instead of inheriting the fenced one.
154
+ persistence.invalidated = true;
155
+ byServerId.delete(serverId);
156
+ if (byServerId.size === 0) persistenceByIdentity.delete(identity);
157
+
158
+ await Promise.allSettled([persistence.cachePromise, persistence.queue]);
159
+ await storage.clearAuthState(serverId);
160
+ };
161
+
60
162
  /**
61
163
  * Builds an OAuthClientProvider for the MCP SDK, backed by MCPStorage.
62
164
  * Token refresh and DCR are handled by the SDK; this provider only mediates
@@ -66,46 +168,66 @@ export function createOAuthProvider(
66
168
  opts: CreateOAuthProviderOptions,
67
169
  ): OAuthClientProvider {
68
170
  const { serverId, config, storage, redirectUri, onAuthorizationUrl } = opts;
171
+ const persistence = getPersistence(storage, serverId);
172
+ let pendingState: string | undefined;
69
173
 
70
- type Cache = {
71
- tokens?: OAuthTokens | undefined;
72
- clientInformation?: OAuthClientInformationFull | undefined;
73
- codeVerifier?: string | undefined;
74
- discoveryState?: OAuthDiscoveryState | undefined;
75
- };
76
- let cached: Cache | null = null;
77
-
78
- const loadCache = async (): Promise<Cache> => {
79
- if (cached) return cached;
80
- const persisted = await storage.loadAuthState(serverId);
81
- const initial: Cache = {};
82
- if (persisted?.tokens) initial.tokens = persisted.tokens;
83
- if (config.clientId) {
84
- const ci: OAuthClientInformationFull = {
85
- client_id: config.clientId,
86
- redirect_uris: [redirectUri],
87
- };
88
- if (config.clientSecret) ci.client_secret = config.clientSecret;
89
- initial.clientInformation = ci;
90
- } else if (persisted?.clientInformation) {
91
- initial.clientInformation = persisted.clientInformation;
92
- }
93
- if (persisted?.codeVerifier) initial.codeVerifier = persisted.codeVerifier;
94
- if (persisted?.discoveryState)
95
- initial.discoveryState = persisted.discoveryState;
96
- cached = initial;
97
- return cached;
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 = (():
179
+ | OAuthClientInformationFull
180
+ | undefined => {
181
+ if (!config.clientId) return undefined;
182
+ const ci: OAuthClientInformationFull = {
183
+ client_id: config.clientId,
184
+ redirect_uris: [redirectUri],
185
+ };
186
+ if (config.clientSecret) ci.client_secret = config.clientSecret;
187
+ return ci;
188
+ })();
189
+
190
+ 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;
98
214
  };
99
215
 
100
- const persist = async () => {
101
- const c = cached;
102
- if (!c) return;
103
- const next: Parameters<typeof storage.saveAuthState>[1] = {};
104
- if (c.tokens) next.tokens = c.tokens;
105
- if (c.clientInformation) next.clientInformation = c.clientInformation;
106
- if (c.codeVerifier) next.codeVerifier = c.codeVerifier;
107
- if (c.discoveryState) next.discoveryState = c.discoveryState;
108
- await storage.saveAuthState(serverId, next);
216
+ const persist = () => {
217
+ const task = persistence.queue.then(async () => {
218
+ if (persistence.invalidated) return;
219
+ const c = persistence.cached;
220
+ if (!c) return;
221
+ const next: Parameters<typeof storage.saveAuthState>[1] = {};
222
+ if (c.tokens) next.tokens = c.tokens;
223
+ if (c.clientInformation) next.clientInformation = c.clientInformation;
224
+ if (c.codeVerifier) next.codeVerifier = c.codeVerifier;
225
+ if (c.state) next.state = c.state;
226
+ if (c.discoveryState) next.discoveryState = c.discoveryState;
227
+ await storage.saveAuthState(serverId, next);
228
+ });
229
+ persistence.queue = task.catch(() => {});
230
+ return task;
109
231
  };
110
232
 
111
233
  const clientMetadata: OAuthClientMetadata = {
@@ -131,11 +253,12 @@ export function createOAuthProvider(
131
253
  typeof crypto !== "undefined" && "randomUUID" in crypto
132
254
  ? crypto.randomUUID()
133
255
  : `${Date.now()}.${Math.random()}`;
134
- return `${encodeServerIdInState(serverId)}.${nonce}`;
256
+ pendingState = `${encodeServerIdInState(serverId)}.${nonce}`;
257
+ return pendingState;
135
258
  },
136
259
  async clientInformation() {
137
260
  const c = await loadCache();
138
- return c.clientInformation;
261
+ return staticClientInformation ?? c.clientInformation;
139
262
  },
140
263
  async saveClientInformation(info) {
141
264
  const c = await loadCache();
@@ -149,6 +272,7 @@ export function createOAuthProvider(
149
272
  async saveTokens(tokens) {
150
273
  const c = await loadCache();
151
274
  c.tokens = tokens;
275
+ delete c.state;
152
276
  await persist();
153
277
  },
154
278
  async redirectToAuthorization(url) {
@@ -157,6 +281,10 @@ export function createOAuthProvider(
157
281
  async saveCodeVerifier(codeVerifier) {
158
282
  const c = await loadCache();
159
283
  c.codeVerifier = codeVerifier;
284
+ if (pendingState) {
285
+ c.state = pendingState;
286
+ pendingState = undefined;
287
+ }
160
288
  await persist();
161
289
  },
162
290
  async codeVerifier() {
@@ -179,7 +307,10 @@ export function createOAuthProvider(
179
307
  const c = await loadCache();
180
308
  if (scope === "all" || scope === "tokens") delete c.tokens;
181
309
  if (scope === "all" || scope === "client") delete c.clientInformation;
182
- if (scope === "all" || scope === "verifier") delete c.codeVerifier;
310
+ if (scope === "all" || scope === "verifier") {
311
+ delete c.codeVerifier;
312
+ delete c.state;
313
+ }
183
314
  if (scope === "all" || scope === "discovery") delete c.discoveryState;
184
315
  await persist();
185
316
  },