@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.
- package/dist/auth/buildHeaders.d.ts +0 -1
- package/dist/auth/buildHeaders.d.ts.map +1 -1
- package/dist/auth/createOAuthProvider.d.ts +4 -3
- package/dist/auth/createOAuthProvider.d.ts.map +1 -1
- package/dist/auth/createOAuthProvider.js +2 -1
- package/dist/auth/createOAuthProvider.js.map +1 -1
- package/dist/auth/types.d.ts +2 -2
- package/dist/auth/types.d.ts.map +1 -1
- package/dist/connector.d.ts +0 -1
- package/dist/connector.d.ts.map +1 -1
- package/dist/context/McpConnectorByIndexProvider.d.ts +0 -1
- package/dist/context/McpConnectorByIndexProvider.d.ts.map +1 -1
- package/dist/context/McpCustomServerByIndexProvider.d.ts +0 -1
- package/dist/context/McpCustomServerByIndexProvider.d.ts.map +1 -1
- package/dist/context/McpServerByIdProvider.d.ts +0 -1
- package/dist/context/McpServerByIdProvider.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.d.ts +4 -3
- package/dist/hooks/useMcpOAuthCallback.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.js +14 -8
- package/dist/hooks/useMcpOAuthCallback.js.map +1 -1
- package/dist/mcp-scope.d.ts +16 -6
- package/dist/mcp-scope.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormAuthFields.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormAuthFields.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormAuthSelect.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormAuthSelect.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormCancel.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormCancel.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormError.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormError.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormNameField.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormNameField.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormRoot.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormRoot.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormSubmit.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormSubmit.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormUrlField.d.ts +0 -1
- package/dist/primitives/addForm/McpAddFormUrlField.d.ts.map +1 -1
- package/dist/primitives/addForm/context.d.ts +0 -1
- package/dist/primitives/addForm/context.d.ts.map +1 -1
- package/dist/primitives/addForm.d.ts +0 -2
- package/dist/primitives/manager/McpManagerAddCustomTrigger.d.ts +0 -1
- package/dist/primitives/manager/McpManagerAddCustomTrigger.d.ts.map +1 -1
- package/dist/primitives/manager/McpManagerConnectors.d.ts +0 -1
- package/dist/primitives/manager/McpManagerConnectors.d.ts.map +1 -1
- package/dist/primitives/manager/McpManagerCustomServers.d.ts +0 -1
- package/dist/primitives/manager/McpManagerCustomServers.d.ts.map +1 -1
- package/dist/primitives/manager/McpManagerRoot.d.ts +0 -1
- package/dist/primitives/manager/McpManagerRoot.d.ts.map +1 -1
- package/dist/primitives/manager.d.ts +0 -2
- package/dist/primitives/server/McpServerConnectButton.d.ts +0 -1
- package/dist/primitives/server/McpServerConnectButton.d.ts.map +1 -1
- package/dist/primitives/server/McpServerDisconnectButton.d.ts +0 -1
- package/dist/primitives/server/McpServerDisconnectButton.d.ts.map +1 -1
- package/dist/primitives/server/McpServerError.d.ts +0 -1
- package/dist/primitives/server/McpServerError.d.ts.map +1 -1
- package/dist/primitives/server/McpServerIcon.d.ts +4 -3
- package/dist/primitives/server/McpServerIcon.d.ts.map +1 -1
- package/dist/primitives/server/McpServerName.d.ts +0 -1
- package/dist/primitives/server/McpServerName.d.ts.map +1 -1
- package/dist/primitives/server/McpServerOAuthLink.d.ts +4 -3
- package/dist/primitives/server/McpServerOAuthLink.d.ts.map +1 -1
- package/dist/primitives/server/McpServerRemoveButton.d.ts +0 -1
- package/dist/primitives/server/McpServerRemoveButton.d.ts.map +1 -1
- package/dist/primitives/server/McpServerRoot.d.ts +0 -1
- package/dist/primitives/server/McpServerRoot.d.ts.map +1 -1
- package/dist/primitives/server/McpServerStatus.d.ts +0 -1
- package/dist/primitives/server/McpServerStatus.d.ts.map +1 -1
- package/dist/primitives/server/McpServerToolName.d.ts +0 -1
- package/dist/primitives/server/McpServerToolName.d.ts.map +1 -1
- package/dist/primitives/server/McpServerTools.d.ts +0 -1
- package/dist/primitives/server/McpServerTools.d.ts.map +1 -1
- package/dist/primitives/server.d.ts +0 -2
- package/dist/resources/McpManagerResource.d.ts +6 -4
- package/dist/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.d.ts +0 -1
- package/dist/resources/McpServerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.js +10 -3
- package/dist/resources/McpServerResource.js.map +1 -1
- package/dist/resources/storage/McpCustomStorage.d.ts +0 -1
- package/dist/resources/storage/McpCustomStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpLocalStorage.d.ts +8 -3
- package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpLocalStorage.js +55 -3
- package/dist/resources/storage/McpLocalStorage.js.map +1 -1
- package/dist/resources/storage/McpMemoryStorage.d.ts +0 -1
- package/dist/resources/storage/McpMemoryStorage.d.ts.map +1 -1
- package/dist/resources/storage/types.d.ts +0 -1
- package/dist/resources/storage/types.d.ts.map +1 -1
- package/dist/utils/serverId.d.ts.map +1 -1
- package/package.json +8 -8
- package/src/hooks/useMcpOAuthCallback.test.ts +24 -0
- package/src/hooks/useMcpOAuthCallback.tsx +20 -5
- package/src/mcp-scope.ts +2 -0
- package/src/resources/McpServerResource.test.ts +152 -9
- package/src/resources/McpServerResource.ts +10 -1
- package/src/resources/storage/McpLocalStorage.test.ts +184 -0
- 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(
|
|
52
|
-
const serverId = decodeServerIdFromState(state);
|
|
67
|
+
if (!state) throw new Error('missing "state" parameter');
|
|
53
68
|
if (!serverId) {
|
|
54
|
-
throw new Error("
|
|
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
|
|
62
|
-
setResult(
|
|
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(
|
|
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
|
|
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
|
-
|
|
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:
|
|
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
|
|
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<
|
|
172
|
+
normalizeCustomServerRecords(read<unknown>(customServersKey, [])),
|
|
63
173
|
saveCustomServers: async (records) => {
|
|
64
174
|
write(customServersKey, records);
|
|
65
175
|
},
|
|
66
176
|
loadAuthState: async (id) =>
|
|
67
|
-
read<
|
|
177
|
+
normalizePersistedAuthState(read<unknown>(authKey(id), null)),
|
|
68
178
|
saveAuthState: async (id, state) => {
|
|
69
179
|
write(authKey(id), state);
|
|
70
180
|
},
|