@assistant-ui/react-mcp 0.0.17 → 0.0.18

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (98) hide show
  1. package/dist/auth/buildHeaders.d.ts +0 -1
  2. package/dist/auth/buildHeaders.d.ts.map +1 -1
  3. package/dist/auth/createOAuthProvider.d.ts +4 -3
  4. package/dist/auth/createOAuthProvider.d.ts.map +1 -1
  5. package/dist/auth/createOAuthProvider.js +2 -1
  6. package/dist/auth/createOAuthProvider.js.map +1 -1
  7. package/dist/auth/types.d.ts +2 -2
  8. package/dist/auth/types.d.ts.map +1 -1
  9. package/dist/connector.d.ts +0 -1
  10. package/dist/connector.d.ts.map +1 -1
  11. package/dist/context/McpConnectorByIndexProvider.d.ts +0 -1
  12. package/dist/context/McpConnectorByIndexProvider.d.ts.map +1 -1
  13. package/dist/context/McpCustomServerByIndexProvider.d.ts +0 -1
  14. package/dist/context/McpCustomServerByIndexProvider.d.ts.map +1 -1
  15. package/dist/context/McpServerByIdProvider.d.ts +0 -1
  16. package/dist/context/McpServerByIdProvider.d.ts.map +1 -1
  17. package/dist/hooks/useMcpOAuthCallback.d.ts +4 -3
  18. package/dist/hooks/useMcpOAuthCallback.d.ts.map +1 -1
  19. package/dist/hooks/useMcpOAuthCallback.js +14 -8
  20. package/dist/hooks/useMcpOAuthCallback.js.map +1 -1
  21. package/dist/mcp-scope.d.ts +16 -6
  22. package/dist/mcp-scope.d.ts.map +1 -1
  23. package/dist/primitives/addForm/McpAddFormAuthFields.d.ts +0 -1
  24. package/dist/primitives/addForm/McpAddFormAuthFields.d.ts.map +1 -1
  25. package/dist/primitives/addForm/McpAddFormAuthSelect.d.ts +0 -1
  26. package/dist/primitives/addForm/McpAddFormAuthSelect.d.ts.map +1 -1
  27. package/dist/primitives/addForm/McpAddFormCancel.d.ts +0 -1
  28. package/dist/primitives/addForm/McpAddFormCancel.d.ts.map +1 -1
  29. package/dist/primitives/addForm/McpAddFormError.d.ts +0 -1
  30. package/dist/primitives/addForm/McpAddFormError.d.ts.map +1 -1
  31. package/dist/primitives/addForm/McpAddFormNameField.d.ts +0 -1
  32. package/dist/primitives/addForm/McpAddFormNameField.d.ts.map +1 -1
  33. package/dist/primitives/addForm/McpAddFormRoot.d.ts +0 -1
  34. package/dist/primitives/addForm/McpAddFormRoot.d.ts.map +1 -1
  35. package/dist/primitives/addForm/McpAddFormSubmit.d.ts +0 -1
  36. package/dist/primitives/addForm/McpAddFormSubmit.d.ts.map +1 -1
  37. package/dist/primitives/addForm/McpAddFormUrlField.d.ts +0 -1
  38. package/dist/primitives/addForm/McpAddFormUrlField.d.ts.map +1 -1
  39. package/dist/primitives/addForm/context.d.ts +0 -1
  40. package/dist/primitives/addForm/context.d.ts.map +1 -1
  41. package/dist/primitives/addForm.d.ts +0 -2
  42. package/dist/primitives/manager/McpManagerAddCustomTrigger.d.ts +0 -1
  43. package/dist/primitives/manager/McpManagerAddCustomTrigger.d.ts.map +1 -1
  44. package/dist/primitives/manager/McpManagerConnectors.d.ts +0 -1
  45. package/dist/primitives/manager/McpManagerConnectors.d.ts.map +1 -1
  46. package/dist/primitives/manager/McpManagerCustomServers.d.ts +0 -1
  47. package/dist/primitives/manager/McpManagerCustomServers.d.ts.map +1 -1
  48. package/dist/primitives/manager/McpManagerRoot.d.ts +0 -1
  49. package/dist/primitives/manager/McpManagerRoot.d.ts.map +1 -1
  50. package/dist/primitives/manager.d.ts +0 -2
  51. package/dist/primitives/server/McpServerConnectButton.d.ts +0 -1
  52. package/dist/primitives/server/McpServerConnectButton.d.ts.map +1 -1
  53. package/dist/primitives/server/McpServerDisconnectButton.d.ts +0 -1
  54. package/dist/primitives/server/McpServerDisconnectButton.d.ts.map +1 -1
  55. package/dist/primitives/server/McpServerError.d.ts +0 -1
  56. package/dist/primitives/server/McpServerError.d.ts.map +1 -1
  57. package/dist/primitives/server/McpServerIcon.d.ts +4 -3
  58. package/dist/primitives/server/McpServerIcon.d.ts.map +1 -1
  59. package/dist/primitives/server/McpServerName.d.ts +0 -1
  60. package/dist/primitives/server/McpServerName.d.ts.map +1 -1
  61. package/dist/primitives/server/McpServerOAuthLink.d.ts +4 -3
  62. package/dist/primitives/server/McpServerOAuthLink.d.ts.map +1 -1
  63. package/dist/primitives/server/McpServerRemoveButton.d.ts +0 -1
  64. package/dist/primitives/server/McpServerRemoveButton.d.ts.map +1 -1
  65. package/dist/primitives/server/McpServerRoot.d.ts +0 -1
  66. package/dist/primitives/server/McpServerRoot.d.ts.map +1 -1
  67. package/dist/primitives/server/McpServerStatus.d.ts +0 -1
  68. package/dist/primitives/server/McpServerStatus.d.ts.map +1 -1
  69. package/dist/primitives/server/McpServerToolName.d.ts +0 -1
  70. package/dist/primitives/server/McpServerToolName.d.ts.map +1 -1
  71. package/dist/primitives/server/McpServerTools.d.ts +0 -1
  72. package/dist/primitives/server/McpServerTools.d.ts.map +1 -1
  73. package/dist/primitives/server.d.ts +0 -2
  74. package/dist/resources/McpManagerResource.d.ts +6 -4
  75. package/dist/resources/McpManagerResource.d.ts.map +1 -1
  76. package/dist/resources/McpServerResource.d.ts +0 -1
  77. package/dist/resources/McpServerResource.d.ts.map +1 -1
  78. package/dist/resources/McpServerResource.js +10 -3
  79. package/dist/resources/McpServerResource.js.map +1 -1
  80. package/dist/resources/storage/McpCustomStorage.d.ts +0 -1
  81. package/dist/resources/storage/McpCustomStorage.d.ts.map +1 -1
  82. package/dist/resources/storage/McpLocalStorage.d.ts +8 -3
  83. package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
  84. package/dist/resources/storage/McpLocalStorage.js +55 -3
  85. package/dist/resources/storage/McpLocalStorage.js.map +1 -1
  86. package/dist/resources/storage/McpMemoryStorage.d.ts +0 -1
  87. package/dist/resources/storage/McpMemoryStorage.d.ts.map +1 -1
  88. package/dist/resources/storage/types.d.ts +0 -1
  89. package/dist/resources/storage/types.d.ts.map +1 -1
  90. package/dist/utils/serverId.d.ts.map +1 -1
  91. package/package.json +8 -8
  92. package/src/hooks/useMcpOAuthCallback.test.ts +24 -0
  93. package/src/hooks/useMcpOAuthCallback.tsx +20 -5
  94. package/src/mcp-scope.ts +2 -0
  95. package/src/resources/McpServerResource.test.ts +152 -9
  96. package/src/resources/McpServerResource.ts +10 -1
  97. package/src/resources/storage/McpLocalStorage.test.ts +184 -0
  98. package/src/resources/storage/McpLocalStorage.ts +113 -3
