@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.
- package/dist/auth/createOAuthProvider.d.ts +7 -1
- package/dist/auth/createOAuthProvider.d.ts.map +1 -1
- package/dist/auth/createOAuthProvider.js +115 -31
- package/dist/auth/createOAuthProvider.js.map +1 -1
- package/dist/auth/types.d.ts +1 -0
- 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/primitives/server/McpServerIcon.js.map +1 -1
- package/dist/primitives/server/McpServerOAuthLink.js.map +1 -1
- package/dist/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpManagerResource.js +4 -2
- 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 +107 -21
- 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 +132 -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/createMcpId.d.ts +9 -0
- package/dist/utils/createMcpId.d.ts.map +1 -0
- package/dist/utils/createMcpId.js +11 -0
- package/dist/utils/createMcpId.js.map +1 -0
- package/dist/utils/invokeMcpCallback.d.ts.map +1 -1
- package/dist/utils/invokeMcpCallback.js +2 -12
- package/dist/utils/invokeMcpCallback.js.map +1 -1
- package/package.json +8 -8
- package/src/auth/createOAuthProvider.test.ts +407 -2
- package/src/auth/createOAuthProvider.ts +171 -40
- package/src/auth/types.ts +1 -0
- package/src/hooks/useMcpOAuthCallback.test.ts +74 -1
- package/src/hooks/useMcpOAuthCallback.tsx +11 -8
- package/src/resources/McpManagerResource.test.ts +128 -0
- package/src/resources/McpManagerResource.ts +4 -5
- package/src/resources/McpServerResource.test.ts +420 -18
- package/src/resources/McpServerResource.ts +148 -27
- package/src/resources/storage/McpLocalStorage.test.ts +71 -1
- package/src/resources/storage/McpLocalStorage.ts +69 -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/createMcpId.test.ts +25 -0
- package/src/utils/createMcpId.ts +10 -0
- 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 {
|
|
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
|
-
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
if (
|
|
83
|
-
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
|
|
95
|
-
|
|
96
|
-
|
|
97
|
-
|
|
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 =
|
|
101
|
-
const
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
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
|
-
|
|
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")
|
|
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
|
},
|