@assistant-ui/react-mcp 0.1.16 → 0.1.18
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 +13 -2
- package/dist/auth/createOAuthProvider.d.ts.map +1 -1
- package/dist/auth/createOAuthProvider.js +109 -31
- package/dist/auth/createOAuthProvider.js.map +1 -1
- package/dist/auth/types.d.ts +5 -1
- package/dist/auth/types.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormAuthFields.js +71 -28
- package/dist/primitives/addForm/McpAddFormAuthFields.js.map +1 -1
- package/dist/primitives/addForm/McpAddFormError.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormError.js +30 -10
- package/dist/primitives/addForm/McpAddFormError.js.map +1 -1
- package/dist/primitives/addForm/McpAddFormNameField.js +31 -18
- package/dist/primitives/addForm/McpAddFormNameField.js.map +1 -1
- package/dist/primitives/addForm/McpAddFormRoot.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormRoot.js +32 -11
- package/dist/primitives/addForm/McpAddFormRoot.js.map +1 -1
- package/dist/primitives/addForm/McpAddFormUrlField.js +31 -18
- package/dist/primitives/addForm/McpAddFormUrlField.js.map +1 -1
- package/dist/primitives/addForm/context.d.ts +9 -1
- package/dist/primitives/addForm/context.d.ts.map +1 -1
- package/dist/primitives/addForm/context.js.map +1 -1
- package/dist/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpManagerResource.js +330 -247
- package/dist/resources/McpManagerResource.js.map +1 -1
- package/dist/resources/McpServerRemovalFence.d.ts +7 -0
- package/dist/resources/McpServerRemovalFence.d.ts.map +1 -0
- package/dist/resources/McpServerRemovalFence.js +11 -0
- package/dist/resources/McpServerRemovalFence.js.map +1 -0
- package/dist/resources/McpServerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.js +41 -16
- package/dist/resources/McpServerResource.js.map +1 -1
- package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpLocalStorage.js +16 -1
- package/dist/resources/storage/McpLocalStorage.js.map +1 -1
- package/dist/resources/storage/types.d.ts +10 -8
- 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 +544 -35
- package/src/auth/createOAuthProvider.ts +201 -46
- package/src/auth/types.ts +5 -1
- package/src/primitives/addForm/McpAddFormAccessibility.test.tsx +261 -0
- package/src/primitives/addForm/McpAddFormAuthFields.tsx +33 -15
- package/src/primitives/addForm/McpAddFormError.tsx +17 -3
- package/src/primitives/addForm/McpAddFormNameField.tsx +13 -1
- package/src/primitives/addForm/McpAddFormRoot.tsx +51 -8
- package/src/primitives/addForm/McpAddFormUrlField.tsx +12 -1
- package/src/primitives/addForm/context.tsx +10 -0
- package/src/resources/McpManagerResource.test.ts +712 -9
- package/src/resources/McpManagerResource.ts +208 -56
- package/src/resources/McpServerRemovalFence.ts +21 -0
- package/src/resources/McpServerResource.test.ts +252 -17
- package/src/resources/McpServerResource.ts +40 -15
- package/src/resources/storage/McpLocalStorage.test.ts +28 -0
- package/src/resources/storage/McpLocalStorage.ts +23 -1
- package/src/resources/storage/types.ts +10 -8
- package/src/utils/serverUrl.test.ts +66 -0
- package/src/utils/serverUrl.ts +23 -0
|
@@ -3,6 +3,7 @@ import type { ClientOutput } from "@assistant-ui/store";
|
|
|
3
3
|
import { useEffect, useState } from "react";
|
|
4
4
|
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
|
5
5
|
import type { MCPAuthConfig } from "../mcp-scope";
|
|
6
|
+
import type { MCPPersistedAuthState } from "../auth/types";
|
|
6
7
|
import type { MCPStorage } from "./storage/types";
|
|
7
8
|
import type { McpServerResourceProps } from "./McpServerResource";
|
|
8
9
|
|
|
@@ -189,9 +190,191 @@ const mount = (
|
|
|
189
190
|
});
|
|
190
191
|
};
|
|
191
192
|
|
|
193
|
+
const unboundAuthMessage =
|
|
194
|
+
'MCP server "docs" has saved authentication for a different URL. Authenticate again to connect to https://example.com/mcp.';
|
|
195
|
+
|
|
192
196
|
describe("McpServerResource automatic authentication", () => {
|
|
193
197
|
beforeEach(resetMocks);
|
|
194
198
|
|
|
199
|
+
it("does not auto-connect with authentication from another server URL", async () => {
|
|
200
|
+
const storage = createStorage();
|
|
201
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
202
|
+
serverUrl: "https://other.example.com/mcp",
|
|
203
|
+
token: "secret",
|
|
204
|
+
});
|
|
205
|
+
const root = mount({
|
|
206
|
+
auth: { type: "bearer" },
|
|
207
|
+
storage,
|
|
208
|
+
autoConnect: true,
|
|
209
|
+
});
|
|
210
|
+
|
|
211
|
+
try {
|
|
212
|
+
await waitForResourceUpdate(
|
|
213
|
+
() => root.getValue().getState().lastError !== null,
|
|
214
|
+
);
|
|
215
|
+
|
|
216
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
217
|
+
connectionState: "disconnected",
|
|
218
|
+
lastError: { message: unboundAuthMessage },
|
|
219
|
+
});
|
|
220
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
221
|
+
} finally {
|
|
222
|
+
root.unmount();
|
|
223
|
+
}
|
|
224
|
+
});
|
|
225
|
+
|
|
226
|
+
it("does not auto-connect with an unbound legacy bearer token", async () => {
|
|
227
|
+
const storage = createStorage();
|
|
228
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({ token: "secret" });
|
|
229
|
+
const root = mount({
|
|
230
|
+
auth: { type: "bearer" },
|
|
231
|
+
storage,
|
|
232
|
+
autoConnect: true,
|
|
233
|
+
});
|
|
234
|
+
|
|
235
|
+
try {
|
|
236
|
+
await waitForResourceUpdate(
|
|
237
|
+
() => root.getValue().getState().lastError !== null,
|
|
238
|
+
);
|
|
239
|
+
|
|
240
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
241
|
+
connectionState: "disconnected",
|
|
242
|
+
lastError: { message: unboundAuthMessage },
|
|
243
|
+
});
|
|
244
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
245
|
+
} finally {
|
|
246
|
+
root.unmount();
|
|
247
|
+
}
|
|
248
|
+
});
|
|
249
|
+
|
|
250
|
+
it("does not auto-connect with unbound legacy OAuth tokens", async () => {
|
|
251
|
+
const storage = createStorage();
|
|
252
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
253
|
+
tokens: { access_token: "secret", token_type: "bearer" },
|
|
254
|
+
});
|
|
255
|
+
const root = mount({
|
|
256
|
+
auth: { type: "oauth" },
|
|
257
|
+
storage,
|
|
258
|
+
autoConnect: true,
|
|
259
|
+
});
|
|
260
|
+
|
|
261
|
+
try {
|
|
262
|
+
await waitForResourceUpdate(
|
|
263
|
+
() => root.getValue().getState().lastError !== null,
|
|
264
|
+
);
|
|
265
|
+
|
|
266
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
267
|
+
connectionState: "disconnected",
|
|
268
|
+
lastError: { message: unboundAuthMessage },
|
|
269
|
+
});
|
|
270
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
271
|
+
} finally {
|
|
272
|
+
root.unmount();
|
|
273
|
+
}
|
|
274
|
+
});
|
|
275
|
+
|
|
276
|
+
it("does not auto-connect OAuth tokens for a different client", async () => {
|
|
277
|
+
const storage = createStorage();
|
|
278
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
279
|
+
serverUrl: "https://example.com/mcp",
|
|
280
|
+
clientInformation: {
|
|
281
|
+
client_id: "client-a",
|
|
282
|
+
redirect_uris: ["https://example.com/callback"],
|
|
283
|
+
},
|
|
284
|
+
clientInformationSource: "registered",
|
|
285
|
+
tokens: { access_token: "secret", token_type: "bearer" },
|
|
286
|
+
tokensClientId: "client-a",
|
|
287
|
+
});
|
|
288
|
+
const root = mount({
|
|
289
|
+
auth: { type: "oauth", clientId: "client-b" },
|
|
290
|
+
storage,
|
|
291
|
+
autoConnect: true,
|
|
292
|
+
});
|
|
293
|
+
|
|
294
|
+
try {
|
|
295
|
+
await waitFor(() => storage.loadAuthState.mock.calls.length > 0);
|
|
296
|
+
await flushMacrotask();
|
|
297
|
+
|
|
298
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
299
|
+
connectionState: "disconnected",
|
|
300
|
+
lastError: null,
|
|
301
|
+
});
|
|
302
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
303
|
+
} finally {
|
|
304
|
+
root.unmount();
|
|
305
|
+
}
|
|
306
|
+
});
|
|
307
|
+
|
|
308
|
+
it("keeps a static bearer token usable when the saved record is unbound", async () => {
|
|
309
|
+
const storage = createStorage();
|
|
310
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({ token: "stale" });
|
|
311
|
+
const root = mount({
|
|
312
|
+
auth: { type: "bearer", token: "static" },
|
|
313
|
+
storage,
|
|
314
|
+
autoConnect: true,
|
|
315
|
+
});
|
|
316
|
+
|
|
317
|
+
try {
|
|
318
|
+
await waitFor(() => mocks.transports.length > 0);
|
|
319
|
+
await flushMacrotask();
|
|
320
|
+
|
|
321
|
+
expect(mocks.StreamableHTTPClientTransport).toHaveBeenCalledWith(
|
|
322
|
+
new URL("https://example.com/mcp"),
|
|
323
|
+
{ requestInit: { headers: { Authorization: "Bearer static" } } },
|
|
324
|
+
);
|
|
325
|
+
expect(root.getValue().getState().lastError).toBeNull();
|
|
326
|
+
} finally {
|
|
327
|
+
root.unmount();
|
|
328
|
+
}
|
|
329
|
+
});
|
|
330
|
+
|
|
331
|
+
it("does not report a record that holds no credentials", async () => {
|
|
332
|
+
const storage = createStorage();
|
|
333
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({});
|
|
334
|
+
const root = mount({
|
|
335
|
+
auth: { type: "bearer" },
|
|
336
|
+
storage,
|
|
337
|
+
autoConnect: true,
|
|
338
|
+
});
|
|
339
|
+
|
|
340
|
+
try {
|
|
341
|
+
await waitFor(() => storage.loadAuthState.mock.calls.length > 0);
|
|
342
|
+
await flushMacrotask();
|
|
343
|
+
|
|
344
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
345
|
+
connectionState: "disconnected",
|
|
346
|
+
lastError: null,
|
|
347
|
+
});
|
|
348
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
349
|
+
} finally {
|
|
350
|
+
root.unmount();
|
|
351
|
+
}
|
|
352
|
+
});
|
|
353
|
+
|
|
354
|
+
it("reports an unbound bearer record on a manual connect", async () => {
|
|
355
|
+
const storage = createStorage();
|
|
356
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
357
|
+
serverUrl: "https://other.example.com/mcp",
|
|
358
|
+
token: "secret",
|
|
359
|
+
});
|
|
360
|
+
const root = mount({ auth: { type: "bearer" }, storage });
|
|
361
|
+
|
|
362
|
+
try {
|
|
363
|
+
await expect(root.getValue().connect()).resolves.toBeUndefined();
|
|
364
|
+
await waitForResourceUpdate(
|
|
365
|
+
() => root.getValue().getState().connectionState === "error",
|
|
366
|
+
);
|
|
367
|
+
|
|
368
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
369
|
+
connectionState: "error",
|
|
370
|
+
lastError: { message: unboundAuthMessage },
|
|
371
|
+
});
|
|
372
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
373
|
+
} finally {
|
|
374
|
+
root.unmount();
|
|
375
|
+
}
|
|
376
|
+
});
|
|
377
|
+
|
|
195
378
|
it("reports auth storage load failures", async () => {
|
|
196
379
|
const storage = createStorage();
|
|
197
380
|
vi.mocked(storage.loadAuthState).mockRejectedValue(
|
|
@@ -562,9 +745,8 @@ describe("McpServerResource completeAuth", () => {
|
|
|
562
745
|
beforeEach(resetMocks);
|
|
563
746
|
|
|
564
747
|
it("lets callback validation win over mount-time auto-connect", async () => {
|
|
565
|
-
const pendingLoads: Array<
|
|
566
|
-
|
|
567
|
-
> = [];
|
|
748
|
+
const pendingLoads: Array<(value: MCPPersistedAuthState | null) => void> =
|
|
749
|
+
[];
|
|
568
750
|
const storage = createStorage();
|
|
569
751
|
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
570
752
|
() =>
|
|
@@ -587,13 +769,25 @@ describe("McpServerResource completeAuth", () => {
|
|
|
587
769
|
await waitFor(() => pendingLoads.length > callbackLoadIndex);
|
|
588
770
|
|
|
589
771
|
for (const resolve of pendingLoads.slice(0, callbackLoadIndex)) {
|
|
590
|
-
resolve({
|
|
772
|
+
resolve({
|
|
773
|
+
serverUrl: "https://example.com/mcp",
|
|
774
|
+
clientInformation: {
|
|
775
|
+
client_id: "registered-client",
|
|
776
|
+
redirect_uris: ["https://example.com/callback"],
|
|
777
|
+
},
|
|
778
|
+
clientInformationSource: "registered",
|
|
779
|
+
tokens: { access_token: "persisted", token_type: "bearer" },
|
|
780
|
+
tokensClientId: "registered-client",
|
|
781
|
+
});
|
|
591
782
|
}
|
|
592
783
|
await flushMacrotask();
|
|
593
784
|
|
|
594
785
|
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
595
786
|
|
|
596
|
-
pendingLoads[callbackLoadIndex]!({
|
|
787
|
+
pendingLoads[callbackLoadIndex]!({
|
|
788
|
+
serverUrl: "https://example.com/mcp",
|
|
789
|
+
state: "expected",
|
|
790
|
+
});
|
|
597
791
|
await expect(completeAuth).resolves.toBeUndefined();
|
|
598
792
|
await flushMacrotask();
|
|
599
793
|
|
|
@@ -606,9 +800,8 @@ describe("McpServerResource completeAuth", () => {
|
|
|
606
800
|
});
|
|
607
801
|
|
|
608
802
|
it("resumes auto-connect when callback validation fails", async () => {
|
|
609
|
-
const pendingLoads: Array<
|
|
610
|
-
|
|
611
|
-
> = [];
|
|
803
|
+
const pendingLoads: Array<(value: MCPPersistedAuthState | null) => void> =
|
|
804
|
+
[];
|
|
612
805
|
const storage = createStorage();
|
|
613
806
|
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
614
807
|
() =>
|
|
@@ -631,13 +824,25 @@ describe("McpServerResource completeAuth", () => {
|
|
|
631
824
|
await waitFor(() => pendingLoads.length > callbackLoadIndex);
|
|
632
825
|
|
|
633
826
|
for (const resolve of pendingLoads.slice(0, callbackLoadIndex)) {
|
|
634
|
-
resolve({
|
|
827
|
+
resolve({
|
|
828
|
+
serverUrl: "https://example.com/mcp",
|
|
829
|
+
clientInformation: {
|
|
830
|
+
client_id: "registered-client",
|
|
831
|
+
redirect_uris: ["https://example.com/callback"],
|
|
832
|
+
},
|
|
833
|
+
clientInformationSource: "registered",
|
|
834
|
+
tokens: { access_token: "persisted", token_type: "bearer" },
|
|
835
|
+
tokensClientId: "registered-client",
|
|
836
|
+
});
|
|
635
837
|
}
|
|
636
838
|
await flushMacrotask();
|
|
637
839
|
|
|
638
840
|
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
639
841
|
|
|
640
|
-
pendingLoads[callbackLoadIndex]!({
|
|
842
|
+
pendingLoads[callbackLoadIndex]!({
|
|
843
|
+
serverUrl: "https://example.com/mcp",
|
|
844
|
+
state: "different",
|
|
845
|
+
});
|
|
641
846
|
await expect(completeAuth).rejects.toThrow(
|
|
642
847
|
"OAuth state does not match the authorization request",
|
|
643
848
|
);
|
|
@@ -653,6 +858,7 @@ describe("McpServerResource completeAuth", () => {
|
|
|
653
858
|
it("completes auth across the StrictMode effect replay", async () => {
|
|
654
859
|
const storage = createStorage();
|
|
655
860
|
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
861
|
+
serverUrl: "https://example.com/mcp",
|
|
656
862
|
state: "aui-mcp:ZG9jcw.nonce",
|
|
657
863
|
});
|
|
658
864
|
let completeAuth: Promise<void> | undefined;
|
|
@@ -682,7 +888,10 @@ describe("McpServerResource completeAuth", () => {
|
|
|
682
888
|
|
|
683
889
|
it("rejects before transport setup without a usable authorization code", async () => {
|
|
684
890
|
const storage = createStorage();
|
|
685
|
-
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
891
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
892
|
+
serverUrl: "https://example.com/mcp",
|
|
893
|
+
state: "abc",
|
|
894
|
+
});
|
|
686
895
|
const root = mount({ auth: { type: "oauth" }, storage });
|
|
687
896
|
|
|
688
897
|
try {
|
|
@@ -705,7 +914,10 @@ describe("McpServerResource completeAuth", () => {
|
|
|
705
914
|
|
|
706
915
|
it("forwards OAuth error callbacks to the transport", async () => {
|
|
707
916
|
const storage = createStorage();
|
|
708
|
-
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
917
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
918
|
+
serverUrl: "https://example.com/mcp",
|
|
919
|
+
state: "expected",
|
|
920
|
+
});
|
|
709
921
|
mocks.finishAuthResults.push(() =>
|
|
710
922
|
Promise.reject(new Error("access_denied: Denied")),
|
|
711
923
|
);
|
|
@@ -758,6 +970,7 @@ describe("McpServerResource completeAuth", () => {
|
|
|
758
970
|
it("rejects callbacks whose state does not match the authorization request", async () => {
|
|
759
971
|
const storage = createStorage();
|
|
760
972
|
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
973
|
+
serverUrl: "https://example.com/mcp",
|
|
761
974
|
state: "aui-mcp:ZG9jcw.expected",
|
|
762
975
|
});
|
|
763
976
|
const root = mount({ auth: { type: "oauth" }, storage });
|
|
@@ -787,7 +1000,10 @@ describe("McpServerResource completeAuth", () => {
|
|
|
787
1000
|
});
|
|
788
1001
|
|
|
789
1002
|
it("does not resume authorization after disconnecting during validation", async () => {
|
|
790
|
-
let resolveAuthState!: (state: {
|
|
1003
|
+
let resolveAuthState!: (state: {
|
|
1004
|
+
serverUrl: string;
|
|
1005
|
+
state: string;
|
|
1006
|
+
}) => void;
|
|
791
1007
|
const storage = createStorage();
|
|
792
1008
|
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
793
1009
|
() =>
|
|
@@ -804,7 +1020,10 @@ describe("McpServerResource completeAuth", () => {
|
|
|
804
1020
|
await waitFor(() => resolveAuthState !== undefined);
|
|
805
1021
|
|
|
806
1022
|
await root.getValue().disconnect();
|
|
807
|
-
resolveAuthState({
|
|
1023
|
+
resolveAuthState({
|
|
1024
|
+
serverUrl: "https://example.com/mcp",
|
|
1025
|
+
state: "expected",
|
|
1026
|
+
});
|
|
808
1027
|
|
|
809
1028
|
await expect(completeAuth).rejects.toThrow(
|
|
810
1029
|
'MCP server "docs" authorization was interrupted before completion.',
|
|
@@ -824,7 +1043,10 @@ describe("McpServerResource completeAuth", () => {
|
|
|
824
1043
|
Promise.reject(new Error("invalid_grant")),
|
|
825
1044
|
);
|
|
826
1045
|
const storage = createStorage();
|
|
827
|
-
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
1046
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
1047
|
+
serverUrl: "https://example.com/mcp",
|
|
1048
|
+
state: "expected",
|
|
1049
|
+
});
|
|
828
1050
|
const root = mount({ auth: { type: "oauth" }, storage });
|
|
829
1051
|
|
|
830
1052
|
try {
|
|
@@ -859,7 +1081,10 @@ describe("McpServerResource completeAuth", () => {
|
|
|
859
1081
|
}),
|
|
860
1082
|
);
|
|
861
1083
|
const storage = createStorage();
|
|
862
|
-
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
1084
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
1085
|
+
serverUrl: "https://example.com/mcp",
|
|
1086
|
+
state: "expected",
|
|
1087
|
+
});
|
|
863
1088
|
const root = mount({ auth: { type: "oauth" }, storage });
|
|
864
1089
|
let didUnmount = false;
|
|
865
1090
|
|
|
@@ -892,7 +1117,10 @@ describe("McpServerResource completeAuth", () => {
|
|
|
892
1117
|
}),
|
|
893
1118
|
);
|
|
894
1119
|
const storage = createStorage();
|
|
895
|
-
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
1120
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
1121
|
+
serverUrl: "https://example.com/mcp",
|
|
1122
|
+
state: "expected",
|
|
1123
|
+
});
|
|
896
1124
|
const root = mount({ auth: { type: "oauth" }, storage });
|
|
897
1125
|
let didUnmount = false;
|
|
898
1126
|
|
|
@@ -1766,7 +1994,14 @@ describe("McpServerResource oauth storage swap", () => {
|
|
|
1766
1994
|
|
|
1767
1995
|
it("reconnects onto the replacement storage when the scope changes", async () => {
|
|
1768
1996
|
const persisted = {
|
|
1997
|
+
serverUrl: "https://example.com/mcp",
|
|
1998
|
+
clientInformation: {
|
|
1999
|
+
client_id: "registered-client",
|
|
2000
|
+
redirect_uris: ["https://example.com/callback"],
|
|
2001
|
+
},
|
|
2002
|
+
clientInformationSource: "registered" as const,
|
|
1769
2003
|
tokens: { access_token: "tok", token_type: "bearer" },
|
|
2004
|
+
tokensClientId: "registered-client",
|
|
1770
2005
|
};
|
|
1771
2006
|
const storageA = {
|
|
1772
2007
|
...createStorage(),
|
|
@@ -1,6 +1,7 @@
|
|
|
1
1
|
import { useState, useRef, useEffect, useMemo, useEffectEvent } from "react";
|
|
2
2
|
import { resource, useResource, withKey } from "@assistant-ui/tap";
|
|
3
3
|
import type { ClientOutput } from "@assistant-ui/store";
|
|
4
|
+
import { shallowEqual } from "@assistant-ui/store/internal";
|
|
4
5
|
import {
|
|
5
6
|
Client,
|
|
6
7
|
StreamableHTTPClientTransport,
|
|
@@ -13,9 +14,14 @@ import {
|
|
|
13
14
|
import {
|
|
14
15
|
clearOAuthProviderAuthState,
|
|
15
16
|
createOAuthProvider,
|
|
17
|
+
hasUsableOAuthTokens,
|
|
16
18
|
} from "../auth/createOAuthProvider";
|
|
17
19
|
import { buildHeaders } from "../auth/buildHeaders";
|
|
18
20
|
import { assertValidServerId } from "../utils/serverId";
|
|
21
|
+
import {
|
|
22
|
+
hasPersistedCredentials,
|
|
23
|
+
isAuthStateForServerUrl,
|
|
24
|
+
} from "../utils/serverUrl";
|
|
19
25
|
import { validateElicitationContent } from "./validateElicitationContent";
|
|
20
26
|
import type { MCPStorage } from "./storage/types";
|
|
21
27
|
import type {
|
|
@@ -28,6 +34,7 @@ import type {
|
|
|
28
34
|
MCPToolInfo,
|
|
29
35
|
} from "../mcp-scope";
|
|
30
36
|
import { createMcpId } from "../utils/createMcpId";
|
|
37
|
+
import { beginMcpServerRemovalFence } from "./McpServerRemovalFence";
|
|
31
38
|
|
|
32
39
|
export type McpServerResourceProps = {
|
|
33
40
|
id: string;
|
|
@@ -80,13 +87,6 @@ export const getConnectionDependencies = (
|
|
|
80
87
|
];
|
|
81
88
|
};
|
|
82
89
|
|
|
83
|
-
const areConnectionDependenciesEqual = (
|
|
84
|
-
left: readonly unknown[],
|
|
85
|
-
right: readonly unknown[],
|
|
86
|
-
) =>
|
|
87
|
-
left.length === right.length &&
|
|
88
|
-
left.every((value, index) => Object.is(value, right[index]));
|
|
89
|
-
|
|
90
90
|
const useMcpServerResourceInstance = (
|
|
91
91
|
props: McpServerResourceInstanceProps,
|
|
92
92
|
): ClientOutput<"mcpServer"> => {
|
|
@@ -242,11 +242,23 @@ const useMcpServerResourceInstance = (
|
|
|
242
242
|
},
|
|
243
243
|
);
|
|
244
244
|
|
|
245
|
+
const unboundAuthMessage = () =>
|
|
246
|
+
`MCP server "${props.id}" has saved authentication for a different URL. Authenticate again to connect to ${props.url}.`;
|
|
247
|
+
|
|
248
|
+
const loadAuthState = useEffectEvent(async () => {
|
|
249
|
+
const state = await props.storage.loadAuthState(props.id);
|
|
250
|
+
if (isAuthStateForServerUrl(state, props.url)) {
|
|
251
|
+
return { state, unbound: false };
|
|
252
|
+
}
|
|
253
|
+
return { state: null, unbound: hasPersistedCredentials(state) };
|
|
254
|
+
});
|
|
255
|
+
|
|
245
256
|
const buildTransport = useEffectEvent(
|
|
246
257
|
async (): Promise<StreamableHTTPClientTransport> => {
|
|
247
258
|
if (props.auth.type === "oauth") {
|
|
248
259
|
const authProvider = createOAuthProvider({
|
|
249
260
|
serverId: props.id,
|
|
261
|
+
serverUrl: props.url,
|
|
250
262
|
config: props.auth,
|
|
251
263
|
storage: props.storage,
|
|
252
264
|
redirectUri: props.redirectUri,
|
|
@@ -257,8 +269,9 @@ const useMcpServerResourceInstance = (
|
|
|
257
269
|
});
|
|
258
270
|
}
|
|
259
271
|
if (props.auth.type === "bearer") {
|
|
260
|
-
const
|
|
261
|
-
const headers = buildHeaders(props.auth,
|
|
272
|
+
const { state, unbound } = await loadAuthState();
|
|
273
|
+
const headers = buildHeaders(props.auth, state);
|
|
274
|
+
if (!headers && unbound) throw new Error(unboundAuthMessage());
|
|
262
275
|
const transportOpts: StreamableHTTPClientTransportOptions = {};
|
|
263
276
|
if (headers) transportOpts.requestInit = { headers };
|
|
264
277
|
return new StreamableHTTPClientTransport(
|
|
@@ -488,10 +501,11 @@ const useMcpServerResourceInstance = (
|
|
|
488
501
|
}
|
|
489
502
|
pendingAuthValidation.count += 1;
|
|
490
503
|
try {
|
|
491
|
-
const persisted = await
|
|
504
|
+
const { state: persisted, unbound } = await loadAuthState();
|
|
492
505
|
if (!isCurrentConnection(validationGeneration)) {
|
|
493
506
|
throw createInterruptedAuthError();
|
|
494
507
|
}
|
|
508
|
+
if (unbound) throw new Error(unboundAuthMessage());
|
|
495
509
|
if (!persisted?.state) {
|
|
496
510
|
throw new Error(
|
|
497
511
|
"no pending OAuth authorization request for this server",
|
|
@@ -563,9 +577,9 @@ const useMcpServerResourceInstance = (
|
|
|
563
577
|
return;
|
|
564
578
|
}
|
|
565
579
|
const generation = connectionGenerationRef.current;
|
|
566
|
-
let
|
|
580
|
+
let loaded: Awaited<ReturnType<typeof loadAuthState>>;
|
|
567
581
|
try {
|
|
568
|
-
|
|
582
|
+
loaded = await loadAuthState();
|
|
569
583
|
} catch (error) {
|
|
570
584
|
if (signal.cancelled || !isCurrentConnection(generation)) return;
|
|
571
585
|
const message = error instanceof Error ? error.message : String(error);
|
|
@@ -576,8 +590,13 @@ const useMcpServerResourceInstance = (
|
|
|
576
590
|
return;
|
|
577
591
|
}
|
|
578
592
|
if (signal.cancelled || !isCurrentConnection(generation)) return;
|
|
593
|
+
if (loaded.unbound) {
|
|
594
|
+
setLastError({ message: unboundAuthMessage() });
|
|
595
|
+
return;
|
|
596
|
+
}
|
|
597
|
+
const persisted = loaded.state;
|
|
579
598
|
if (props.auth.type === "oauth") {
|
|
580
|
-
if (!persisted
|
|
599
|
+
if (!hasUsableOAuthTokens(persisted, props.auth)) return;
|
|
581
600
|
} else if (!persisted?.token) {
|
|
582
601
|
return;
|
|
583
602
|
}
|
|
@@ -599,6 +618,9 @@ const useMcpServerResourceInstance = (
|
|
|
599
618
|
pendingDisposalRef.current = pendingDisposal;
|
|
600
619
|
mountedRef.current = true;
|
|
601
620
|
const signal = { cancelled: false };
|
|
621
|
+
// Auto-connect opens a transport, so it belongs to the same effect as the
|
|
622
|
+
// disposal that closes it.
|
|
623
|
+
// eslint-disable-next-line react-hooks/set-state-in-effect
|
|
602
624
|
void tryAutoConnect(signal);
|
|
603
625
|
return () => {
|
|
604
626
|
mountedRef.current = false;
|
|
@@ -645,16 +667,19 @@ const useMcpServerResourceInstance = (
|
|
|
645
667
|
connect: doConnect,
|
|
646
668
|
disconnect: doDisconnect,
|
|
647
669
|
remove: async () => {
|
|
648
|
-
|
|
670
|
+
const releaseRemovalFence = beginMcpServerRemovalFence(props);
|
|
649
671
|
try {
|
|
672
|
+
await doDisconnect();
|
|
650
673
|
await clearOAuthProviderAuthState(props.storage, props.id);
|
|
651
674
|
await props.onRemove();
|
|
652
675
|
} catch (err) {
|
|
676
|
+
releaseRemovalFence?.();
|
|
653
677
|
setLastError({
|
|
654
678
|
message: err instanceof Error ? err.message : String(err),
|
|
655
679
|
});
|
|
656
680
|
throw err;
|
|
657
681
|
}
|
|
682
|
+
releaseRemovalFence?.();
|
|
658
683
|
},
|
|
659
684
|
callTool: async (name, args) => {
|
|
660
685
|
const client = clientRef.current;
|
|
@@ -744,7 +769,7 @@ export const McpServerResource = resource(function useMcpServerResource(
|
|
|
744
769
|
const dependencies = getConnectionDependencies(props);
|
|
745
770
|
const [connection, setConnection] = useState({ dependencies, generation: 0 });
|
|
746
771
|
let currentConnection = connection;
|
|
747
|
-
if (!
|
|
772
|
+
if (!shallowEqual(connection.dependencies, dependencies)) {
|
|
748
773
|
currentConnection = {
|
|
749
774
|
dependencies,
|
|
750
775
|
generation: connection.generation + 1,
|
|
@@ -136,6 +136,27 @@ describe("normalizePersistedAuthState", () => {
|
|
|
136
136
|
});
|
|
137
137
|
});
|
|
138
138
|
|
|
139
|
+
it("keeps valid server URL bindings", () => {
|
|
140
|
+
expect(
|
|
141
|
+
normalizePersistedAuthState({
|
|
142
|
+
serverUrl: "http://mcp.example.com/docs",
|
|
143
|
+
token: "bearer-token",
|
|
144
|
+
}),
|
|
145
|
+
).toEqual({
|
|
146
|
+
serverUrl: "http://mcp.example.com/docs",
|
|
147
|
+
token: "bearer-token",
|
|
148
|
+
});
|
|
149
|
+
});
|
|
150
|
+
|
|
151
|
+
it("rejects auth state with an unsafe server URL binding", () => {
|
|
152
|
+
expect(
|
|
153
|
+
normalizePersistedAuthState({
|
|
154
|
+
serverUrl: "javascript:alert(1)",
|
|
155
|
+
token: "bearer-token",
|
|
156
|
+
}),
|
|
157
|
+
).toBeNull();
|
|
158
|
+
});
|
|
159
|
+
|
|
139
160
|
it("keeps valid OAuth tokens and client information", () => {
|
|
140
161
|
const tokens = {
|
|
141
162
|
access_token: "access-token",
|
|
@@ -153,11 +174,15 @@ describe("normalizePersistedAuthState", () => {
|
|
|
153
174
|
expect(
|
|
154
175
|
normalizePersistedAuthState({
|
|
155
176
|
tokens,
|
|
177
|
+
tokensClientId: "client-id",
|
|
156
178
|
clientInformation,
|
|
179
|
+
clientInformationSource: "registered",
|
|
157
180
|
}),
|
|
158
181
|
).toEqual({
|
|
159
182
|
tokens,
|
|
183
|
+
tokensClientId: "client-id",
|
|
160
184
|
clientInformation,
|
|
185
|
+
clientInformationSource: "registered",
|
|
161
186
|
});
|
|
162
187
|
});
|
|
163
188
|
|
|
@@ -214,6 +239,8 @@ describe("normalizePersistedAuthState", () => {
|
|
|
214
239
|
|
|
215
240
|
it.each([
|
|
216
241
|
"http://auth.example.com",
|
|
242
|
+
"http://127.example.com",
|
|
243
|
+
"http://127.0.0.1.example.com",
|
|
217
244
|
"data:text/plain,auth",
|
|
218
245
|
"file:///tmp/auth",
|
|
219
246
|
])("drops discovery state with an unsafe URL: %s", (url) => {
|
|
@@ -377,6 +404,7 @@ describe("McpLocalStorage auth state", () => {
|
|
|
377
404
|
const createProvider = () =>
|
|
378
405
|
createOAuthProvider({
|
|
379
406
|
serverId: "docs",
|
|
407
|
+
serverUrl: "https://mcp.example.com/mcp",
|
|
380
408
|
config: { type: "oauth", clientId: "client-id" },
|
|
381
409
|
storage: loadStorage(storage),
|
|
382
410
|
redirectUri: "http://localhost/callback",
|
|
@@ -6,6 +6,7 @@ import {
|
|
|
6
6
|
OAuthProtectedResourceMetadataSchema,
|
|
7
7
|
OAuthTokensSchema,
|
|
8
8
|
} from "@modelcontextprotocol/core";
|
|
9
|
+
import { normalizeMcpServerUrl } from "../../utils/serverUrl";
|
|
9
10
|
import type { MCPAuthConfig, MCPCustomServerRecord } from "../../mcp-scope";
|
|
10
11
|
import type { MCPPersistedAuthState } from "../../auth/types";
|
|
11
12
|
import { assertValidServerId } from "../../utils/serverId";
|
|
@@ -147,12 +148,13 @@ const isSecureNetworkUrl = (value: unknown): value is string => {
|
|
|
147
148
|
if (!isNonEmptyString(value)) return false;
|
|
148
149
|
try {
|
|
149
150
|
const url = new URL(value);
|
|
151
|
+
const isIpv4Loopback = /^127(?:\.\d{1,3}){3}$/.test(url.hostname);
|
|
150
152
|
return (
|
|
151
153
|
url.protocol === "https:" ||
|
|
152
154
|
(url.protocol === "http:" &&
|
|
153
155
|
(url.hostname === "localhost" ||
|
|
154
156
|
url.hostname.endsWith(".localhost") ||
|
|
155
|
-
|
|
157
|
+
isIpv4Loopback ||
|
|
156
158
|
url.hostname === "[::1]"))
|
|
157
159
|
);
|
|
158
160
|
} catch {
|
|
@@ -160,6 +162,16 @@ const isSecureNetworkUrl = (value: unknown): value is string => {
|
|
|
160
162
|
}
|
|
161
163
|
};
|
|
162
164
|
|
|
165
|
+
const isMcpServerUrl = (value: unknown): value is string => {
|
|
166
|
+
if (typeof value !== "string") return false;
|
|
167
|
+
try {
|
|
168
|
+
const url = new URL(value);
|
|
169
|
+
return url.protocol === "https:" || url.protocol === "http:";
|
|
170
|
+
} catch {
|
|
171
|
+
return false;
|
|
172
|
+
}
|
|
173
|
+
};
|
|
174
|
+
|
|
163
175
|
const normalizeDiscoveryState = (
|
|
164
176
|
value: unknown,
|
|
165
177
|
): MCPPersistedAuthState["discoveryState"] | undefined => {
|
|
@@ -195,9 +207,16 @@ export const normalizePersistedAuthState = (
|
|
|
195
207
|
value: unknown,
|
|
196
208
|
): MCPPersistedAuthState | null => {
|
|
197
209
|
if (!isRecord(value)) return null;
|
|
210
|
+
if ("serverUrl" in value && !isMcpServerUrl(value.serverUrl)) return null;
|
|
198
211
|
|
|
199
212
|
const state: MCPPersistedAuthState = {};
|
|
213
|
+
if (isMcpServerUrl(value.serverUrl)) {
|
|
214
|
+
state.serverUrl = normalizeMcpServerUrl(value.serverUrl);
|
|
215
|
+
}
|
|
200
216
|
if (isNonEmptyString(value.token)) state.token = value.token;
|
|
217
|
+
if (isNonEmptyString(value.tokensClientId)) {
|
|
218
|
+
state.tokensClientId = value.tokensClientId;
|
|
219
|
+
}
|
|
201
220
|
if (isNonEmptyString(value.codeVerifier)) {
|
|
202
221
|
state.codeVerifier = value.codeVerifier;
|
|
203
222
|
}
|
|
@@ -208,6 +227,9 @@ export const normalizePersistedAuthState = (
|
|
|
208
227
|
|
|
209
228
|
const clientInformation = normalizeClientInformation(value.clientInformation);
|
|
210
229
|
if (clientInformation) state.clientInformation = clientInformation;
|
|
230
|
+
if (value.clientInformationSource === "registered") {
|
|
231
|
+
state.clientInformationSource = value.clientInformationSource;
|
|
232
|
+
}
|
|
211
233
|
|
|
212
234
|
const discoveryState = normalizeDiscoveryState(value.discoveryState);
|
|
213
235
|
if (discoveryState) state.discoveryState = discoveryState;
|