@assistant-ui/react-mcp 0.1.15 → 0.1.16
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- 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/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpManagerResource.js +2 -1
- package/dist/resources/McpManagerResource.js.map +1 -1
- package/dist/resources/McpServerResource.d.ts +2 -1
- package/dist/resources/McpServerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.js +50 -11
- 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/package.json +7 -7
- 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.ts +2 -1
- package/src/resources/McpServerResource.test.ts +377 -16
- package/src/resources/McpServerResource.ts +61 -10
- 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
|
@@ -4,6 +4,7 @@ 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
6
|
import type { MCPStorage } from "./storage/types";
|
|
7
|
+
import type { McpServerResourceProps } from "./McpServerResource";
|
|
7
8
|
|
|
8
9
|
const mocks = vi.hoisted(() => {
|
|
9
10
|
const clients: any[] = [];
|
|
@@ -70,7 +71,8 @@ vi.mock("@modelcontextprotocol/client", async (importOriginal) => ({
|
|
|
70
71
|
StreamableHTTPClientTransport: mocks.StreamableHTTPClientTransport,
|
|
71
72
|
}));
|
|
72
73
|
|
|
73
|
-
const { McpServerResource } =
|
|
74
|
+
const { McpServerResource, getConnectionDependencies } =
|
|
75
|
+
await import("./McpServerResource");
|
|
74
76
|
|
|
75
77
|
const never = <T>() => new Promise<T>(() => {});
|
|
76
78
|
|
|
@@ -559,14 +561,107 @@ describe("McpServerResource connection lifecycle", () => {
|
|
|
559
561
|
describe("McpServerResource completeAuth", () => {
|
|
560
562
|
beforeEach(resetMocks);
|
|
561
563
|
|
|
564
|
+
it("lets callback validation win over mount-time auto-connect", async () => {
|
|
565
|
+
const pendingLoads: Array<
|
|
566
|
+
(value: { state?: string; tokens?: { access_token: string } }) => void
|
|
567
|
+
> = [];
|
|
568
|
+
const storage = createStorage();
|
|
569
|
+
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
570
|
+
() =>
|
|
571
|
+
new Promise((resolve) => {
|
|
572
|
+
pendingLoads.push(resolve);
|
|
573
|
+
}),
|
|
574
|
+
);
|
|
575
|
+
const root = mount({
|
|
576
|
+
auth: { type: "oauth" },
|
|
577
|
+
storage,
|
|
578
|
+
autoConnect: true,
|
|
579
|
+
});
|
|
580
|
+
|
|
581
|
+
try {
|
|
582
|
+
await waitFor(() => pendingLoads.length > 1);
|
|
583
|
+
const callbackLoadIndex = pendingLoads.length;
|
|
584
|
+
const completeAuth = root
|
|
585
|
+
.getValue()
|
|
586
|
+
.completeAuth("https://example.com/callback?code=abc&state=expected");
|
|
587
|
+
await waitFor(() => pendingLoads.length > callbackLoadIndex);
|
|
588
|
+
|
|
589
|
+
for (const resolve of pendingLoads.slice(0, callbackLoadIndex)) {
|
|
590
|
+
resolve({ tokens: { access_token: "persisted" } });
|
|
591
|
+
}
|
|
592
|
+
await flushMacrotask();
|
|
593
|
+
|
|
594
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
595
|
+
|
|
596
|
+
pendingLoads[callbackLoadIndex]!({ state: "expected" });
|
|
597
|
+
await expect(completeAuth).resolves.toBeUndefined();
|
|
598
|
+
await flushMacrotask();
|
|
599
|
+
|
|
600
|
+
expect(mocks.transports).toHaveLength(1);
|
|
601
|
+
expect(mocks.transports[0].finishAuth).toHaveBeenCalledTimes(1);
|
|
602
|
+
expect(root.getValue().getState().connectionState).toBe("connected");
|
|
603
|
+
} finally {
|
|
604
|
+
root.unmount();
|
|
605
|
+
}
|
|
606
|
+
});
|
|
607
|
+
|
|
608
|
+
it("resumes auto-connect when callback validation fails", async () => {
|
|
609
|
+
const pendingLoads: Array<
|
|
610
|
+
(value: { state?: string; tokens?: { access_token: string } }) => void
|
|
611
|
+
> = [];
|
|
612
|
+
const storage = createStorage();
|
|
613
|
+
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
614
|
+
() =>
|
|
615
|
+
new Promise((resolve) => {
|
|
616
|
+
pendingLoads.push(resolve);
|
|
617
|
+
}),
|
|
618
|
+
);
|
|
619
|
+
const root = mount({
|
|
620
|
+
auth: { type: "oauth" },
|
|
621
|
+
storage,
|
|
622
|
+
autoConnect: true,
|
|
623
|
+
});
|
|
624
|
+
|
|
625
|
+
try {
|
|
626
|
+
await waitFor(() => pendingLoads.length > 1);
|
|
627
|
+
const callbackLoadIndex = pendingLoads.length;
|
|
628
|
+
const completeAuth = root
|
|
629
|
+
.getValue()
|
|
630
|
+
.completeAuth("https://example.com/callback?code=abc&state=expected");
|
|
631
|
+
await waitFor(() => pendingLoads.length > callbackLoadIndex);
|
|
632
|
+
|
|
633
|
+
for (const resolve of pendingLoads.slice(0, callbackLoadIndex)) {
|
|
634
|
+
resolve({ tokens: { access_token: "persisted" } });
|
|
635
|
+
}
|
|
636
|
+
await flushMacrotask();
|
|
637
|
+
|
|
638
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
639
|
+
|
|
640
|
+
pendingLoads[callbackLoadIndex]!({ state: "different" });
|
|
641
|
+
await expect(completeAuth).rejects.toThrow(
|
|
642
|
+
"OAuth state does not match the authorization request",
|
|
643
|
+
);
|
|
644
|
+
await waitFor(() => mocks.transports.length === 1);
|
|
645
|
+
await waitForResourceUpdate(
|
|
646
|
+
() => root.getValue().getState().connectionState === "connected",
|
|
647
|
+
);
|
|
648
|
+
} finally {
|
|
649
|
+
root.unmount();
|
|
650
|
+
}
|
|
651
|
+
});
|
|
652
|
+
|
|
562
653
|
it("completes auth across the StrictMode effect replay", async () => {
|
|
654
|
+
const storage = createStorage();
|
|
655
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
656
|
+
state: "aui-mcp:ZG9jcw.nonce",
|
|
657
|
+
});
|
|
563
658
|
let completeAuth: Promise<void> | undefined;
|
|
564
659
|
let started = false;
|
|
565
|
-
const root = mount({ auth: { type: "oauth" } }, (server) => {
|
|
660
|
+
const root = mount({ auth: { type: "oauth" }, storage }, (server) => {
|
|
566
661
|
if (started) return;
|
|
567
662
|
started = true;
|
|
568
663
|
completeAuth = server.completeAuth(
|
|
569
|
-
"https://example.com/callback?code=abc",
|
|
664
|
+
"https://example.com/callback?code=abc&state=aui-mcp%3AZG9jcw.nonce&iss=https%3A%2F%2Fauth.example.com",
|
|
570
665
|
);
|
|
571
666
|
});
|
|
572
667
|
|
|
@@ -575,26 +670,149 @@ describe("McpServerResource completeAuth", () => {
|
|
|
575
670
|
await flushMacrotask();
|
|
576
671
|
|
|
577
672
|
expect(mocks.transports[0].finishAuth).toHaveBeenCalledTimes(1);
|
|
673
|
+
const params = mocks.transports[0].finishAuth.mock.calls[0][0];
|
|
674
|
+
expect(params).toBeInstanceOf(URLSearchParams);
|
|
675
|
+
expect(params.get("code")).toBe("abc");
|
|
676
|
+
expect(params.get("iss")).toBe("https://auth.example.com");
|
|
578
677
|
expect(root.getValue().getState().connectionState).toBe("connected");
|
|
579
678
|
} finally {
|
|
580
679
|
root.unmount();
|
|
581
680
|
}
|
|
582
681
|
});
|
|
583
682
|
|
|
584
|
-
it("rejects
|
|
585
|
-
const
|
|
683
|
+
it("rejects before transport setup without a usable authorization code", async () => {
|
|
684
|
+
const storage = createStorage();
|
|
685
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({ state: "abc" });
|
|
686
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
586
687
|
|
|
587
688
|
try {
|
|
588
689
|
await expect(
|
|
589
|
-
root
|
|
690
|
+
root
|
|
691
|
+
.getValue()
|
|
692
|
+
.completeAuth("https://example.com/callback?state=abc&code="),
|
|
590
693
|
).rejects.toThrow("missing authorization code in callback URL");
|
|
591
694
|
await flushMacrotask();
|
|
592
695
|
|
|
696
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
697
|
+
connectionState: "disconnected",
|
|
698
|
+
lastError: null,
|
|
699
|
+
});
|
|
700
|
+
expect(mocks.transports).toHaveLength(0);
|
|
701
|
+
} finally {
|
|
702
|
+
root.unmount();
|
|
703
|
+
}
|
|
704
|
+
});
|
|
705
|
+
|
|
706
|
+
it("forwards OAuth error callbacks to the transport", async () => {
|
|
707
|
+
const storage = createStorage();
|
|
708
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({ state: "expected" });
|
|
709
|
+
mocks.finishAuthResults.push(() =>
|
|
710
|
+
Promise.reject(new Error("access_denied: Denied")),
|
|
711
|
+
);
|
|
712
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
713
|
+
|
|
714
|
+
try {
|
|
715
|
+
await expect(
|
|
716
|
+
root
|
|
717
|
+
.getValue()
|
|
718
|
+
.completeAuth(
|
|
719
|
+
"https://example.com/callback?error=access_denied&error_description=Denied&state=expected&iss=https%3A%2F%2Fauth.example.com",
|
|
720
|
+
),
|
|
721
|
+
).rejects.toThrow("access_denied: Denied");
|
|
722
|
+
await flushMacrotask();
|
|
723
|
+
|
|
724
|
+
const params = mocks.transports[0].finishAuth.mock.calls[0][0];
|
|
725
|
+
expect(params).toBeInstanceOf(URLSearchParams);
|
|
726
|
+
expect(params.get("error")).toBe("access_denied");
|
|
727
|
+
expect(params.get("error_description")).toBe("Denied");
|
|
728
|
+
expect(params.get("iss")).toBe("https://auth.example.com");
|
|
593
729
|
expect(root.getValue().getState()).toMatchObject({
|
|
594
730
|
connectionState: "error",
|
|
595
|
-
lastError: {
|
|
596
|
-
|
|
597
|
-
|
|
731
|
+
lastError: { message: "access_denied: Denied" },
|
|
732
|
+
});
|
|
733
|
+
} finally {
|
|
734
|
+
root.unmount();
|
|
735
|
+
}
|
|
736
|
+
});
|
|
737
|
+
|
|
738
|
+
it("rejects callbacks without a pending authorization request", async () => {
|
|
739
|
+
const storage = createStorage();
|
|
740
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue(null);
|
|
741
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
742
|
+
|
|
743
|
+
try {
|
|
744
|
+
await expect(
|
|
745
|
+
root
|
|
746
|
+
.getValue()
|
|
747
|
+
.completeAuth("https://example.com/callback?code=abc&state=expected"),
|
|
748
|
+
).rejects.toThrow(
|
|
749
|
+
"no pending OAuth authorization request for this server",
|
|
750
|
+
);
|
|
751
|
+
|
|
752
|
+
expect(mocks.transports).toHaveLength(0);
|
|
753
|
+
} finally {
|
|
754
|
+
root.unmount();
|
|
755
|
+
}
|
|
756
|
+
});
|
|
757
|
+
|
|
758
|
+
it("rejects callbacks whose state does not match the authorization request", async () => {
|
|
759
|
+
const storage = createStorage();
|
|
760
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({
|
|
761
|
+
state: "aui-mcp:ZG9jcw.expected",
|
|
762
|
+
});
|
|
763
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
764
|
+
|
|
765
|
+
try {
|
|
766
|
+
await root.getValue().connect();
|
|
767
|
+
const transport = mocks.transports[0];
|
|
768
|
+
|
|
769
|
+
await expect(
|
|
770
|
+
root
|
|
771
|
+
.getValue()
|
|
772
|
+
.completeAuth(
|
|
773
|
+
"https://example.com/callback?code=abc&state=aui-mcp%3AZG9jcw.forged",
|
|
774
|
+
),
|
|
775
|
+
).rejects.toThrow("OAuth state does not match the authorization request");
|
|
776
|
+
await flushMacrotask();
|
|
777
|
+
|
|
778
|
+
expect(mocks.transports).toHaveLength(1);
|
|
779
|
+
expect(transport.close).not.toHaveBeenCalled();
|
|
780
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
781
|
+
connectionState: "connected",
|
|
782
|
+
lastError: null,
|
|
783
|
+
});
|
|
784
|
+
} finally {
|
|
785
|
+
root.unmount();
|
|
786
|
+
}
|
|
787
|
+
});
|
|
788
|
+
|
|
789
|
+
it("does not resume authorization after disconnecting during validation", async () => {
|
|
790
|
+
let resolveAuthState!: (state: { state: string }) => void;
|
|
791
|
+
const storage = createStorage();
|
|
792
|
+
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
793
|
+
() =>
|
|
794
|
+
new Promise((resolve) => {
|
|
795
|
+
resolveAuthState = resolve;
|
|
796
|
+
}),
|
|
797
|
+
);
|
|
798
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
799
|
+
|
|
800
|
+
try {
|
|
801
|
+
const completeAuth = root
|
|
802
|
+
.getValue()
|
|
803
|
+
.completeAuth("https://example.com/callback?code=abc&state=expected");
|
|
804
|
+
await waitFor(() => resolveAuthState !== undefined);
|
|
805
|
+
|
|
806
|
+
await root.getValue().disconnect();
|
|
807
|
+
resolveAuthState({ state: "expected" });
|
|
808
|
+
|
|
809
|
+
await expect(completeAuth).rejects.toThrow(
|
|
810
|
+
'MCP server "docs" authorization was interrupted before completion.',
|
|
811
|
+
);
|
|
812
|
+
expect(mocks.transports).toHaveLength(0);
|
|
813
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
814
|
+
connectionState: "disconnected",
|
|
815
|
+
lastError: null,
|
|
598
816
|
});
|
|
599
817
|
} finally {
|
|
600
818
|
root.unmount();
|
|
@@ -605,11 +823,15 @@ describe("McpServerResource completeAuth", () => {
|
|
|
605
823
|
mocks.finishAuthResults.push(() =>
|
|
606
824
|
Promise.reject(new Error("invalid_grant")),
|
|
607
825
|
);
|
|
608
|
-
const
|
|
826
|
+
const storage = createStorage();
|
|
827
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({ state: "expected" });
|
|
828
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
609
829
|
|
|
610
830
|
try {
|
|
611
831
|
await expect(
|
|
612
|
-
root
|
|
832
|
+
root
|
|
833
|
+
.getValue()
|
|
834
|
+
.completeAuth("https://example.com/callback?code=abc&state=expected"),
|
|
613
835
|
).rejects.toThrow("invalid_grant");
|
|
614
836
|
await flushMacrotask();
|
|
615
837
|
|
|
@@ -619,7 +841,9 @@ describe("McpServerResource completeAuth", () => {
|
|
|
619
841
|
message: "invalid_grant",
|
|
620
842
|
},
|
|
621
843
|
});
|
|
622
|
-
expect(mocks.transports[0].finishAuth).toHaveBeenCalledWith(
|
|
844
|
+
expect(mocks.transports[0].finishAuth).toHaveBeenCalledWith(
|
|
845
|
+
expect.any(URLSearchParams),
|
|
846
|
+
);
|
|
623
847
|
expect(mocks.transports[0].close).toHaveBeenCalledTimes(1);
|
|
624
848
|
} finally {
|
|
625
849
|
root.unmount();
|
|
@@ -634,13 +858,15 @@ describe("McpServerResource completeAuth", () => {
|
|
|
634
858
|
resolveFinishAuth = resolve;
|
|
635
859
|
}),
|
|
636
860
|
);
|
|
637
|
-
const
|
|
861
|
+
const storage = createStorage();
|
|
862
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({ state: "expected" });
|
|
863
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
638
864
|
let didUnmount = false;
|
|
639
865
|
|
|
640
866
|
try {
|
|
641
867
|
const completeAuth = root
|
|
642
868
|
.getValue()
|
|
643
|
-
.completeAuth("https://example.com/callback?code=abc");
|
|
869
|
+
.completeAuth("https://example.com/callback?code=abc&state=expected");
|
|
644
870
|
await waitFor(
|
|
645
871
|
() => mocks.transports[0]?.finishAuth.mock.calls.length === 1,
|
|
646
872
|
);
|
|
@@ -665,13 +891,15 @@ describe("McpServerResource completeAuth", () => {
|
|
|
665
891
|
rejectFinishAuth = reject;
|
|
666
892
|
}),
|
|
667
893
|
);
|
|
668
|
-
const
|
|
894
|
+
const storage = createStorage();
|
|
895
|
+
vi.mocked(storage.loadAuthState).mockResolvedValue({ state: "expected" });
|
|
896
|
+
const root = mount({ auth: { type: "oauth" }, storage });
|
|
669
897
|
let didUnmount = false;
|
|
670
898
|
|
|
671
899
|
try {
|
|
672
900
|
const completeAuth = root
|
|
673
901
|
.getValue()
|
|
674
|
-
.completeAuth("https://example.com/callback?code=abc");
|
|
902
|
+
.completeAuth("https://example.com/callback?code=abc&state=expected");
|
|
675
903
|
await waitFor(
|
|
676
904
|
() => mocks.transports[0]?.finishAuth.mock.calls.length === 1,
|
|
677
905
|
);
|
|
@@ -1481,3 +1709,136 @@ describe("McpServerResource resource methods", () => {
|
|
|
1481
1709
|
}
|
|
1482
1710
|
});
|
|
1483
1711
|
});
|
|
1712
|
+
|
|
1713
|
+
describe("getConnectionDependencies storage scope", () => {
|
|
1714
|
+
const propsWith = (storage: MCPStorage): McpServerResourceProps => ({
|
|
1715
|
+
id: "docs",
|
|
1716
|
+
kind: "connector",
|
|
1717
|
+
name: "Docs",
|
|
1718
|
+
url: "https://example.com/mcp",
|
|
1719
|
+
auth: { type: "oauth" },
|
|
1720
|
+
storage,
|
|
1721
|
+
redirectUri: "https://example.com/callback",
|
|
1722
|
+
autoConnect: false,
|
|
1723
|
+
onRemove: async () => {},
|
|
1724
|
+
});
|
|
1725
|
+
|
|
1726
|
+
it("keys the connection on a declared storage scopeId", () => {
|
|
1727
|
+
const a = { ...createStorage(), scopeId: "local-storage:a" };
|
|
1728
|
+
const b = { ...createStorage(), scopeId: "local-storage:b" };
|
|
1729
|
+
|
|
1730
|
+
expect(getConnectionDependencies(propsWith(a))).not.toEqual(
|
|
1731
|
+
getConnectionDependencies(propsWith(b)),
|
|
1732
|
+
);
|
|
1733
|
+
});
|
|
1734
|
+
|
|
1735
|
+
it("treats storages sharing a scopeId as the same connection target", () => {
|
|
1736
|
+
const a = { ...createStorage(), scopeId: "local-storage:same" };
|
|
1737
|
+
const b = { ...createStorage(), scopeId: "local-storage:same" };
|
|
1738
|
+
|
|
1739
|
+
expect(getConnectionDependencies(propsWith(a))).toEqual(
|
|
1740
|
+
getConnectionDependencies(propsWith(b)),
|
|
1741
|
+
);
|
|
1742
|
+
});
|
|
1743
|
+
|
|
1744
|
+
it("does not key the connection on storage identity when no scopeId is declared", () => {
|
|
1745
|
+
expect(getConnectionDependencies(propsWith(createStorage()))).toEqual(
|
|
1746
|
+
getConnectionDependencies(propsWith(createStorage())),
|
|
1747
|
+
);
|
|
1748
|
+
});
|
|
1749
|
+
|
|
1750
|
+
it("ignores the storage scope for none-auth servers", () => {
|
|
1751
|
+
const a = { ...createStorage(), scopeId: "local-storage:a" };
|
|
1752
|
+
const b = { ...createStorage(), scopeId: "local-storage:b" };
|
|
1753
|
+
const noneProps = (storage: MCPStorage): McpServerResourceProps => ({
|
|
1754
|
+
...propsWith(storage),
|
|
1755
|
+
auth: { type: "none" },
|
|
1756
|
+
});
|
|
1757
|
+
|
|
1758
|
+
expect(getConnectionDependencies(noneProps(a))).toEqual(
|
|
1759
|
+
getConnectionDependencies(noneProps(b)),
|
|
1760
|
+
);
|
|
1761
|
+
});
|
|
1762
|
+
});
|
|
1763
|
+
|
|
1764
|
+
describe("McpServerResource oauth storage swap", () => {
|
|
1765
|
+
beforeEach(resetMocks);
|
|
1766
|
+
|
|
1767
|
+
it("reconnects onto the replacement storage when the scope changes", async () => {
|
|
1768
|
+
const persisted = {
|
|
1769
|
+
tokens: { access_token: "tok", token_type: "bearer" },
|
|
1770
|
+
};
|
|
1771
|
+
const storageA = {
|
|
1772
|
+
...createStorage(),
|
|
1773
|
+
scopeId: "scope:a",
|
|
1774
|
+
loadAuthState: vi.fn(async () => persisted),
|
|
1775
|
+
};
|
|
1776
|
+
const storageB = {
|
|
1777
|
+
...createStorage(),
|
|
1778
|
+
scopeId: "scope:b",
|
|
1779
|
+
loadAuthState: vi.fn(async () => persisted),
|
|
1780
|
+
};
|
|
1781
|
+
let setStorage!: (s: MCPStorage) => void;
|
|
1782
|
+
|
|
1783
|
+
const Host = resource(function useHost() {
|
|
1784
|
+
const [storage, set] = useState<MCPStorage>(storageA);
|
|
1785
|
+
setStorage = set;
|
|
1786
|
+
return useResource(
|
|
1787
|
+
McpServerResource({
|
|
1788
|
+
id: "docs",
|
|
1789
|
+
kind: "connector",
|
|
1790
|
+
name: "Docs",
|
|
1791
|
+
url: "https://example.com/mcp",
|
|
1792
|
+
auth: { type: "oauth" },
|
|
1793
|
+
storage,
|
|
1794
|
+
redirectUri: "https://example.com/callback",
|
|
1795
|
+
autoConnect: true,
|
|
1796
|
+
connectionTimeout: 10_000,
|
|
1797
|
+
onRemove: vi.fn(async () => {}),
|
|
1798
|
+
}),
|
|
1799
|
+
);
|
|
1800
|
+
});
|
|
1801
|
+
|
|
1802
|
+
const root = createTapRoot(function SwapRoot() {
|
|
1803
|
+
return useResource(Host());
|
|
1804
|
+
});
|
|
1805
|
+
|
|
1806
|
+
try {
|
|
1807
|
+
await waitFor(() => mocks.transports.length === 1);
|
|
1808
|
+
const authProviderA =
|
|
1809
|
+
mocks.StreamableHTTPClientTransport.mock.calls[0]?.[1]?.authProvider;
|
|
1810
|
+
await authProviderA.tokens();
|
|
1811
|
+
expect(storageA.loadAuthState).toHaveBeenCalledWith("docs");
|
|
1812
|
+
expect(storageB.loadAuthState).not.toHaveBeenCalled();
|
|
1813
|
+
|
|
1814
|
+
setStorage(storageB);
|
|
1815
|
+
await waitForResourceUpdate(() => mocks.transports.length === 2);
|
|
1816
|
+
await waitForResourceUpdate(
|
|
1817
|
+
() => vi.mocked(mocks.transports[0]!.close).mock.calls.length > 0,
|
|
1818
|
+
);
|
|
1819
|
+
|
|
1820
|
+
const authProviderB =
|
|
1821
|
+
mocks.StreamableHTTPClientTransport.mock.calls[1]?.[1]?.authProvider;
|
|
1822
|
+
const callsBefore = vi.mocked(storageA.loadAuthState).mock.calls.length;
|
|
1823
|
+
await authProviderB.tokens();
|
|
1824
|
+
expect(storageB.loadAuthState).toHaveBeenCalledWith("docs");
|
|
1825
|
+
expect(vi.mocked(storageA.loadAuthState).mock.calls.length).toBe(
|
|
1826
|
+
callsBefore,
|
|
1827
|
+
);
|
|
1828
|
+
|
|
1829
|
+
await authProviderB.saveTokens({
|
|
1830
|
+
access_token: "fresh",
|
|
1831
|
+
token_type: "bearer",
|
|
1832
|
+
});
|
|
1833
|
+
expect(storageB.saveAuthState).toHaveBeenCalledWith(
|
|
1834
|
+
"docs",
|
|
1835
|
+
expect.objectContaining({
|
|
1836
|
+
tokens: expect.objectContaining({ access_token: "fresh" }),
|
|
1837
|
+
}),
|
|
1838
|
+
);
|
|
1839
|
+
expect(storageA.saveAuthState).not.toHaveBeenCalled();
|
|
1840
|
+
} finally {
|
|
1841
|
+
root.unmount();
|
|
1842
|
+
}
|
|
1843
|
+
});
|
|
1844
|
+
});
|
|
@@ -10,7 +10,10 @@ import {
|
|
|
10
10
|
type ElicitResult,
|
|
11
11
|
type StreamableHTTPClientTransportOptions,
|
|
12
12
|
} from "@modelcontextprotocol/client";
|
|
13
|
-
import {
|
|
13
|
+
import {
|
|
14
|
+
clearOAuthProviderAuthState,
|
|
15
|
+
createOAuthProvider,
|
|
16
|
+
} from "../auth/createOAuthProvider";
|
|
14
17
|
import { buildHeaders } from "../auth/buildHeaders";
|
|
15
18
|
import { assertValidServerId } from "../utils/serverId";
|
|
16
19
|
import { validateElicitationContent } from "./validateElicitationContent";
|
|
@@ -46,13 +49,13 @@ type McpServerResourceInstanceProps = McpServerResourceProps & {
|
|
|
46
49
|
transportCloseQueueRef: { current: Promise<void> };
|
|
47
50
|
};
|
|
48
51
|
|
|
49
|
-
const getConnectionDependencies = (
|
|
52
|
+
export const getConnectionDependencies = (
|
|
50
53
|
props: McpServerResourceProps,
|
|
51
54
|
): readonly unknown[] => {
|
|
52
55
|
const auth = props.auth;
|
|
53
56
|
const authDependencies =
|
|
54
57
|
auth.type === "bearer"
|
|
55
|
-
? [auth.type, auth.token]
|
|
58
|
+
? [auth.type, auth.token, props.storage.scopeId]
|
|
56
59
|
: auth.type === "oauth"
|
|
57
60
|
? [
|
|
58
61
|
auth.type,
|
|
@@ -63,6 +66,7 @@ const getConnectionDependencies = (
|
|
|
63
66
|
auth.registrationEndpoint,
|
|
64
67
|
auth.clientId,
|
|
65
68
|
auth.clientSecret,
|
|
69
|
+
props.storage.scopeId,
|
|
66
70
|
]
|
|
67
71
|
: [auth.type];
|
|
68
72
|
|
|
@@ -102,6 +106,11 @@ const useMcpServerResourceInstance = (
|
|
|
102
106
|
null,
|
|
103
107
|
);
|
|
104
108
|
const connectionGenerationRef = useRef(0);
|
|
109
|
+
const pendingAuthValidationRef = useRef<{
|
|
110
|
+
count: number;
|
|
111
|
+
promise: Promise<void>;
|
|
112
|
+
resolve: () => void;
|
|
113
|
+
} | null>(null);
|
|
105
114
|
const elicitationResolversRef = useRef(
|
|
106
115
|
new Map<
|
|
107
116
|
string,
|
|
@@ -464,6 +473,45 @@ const useMcpServerResourceInstance = (
|
|
|
464
473
|
});
|
|
465
474
|
|
|
466
475
|
const doCompleteAuth = useEffectEvent(async (callbackUrl: string) => {
|
|
476
|
+
const validationGeneration = connectionGenerationRef.current;
|
|
477
|
+
const url = new URL(callbackUrl);
|
|
478
|
+
const state = url.searchParams.get("state");
|
|
479
|
+
if (!state) throw new Error('missing "state" parameter');
|
|
480
|
+
let pendingAuthValidation = pendingAuthValidationRef.current;
|
|
481
|
+
if (!pendingAuthValidation) {
|
|
482
|
+
let resolve!: () => void;
|
|
483
|
+
const promise = new Promise<void>((resolvePromise) => {
|
|
484
|
+
resolve = resolvePromise;
|
|
485
|
+
});
|
|
486
|
+
pendingAuthValidation = { count: 0, promise, resolve };
|
|
487
|
+
pendingAuthValidationRef.current = pendingAuthValidation;
|
|
488
|
+
}
|
|
489
|
+
pendingAuthValidation.count += 1;
|
|
490
|
+
try {
|
|
491
|
+
const persisted = await props.storage.loadAuthState(props.id);
|
|
492
|
+
if (!isCurrentConnection(validationGeneration)) {
|
|
493
|
+
throw createInterruptedAuthError();
|
|
494
|
+
}
|
|
495
|
+
if (!persisted?.state) {
|
|
496
|
+
throw new Error(
|
|
497
|
+
"no pending OAuth authorization request for this server",
|
|
498
|
+
);
|
|
499
|
+
}
|
|
500
|
+
if (persisted.state !== state) {
|
|
501
|
+
throw new Error("OAuth state does not match the authorization request");
|
|
502
|
+
}
|
|
503
|
+
if (!url.searchParams.get("code") && !url.searchParams.get("error")) {
|
|
504
|
+
throw new Error("missing authorization code in callback URL");
|
|
505
|
+
}
|
|
506
|
+
} finally {
|
|
507
|
+
pendingAuthValidation.count -= 1;
|
|
508
|
+
if (pendingAuthValidation.count === 0) {
|
|
509
|
+
pendingAuthValidationRef.current = null;
|
|
510
|
+
pendingAuthValidation.resolve();
|
|
511
|
+
}
|
|
512
|
+
}
|
|
513
|
+
|
|
514
|
+
// Claim the generation before a waiting auto-connect can resume.
|
|
467
515
|
const generation = ++connectionGenerationRef.current;
|
|
468
516
|
cancelPendingElicitations();
|
|
469
517
|
await closePendingTransport();
|
|
@@ -472,9 +520,6 @@ const useMcpServerResourceInstance = (
|
|
|
472
520
|
setConnectionState("authPending");
|
|
473
521
|
setLastError(null);
|
|
474
522
|
try {
|
|
475
|
-
const url = new URL(callbackUrl);
|
|
476
|
-
const code = url.searchParams.get("code");
|
|
477
|
-
if (!code) throw new Error("missing authorization code in callback URL");
|
|
478
523
|
let transport = transportRef.current;
|
|
479
524
|
if (!transport) {
|
|
480
525
|
transport = await buildTransport();
|
|
@@ -486,7 +531,7 @@ const useMcpServerResourceInstance = (
|
|
|
486
531
|
transportRef.current = null;
|
|
487
532
|
clientRef.current = null;
|
|
488
533
|
pendingTransportRef.current = transport;
|
|
489
|
-
await transport.finishAuth(
|
|
534
|
+
await transport.finishAuth(url.searchParams);
|
|
490
535
|
if (!isCurrentConnection(generation)) throw createInterruptedAuthError();
|
|
491
536
|
setAuthorizationUrl(null);
|
|
492
537
|
const connected = await finalizeConnect(transport, generation);
|
|
@@ -536,6 +581,11 @@ const useMcpServerResourceInstance = (
|
|
|
536
581
|
} else if (!persisted?.token) {
|
|
537
582
|
return;
|
|
538
583
|
}
|
|
584
|
+
const pendingAuthValidation = pendingAuthValidationRef.current;
|
|
585
|
+
if (pendingAuthValidation) {
|
|
586
|
+
await pendingAuthValidation.promise;
|
|
587
|
+
if (signal.cancelled || !isCurrentConnection(generation)) return;
|
|
588
|
+
}
|
|
539
589
|
void doConnect();
|
|
540
590
|
},
|
|
541
591
|
);
|
|
@@ -597,7 +647,7 @@ const useMcpServerResourceInstance = (
|
|
|
597
647
|
remove: async () => {
|
|
598
648
|
await doDisconnect();
|
|
599
649
|
try {
|
|
600
|
-
await props.storage
|
|
650
|
+
await clearOAuthProviderAuthState(props.storage, props.id);
|
|
601
651
|
await props.onRemove();
|
|
602
652
|
} catch (err) {
|
|
603
653
|
setLastError({
|
|
@@ -702,8 +752,9 @@ export const McpServerResource = resource(function useMcpServerResource(
|
|
|
702
752
|
setConnection(currentConnection);
|
|
703
753
|
}
|
|
704
754
|
|
|
705
|
-
// Storage
|
|
706
|
-
//
|
|
755
|
+
// Storage keys remounts through its optional scopeId rather than object
|
|
756
|
+
// identity, because a defaulted storage element is rebuilt on ordinary
|
|
757
|
+
// renders.
|
|
707
758
|
return useResource(
|
|
708
759
|
withKey(
|
|
709
760
|
currentConnection.generation,
|