@assistant-ui/react-mcp 0.1.11 → 0.1.13
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/dist/auth/createOAuthProvider.d.ts.map +1 -1
- package/dist/auth/createOAuthProvider.js +11 -0
- package/dist/auth/createOAuthProvider.js.map +1 -1
- package/dist/auth/types.d.ts +2 -1
- package/dist/auth/types.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.js +3 -16
- package/dist/hooks/useMcpOAuthCallback.js.map +1 -1
- package/dist/primitives/addForm/McpAddFormRoot.d.ts.map +1 -1
- package/dist/primitives/addForm/McpAddFormRoot.js +2 -1
- package/dist/primitives/addForm/McpAddFormRoot.js.map +1 -1
- package/dist/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpManagerResource.js +47 -46
- package/dist/resources/McpManagerResource.js.map +1 -1
- package/dist/resources/McpServerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.js +13 -3
- package/dist/resources/McpServerResource.js.map +1 -1
- package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpLocalStorage.js +22 -1
- package/dist/resources/storage/McpLocalStorage.js.map +1 -1
- package/dist/utils/invokeMcpCallback.d.ts +5 -0
- package/dist/utils/invokeMcpCallback.d.ts.map +1 -0
- package/dist/utils/invokeMcpCallback.js +19 -0
- package/dist/utils/invokeMcpCallback.js.map +1 -0
- package/package.json +6 -6
- package/src/auth/createOAuthProvider.test.ts +86 -0
- package/src/auth/createOAuthProvider.ts +15 -0
- package/src/auth/types.ts +2 -0
- package/src/hooks/useMcpOAuthCallback.tsx +3 -36
- package/src/primitives/addForm/McpAddFormRoot.test.tsx +80 -0
- package/src/primitives/addForm/McpAddFormRoot.tsx +2 -1
- package/src/resources/McpManagerResource.test.ts +21 -0
- package/src/resources/McpManagerResource.ts +12 -7
- package/src/resources/McpServerResource.test.ts +141 -1
- package/src/resources/McpServerResource.ts +14 -2
- package/src/resources/storage/McpLocalStorage.test.ts +180 -0
- package/src/resources/storage/McpLocalStorage.ts +53 -0
- package/src/utils/invokeMcpCallback.ts +27 -0
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
// @vitest-environment jsdom
|
|
2
|
+
|
|
3
|
+
import {
|
|
4
|
+
cleanup,
|
|
5
|
+
fireEvent,
|
|
6
|
+
render,
|
|
7
|
+
screen,
|
|
8
|
+
waitFor,
|
|
9
|
+
} from "@testing-library/react";
|
|
10
|
+
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
|
11
|
+
|
|
12
|
+
const mocks = vi.hoisted(() => ({
|
|
13
|
+
addCustomServer: vi.fn(),
|
|
14
|
+
}));
|
|
15
|
+
|
|
16
|
+
vi.mock("@assistant-ui/store", async (importOriginal) => ({
|
|
17
|
+
...(await importOriginal()),
|
|
18
|
+
useAui: () => ({
|
|
19
|
+
mcp: { addCustomServer: mocks.addCustomServer },
|
|
20
|
+
}),
|
|
21
|
+
}));
|
|
22
|
+
|
|
23
|
+
import { McpAddFormPrimitiveError } from "./McpAddFormError";
|
|
24
|
+
import { McpAddFormPrimitiveNameField } from "./McpAddFormNameField";
|
|
25
|
+
import { McpAddFormPrimitiveRoot } from "./McpAddFormRoot";
|
|
26
|
+
import { McpAddFormPrimitiveSubmit } from "./McpAddFormSubmit";
|
|
27
|
+
import { McpAddFormPrimitiveUrlField } from "./McpAddFormUrlField";
|
|
28
|
+
|
|
29
|
+
describe("McpAddFormPrimitiveRoot", () => {
|
|
30
|
+
beforeEach(() => {
|
|
31
|
+
mocks.addCustomServer.mockReset();
|
|
32
|
+
mocks.addCustomServer.mockResolvedValue("server-1");
|
|
33
|
+
});
|
|
34
|
+
|
|
35
|
+
afterEach(() => {
|
|
36
|
+
cleanup();
|
|
37
|
+
vi.restoreAllMocks();
|
|
38
|
+
});
|
|
39
|
+
|
|
40
|
+
it.each(["throws", "rejects"] as const)(
|
|
41
|
+
"does not turn a successful add into an error when onSubmitted %s",
|
|
42
|
+
async (mode) => {
|
|
43
|
+
const callbackError = new Error("navigation failed");
|
|
44
|
+
const consoleError = vi
|
|
45
|
+
.spyOn(console, "error")
|
|
46
|
+
.mockImplementation(() => undefined);
|
|
47
|
+
|
|
48
|
+
render(
|
|
49
|
+
<McpAddFormPrimitiveRoot
|
|
50
|
+
onSubmitted={() => {
|
|
51
|
+
if (mode === "throws") throw callbackError;
|
|
52
|
+
return Promise.reject(callbackError);
|
|
53
|
+
}}
|
|
54
|
+
>
|
|
55
|
+
<McpAddFormPrimitiveNameField aria-label="Name" />
|
|
56
|
+
<McpAddFormPrimitiveUrlField aria-label="URL" />
|
|
57
|
+
<McpAddFormPrimitiveError />
|
|
58
|
+
<McpAddFormPrimitiveSubmit>Submit</McpAddFormPrimitiveSubmit>
|
|
59
|
+
</McpAddFormPrimitiveRoot>,
|
|
60
|
+
);
|
|
61
|
+
|
|
62
|
+
fireEvent.change(screen.getByLabelText("Name"), {
|
|
63
|
+
target: { value: "Docs" },
|
|
64
|
+
});
|
|
65
|
+
fireEvent.change(screen.getByLabelText("URL"), {
|
|
66
|
+
target: { value: "https://example.com/mcp" },
|
|
67
|
+
});
|
|
68
|
+
fireEvent.click(screen.getByRole("button", { name: "Submit" }));
|
|
69
|
+
|
|
70
|
+
await waitFor(() => expect(mocks.addCustomServer).toHaveBeenCalledOnce());
|
|
71
|
+
expect(screen.queryByText(callbackError.message)).toBeNull();
|
|
72
|
+
await waitFor(() => {
|
|
73
|
+
expect(consoleError).toHaveBeenCalledWith(
|
|
74
|
+
"[react-mcp] onSubmitted callback threw an error",
|
|
75
|
+
callbackError,
|
|
76
|
+
);
|
|
77
|
+
});
|
|
78
|
+
},
|
|
79
|
+
);
|
|
80
|
+
});
|
|
@@ -11,6 +11,7 @@ import { Primitive } from "@radix-ui/react-primitive";
|
|
|
11
11
|
import { useAui } from "@assistant-ui/store";
|
|
12
12
|
import { AddFormContext, type AddFormState } from "./context";
|
|
13
13
|
import type { MCPAuthConfig } from "../../mcp-scope";
|
|
14
|
+
import { invokeMcpCallback } from "../../utils/invokeMcpCallback";
|
|
14
15
|
|
|
15
16
|
const INITIAL: AddFormState = {
|
|
16
17
|
name: "",
|
|
@@ -106,7 +107,7 @@ export const McpAddFormPrimitiveRoot = forwardRef<
|
|
|
106
107
|
auth: buildAuth(),
|
|
107
108
|
});
|
|
108
109
|
setState(INITIAL);
|
|
109
|
-
onSubmitted
|
|
110
|
+
invokeMcpCallback("onSubmitted", onSubmitted, id);
|
|
110
111
|
} catch (err) {
|
|
111
112
|
setState((p) => ({
|
|
112
113
|
...p,
|
|
@@ -36,6 +36,27 @@ vi.mock("@assistant-ui/store", async (importOriginal) => ({
|
|
|
36
36
|
useAssistantClientRef: () => ({ current: null }),
|
|
37
37
|
}));
|
|
38
38
|
|
|
39
|
+
vi.mock("@assistant-ui/store/client", async (importOriginal) => {
|
|
40
|
+
const actual =
|
|
41
|
+
await importOriginal<typeof import("@assistant-ui/store/client")>();
|
|
42
|
+
const { useEffect } = await import("react");
|
|
43
|
+
const useScopeEffectShim = (
|
|
44
|
+
_scope: string,
|
|
45
|
+
effect: () => (() => void) | void,
|
|
46
|
+
deps: readonly unknown[],
|
|
47
|
+
) => {
|
|
48
|
+
useEffect(() => {
|
|
49
|
+
const cleanup = effect();
|
|
50
|
+
return typeof cleanup === "function" ? cleanup : undefined;
|
|
51
|
+
// oxlint-disable-next-line react-hooks/exhaustive-deps -- caller-provided deps, mirrors the real hook
|
|
52
|
+
}, deps);
|
|
53
|
+
};
|
|
54
|
+
return {
|
|
55
|
+
...actual,
|
|
56
|
+
useAssistantScopeEffect: useScopeEffectShim,
|
|
57
|
+
};
|
|
58
|
+
});
|
|
59
|
+
|
|
39
60
|
const connector = (id: string, name = id): MCPConnector =>
|
|
40
61
|
defineConnector({
|
|
41
62
|
id,
|
|
@@ -6,6 +6,7 @@ import {
|
|
|
6
6
|
attachTransformScopes,
|
|
7
7
|
type ClientOutput,
|
|
8
8
|
} from "@assistant-ui/store";
|
|
9
|
+
import { useAssistantScopeEffect } from "@assistant-ui/store/client";
|
|
9
10
|
import { ModelContext } from "@assistant-ui/core/store";
|
|
10
11
|
import type { Tool } from "assistant-stream";
|
|
11
12
|
import { McpServerResource } from "./McpServerResource";
|
|
@@ -234,13 +235,17 @@ const useMcpManagerResource = (
|
|
|
234
235
|
|
|
235
236
|
const clientRef = useAssistantClientRef();
|
|
236
237
|
|
|
237
|
-
|
|
238
|
-
|
|
239
|
-
|
|
240
|
-
|
|
241
|
-
|
|
242
|
-
|
|
243
|
-
|
|
238
|
+
useAssistantScopeEffect(
|
|
239
|
+
"modelContext",
|
|
240
|
+
() => {
|
|
241
|
+
const client = clientRef.current;
|
|
242
|
+
if (!client) return;
|
|
243
|
+
return client.modelContext.register({
|
|
244
|
+
getModelContext: () => ({ tools: toolkit }),
|
|
245
|
+
});
|
|
246
|
+
},
|
|
247
|
+
[toolkit],
|
|
248
|
+
);
|
|
244
249
|
|
|
245
250
|
const serverByKind = (kind: "connector" | "custom", index: number) => {
|
|
246
251
|
const list = kind === "connector" ? state.connectors : state.customServers;
|
|
@@ -147,6 +147,8 @@ const resetMocks = () => {
|
|
|
147
147
|
const mount = (
|
|
148
148
|
props?: {
|
|
149
149
|
auth?: MCPAuthConfig | undefined;
|
|
150
|
+
storage?: MCPStorage | undefined;
|
|
151
|
+
autoConnect?: boolean | undefined;
|
|
150
152
|
connectionTimeout?: number | undefined;
|
|
151
153
|
cache?: { readonly defaultTtlMs?: number } | undefined;
|
|
152
154
|
elicitation?: boolean | undefined;
|
|
@@ -169,7 +171,7 @@ const mount = (
|
|
|
169
171
|
auth: props?.auth ?? { type: "none" },
|
|
170
172
|
storage: props?.storage ?? createStorage(),
|
|
171
173
|
redirectUri: "https://example.com/callback",
|
|
172
|
-
autoConnect: false,
|
|
174
|
+
autoConnect: props?.autoConnect ?? false,
|
|
173
175
|
connectionTimeout,
|
|
174
176
|
cache: props?.cache,
|
|
175
177
|
...(props?.elicitation !== undefined
|
|
@@ -185,6 +187,144 @@ const mount = (
|
|
|
185
187
|
});
|
|
186
188
|
};
|
|
187
189
|
|
|
190
|
+
describe("McpServerResource automatic authentication", () => {
|
|
191
|
+
beforeEach(resetMocks);
|
|
192
|
+
|
|
193
|
+
it("reports auth storage load failures", async () => {
|
|
194
|
+
const storage = createStorage();
|
|
195
|
+
vi.mocked(storage.loadAuthState).mockRejectedValue(
|
|
196
|
+
new Error("auth storage unavailable"),
|
|
197
|
+
);
|
|
198
|
+
const root = mount({
|
|
199
|
+
auth: { type: "oauth" },
|
|
200
|
+
storage,
|
|
201
|
+
autoConnect: true,
|
|
202
|
+
});
|
|
203
|
+
|
|
204
|
+
try {
|
|
205
|
+
await waitForResourceUpdate(
|
|
206
|
+
() => root.getValue().getState().connectionState === "error",
|
|
207
|
+
);
|
|
208
|
+
|
|
209
|
+
expect(storage.loadAuthState).toHaveBeenCalledWith("docs");
|
|
210
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
211
|
+
connectionState: "error",
|
|
212
|
+
tools: [],
|
|
213
|
+
lastError: {
|
|
214
|
+
message:
|
|
215
|
+
'MCP server "docs" failed to load saved authentication: auth storage unavailable',
|
|
216
|
+
},
|
|
217
|
+
});
|
|
218
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
219
|
+
} finally {
|
|
220
|
+
root.unmount();
|
|
221
|
+
}
|
|
222
|
+
});
|
|
223
|
+
|
|
224
|
+
it("ignores auth storage failures from the cancelled StrictMode setup", async () => {
|
|
225
|
+
let rejectCancelledLoad!: (error: Error) => void;
|
|
226
|
+
const storage = createStorage();
|
|
227
|
+
vi.mocked(storage.loadAuthState)
|
|
228
|
+
.mockImplementationOnce(
|
|
229
|
+
() =>
|
|
230
|
+
new Promise((_, reject) => {
|
|
231
|
+
rejectCancelledLoad = reject;
|
|
232
|
+
}),
|
|
233
|
+
)
|
|
234
|
+
.mockResolvedValueOnce(null);
|
|
235
|
+
const root = mount({
|
|
236
|
+
auth: { type: "oauth" },
|
|
237
|
+
storage,
|
|
238
|
+
autoConnect: true,
|
|
239
|
+
});
|
|
240
|
+
|
|
241
|
+
try {
|
|
242
|
+
await waitFor(
|
|
243
|
+
() => vi.mocked(storage.loadAuthState).mock.calls.length > 1,
|
|
244
|
+
);
|
|
245
|
+
rejectCancelledLoad(new Error("cancelled auth storage failure"));
|
|
246
|
+
await flushMacrotask();
|
|
247
|
+
|
|
248
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
249
|
+
connectionState: "disconnected",
|
|
250
|
+
lastError: null,
|
|
251
|
+
});
|
|
252
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
253
|
+
} finally {
|
|
254
|
+
root.unmount();
|
|
255
|
+
}
|
|
256
|
+
});
|
|
257
|
+
|
|
258
|
+
it("ignores auth storage failures superseded by a manual connection", async () => {
|
|
259
|
+
let rejectLoad!: (error: Error) => void;
|
|
260
|
+
const storage = createStorage();
|
|
261
|
+
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
262
|
+
() =>
|
|
263
|
+
new Promise((_, reject) => {
|
|
264
|
+
rejectLoad = reject;
|
|
265
|
+
}),
|
|
266
|
+
);
|
|
267
|
+
const root = mount({
|
|
268
|
+
auth: { type: "oauth" },
|
|
269
|
+
storage,
|
|
270
|
+
autoConnect: true,
|
|
271
|
+
});
|
|
272
|
+
|
|
273
|
+
try {
|
|
274
|
+
await waitFor(
|
|
275
|
+
() => vi.mocked(storage.loadAuthState).mock.calls.length > 0,
|
|
276
|
+
);
|
|
277
|
+
await root.getValue().connect();
|
|
278
|
+
await waitForResourceUpdate(
|
|
279
|
+
() => root.getValue().getState().connectionState === "connected",
|
|
280
|
+
);
|
|
281
|
+
|
|
282
|
+
rejectLoad(new Error("superseded auth storage failure"));
|
|
283
|
+
await flushMacrotask();
|
|
284
|
+
|
|
285
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
286
|
+
connectionState: "connected",
|
|
287
|
+
lastError: null,
|
|
288
|
+
});
|
|
289
|
+
} finally {
|
|
290
|
+
root.unmount();
|
|
291
|
+
}
|
|
292
|
+
});
|
|
293
|
+
|
|
294
|
+
it("ignores successful auth storage loads superseded by disconnect", async () => {
|
|
295
|
+
let resolveLoad!: (value: { token: string }) => void;
|
|
296
|
+
const storage = createStorage();
|
|
297
|
+
vi.mocked(storage.loadAuthState).mockImplementation(
|
|
298
|
+
() =>
|
|
299
|
+
new Promise((resolve) => {
|
|
300
|
+
resolveLoad = resolve;
|
|
301
|
+
}),
|
|
302
|
+
);
|
|
303
|
+
const root = mount({
|
|
304
|
+
auth: { type: "bearer" },
|
|
305
|
+
storage,
|
|
306
|
+
autoConnect: true,
|
|
307
|
+
});
|
|
308
|
+
|
|
309
|
+
try {
|
|
310
|
+
await waitFor(
|
|
311
|
+
() => vi.mocked(storage.loadAuthState).mock.calls.length > 0,
|
|
312
|
+
);
|
|
313
|
+
await root.getValue().disconnect();
|
|
314
|
+
resolveLoad({ token: "secret" });
|
|
315
|
+
await flushMacrotask();
|
|
316
|
+
|
|
317
|
+
expect(root.getValue().getState()).toMatchObject({
|
|
318
|
+
connectionState: "disconnected",
|
|
319
|
+
lastError: null,
|
|
320
|
+
});
|
|
321
|
+
expect(mocks.StreamableHTTPClientTransport).not.toHaveBeenCalled();
|
|
322
|
+
} finally {
|
|
323
|
+
root.unmount();
|
|
324
|
+
}
|
|
325
|
+
});
|
|
326
|
+
});
|
|
327
|
+
|
|
188
328
|
describe("McpServerResource connectionTimeout", () => {
|
|
189
329
|
beforeEach(() => {
|
|
190
330
|
vi.useFakeTimers();
|
|
@@ -474,8 +474,20 @@ const useMcpServerResource = (
|
|
|
474
474
|
void doConnect();
|
|
475
475
|
return;
|
|
476
476
|
}
|
|
477
|
-
const
|
|
478
|
-
|
|
477
|
+
const generation = connectionGenerationRef.current;
|
|
478
|
+
let persisted: Awaited<ReturnType<MCPStorage["loadAuthState"]>>;
|
|
479
|
+
try {
|
|
480
|
+
persisted = await props.storage.loadAuthState(props.id);
|
|
481
|
+
} catch (error) {
|
|
482
|
+
if (signal.cancelled || !isCurrentConnection(generation)) return;
|
|
483
|
+
const message = error instanceof Error ? error.message : String(error);
|
|
484
|
+
setLastError({
|
|
485
|
+
message: `MCP server "${props.id}" failed to load saved authentication: ${message}`,
|
|
486
|
+
});
|
|
487
|
+
setConnectionState("error");
|
|
488
|
+
return;
|
|
489
|
+
}
|
|
490
|
+
if (signal.cancelled || !isCurrentConnection(generation)) return;
|
|
479
491
|
if (props.auth.type === "oauth") {
|
|
480
492
|
if (!persisted?.tokens) return;
|
|
481
493
|
} else if (!persisted?.token) {
|
|
@@ -1,5 +1,7 @@
|
|
|
1
1
|
import { createTapRoot, useResource } from "@assistant-ui/tap";
|
|
2
|
+
import { auth, type FetchLike } from "@modelcontextprotocol/client";
|
|
2
3
|
import { describe, expect, it } from "vitest";
|
|
4
|
+
import { createOAuthProvider } from "../../auth/createOAuthProvider";
|
|
3
5
|
|
|
4
6
|
import {
|
|
5
7
|
McpLocalStorage,
|
|
@@ -155,6 +157,103 @@ describe("normalizePersistedAuthState", () => {
|
|
|
155
157
|
});
|
|
156
158
|
});
|
|
157
159
|
|
|
160
|
+
it("keeps valid OAuth discovery state", () => {
|
|
161
|
+
const discoveryState = {
|
|
162
|
+
authorizationServerUrl: "https://auth.example.com",
|
|
163
|
+
resourceMetadataUrl:
|
|
164
|
+
"https://mcp.example.com/.well-known/oauth-protected-resource",
|
|
165
|
+
authorizationServerMetadata: {
|
|
166
|
+
issuer: "https://auth.example.com",
|
|
167
|
+
authorization_endpoint: "https://auth.example.com/authorize",
|
|
168
|
+
token_endpoint: "https://auth.example.com/token",
|
|
169
|
+
response_types_supported: ["code"],
|
|
170
|
+
},
|
|
171
|
+
resourceMetadata: {
|
|
172
|
+
resource: "https://mcp.example.com",
|
|
173
|
+
authorization_servers: ["https://auth.example.com"],
|
|
174
|
+
},
|
|
175
|
+
};
|
|
176
|
+
|
|
177
|
+
expect(normalizePersistedAuthState({ discoveryState })).toEqual({
|
|
178
|
+
discoveryState,
|
|
179
|
+
});
|
|
180
|
+
});
|
|
181
|
+
|
|
182
|
+
it("drops malformed OAuth discovery state", () => {
|
|
183
|
+
expect(
|
|
184
|
+
normalizePersistedAuthState({
|
|
185
|
+
token: "bearer-token",
|
|
186
|
+
discoveryState: {
|
|
187
|
+
authorizationServerUrl: "not-a-url",
|
|
188
|
+
},
|
|
189
|
+
}),
|
|
190
|
+
).toEqual({ token: "bearer-token" });
|
|
191
|
+
});
|
|
192
|
+
|
|
193
|
+
it("keeps the URL binding when optional discovery fields are malformed", () => {
|
|
194
|
+
expect(
|
|
195
|
+
normalizePersistedAuthState({
|
|
196
|
+
token: "bearer-token",
|
|
197
|
+
discoveryState: {
|
|
198
|
+
authorizationServerUrl: "https://auth.example.com",
|
|
199
|
+
authorizationServerMetadata: { issuer: 123 },
|
|
200
|
+
resourceMetadata: { resource: 42 },
|
|
201
|
+
},
|
|
202
|
+
}),
|
|
203
|
+
).toEqual({
|
|
204
|
+
token: "bearer-token",
|
|
205
|
+
discoveryState: {
|
|
206
|
+
authorizationServerUrl: "https://auth.example.com",
|
|
207
|
+
},
|
|
208
|
+
});
|
|
209
|
+
});
|
|
210
|
+
|
|
211
|
+
it.each([
|
|
212
|
+
"http://auth.example.com",
|
|
213
|
+
"data:text/plain,auth",
|
|
214
|
+
"file:///tmp/auth",
|
|
215
|
+
])("drops discovery state with an unsafe URL: %s", (url) => {
|
|
216
|
+
expect(
|
|
217
|
+
normalizePersistedAuthState({
|
|
218
|
+
token: "bearer-token",
|
|
219
|
+
discoveryState: { authorizationServerUrl: url },
|
|
220
|
+
}),
|
|
221
|
+
).toEqual({ token: "bearer-token" });
|
|
222
|
+
});
|
|
223
|
+
|
|
224
|
+
it.each([
|
|
225
|
+
"http://localhost:3000",
|
|
226
|
+
"http://foo.localhost:3000",
|
|
227
|
+
"http://127.0.0.1:3000",
|
|
228
|
+
"http://127.0.0.2:3000",
|
|
229
|
+
"http://[::1]:3000",
|
|
230
|
+
])("keeps loopback HTTP discovery URLs: %s", (url) => {
|
|
231
|
+
expect(
|
|
232
|
+
normalizePersistedAuthState({
|
|
233
|
+
discoveryState: { authorizationServerUrl: url },
|
|
234
|
+
}),
|
|
235
|
+
).toEqual({
|
|
236
|
+
discoveryState: { authorizationServerUrl: url },
|
|
237
|
+
});
|
|
238
|
+
});
|
|
239
|
+
|
|
240
|
+
it("drops an unsafe resource metadata URL but keeps the binding", () => {
|
|
241
|
+
expect(
|
|
242
|
+
normalizePersistedAuthState({
|
|
243
|
+
token: "bearer-token",
|
|
244
|
+
discoveryState: {
|
|
245
|
+
authorizationServerUrl: "https://auth.example.com",
|
|
246
|
+
resourceMetadataUrl: "http://mcp.example.com/oauth-resource",
|
|
247
|
+
},
|
|
248
|
+
}),
|
|
249
|
+
).toEqual({
|
|
250
|
+
token: "bearer-token",
|
|
251
|
+
discoveryState: {
|
|
252
|
+
authorizationServerUrl: "https://auth.example.com",
|
|
253
|
+
},
|
|
254
|
+
});
|
|
255
|
+
});
|
|
256
|
+
|
|
158
257
|
it("keeps valid fields when neighboring fields are malformed", () => {
|
|
159
258
|
expect(
|
|
160
259
|
normalizePersistedAuthState({
|
|
@@ -237,4 +336,85 @@ describe("McpLocalStorage auth state", () => {
|
|
|
237
336
|
token: "bearer-token",
|
|
238
337
|
});
|
|
239
338
|
});
|
|
339
|
+
|
|
340
|
+
it("completes an SDK callback with discovery state from localStorage", async () => {
|
|
341
|
+
const storage = createStorage();
|
|
342
|
+
const requests: Array<{ method: string; url: string }> = [];
|
|
343
|
+
const fetchFn: FetchLike = async (input, init) => {
|
|
344
|
+
const url = new URL(
|
|
345
|
+
input instanceof Request ? input.url : input.toString(),
|
|
346
|
+
);
|
|
347
|
+
requests.push({ method: init?.method ?? "GET", url: url.toString() });
|
|
348
|
+
|
|
349
|
+
if (url.pathname === "/.well-known/oauth-protected-resource") {
|
|
350
|
+
return Response.json({
|
|
351
|
+
resource: "https://mcp.example.com/mcp",
|
|
352
|
+
authorization_servers: ["https://auth.example.com"],
|
|
353
|
+
});
|
|
354
|
+
}
|
|
355
|
+
if (url.pathname === "/.well-known/oauth-authorization-server") {
|
|
356
|
+
return Response.json({
|
|
357
|
+
issuer: "https://auth.example.com",
|
|
358
|
+
authorization_endpoint: "https://auth.example.com/authorize",
|
|
359
|
+
token_endpoint: "https://auth.example.com/token",
|
|
360
|
+
response_types_supported: ["code"],
|
|
361
|
+
code_challenge_methods_supported: ["S256"],
|
|
362
|
+
});
|
|
363
|
+
}
|
|
364
|
+
if (url.pathname === "/token") {
|
|
365
|
+
return Response.json({
|
|
366
|
+
access_token: "access-token",
|
|
367
|
+
token_type: "Bearer",
|
|
368
|
+
});
|
|
369
|
+
}
|
|
370
|
+
throw new Error(`Unexpected OAuth request: ${url}`);
|
|
371
|
+
};
|
|
372
|
+
const authorizationUrls: URL[] = [];
|
|
373
|
+
const createProvider = () =>
|
|
374
|
+
createOAuthProvider({
|
|
375
|
+
serverId: "docs",
|
|
376
|
+
config: { type: "oauth", clientId: "client-id" },
|
|
377
|
+
storage: loadStorage(storage),
|
|
378
|
+
redirectUri: "http://localhost/callback",
|
|
379
|
+
onAuthorizationUrl: (url) => authorizationUrls.push(url),
|
|
380
|
+
});
|
|
381
|
+
const redirectProvider = createProvider();
|
|
382
|
+
|
|
383
|
+
await expect(
|
|
384
|
+
auth(redirectProvider, {
|
|
385
|
+
serverUrl: "https://mcp.example.com/mcp",
|
|
386
|
+
resourceMetadataUrl: new URL(
|
|
387
|
+
"https://mcp.example.com/.well-known/oauth-protected-resource",
|
|
388
|
+
),
|
|
389
|
+
fetchFn,
|
|
390
|
+
}),
|
|
391
|
+
).resolves.toBe("REDIRECT");
|
|
392
|
+
expect(authorizationUrls).toHaveLength(1);
|
|
393
|
+
expect(
|
|
394
|
+
JSON.parse(storage.getItem("test-mcp:auth:docs") ?? "null"),
|
|
395
|
+
).toMatchObject({
|
|
396
|
+
codeVerifier: expect.any(String),
|
|
397
|
+
discoveryState: {
|
|
398
|
+
authorizationServerUrl: "https://auth.example.com",
|
|
399
|
+
},
|
|
400
|
+
});
|
|
401
|
+
|
|
402
|
+
requests.length = 0;
|
|
403
|
+
const callbackProvider = createProvider();
|
|
404
|
+
await expect(
|
|
405
|
+
auth(callbackProvider, {
|
|
406
|
+
serverUrl: "https://mcp.example.com/mcp",
|
|
407
|
+
authorizationCode: "authorization-code",
|
|
408
|
+
fetchFn,
|
|
409
|
+
}),
|
|
410
|
+
).resolves.toBe("AUTHORIZED");
|
|
411
|
+
expect(requests).toEqual([
|
|
412
|
+
{ method: "POST", url: "https://auth.example.com/token" },
|
|
413
|
+
]);
|
|
414
|
+
await expect(
|
|
415
|
+
loadStorage(storage).loadAuthState("docs"),
|
|
416
|
+
).resolves.toMatchObject({
|
|
417
|
+
tokens: { access_token: "access-token" },
|
|
418
|
+
});
|
|
419
|
+
});
|
|
240
420
|
});
|
|
@@ -1,6 +1,8 @@
|
|
|
1
1
|
import { resource } from "@assistant-ui/tap";
|
|
2
2
|
import {
|
|
3
|
+
OAuthMetadataSchema,
|
|
3
4
|
OAuthClientInformationFullSchema,
|
|
5
|
+
OAuthProtectedResourceMetadataSchema,
|
|
4
6
|
OAuthTokensSchema,
|
|
5
7
|
} from "@modelcontextprotocol/core";
|
|
6
8
|
import type { MCPAuthConfig, MCPCustomServerRecord } from "../../mcp-scope";
|
|
@@ -133,6 +135,54 @@ const normalizeClientInformation = (
|
|
|
133
135
|
return result.success ? result.data : undefined;
|
|
134
136
|
};
|
|
135
137
|
|
|
138
|
+
const isSecureNetworkUrl = (value: unknown): value is string => {
|
|
139
|
+
if (!isNonEmptyString(value)) return false;
|
|
140
|
+
try {
|
|
141
|
+
const url = new URL(value);
|
|
142
|
+
return (
|
|
143
|
+
url.protocol === "https:" ||
|
|
144
|
+
(url.protocol === "http:" &&
|
|
145
|
+
(url.hostname === "localhost" ||
|
|
146
|
+
url.hostname.endsWith(".localhost") ||
|
|
147
|
+
url.hostname.startsWith("127.") ||
|
|
148
|
+
url.hostname === "[::1]"))
|
|
149
|
+
);
|
|
150
|
+
} catch {
|
|
151
|
+
return false;
|
|
152
|
+
}
|
|
153
|
+
};
|
|
154
|
+
|
|
155
|
+
const normalizeDiscoveryState = (
|
|
156
|
+
value: unknown,
|
|
157
|
+
): MCPPersistedAuthState["discoveryState"] | undefined => {
|
|
158
|
+
if (!isRecord(value) || !isSecureNetworkUrl(value.authorizationServerUrl)) {
|
|
159
|
+
return undefined;
|
|
160
|
+
}
|
|
161
|
+
|
|
162
|
+
// A malformed optional field is dropped alone: keeping the validated
|
|
163
|
+
// authorization server URL preserves the redirect-time binding, and the SDK
|
|
164
|
+
// re-discovers whatever metadata is missing.
|
|
165
|
+
const state: NonNullable<MCPPersistedAuthState["discoveryState"]> = {
|
|
166
|
+
authorizationServerUrl: value.authorizationServerUrl,
|
|
167
|
+
};
|
|
168
|
+
|
|
169
|
+
if (isSecureNetworkUrl(value.resourceMetadataUrl)) {
|
|
170
|
+
state.resourceMetadataUrl = value.resourceMetadataUrl;
|
|
171
|
+
}
|
|
172
|
+
|
|
173
|
+
const metadata = OAuthMetadataSchema.safeParse(
|
|
174
|
+
value.authorizationServerMetadata,
|
|
175
|
+
);
|
|
176
|
+
if (metadata.success) state.authorizationServerMetadata = metadata.data;
|
|
177
|
+
|
|
178
|
+
const resourceMetadata = OAuthProtectedResourceMetadataSchema.safeParse(
|
|
179
|
+
value.resourceMetadata,
|
|
180
|
+
);
|
|
181
|
+
if (resourceMetadata.success) state.resourceMetadata = resourceMetadata.data;
|
|
182
|
+
|
|
183
|
+
return state;
|
|
184
|
+
};
|
|
185
|
+
|
|
136
186
|
export const normalizePersistedAuthState = (
|
|
137
187
|
value: unknown,
|
|
138
188
|
): MCPPersistedAuthState | null => {
|
|
@@ -150,6 +200,9 @@ export const normalizePersistedAuthState = (
|
|
|
150
200
|
const clientInformation = normalizeClientInformation(value.clientInformation);
|
|
151
201
|
if (clientInformation) state.clientInformation = clientInformation;
|
|
152
202
|
|
|
203
|
+
const discoveryState = normalizeDiscoveryState(value.discoveryState);
|
|
204
|
+
if (discoveryState) state.discoveryState = discoveryState;
|
|
205
|
+
|
|
153
206
|
return Object.keys(state).length > 0 ? state : null;
|
|
154
207
|
};
|
|
155
208
|
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
const reportCallbackError = (name: string, error: unknown) => {
|
|
2
|
+
console.error(`[react-mcp] ${name} callback threw an error`, error);
|
|
3
|
+
};
|
|
4
|
+
|
|
5
|
+
export const invokeMcpCallback = <TArgs extends unknown[]>(
|
|
6
|
+
name: string,
|
|
7
|
+
callback: ((...args: TArgs) => void) | undefined,
|
|
8
|
+
...args: TArgs
|
|
9
|
+
) => {
|
|
10
|
+
if (!callback) return;
|
|
11
|
+
|
|
12
|
+
try {
|
|
13
|
+
const result = callback(...args) as unknown;
|
|
14
|
+
if (
|
|
15
|
+
result !== null &&
|
|
16
|
+
(typeof result === "object" || typeof result === "function") &&
|
|
17
|
+
"then" in result &&
|
|
18
|
+
typeof result.then === "function"
|
|
19
|
+
) {
|
|
20
|
+
void Promise.resolve(result).catch((error) => {
|
|
21
|
+
reportCallbackError(name, error);
|
|
22
|
+
});
|
|
23
|
+
}
|
|
24
|
+
} catch (error) {
|
|
25
|
+
reportCallbackError(name, error);
|
|
26
|
+
}
|
|
27
|
+
};
|