@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.
@@ -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
- (value: { state?: string; tokens?: { access_token: string } }) => void
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({ tokens: { access_token: "persisted" } });
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]!({ state: "expected" });
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
- (value: { state?: string; tokens?: { access_token: string } }) => void
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({ tokens: { access_token: "persisted" } });
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]!({ state: "different" });
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({ state: "abc" });
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({ state: "expected" });
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: { state: string }) => void;
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({ state: "expected" });
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({ state: "expected" });
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({ state: "expected" });
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({ state: "expected" });
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 persisted = await props.storage.loadAuthState(props.id);
261
- const headers = buildHeaders(props.auth, persisted);
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 props.storage.loadAuthState(props.id);
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 persisted: Awaited<ReturnType<MCPStorage["loadAuthState"]>>;
579
+ let loaded: Awaited<ReturnType<typeof loadAuthState>>;
567
580
  try {
568
- persisted = await props.storage.loadAuthState(props.id);
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?.tokens) return;
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 (!areConnectionDependenciesEqual(connection.dependencies, dependencies)) {
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
+ });