@assistant-ui/react-mcp 0.1.15 → 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.
Files changed (45) hide show
  1. package/dist/auth/createOAuthProvider.d.ts +19 -2
  2. package/dist/auth/createOAuthProvider.d.ts.map +1 -1
  3. package/dist/auth/createOAuthProvider.js +197 -35
  4. package/dist/auth/createOAuthProvider.js.map +1 -1
  5. package/dist/auth/types.d.ts +6 -1
  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/resources/McpManagerResource.d.ts.map +1 -1
  11. package/dist/resources/McpManagerResource.js +2 -1
  12. package/dist/resources/McpManagerResource.js.map +1 -1
  13. package/dist/resources/McpServerResource.d.ts +2 -1
  14. package/dist/resources/McpServerResource.d.ts.map +1 -1
  15. package/dist/resources/McpServerResource.js +79 -19
  16. package/dist/resources/McpServerResource.js.map +1 -1
  17. package/dist/resources/storage/McpLocalStorage.d.ts +7 -0
  18. package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
  19. package/dist/resources/storage/McpLocalStorage.js +146 -37
  20. package/dist/resources/storage/McpLocalStorage.js.map +1 -1
  21. package/dist/resources/storage/McpMemoryStorage.d.ts.map +1 -1
  22. package/dist/resources/storage/McpMemoryStorage.js +9 -6
  23. package/dist/resources/storage/McpMemoryStorage.js.map +1 -1
  24. package/dist/resources/storage/types.d.ts +12 -0
  25. package/dist/resources/storage/types.d.ts.map +1 -1
  26. package/dist/utils/serverUrl.d.ts +8 -0
  27. package/dist/utils/serverUrl.d.ts.map +1 -0
  28. package/dist/utils/serverUrl.js +15 -0
  29. package/dist/utils/serverUrl.js.map +1 -0
  30. package/package.json +7 -7
  31. package/src/auth/createOAuthProvider.test.ts +919 -5
  32. package/src/auth/createOAuthProvider.ts +328 -42
  33. package/src/auth/types.ts +6 -1
  34. package/src/hooks/useMcpOAuthCallback.test.ts +74 -1
  35. package/src/hooks/useMcpOAuthCallback.tsx +11 -8
  36. package/src/resources/McpManagerResource.ts +5 -1
  37. package/src/resources/McpServerResource.test.ts +612 -16
  38. package/src/resources/McpServerResource.ts +95 -23
  39. package/src/resources/storage/McpLocalStorage.test.ts +97 -1
  40. package/src/resources/storage/McpLocalStorage.ts +90 -47
  41. package/src/resources/storage/McpMemoryStorage.test.ts +99 -0
  42. package/src/resources/storage/McpMemoryStorage.ts +23 -17
  43. package/src/resources/storage/types.ts +12 -0
  44. package/src/utils/serverUrl.test.ts +66 -0
  45. package/src/utils/serverUrl.ts +23 -0
@@ -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;
@@ -57,6 +63,174 @@ export type CreateOAuthProviderOptions = {
57
63
  onAuthorizationUrl: (url: URL) => void;
58
64
  };
59
65
 