@@ -2,6 +2,20 @@ import { type FC, type ReactNode, useEffect, useRef, useState } from "react";
2
2
  import { useAui } from "@assistant-ui/store";
3
3
  import { decodeServerIdFromState } from "../auth/createOAuthProvider";
4
4
 
5
+ export const createMcpOAuthCallbackError = (
6
+ err: unknown,
7
+ serverId: string | null,
8
+ ): Error => {
9
+ const message = err instanceof Error ? err.message : String(err);
10
+ if (serverId) {
11
+ return new Error(
12
+ `MCP OAuth callback for server "${serverId}" failed: ${message}`,
13
+ { cause: err },
14
+ );
15
+ }
16
+ return new Error(`MCP OAuth callback failed: ${message}`, { cause: err });
17
+ };
18
+
5
19
  export type UseMcpOAuthCallbackOptions = {
6
20
  /** Defaults to `window.location.href`. */
7
21
  url?: string;
@@ -39,27 +53,28 @@ export function useMcpOAuthCallback(
39
53
  startedRef.current = url;
40
54
 
41
55
  (async () => {
56
+ let serverId: string | null = null;
42
57
  try {
43
58
  const parsed = new URL(url);
44
59
  const state = parsed.searchParams.get("state");
60
+ if (state) serverId = decodeServerIdFromState(state);
45
61
  const error = parsed.searchParams.get("error");
46
62
  if (error) {
47
63
  throw new Error(
48
64
  parsed.searchParams.get("error_description") ?? error,
49
65
  );
50
66
  }
51
- if (!state) throw new Error("missing state parameter in callback URL");
52
- const serverId = decodeServerIdFromState(state);
67
+ if (!state) throw new Error('missing "state" parameter');
53
68
  if (!serverId) {
54
- throw new Error("callback state does not match an MCP server");
69
+ throw new Error("state was not created by assistant-ui MCP");
55
70
  }
56
71
  setResult({ status: "running", serverId, error: null });
57
72
  await aui.mcp().server({ id: serverId }).completeAuth(url);
58
73
  setResult({ status: "done", serverId, error: null });
59
74
  optsRef.current.onComplete?.(serverId);
60
75
  } catch (err) {
61
- const e = err instanceof Error ? err : new Error(String(err));
62
- setResult((prev) => ({ ...prev, status: "error", error: e }));
76
+ const e = createMcpOAuthCallbackError(err, serverId);
77
+ setResult({ status: "error", serverId, error: e });
63
78
  optsRef.current.onError?.(e);
64
79
  }
65
80
  })();
package/src/mcp-scope.ts CHANGED
@@ -72,6 +72,8 @@ export type MCPServerMethods = {
72
72
  disconnect: () => Promise<void>;
73
73
  remove: () => Promise<void>;
74
74
  callTool: (name: string, args: unknown) => Promise<unknown>;
75
+ /** List resources exposed by the server. Returns the raw MCP `ListResourcesResult`. */
76
+ listResources: (params?: { cursor?: string | undefined }) => Promise<unknown>;
75
77
  /** Read a resource by URI. Returns the raw MCP `ReadResourceResult`. */
76
78
  readResource: (uri: string) => Promise<unknown>;
77
79
  /** OAuth only: pass full callback URL (e.g. window.location.href) */
@@ -1,5 +1,6 @@
1
1
  import { createTapRoot, useResource } from "@assistant-ui/tap";
2
2
  import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
3
+ import type { MCPAuthConfig } from "../mcp-scope";
3
4
  import type { MCPStorage } from "./storage/types";
4
5
 
5
6
  const mocks = vi.hoisted(() => {
@@ -15,6 +16,7 @@ const mocks = vi.hoisted(() => {
15
16
  }>;
16
17
  }>
17
18
  > = [];
19
+ const finishAuthResults: Array<() => Promise<void>> = [];
18
20
 
19
21
  const Client = vi.fn().mockImplementation(function Client(this: any) {
20
22
  const index = clients.length;
@@ -23,6 +25,7 @@ const mocks = vi.hoisted(() => {
23
25
  () => listToolsResults[index]?.() ?? Promise.resolve({ tools: [] }),
24
26
  );
25
27
  this.callTool = vi.fn();
28
+ this.listResources = vi.fn(() => Promise.resolve({ resources: [] }));
26
29
  this.readResource = vi.fn();
27
30
  clients.push(this);
28
31
  });
@@ -30,8 +33,11 @@ const mocks = vi.hoisted(() => {
30
33
  const StreamableHTTPClientTransport = vi
31
34
  .fn()
32
35
  .mockImplementation(function StreamableHTTPClientTransport(this: any) {
36
+ const index = transports.length;
33
37
  this.close = vi.fn(() => Promise.resolve());
34
- this.finishAuth = vi.fn(() => Promise.resolve());
38
+ this.finishAuth = vi.fn(
39
+ () => finishAuthResults[index]?.() ?? Promise.resolve(),
40
+ );
35
41
  transports.push(this);
36
42
  });
37
43
 
@@ -42,6 +48,7 @@ const mocks = vi.hoisted(() => {
42
48
  transports,
43
49
  connectResults,
44
50
  listToolsResults,
51
+ finishAuthResults,
45
52
  };
46
53
  });
47
54
 
@@ -61,6 +68,18 @@ const tick = async () => {
61
68
  await Promise.resolve();
62
69
  };
63
70
 
71
+ const flushMacrotask = async () => {
72
+ await new Promise<void>((resolve) => {
73
+ const channel = new MessageChannel();
74
+ channel.port1.onmessage = () => {
75
+ channel.port1.close();
76
+ channel.port2.close();
77
+ resolve();
78
+ };
79
+ channel.port2.postMessage(null);
80
+ });
81
+ };
82
+
64
83
  const waitFor = async (predicate: () => boolean) => {
65
84
  for (let i = 0; i < 20; i++) {
66
85
  if (predicate()) return;
@@ -77,7 +96,20 @@ const createStorage = (): MCPStorage => ({
77
96
  clearAuthState: vi.fn(async () => {}),
78
97
  });
79
98
 
80
- const mount = (props?: { connectionTimeout?: number | undefined }) => {
99
+ const resetMocks = () => {
100
+ mocks.clients.length = 0;
101
+ mocks.transports.length = 0;
102
+ mocks.connectResults.length = 0;
103
+ mocks.listToolsResults.length = 0;
104
+ mocks.finishAuthResults.length = 0;
105
+ mocks.Client.mockClear();
106
+ mocks.StreamableHTTPClientTransport.mockClear();
107
+ };
108
+
109
+ const mount = (props?: {
110
+ auth?: MCPAuthConfig | undefined;
111
+ connectionTimeout?: number | undefined;
112
+ }) => {
81
113
  const connectionTimeout =
82
114
  props && "connectionTimeout" in props ? props.connectionTimeout : 10_000;
83
115
 
@@ -88,7 +120,7 @@ const mount = (props?: { connectionTimeout?: number | undefined }) => {
88
120
  kind: "connector",
89
121
  name: "Docs",
90
122
  url: "https://example.com/mcp",
91
- auth: { type: "none" },
123
+ auth: props?.auth ?? { type: "none" },
92
124
  storage: createStorage(),
93
125
  redirectUri: "https://example.com/callback",
94
126
  autoConnect: false,
@@ -102,12 +134,7 @@ const mount = (props?: { connectionTimeout?: number | undefined }) => {
102
134
  describe("McpServerResource connectionTimeout", () => {
103
135
  beforeEach(() => {
104
136
  vi.useFakeTimers();
105
- mocks.clients.length = 0;
106
- mocks.transports.length = 0;
107
- mocks.connectResults.length = 0;
108
- mocks.listToolsResults.length = 0;
109
- mocks.Client.mockClear();
110
- mocks.StreamableHTTPClientTransport.mockClear();
137
+ resetMocks();
111
138
  });
112
139
 
113
140
  afterEach(() => {
@@ -219,3 +246,119 @@ describe("McpServerResource connectionTimeout", () => {
219
246
  }
220
247
  });
221
248
  });
249
+
250
+ describe("McpServerResource completeAuth", () => {
251
+ beforeEach(resetMocks);
252
+
253
+ it("rejects when the callback URL has no authorization code", async () => {
254
+ const root = mount();
255
+
256
+ try {
257
+ await expect(
258
+ root.getValue().completeAuth("https://example.com/callback?state=abc"),
259
+ ).rejects.toThrow("missing authorization code in callback URL");
260
+ await flushMacrotask();
261
+
262
+ expect(root.getValue().getState()).toMatchObject({
263
+ connectionState: "error",
264
+ lastError: {
265
+ message: "missing authorization code in callback URL",
266
+ },
267
+ });
268
+ } finally {
269
+ root.unmount();
270
+ }
271
+ });
272
+
273
+ it("rejects after storing finishAuth failures on the server state", async () => {
274
+ mocks.finishAuthResults.push(() =>
275
+ Promise.reject(new Error("invalid_grant")),
276
+ );
277
+ const root = mount({ auth: { type: "oauth" } });
278
+
279
+ try {
280
+ await expect(
281
+ root.getValue().completeAuth("https://example.com/callback?code=abc"),
282
+ ).rejects.toThrow("invalid_grant");
283
+ await flushMacrotask();
284
+
285
+ expect(root.getValue().getState()).toMatchObject({
286
+ connectionState: "error",
287
+ lastError: {
288
+ message: "invalid_grant",
289
+ },
290
+ });
291
+ expect(mocks.transports[0].finishAuth).toHaveBeenCalledWith("abc");
292
+ expect(mocks.transports[0].close).toHaveBeenCalledTimes(1);
293
+ } finally {
294
+ root.unmount();
295
+ }
296
+ });
297
+ });
298
+
299
+ describe("McpServerResource resource methods", () => {
300
+ beforeEach(() => {
301
+ mocks.clients.length = 0;
302
+ mocks.transports.length = 0;
303
+ mocks.connectResults.length = 0;
304
+ mocks.listToolsResults.length = 0;
305
+ mocks.Client.mockClear();
306
+ mocks.StreamableHTTPClientTransport.mockClear();
307
+ });
308
+
309
+ it("lists resources from a connected server", async () => {
310
+ const result = {
311
+ resources: [
312
+ {
313
+ uri: "docs://intro",
314
+ name: "Intro",
315
+ mimeType: "text/markdown",
316
+ },
317
+ ],
318
+ };
319
+ const root = mount();
320
+
321
+ try {
322
+ await root.getValue().connect();
323
+ mocks.clients[0].listResources.mockResolvedValueOnce(result);
324
+
325
+ await expect(root.getValue().listResources()).resolves.toBe(result);
326
+ expect(mocks.clients[0].listResources).toHaveBeenCalledTimes(1);
327
+ } finally {
328
+ root.unmount();
329
+ }
330
+ });
331
+
332
+ it("forwards the resource pagination cursor", async () => {
333
+ const result = {
334
+ resources: [{ uri: "docs://page-two", name: "Page two" }],
335
+ };
336
+ const root = mount();
337
+
338
+ try {
339
+ await root.getValue().connect();
340
+ mocks.clients[0].listResources.mockResolvedValueOnce(result);
341
+
342
+ await expect(
343
+ root.getValue().listResources({ cursor: "next-page" }),
344
+ ).resolves.toBe(result);
345
+ expect(mocks.clients[0].listResources).toHaveBeenCalledWith({
346
+ cursor: "next-page",
347
+ });
348
+ } finally {
349
+ root.unmount();
350
+ }
351
+ });
352
+
353
+ it("rejects listResources when the server is disconnected", async () => {
354
+ const root = mount();
355
+
356
+ try {
357
+ await expect(root.getValue().listResources()).rejects.toThrow(
358
+ 'MCP server "docs" is not connected',
359
+ );
360
+ } finally {
361
+ root.unmount();
362
+ }
363
+ });
364
+ });
@@ -221,10 +221,12 @@ const useMcpServerResource = (
221
221
  await finalizeConnect(transport);
222
222
  } catch (err) {
223
223
  await closeTransport();
224
+ const error = err instanceof Error ? err : new Error(String(err));
224
225
  setLastError({
225
- message: err instanceof Error ? err.message : String(err),
226
+ message: error.message,
226
227
  });
227
228
  setConnectionState("error");
229
+ throw error;
228
230
  }
229
231
  });
230
232
 
@@ -308,6 +310,13 @@ const useMcpServerResource = (
308
310
  arguments: args as Record<string, unknown> | undefined,
309
311
  });
310
312
  },
313
+ listResources: async (params) => {
314
+ const client = clientRef.current;
315
+ if (!client) {
316
+ throw new Error(`MCP server "${props.id}" is not connected`);
317
+ }
318
+ return await client.listResources(params);
319
+ },
311
320
  readResource: async (uri) => {
312
321
  const client = clientRef.current;
313
322
  if (!client) {
@@ -0,0 +1,184 @@
1
+ import { createTapRoot, useResource } from "@assistant-ui/tap";
2
+ import { describe, expect, it } from "vitest";
3
+
4
+ import {
5
+ McpLocalStorage,
6
+ normalizeCustomServerRecords,
7
+ normalizePersistedAuthState,
8
+ } from "./McpLocalStorage";
9
+
10
+ const validRecord = {
11
+ id: "docs",
12
+ name: "Docs",
13
+ url: "https://docs.example.com/mcp",
14
+ auth: { type: "none" },
15
+ createdAt: 1,
16
+ } as const;
17
+
18
+ describe("normalizeCustomServerRecords", () => {
19
+ it("returns an empty list when the persisted value is not an array", () => {
20
+ expect(normalizeCustomServerRecords({ bad: true })).toEqual([]);
21
+ expect(normalizeCustomServerRecords(null)).toEqual([]);
22
+ });
23
+
24
+ it("filters malformed custom server entries", () => {
25
+ expect(
26
+ normalizeCustomServerRecords([
27
+ validRecord,
28
+ null,
29
+ { ...validRecord, id: "" },
30
+ { ...validRecord, id: "docs__search" },
31
+ { ...validRecord, name: "" },
32
+ { ...validRecord, url: " " },
33
+ { ...validRecord, name: 123 },
34
+ { ...validRecord, url: null },
35
+ { ...validRecord, createdAt: Number.NaN },
36
+ { ...validRecord, auth: { type: "bearer", token: "" } },
37
+ { ...validRecord, auth: { type: "bearer", token: 123 } },
38
+ ]),
39
+ ).toEqual([validRecord]);
40
+ });
41
+
42
+ it("accepts persisted bearer and oauth auth configs", () => {
43
+ expect(
44
+ normalizeCustomServerRecords([
45
+ {
46
+ ...validRecord,
47
+ id: "private-docs",
48
+ auth: { type: "bearer", token: "token" },
49
+ },
50
+ {
51
+ ...validRecord,
52
+ id: "oauth-docs",
53
+ auth: {
54
+ type: "oauth",
55
+ scopes: ["docs.read"],
56
+ authorizationEndpoint: "https://docs.example.com/oauth/authorize",
57
+ tokenEndpoint: "https://docs.example.com/oauth/token",
58
+ },
59
+ },
60
+ ]),
61
+ ).toHaveLength(2);
62
+ });
63
+ });
64
+
65
+ describe("normalizePersistedAuthState", () => {
66
+ it("returns null when the persisted value is not an object", () => {
67
+ expect(normalizePersistedAuthState(null)).toBeNull();
68
+ expect(normalizePersistedAuthState("bad")).toBeNull();
69
+ expect(normalizePersistedAuthState([])).toBeNull();
70
+ });
71
+
72
+ it("drops malformed auth state fields", () => {
73
+ expect(
74
+ normalizePersistedAuthState({
75
+ token: "",
76
+ codeVerifier: 123,
77
+ tokens: {
78
+ access_token: "access-token",
79
+ token_type: "Bearer",
80
+ expires_in: Number.POSITIVE_INFINITY,
81
+ },
82
+ clientInformation: { client_id: "client-id" },
83
+ }),
84
+ ).toBeNull();
85
+ });
86
+
87
+ it("keeps valid bearer and OAuth callback state", () => {
88
+ expect(
89
+ normalizePersistedAuthState({
90
+ token: "bearer-token",
91
+ codeVerifier: "pkce-verifier",
92
+ }),
93
+ ).toEqual({
94
+ token: "bearer-token",
95
+ codeVerifier: "pkce-verifier",
96
+ });
97
+ });
98
+
99
+ it("keeps valid OAuth tokens and client information", () => {
100
+ const tokens = {
101
+ access_token: "access-token",
102
+ token_type: "Bearer",
103
+ refresh_token: "refresh-token",
104
+ expires_in: 3600,
105
+ scope: "docs.read",
106
+ };
107
+ const clientInformation = {
108
+ client_id: "client-id",
109
+ client_secret: "client-secret",
110
+ redirect_uris: ["http://localhost/callback"],
111
+ };
112
+
113
+ expect(
114
+ normalizePersistedAuthState({
115
+ tokens,
116
+ clientInformation,
117
+ }),
118
+ ).toEqual({
119
+ tokens,
120
+ clientInformation,
121
+ });
122
+ });
123
+
124
+ it("keeps valid fields when neighboring fields are malformed", () => {
125
+ expect(
126
+ normalizePersistedAuthState({
127
+ token: "bearer-token",
128
+ codeVerifier: 123,
129
+ tokens: { access_token: "access-token" },
130
+ clientInformation: "not-client-info",
131
+ }),
132
+ ).toEqual({
133
+ token: "bearer-token",
134
+ });
135
+ });
136
+ });
137
+
138
+ const createStorage = (): Storage => {
139
+ const data = new Map<string, string>();
140
+ return {
141
+ get length() {
142
+ return data.size;
143
+ },
144
+ clear: () => {
145
+ data.clear();
146
+ },
147
+ getItem: (key) => data.get(key) ?? null,
148
+ key: (index) => Array.from(data.keys())[index] ?? null,
149
+ removeItem: (key) => {
150
+ data.delete(key);
151
+ },
152
+ setItem: (key, value) => {
153
+ data.set(key, value);
154
+ },
155
+ };
156
+ };
157
+
158
+ const loadStorage = (storage: Storage) =>
159
+ createTapRoot(function McpStorageRoot() {
160
+ return useResource(
161
+ McpLocalStorage({
162
+ keyPrefix: "test-mcp",
163
+ storage,
164
+ }),
165
+ );
166
+ }).getValue();
167
+
168
+ describe("McpLocalStorage auth state", () => {
169
+ it("normalizes loaded auth state from localStorage", async () => {
170
+ const storage = createStorage();
171
+ storage.setItem(
172
+ "test-mcp:auth:docs",
173
+ JSON.stringify({
174
+ token: "bearer-token",
175
+ codeVerifier: 123,
176
+ tokens: "not-tokens",
177
+ }),
178
+ );
179
+
180
+ await expect(loadStorage(storage).loadAuthState("docs")).resolves.toEqual({
181
+ token: "bearer-token",
182
+ });
183
+ });
184
+ });
@@ -1,6 +1,11 @@
1
1
  import { resource } from "@assistant-ui/tap";
2
- import type { MCPCustomServerRecord } from "../../mcp-scope";
2
+ import {
3
+ OAuthClientInformationFullSchema,
4
+ OAuthTokensSchema,
5
+ } from "@modelcontextprotocol/sdk/shared/auth.js";
6
+ import type { MCPAuthConfig, MCPCustomServerRecord } from "../../mcp-scope";
3
7
  import type { MCPPersistedAuthState } from "../../auth/types";
8
+ import { assertValidServerId } from "../../utils/serverId";
4
9
  import type { MCPStorage } from "./types";
5
10
 
6
11
  export type McpLocalStorageOptions = {
@@ -22,6 +27,111 @@ function resolveStorage(opts: McpLocalStorageOptions): Storage | null {
22
27
  return null;
23
28
  }
24
29
 
30
+ const isRecord = (value: unknown): value is Record<string, unknown> =>
31
+ typeof value === "object" && value !== null && !Array.isArray(value);
32
+
33
+ const isOptionalString = (value: unknown): value is string | undefined =>
34
+ value === undefined || typeof value === "string";
35
+
36
+ const isNonEmptyString = (value: unknown): value is string =>
37
+ typeof value === "string" && value.trim().length > 0;
38
+
39
+ const isOptionalNonEmptyString = (
40
+ value: unknown,
41
+ ): value is string | undefined =>
42
+ value === undefined || isNonEmptyString(value);
43
+
44
+ const isOptionalStringArray = (value: unknown): value is string[] | undefined =>
45
+ value === undefined ||
46
+ (Array.isArray(value) && value.every((item) => typeof item === "string"));
47
+
48
+ const isValidServerId = (id: string): boolean => {
49
+ try {
50
+ assertValidServerId(id);
51
+ return true;
52
+ } catch {
53
+ return false;
54
+ }
55
+ };
56
+
57
+ const isMCPAuthConfig = (auth: unknown): auth is MCPAuthConfig => {
58
+ if (!isRecord(auth)) return false;
59
+
60
+ switch (auth.type) {
61
+ case "none":
62
+ return true;
63
+ case "bearer":
64
+ return isOptionalNonEmptyString(auth.token);
65
+ case "oauth":
66
+ return (
67
+ isOptionalStringArray(auth.scopes) &&
68
+ isOptionalString(auth.authorizationEndpoint) &&
69
+ isOptionalString(auth.tokenEndpoint) &&
70
+ isOptionalString(auth.registrationEndpoint) &&
71
+ isOptionalString(auth.clientId) &&
72
+ isOptionalString(auth.clientSecret)
73
+ );
74
+ default:
75
+ return false;
76
+ }
77
+ };
78
+
79
+ const isCustomServerRecord = (
80
+ value: unknown,
81
+ ): value is MCPCustomServerRecord => {
82
+ if (!isRecord(value)) return false;
83
+ if (typeof value.id !== "string" || !isValidServerId(value.id)) {
84
+ return false;
85
+ }
86
+ return (
87
+ isNonEmptyString(value.name) &&
88
+ isNonEmptyString(value.url) &&
89
+ Number.isFinite(value.createdAt) &&
90
+ isMCPAuthConfig(value.auth)
91
+ );
92
+ };
93
+
94
+ export const normalizeCustomServerRecords = (
95
+ value: unknown,
96
+ ): MCPCustomServerRecord[] => {
97
+ if (!Array.isArray(value)) return [];
98
+ return value.filter(isCustomServerRecord);
99
+ };
100
+
101
+ const normalizeOAuthTokens = (
102
+ value: unknown,
103
+ ): MCPPersistedAuthState["tokens"] | undefined => {
104
+ const result = OAuthTokensSchema.safeParse(value);
105
+ return result.success ? result.data : undefined;
106
+ };
107
+
108
+ const normalizeClientInformation = (
109
+ value: unknown,
110
+ ): MCPPersistedAuthState["clientInformation"] | undefined => {
111
+ const result = OAuthClientInformationFullSchema.safeParse(value);
112
+ return result.success ? result.data : undefined;
113
+ };
114
+
115
+ export const normalizePersistedAuthState = (
116
+ value: unknown,
117
+ ): MCPPersistedAuthState | null => {
118
+ if (!isRecord(value)) return null;
119
+
120
+ const state: MCPPersistedAuthState = {};
121
+ if (isNonEmptyString(value.token)) state.token = value.token;
122
+ if (isNonEmptyString(value.codeVerifier)) {
123
+ state.codeVerifier = value.codeVerifier;
124
+ }
125
+
126
+ const tokens = normalizeOAuthTokens(value.tokens);
127
+ if (tokens) state.tokens = tokens;
128
+
129
+ const clientInformation = normalizeClientInformation(value.clientInformation);
130
+ if (clientInformation) state.clientInformation = clientInformation;
131
+
132
+ return Object.keys(state).length > 0 ? state : null;
133
+ };
134
+
25
135
  const useMcpLocalStorage = (opts: McpLocalStorageOptions = {}): MCPStorage => {
26
136
  const prefix = opts.keyPrefix ?? "aui-mcp";
27
137
  const customServersKey = `${prefix}:custom-servers`;
@@ -59,12 +169,12 @@ const useMcpLocalStorage = (opts: McpLocalStorageOptions = {}): MCPStorage => {
59
169
 
60
170
  return {
61
171
  loadCustomServers: async () =>
62
- read<MCPCustomServerRecord[]>(customServersKey, []),
172
+ normalizeCustomServerRecords(read<unknown>(customServersKey, [])),
63
173
  saveCustomServers: async (records) => {
64
174
  write(customServersKey, records);
65
175
  },
66
176
  loadAuthState: async (id) =>
67
- read<MCPPersistedAuthState | null>(authKey(id), null),
177
+ normalizePersistedAuthState(read<unknown>(authKey(id), null)),
68
178
  saveAuthState: async (id, state) => {
69
179
  write(authKey(id), state);
70
180
  },