@assistant-ui/react-mcp 0.1.16 → 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 +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/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpManagerResource.js.map +1 -1
- package/dist/resources/McpServerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.js +36 -15
- package/dist/resources/McpServerResource.js.map +1 -1
- package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpLocalStorage.js +14 -0
- package/dist/resources/storage/McpLocalStorage.js.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 +6 -6
- package/src/auth/createOAuthProvider.test.ts +544 -35
- package/src/auth/createOAuthProvider.ts +201 -46
- package/src/auth/types.ts +5 -1
- package/src/resources/McpManagerResource.ts +3 -0
- package/src/resources/McpServerResource.test.ts +252 -17
- package/src/resources/McpServerResource.ts +35 -14
- package/src/resources/storage/McpLocalStorage.test.ts +26 -0
- package/src/resources/storage/McpLocalStorage.ts +21 -0
- 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 {
|
|
@@ -80,13 +86,6 @@ export const getConnectionDependencies = (
|
|
|
80
86
|
];
|
|
81
87
|
};
|
|
82
88
|
|
|
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
89
|
const useMcpServerResourceInstance = (
|
|
91
90
|
props: McpServerResourceInstanceProps,
|
|
92
91
|
): ClientOutput<"mcpServer"> => {
|
|
@@ -242,11 +241,23 @@ const useMcpServerResourceInstance = (
|
|
|
242
241
|
},
|
|
243
242
|
);
|
|
244
243
|
|
|
244
|
+
const unboundAuthMessage = () =>
|
|
245
|
+
`MCP server "${props.id}" has saved authentication for a different URL. Authenticate again to connect to ${props.url}.`;
|
|
246
|
+
|
|
247
|
+
const loadAuthState = useEffectEvent(async () => {
|
|
248
|
+
const state = await props.storage.loadAuthState(props.id);
|
|
249
|
+
if (isAuthStateForServerUrl(state, props.url)) {
|
|
250
|
+
return { state, unbound: false };
|
|
251
|
+
}
|
|
252
|
+
return { state: null, unbound: hasPersistedCredentials(state) };
|
|
253
|
+
});
|
|
254
|
+
|
|
245
255
|
const buildTransport = useEffectEvent(
|
|
246
256
|
async (): Promise<StreamableHTTPClientTransport> => {
|
|
247
257
|
if (props.auth.type === "oauth") {
|
|
248
258
|
const authProvider = createOAuthProvider({
|
|
249
259
|
serverId: props.id,
|
|
260
|
+
serverUrl: props.url,
|
|
250
261
|
config: props.auth,
|
|
251
262
|
storage: props.storage,
|
|
252
263
|
redirectUri: props.redirectUri,
|
|
@@ -257,8 +268,9 @@ const useMcpServerResourceInstance = (
|
|
|
257
268
|
});
|
|
258
269
|
}
|
|
259
270
|
if (props.auth.type === "bearer") {
|
|
260
|
-
const
|
|
261
|
-
const headers = buildHeaders(props.auth,
|
|
271
|
+
const { state, unbound } = await loadAuthState();
|
|
272
|
+
const headers = buildHeaders(props.auth, state);
|
|
273
|
+
if (!headers && unbound) throw new Error(unboundAuthMessage());
|
|
262
274
|
const transportOpts: StreamableHTTPClientTransportOptions = {};
|
|
263
275
|
if (headers) transportOpts.requestInit = { headers };
|
|
264
276
|
return new StreamableHTTPClientTransport(
|
|
@@ -488,10 +500,11 @@ const useMcpServerResourceInstance = (
|
|
|
488
500
|
}
|
|
489
501
|
pendingAuthValidation.count += 1;
|
|
490
502
|
try {
|
|
491
|
-
const persisted = await
|
|
503
|
+
const { state: persisted, unbound } = await loadAuthState();
|
|
492
504
|
if (!isCurrentConnection(validationGeneration)) {
|
|
493
505
|
throw createInterruptedAuthError();
|
|
494
506
|
}
|
|
507
|
+
if (unbound) throw new Error(unboundAuthMessage());
|
|
495
508
|
if (!persisted?.state) {
|
|
496
509
|
throw new Error(
|
|
497
510
|
"no pending OAuth authorization request for this server",
|
|
@@ -563,9 +576,9 @@ const useMcpServerResourceInstance = (
|
|
|
563
576
|
return;
|
|
564
577
|
}
|
|
565
578
|
const generation = connectionGenerationRef.current;
|
|
566
|
-
let
|
|
579
|
+
let loaded: Awaited<ReturnType<typeof loadAuthState>>;
|
|
567
580
|
try {
|
|
568
|
-
|
|
581
|
+
loaded = await loadAuthState();
|
|
569
582
|
} catch (error) {
|
|
570
583
|
if (signal.cancelled || !isCurrentConnection(generation)) return;
|
|
571
584
|
const message = error instanceof Error ? error.message : String(error);
|
|
@@ -576,8 +589,13 @@ const useMcpServerResourceInstance = (
|
|
|
576
589
|
return;
|
|
577
590
|
}
|
|
578
591
|
if (signal.cancelled || !isCurrentConnection(generation)) return;
|
|
592
|
+
if (loaded.unbound) {
|
|
593
|
+
setLastError({ message: unboundAuthMessage() });
|
|
594
|
+
return;
|
|
595
|
+
}
|
|
596
|
+
const persisted = loaded.state;
|
|
579
597
|
if (props.auth.type === "oauth") {
|
|
580
|
-
if (!persisted
|
|
598
|
+
if (!hasUsableOAuthTokens(persisted, props.auth)) return;
|
|
581
599
|
} else if (!persisted?.token) {
|
|
582
600
|
return;
|
|
583
601
|
}
|
|
@@ -599,6 +617,9 @@ const useMcpServerResourceInstance = (
|
|
|
599
617
|
pendingDisposalRef.current = pendingDisposal;
|
|
600
618
|
mountedRef.current = true;
|
|
601
619
|
const signal = { cancelled: false };
|
|
620
|
+
// Auto-connect opens a transport, so it belongs to the same effect as the
|
|
621
|
+
// disposal that closes it.
|
|
622
|
+
// eslint-disable-next-line react-hooks/set-state-in-effect
|
|
602
623
|
void tryAutoConnect(signal);
|
|
603
624
|
return () => {
|
|
604
625
|
mountedRef.current = false;
|
|
@@ -744,7 +765,7 @@ export const McpServerResource = resource(function useMcpServerResource(
|
|
|
744
765
|
const dependencies = getConnectionDependencies(props);
|
|
745
766
|
const [connection, setConnection] = useState({ dependencies, generation: 0 });
|
|
746
767
|
let currentConnection = connection;
|
|
747
|
-
if (!
|
|
768
|
+
if (!shallowEqual(connection.dependencies, dependencies)) {
|
|
748
769
|
currentConnection = {
|
|
749
770
|
dependencies,
|
|
750
771
|
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
|
|
|
@@ -377,6 +402,7 @@ describe("McpLocalStorage auth state", () => {
|
|
|
377
402
|
const createProvider = () =>
|
|
378
403
|
createOAuthProvider({
|
|
379
404
|
serverId: "docs",
|
|
405
|
+
serverUrl: "https://mcp.example.com/mcp",
|
|
380
406
|
config: { type: "oauth", clientId: "client-id" },
|
|
381
407
|
storage: loadStorage(storage),
|
|
382
408
|
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";
|
|
@@ -160,6 +161,16 @@ const isSecureNetworkUrl = (value: unknown): value is string => {
|
|
|
160
161
|
}
|
|
161
162
|
};
|
|
162
163
|
|
|
164
|
+
const isMcpServerUrl = (value: unknown): value is string => {
|
|
165
|
+
if (typeof value !== "string") return false;
|
|
166
|
+
try {
|
|
167
|
+
const url = new URL(value);
|
|
168
|
+
return url.protocol === "https:" || url.protocol === "http:";
|
|
169
|
+
} catch {
|
|
170
|
+
return false;
|
|
171
|
+
}
|
|
172
|
+
};
|
|
173
|
+
|
|
163
174
|
const normalizeDiscoveryState = (
|
|
164
175
|
value: unknown,
|
|
165
176
|
): MCPPersistedAuthState["discoveryState"] | undefined => {
|
|
@@ -195,9 +206,16 @@ export const normalizePersistedAuthState = (
|
|
|
195
206
|
value: unknown,
|
|
196
207
|
): MCPPersistedAuthState | null => {
|
|
197
208
|
if (!isRecord(value)) return null;
|
|
209
|
+
if ("serverUrl" in value && !isMcpServerUrl(value.serverUrl)) return null;
|
|
198
210
|
|
|
199
211
|
const state: MCPPersistedAuthState = {};
|
|
212
|
+
if (isMcpServerUrl(value.serverUrl)) {
|
|
213
|
+
state.serverUrl = normalizeMcpServerUrl(value.serverUrl);
|
|
214
|
+
}
|
|
200
215
|
if (isNonEmptyString(value.token)) state.token = value.token;
|
|
216
|
+
if (isNonEmptyString(value.tokensClientId)) {
|
|
217
|
+
state.tokensClientId = value.tokensClientId;
|
|
218
|
+
}
|
|
201
219
|
if (isNonEmptyString(value.codeVerifier)) {
|
|
202
220
|
state.codeVerifier = value.codeVerifier;
|
|
203
221
|
}
|
|
@@ -208,6 +226,9 @@ export const normalizePersistedAuthState = (
|
|
|
208
226
|
|
|
209
227
|
const clientInformation = normalizeClientInformation(value.clientInformation);
|
|
210
228
|
if (clientInformation) state.clientInformation = clientInformation;
|
|
229
|
+
if (value.clientInformationSource === "registered") {
|
|
230
|
+
state.clientInformationSource = value.clientInformationSource;
|
|
231
|
+
}
|
|
211
232
|
|
|
212
233
|
const discoveryState = normalizeDiscoveryState(value.discoveryState);
|
|
213
234
|
if (discoveryState) state.discoveryState = discoveryState;
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
import { describe, expect, it } from "vitest";
|
|
2
|
+
import {
|
|
3
|
+
hasPersistedCredentials,
|
|
4
|
+
isAuthStateForServerUrl,
|
|
5
|
+
normalizeMcpServerUrl,
|
|
6
|
+
} from "./serverUrl";
|
|
7
|
+
|
|
8
|
+
describe("normalizeMcpServerUrl", () => {
|
|
9
|
+
it("normalizes host case and the default port", () => {
|
|
10
|
+
expect(normalizeMcpServerUrl("https://MCP.Example.com:443/mcp")).toBe(
|
|
11
|
+
"https://mcp.example.com/mcp",
|
|
12
|
+
);
|
|
13
|
+
});
|
|
14
|
+
});
|
|
15
|
+
|
|
16
|
+
describe("isAuthStateForServerUrl", () => {
|
|
17
|
+
it("matches equivalent spellings of the same endpoint", () => {
|
|
18
|
+
expect(
|
|
19
|
+
isAuthStateForServerUrl(
|
|
20
|
+
{ serverUrl: "https://MCP.Example.com/mcp" },
|
|
21
|
+
"https://mcp.example.com:443/mcp",
|
|
22
|
+
),
|
|
23
|
+
).toBe(true);
|
|
24
|
+
});
|
|
25
|
+
|
|
26
|
+
it("rejects a different endpoint, an unbound record, and no record", () => {
|
|
27
|
+
expect(
|
|
28
|
+
isAuthStateForServerUrl(
|
|
29
|
+
{ serverUrl: "https://a.example.com/mcp" },
|
|
30
|
+
"https://b.example.com/mcp",
|
|
31
|
+
),
|
|
32
|
+
).toBe(false);
|
|
33
|
+
expect(
|
|
34
|
+
isAuthStateForServerUrl({ token: "t" }, "https://a.example.com/mcp"),
|
|
35
|
+
).toBe(false);
|
|
36
|
+
expect(isAuthStateForServerUrl(null, "https://a.example.com/mcp")).toBe(
|
|
37
|
+
false,
|
|
38
|
+
);
|
|
39
|
+
});
|
|
40
|
+
|
|
41
|
+
it("rejects an unparsable URL on either side", () => {
|
|
42
|
+
expect(
|
|
43
|
+
isAuthStateForServerUrl({ serverUrl: "not a url" }, "https://a.test/mcp"),
|
|
44
|
+
).toBe(false);
|
|
45
|
+
expect(
|
|
46
|
+
isAuthStateForServerUrl({ serverUrl: "https://a.test/mcp" }, "not a url"),
|
|
47
|
+
).toBe(false);
|
|
48
|
+
});
|
|
49
|
+
});
|
|
50
|
+
|
|
51
|
+
describe("hasPersistedCredentials", () => {
|
|
52
|
+
it("counts bearer and OAuth tokens, not flow state", () => {
|
|
53
|
+
expect(hasPersistedCredentials({ token: "t" })).toBe(true);
|
|
54
|
+
expect(hasPersistedCredentials({ token: "" })).toBe(false);
|
|
55
|
+
expect(
|
|
56
|
+
hasPersistedCredentials({
|
|
57
|
+
tokens: { access_token: "a", token_type: "bearer" },
|
|
58
|
+
}),
|
|
59
|
+
).toBe(true);
|
|
60
|
+
expect(hasPersistedCredentials({ codeVerifier: "v", state: "s" })).toBe(
|
|
61
|
+
false,
|
|
62
|
+
);
|
|
63
|
+
expect(hasPersistedCredentials({})).toBe(false);
|
|
64
|
+
expect(hasPersistedCredentials(null)).toBe(false);
|
|
65
|
+
});
|
|
66
|
+
});
|