66
+ type OAuthProviderCache = {
67
+ token?: string | undefined;
68
+ tokens?: OAuthTokens | undefined;
69
+ tokensClientId?: string | undefined;
70
+ clientInformation?: OAuthClientInformationFull | undefined;
71
+ clientInformationSource?: MCPPersistedAuthState["clientInformationSource"];
72
+ codeVerifier?: string | undefined;
73
+ state?: string | undefined;
74
+ discoveryState?: OAuthDiscoveryState | undefined;
75
+ };
76
+
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;
118
+ cached: OAuthProviderCache | null;
119
+ cachePromise: Promise<OAuthProviderCache> | null;
120
+ invalidated: boolean;
121
+ };
122
+
123
+ type OAuthProviderPersistence = {
124
+ endpoint: OAuthProviderEndpointCache | null;
125
+ queue: Promise<void>;
126
+ invalidated: boolean;
127
+ };
128
+
129
+ // scopeId, not object identity, is what addresses the same persisted data, so
130
+ // storages sharing one share an anchor and an unscoped storage is its own
131
+ // identity. Every storage declaring a scope holds that scope's anchor, so the
132
+ // coordination state below is collected once the last of them is gone.
133
+ const anchorByStorage = new WeakMap<MCPStorage, object>();
134
+ const anchorByScope = new Map<string, WeakRef<object>>();
135
+ const anchorRegistry = new FinalizationRegistry<string>((scopeId) => {
136
+ if (!anchorByScope.get(scopeId)?.deref()) anchorByScope.delete(scopeId);
137
+ });
138
+
139
+ const getStorageIdentity = (storage: MCPStorage): object => {
140
+ const existing = anchorByStorage.get(storage);
141
+ if (existing) return existing;
142
+
143
+ const { scopeId } = storage;
144
+ if (scopeId === undefined) return storage;
145
+
146
+ let anchor = anchorByScope.get(scopeId)?.deref();
147
+ if (!anchor) {
148
+ anchor = {};
149
+ anchorByScope.set(scopeId, new WeakRef(anchor));
150
+ anchorRegistry.register(anchor, scopeId);
151
+ }
152
+ anchorByStorage.set(storage, anchor);
153
+ return anchor;
154
+ };
155
+
156
+ // McpServerResource builds a fresh provider for every transport, so the cache,
157
+ // the in-flight load, and the write queue have to outlive any one provider.
158
+ // saveAuthState replaces the whole record, so two providers writing their own
159
+ // snapshots concurrently would drop whichever field the loser had added.
160
+ const persistenceByIdentity = new WeakMap<
161
+ object,
162
+ Map<string, OAuthProviderPersistence>
163
+ >();
164
+
165
+ const getPersistence = (
166
+ storage: MCPStorage,
167
+ serverId: string,
168
+ serverUrl: string,
169
+ ): {
170
+ persistence: OAuthProviderPersistence;
171
+ endpoint: OAuthProviderEndpointCache;
172
+ } => {
173
+ const identity = getStorageIdentity(storage);
174
+ let byServerId = persistenceByIdentity.get(identity);
175
+ if (!byServerId) {
176
+ byServerId = new Map();
177
+ persistenceByIdentity.set(identity, byServerId);
178
+ }
179
+
180
+ let persistence = byServerId.get(serverId);
181
+ if (!persistence) {
182
+ persistence = {
183
+ endpoint: null,
184
+ queue: Promise.resolve(),
185
+ invalidated: false,
186
+ };
187
+ byServerId.set(serverId, persistence);
188
+ }
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 };
202
+ };
203
+
204
+ /**
205
+ * Clears persisted OAuth state after the in-flight load and every queued write
206
+ * for that server have settled, so a discarded provider cannot recreate the
207
+ * record it was mid-save on.
208
+ */
209
+ export const clearOAuthProviderAuthState = async (
210
+ storage: MCPStorage,
211
+ serverId: string,
212
+ ): Promise<void> => {
213
+ const identity = getStorageIdentity(storage);
214
+ const byServerId = persistenceByIdentity.get(identity);
215
+ const persistence = byServerId?.get(serverId);
216
+ if (!byServerId || !persistence) {
217
+ await storage.clearAuthState(serverId);
218
+ return;
219
+ }
220
+
221
+ // Detaching the entry before awaiting keeps a provider built during the clear
222
+ // on a fresh generation instead of inheriting the fenced one.
223
+ persistence.invalidated = true;
224
+ byServerId.delete(serverId);
225
+ if (byServerId.size === 0) persistenceByIdentity.delete(identity);
226
+
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;
231
+ await storage.clearAuthState(serverId);
232
+ };
233
+
60
234
  /**
61
235
  * Builds an OAuthClientProvider for the MCP SDK, backed by MCPStorage.
62
236
  * Token refresh and DCR are handled by the SDK; this provider only mediates
@@ -65,49 +239,124 @@ export type CreateOAuthProviderOptions = {
65
239
  export function createOAuthProvider(
66
240
  opts: CreateOAuthProviderOptions,
67
241
  ): OAuthClientProvider {
68
- const { serverId, config, storage, redirectUri, onAuthorizationUrl } = opts;
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
+ );
256
+ let pendingState: string | undefined;
69
257
 
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;
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 = ():
264
+ | OAuthClientInformationFull
265
+ | undefined => {
266
+ if (!config.clientId) return undefined;
267
+ const ci: OAuthClientInformationFull = {
268
+ client_id: config.clientId,
269
+ redirect_uris: [redirectUri],
270
+ };
271
+ if (config.clientSecret) ci.client_secret = config.clientSecret;
272
+ return ci;
98
273
  };
274
+ let clientInformationOverlay = configuredClientInformation();
99
275
 
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);
276
+ const activeClientId = (cache: OAuthProviderCache): string | undefined =>
277
+ clientInformationOverlay?.client_id ?? registeredClientId(cache);
278
+
279
+ const loadCache = (): Promise<OAuthProviderCache> => {
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;
109
333
  };
110
334
 
335
+ function persist() {
336
+ const task = persistence.queue.then(async () => {
337
+ if (persistence.invalidated || endpoint.invalidated) return;
338
+ const c = endpoint.cached;
339
+ if (!c) return;
340
+ const next: Parameters<typeof storage.saveAuthState>[1] = {};
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;
350
+ if (c.codeVerifier) next.codeVerifier = c.codeVerifier;
351
+ if (c.state) next.state = c.state;
352
+ if (c.discoveryState) next.discoveryState = c.discoveryState;
353
+ next.serverUrl = normalizedServerUrl;
354
+ await storage.saveAuthState(serverId, next);
355
+ });
356
+ persistence.queue = task.catch(() => {});
357
+ return task;
358
+ }
359
+
111
360
  const clientMetadata: OAuthClientMetadata = {
112
361
  client_name: "assistant-ui",
113
362
  redirect_uris: [redirectUri],
@@ -131,24 +380,43 @@ export function createOAuthProvider(
131
380
  typeof crypto !== "undefined" && "randomUUID" in crypto
132
381
  ? crypto.randomUUID()
133
382
  : `${Date.now()}.${Math.random()}`;
134
- return `${encodeServerIdInState(serverId)}.${nonce}`;
383
+ pendingState = `${encodeServerIdInState(serverId)}.${nonce}`;
384
+ return pendingState;
135
385
  },
136
386
  async clientInformation() {
137
387
  const c = await loadCache();
388
+ if (clientInformationOverlay) return clientInformationOverlay;
389
+ if (c.clientInformationSource !== "registered") return undefined;
138
390
  return c.clientInformation;
139
391
  },
140
392
  async saveClientInformation(info) {
393
+ if (clientInformationOverlay) {
394
+ clientInformationOverlay = info as OAuthClientInformationFull;
395
+ return;
396
+ }
141
397
  const c = await loadCache();
142
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
+ }
143
404
  await persist();
144
405
  },
145
406
  async tokens() {
146
407
  const c = await loadCache();
408
+ const clientId = activeClientId(c);
409
+ if (clientId === undefined || c.tokensClientId !== clientId)
410
+ return undefined;
147
411
  return c.tokens;
148
412
  },
149
413
  async saveTokens(tokens) {
150
414
  const c = await loadCache();
151
415
  c.tokens = tokens;
416
+ const clientId = activeClientId(c);
417
+ if (clientId) c.tokensClientId = clientId;
418
+ else delete c.tokensClientId;
419
+ delete c.state;
152
420
  await persist();
153
421
  },
154
422
  async redirectToAuthorization(url) {
@@ -157,6 +425,10 @@ export function createOAuthProvider(
157
425
  async saveCodeVerifier(codeVerifier) {
158
426
  const c = await loadCache();
159
427
  c.codeVerifier = codeVerifier;
428
+ if (pendingState) {
429
+ c.state = pendingState;
430
+ pendingState = undefined;
431
+ }
160
432
  await persist();
161
433
  },
162
434
  async codeVerifier() {
@@ -177,9 +449,23 @@ export function createOAuthProvider(
177
449
  },
178
450
  async invalidateCredentials(scope) {
179
451
  const c = await loadCache();
180
- if (scope === "all" || scope === "tokens") delete c.tokens;
181
- if (scope === "all" || scope === "client") delete c.clientInformation;
182
- if (scope === "all" || scope === "verifier") delete c.codeVerifier;
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
+ }
465
+ if (scope === "all" || scope === "verifier") {
466
+ delete c.codeVerifier;
467
+ delete c.state;
468
+ }
183
469
  if (scope === "all" || scope === "discovery") delete c.discoveryState;
184
470
  await persist();
185
471
  },
package/src/auth/types.ts CHANGED
@@ -5,10 +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;
15
+ state?: string;
11
16
  discoveryState?: OAuthDiscoveryState;
12
- /** Bearer token (entered at add-form time). */
17
+ /** Host-persisted bearer token. Must be paired with serverUrl. */
13
18
  token?: string;
