@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.
Files changed (39) hide show
  1. package/dist/auth/createOAuthProvider.d.ts +7 -1
  2. package/dist/auth/createOAuthProvider.d.ts.map +1 -1
  3. package/dist/auth/createOAuthProvider.js +115 -31
  4. package/dist/auth/createOAuthProvider.js.map +1 -1
  5. package/dist/auth/types.d.ts +1 -0
  6. package/dist/auth/types.d.ts.map +1 -1
  7. package/dist/hooks/useMcpOAuthCallback.d.ts.map +1 -1
  8. package/dist/hooks/useMcpOAuthCallback.js +5 -6
  9. package/dist/hooks/useMcpOAuthCallback.js.map +1 -1
  10. package/dist/resources/McpManagerResource.d.ts.map +1 -1
  11. package/dist/resources/McpManagerResource.js +2 -1
  12. package/dist/resources/McpManagerResource.js.map +1 -1
  13. package/dist/resources/McpServerResource.d.ts +2 -1
  14. package/dist/resources/McpServerResource.d.ts.map +1 -1
  15. package/dist/resources/McpServerResource.js +50 -11
  16. package/dist/resources/McpServerResource.js.map +1 -1
  17. package/dist/resources/storage/McpLocalStorage.d.ts +7 -0
  18. package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
  19. package/dist/resources/storage/McpLocalStorage.js +132 -37
  20. package/dist/resources/storage/McpLocalStorage.js.map +1 -1
  21. package/dist/resources/storage/McpMemoryStorage.d.ts.map +1 -1
  22. package/dist/resources/storage/McpMemoryStorage.js +9 -6
  23. package/dist/resources/storage/McpMemoryStorage.js.map +1 -1
  24. package/dist/resources/storage/types.d.ts +12 -0
  25. package/dist/resources/storage/types.d.ts.map +1 -1
  26. package/package.json +7 -7
  27. package/src/auth/createOAuthProvider.test.ts +407 -2
  28. package/src/auth/createOAuthProvider.ts +171 -40
  29. package/src/auth/types.ts +1 -0
  30. package/src/hooks/useMcpOAuthCallback.test.ts +74 -1
  31. package/src/hooks/useMcpOAuthCallback.tsx +11 -8
  32. package/src/resources/McpManagerResource.ts +2 -1
  33. package/src/resources/McpServerResource.test.ts +377 -16
  34. package/src/resources/McpServerResource.ts +61 -10
  35. package/src/resources/storage/McpLocalStorage.test.ts +71 -1
  36. package/src/resources/storage/McpLocalStorage.ts +69 -47
  37. package/src/resources/storage/McpMemoryStorage.test.ts +99 -0
  38. package/src/resources/storage/McpMemoryStorage.ts +23 -17
  39. 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 } = await import("./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 when the callback URL has no authorization code", async () => {
585
- const root = mount();
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.getValue().completeAuth("https://example.com/callback?state=abc"),
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
- message: "missing authorization code in callback URL",
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 root = mount({ auth: { type: "oauth" } });
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.getValue().completeAuth("https://example.com/callback?code=abc"),
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("abc");
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 root = mount({ auth: { type: "oauth" } });
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 root = mount({ auth: { type: "oauth" } });
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 { createOAuthProvider } from "../auth/createOAuthProvider";
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(code);
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.clearAuthState(props.id);
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 resources do not expose a stable scope identity and may return a
706
- // fresh client on ordinary renders, so storage changes cannot key remounts.
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,