@assistant-ui/react-mcp 0.1.15 → 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 (39) 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/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 +50 -11
  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 +132 -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/package.json +7 -7
  27. package/src/auth/createOAuthProvider.test.ts +407 -2
  28. package/src/auth/createOAuthProvider.ts +171 -40
  29. package/src/auth/types.ts +1 -0
  30. package/src/hooks/useMcpOAuthCallback.test.ts +74 -1
  31. package/src/hooks/useMcpOAuthCallback.tsx +11 -8
  32. package/src/resources/McpManagerResource.ts +2 -1
  33. package/src/resources/McpServerResource.test.ts +377 -16
  34. package/src/resources/McpServerResource.ts +61 -10
  35. package/src/resources/storage/McpLocalStorage.test.ts +71 -1
  36. package/src/resources/storage/McpLocalStorage.ts +69 -47
  37. package/src/resources/storage/McpMemoryStorage.test.ts +99 -0
  38. package/src/resources/storage/McpMemoryStorage.ts +23 -17
  39. package/src/resources/storage/types.ts +12 -0
@@ -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
  },
package/src/auth/types.ts CHANGED
@@ -8,6 +8,7 @@ export type MCPPersistedAuthState = {
8
8
  tokens?: OAuthTokens;
9
9
  clientInformation?: OAuthClientInformationFull;
10
10
  codeVerifier?: string;
11
+ state?: string;
11
12
  discoveryState?: OAuthDiscoveryState;
12
13
  /** Bearer token (entered at add-form time). */
13
14
  token?: string;
@@ -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";
@@ -303,7 +304,7 @@ const useMcpManagerResource = (
303
304
  try {
304
305
  await lookup.get({ key: id }).remove();
305
306
  } catch {
306
- await storage.clearAuthState(id);
307
+ await clearOAuthProviderAuthState(storage, id);
307
308
  setCustomServers((prev) => prev.filter((s) => s.id !== id));
308
309
  }
309
310
  },