14
19
  };
@@ -1,6 +1,12 @@
1
1
  // @vitest-environment jsdom
2
2
 
3
- import { renderHook, waitFor } from "@testing-library/react";
3
+ import {
4
+ createElement,
5
+ Suspense,
6
+ startTransition,
7
+ type PropsWithChildren,
8
+ } from "react";
9
+ import { act, renderHook, waitFor } from "@testing-library/react";
4
10
  import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
5
11
 
6
12
  const mocks = vi.hoisted(() => {
@@ -98,6 +104,73 @@ describe("useMcpOAuthCallback", () => {
98
104
  });
99
105
  },
100
106
  );
107
+
108
+ it("lets the server validate error callback parameters", async () => {
109
+ const authError = new Error("authorization response issuer mismatch");
110
+ const url = `${callbackUrl}&error=access_denied&error_description=untrusted`;
111
+ mocks.completeAuth.mockRejectedValueOnce(authError);
112
+
113
+ const { result } = renderHook(() => useMcpOAuthCallback({ url }));
114
+
115
+ await waitFor(() => expect(result.current.status).toBe("error"));
116
+ expect(mocks.completeAuth).toHaveBeenCalledWith(url);
117
+ expect(result.current.error).toMatchObject({
118
+ message:
119
+ 'MCP OAuth callback for server "docs" failed: authorization response issuer mismatch',
120
+ cause: authError,
121
+ });
122
+ });
123
+
124
+ it("keeps callbacks scoped to committed renders", async () => {
125
+ let resolveAuth!: () => void;
126
+ mocks.completeAuth.mockReturnValueOnce(
127
+ new Promise<void>((resolve) => {
128
+ resolveAuth = resolve;
129
+ }),
130
+ );
131
+ const onCompleteA = vi.fn();
132
+ const onCompleteB = vi.fn();
133
+ const interruptedRender = vi.fn();
134
+ const pending = new Promise<never>(() => {});
135
+ let blocked = false;
136
+ const Blocker = () => {
137
+ if (blocked) {
138
+ interruptedRender();
139
+ throw pending;
140
+ }
141
+ return null;
142
+ };
143
+ const Wrapper = ({ children }: PropsWithChildren) =>
144
+ createElement(
145
+ Suspense,
146
+ { fallback: null },
147
+ children,
148
+ createElement(Blocker),
149
+ );
150
+
151
+ const { rerender } = renderHook(
152
+ ({ onComplete }) => useMcpOAuthCallback({ url: callbackUrl, onComplete }),
153
+ {
154
+ initialProps: { onComplete: onCompleteA },
155
+ wrapper: Wrapper,
156
+ },
157
+ );
158
+ await waitFor(() => expect(mocks.completeAuth).toHaveBeenCalledOnce());
159
+
160
+ act(() => {
161
+ blocked = true;
162
+ startTransition(() => rerender({ onComplete: onCompleteB }));
163
+ });
164
+ expect(interruptedRender).toHaveBeenCalled();
165
+
166
+ await act(async () => {
167
+ resolveAuth();
168
+ await Promise.resolve();
169
+ });
170
+
171
+ expect(onCompleteB).not.toHaveBeenCalled();
172
+ expect(onCompleteA).toHaveBeenCalledWith("docs");
173
+ });
101
174
  });
