@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.
- package/dist/auth/createOAuthProvider.d.ts +19 -2
- package/dist/auth/createOAuthProvider.d.ts.map +1 -1
- package/dist/auth/createOAuthProvider.js +197 -35
- package/dist/auth/createOAuthProvider.js.map +1 -1
- package/dist/auth/types.d.ts +6 -1
- package/dist/auth/types.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.js +5 -6
- package/dist/hooks/useMcpOAuthCallback.js.map +1 -1
- package/dist/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpManagerResource.js +2 -1
- package/dist/resources/McpManagerResource.js.map +1 -1
- package/dist/resources/McpServerResource.d.ts +2 -1
- package/dist/resources/McpServerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.js +79 -19
- package/dist/resources/McpServerResource.js.map +1 -1
- package/dist/resources/storage/McpLocalStorage.d.ts +7 -0
- package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpLocalStorage.js +146 -37
- package/dist/resources/storage/McpLocalStorage.js.map +1 -1
- package/dist/resources/storage/McpMemoryStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpMemoryStorage.js +9 -6
- package/dist/resources/storage/McpMemoryStorage.js.map +1 -1
- package/dist/resources/storage/types.d.ts +12 -0
- package/dist/resources/storage/types.d.ts.map +1 -1
- package/dist/utils/serverUrl.d.ts +8 -0
- package/dist/utils/serverUrl.d.ts.map +1 -0
- package/dist/utils/serverUrl.js +15 -0
- package/dist/utils/serverUrl.js.map +1 -0
- package/package.json +7 -7
- package/src/auth/createOAuthProvider.test.ts +919 -5
- package/src/auth/createOAuthProvider.ts +328 -42
- package/src/auth/types.ts +6 -1
- package/src/hooks/useMcpOAuthCallback.test.ts +74 -1
- package/src/hooks/useMcpOAuthCallback.tsx +11 -8
- package/src/resources/McpManagerResource.ts +5 -1
- package/src/resources/McpServerResource.test.ts +612 -16
- package/src/resources/McpServerResource.ts +95 -23
- package/src/resources/storage/McpLocalStorage.test.ts +97 -1
- package/src/resources/storage/McpLocalStorage.ts +90 -47
- package/src/resources/storage/McpMemoryStorage.test.ts +99 -0
- package/src/resources/storage/McpMemoryStorage.ts +23 -17
- package/src/resources/storage/types.ts +12 -0
- package/src/utils/serverUrl.test.ts +66 -0
- 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 {
|
|
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
|
-
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
|
|
83
|
-
if (config.
|
|
84
|
-
|
|
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
|
|
101
|
-
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
if (
|
|
105
|
-
if (
|
|
106
|
-
if (
|
|
107
|
-
|
|
108
|
-
|
|
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
|
-
|
|
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")
|
|
181
|
-
|
|
182
|
-
|
|
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
|
-
/**
|
|
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 {
|
|
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 {
|
|
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
|
-
|
|
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
|
|
310
|
+
await clearOAuthProviderAuthState(storage, id);
|
|
307
311
|
setCustomServers((prev) => prev.filter((s) => s.id !== id));
|
|
308
312
|
}
|
|
309
313
|
},
|