102
175
 
103
176
  describe("createMcpOAuthCallbackError", () => {
@@ -1,4 +1,11 @@
1
- import { type FC, type ReactNode, useEffect, useRef, useState } from "react";
1
+ import {
2
+ type FC,
3
+ type ReactNode,
4
+ useEffect,
5
+ useInsertionEffect,
6
+ useRef,
7
+ useState,
8
+ } from "react";
2
9
  import { useAui } from "@assistant-ui/store";
3
10
  import { decodeServerIdFromState } from "../auth/createOAuthProvider";
4
11
  import { invokeMcpCallback } from "../utils/invokeMcpCallback";
@@ -44,7 +51,9 @@ export function useMcpOAuthCallback(
44
51
  // single-use OAuth code is double-redeemed and the second attempt 4xxs.
45
52
  const startedRef = useRef<string | null>(null);
46
53
  const optsRef = useRef(opts);
47
- optsRef.current = opts;
54
+ useInsertionEffect(() => {
55
+ optsRef.current = opts;
56
+ });
48
57
 
49
58
  useEffect(() => {
50
59
  const url =
@@ -59,12 +68,6 @@ export function useMcpOAuthCallback(
59
68
  const parsed = new URL(url);
60
69
  const state = parsed.searchParams.get("state");
61
70
  if (state) serverId = decodeServerIdFromState(state);
62
- const error = parsed.searchParams.get("error");
63
- if (error) {
64
- throw new Error(
65
- parsed.searchParams.get("error_description") ?? error,
66
- );
67
- }
68
71
  if (!state) throw new Error('missing "state" parameter');
69
72
  if (!serverId) {
70
73
  throw new Error("state was not created by assistant-ui MCP");
@@ -9,6 +9,7 @@ import {
9
9
  import { useAssistantScopeEffect } from "@assistant-ui/store/client";
10
10
  import { ModelContext } from "@assistant-ui/core/store";
11
11
  import { createMcpId } from "../utils/createMcpId";
12
+ import { clearOAuthProviderAuthState } from "../auth/createOAuthProvider";
12
13
  import type { Tool } from "assistant-stream";
13
14
  import { McpServerResource } from "./McpServerResource";
14
15
  import { McpLocalStorage } from "./storage/McpLocalStorage";
@@ -119,6 +120,9 @@ const useMcpManagerResource = (
119
120
 
120
121
  useEffect(() => {
121
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
122
126
  void hydrate(signal);
123
127
  return () => {
124
128
  signal.cancelled = true;
@@ -303,7 +307,7 @@ const useMcpManagerResource = (
303
307
  try {
304
308
  await lookup.get({ key: id }).remove();
305
309
  } catch {
306
- await storage.clearAuthState(id);
310
+ await clearOAuthProviderAuthState(storage, id);
307
311
  setCustomServers((prev) => prev.filter((s) => s.id !== id));
308
312
  }
309
313
  },