@assistant-ui/react-mcp 0.1.15 → 0.1.16
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/dist/auth/createOAuthProvider.d.ts +7 -1
- package/dist/auth/createOAuthProvider.d.ts.map +1 -1
- package/dist/auth/createOAuthProvider.js +115 -31
- package/dist/auth/createOAuthProvider.js.map +1 -1
- package/dist/auth/types.d.ts +1 -0
- package/dist/auth/types.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.d.ts.map +1 -1
- package/dist/hooks/useMcpOAuthCallback.js +5 -6
- package/dist/hooks/useMcpOAuthCallback.js.map +1 -1
- package/dist/resources/McpManagerResource.d.ts.map +1 -1
- package/dist/resources/McpManagerResource.js +2 -1
- package/dist/resources/McpManagerResource.js.map +1 -1
- package/dist/resources/McpServerResource.d.ts +2 -1
- package/dist/resources/McpServerResource.d.ts.map +1 -1
- package/dist/resources/McpServerResource.js +50 -11
- package/dist/resources/McpServerResource.js.map +1 -1
- package/dist/resources/storage/McpLocalStorage.d.ts +7 -0
- package/dist/resources/storage/McpLocalStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpLocalStorage.js +132 -37
- package/dist/resources/storage/McpLocalStorage.js.map +1 -1
- package/dist/resources/storage/McpMemoryStorage.d.ts.map +1 -1
- package/dist/resources/storage/McpMemoryStorage.js +9 -6
- package/dist/resources/storage/McpMemoryStorage.js.map +1 -1
- package/dist/resources/storage/types.d.ts +12 -0
- package/dist/resources/storage/types.d.ts.map +1 -1
- package/package.json +7 -7
- package/src/auth/createOAuthProvider.test.ts +407 -2
- package/src/auth/createOAuthProvider.ts +171 -40
- package/src/auth/types.ts +1 -0
- package/src/hooks/useMcpOAuthCallback.test.ts +74 -1
- package/src/hooks/useMcpOAuthCallback.tsx +11 -8
- package/src/resources/McpManagerResource.ts +2 -1
- package/src/resources/McpServerResource.test.ts +377 -16
- package/src/resources/McpServerResource.ts +61 -10
- package/src/resources/storage/McpLocalStorage.test.ts +71 -1
- package/src/resources/storage/McpLocalStorage.ts +69 -47
- package/src/resources/storage/McpMemoryStorage.test.ts +99 -0
- package/src/resources/storage/McpMemoryStorage.ts +23 -17
- package/src/resources/storage/types.ts +12 -0
|
@@ -57,6 +57,108 @@ export type CreateOAuthProviderOptions = {
|
|
|
57
57
|
onAuthorizationUrl: (url: URL) => void;
|
|
58
58
|
};
|
|
59
59
|
|
|
60
|
+
type OAuthProviderCache = {
|
|
61
|
+
tokens?: OAuthTokens | undefined;
|
|
62
|
+
clientInformation?: OAuthClientInformationFull | undefined;
|
|
63
|
+
codeVerifier?: string | undefined;
|
|
64
|
+
state?: string | undefined;
|
|
65
|
+
discoveryState?: OAuthDiscoveryState | undefined;
|
|
66
|
+
};
|
|
67
|
+
|
|
68
|
+
type OAuthProviderPersistence = {
|
|
69
|
+
cached: OAuthProviderCache | null;
|
|
70
|
+
cachePromise: Promise<OAuthProviderCache> | null;
|
|
71
|
+
queue: Promise<void>;
|
|
72
|
+
invalidated: boolean;
|
|
73
|
+
};
|
|
74
|
+
|
|
75
|
+
// scopeId, not object identity, is what addresses the same persisted data, so
|
|
76
|
+
// storages sharing one share an anchor and an unscoped storage is its own
|
|
77
|
+
// identity. Every storage declaring a scope holds that scope's anchor, so the
|
|
78
|
+
// coordination state below is collected once the last of them is gone.
|
|
79
|
+
const anchorByStorage = new WeakMap<MCPStorage, object>();
|
|
80
|
+
const anchorByScope = new Map<string, WeakRef<object>>();
|
|
81
|
+
const anchorRegistry = new FinalizationRegistry<string>((scopeId) => {
|
|
82
|
+
if (!anchorByScope.get(scopeId)?.deref()) anchorByScope.delete(scopeId);
|
|
83
|
+
});
|
|
84
|
+
|
|
85
|
+
const getStorageIdentity = (storage: MCPStorage): object => {
|
|
86
|
+
const existing = anchorByStorage.get(storage);
|
|
87
|
+
if (existing) return existing;
|
|
88
|
+
|
|
89
|
+
const { scopeId } = storage;
|
|
90
|
+
if (scopeId === undefined) return storage;
|
|
91
|
+
|
|
92
|
+
let anchor = anchorByScope.get(scopeId)?.deref();
|
|
93
|
+
if (!anchor) {
|
|
94
|
+
anchor = {};
|
|
95
|
+
anchorByScope.set(scopeId, new WeakRef(anchor));
|
|
96
|
+
anchorRegistry.register(anchor, scopeId);
|
|
97
|
+
}
|
|
98
|
+
anchorByStorage.set(storage, anchor);
|
|
99
|
+
return anchor;
|
|
100
|
+
};
|
|
101
|
+
|
|
102
|
+
// McpServerResource builds a fresh provider for every transport, so the cache,
|
|
103
|
+
// the in-flight load, and the write queue have to outlive any one provider.
|
|
104
|
+
// saveAuthState replaces the whole record, so two providers writing their own
|
|
105
|
+
// snapshots concurrently would drop whichever field the loser had added.
|
|
106
|
+
const persistenceByIdentity = new WeakMap<
|
|
107
|
+
object,
|
|
108
|
+
Map<string, OAuthProviderPersistence>
|
|
109
|
+
>();
|
|
110
|
+
|
|
111
|
+
const getPersistence = (
|
|
112
|
+
storage: MCPStorage,
|
|
113
|
+
serverId: string,
|
|
114
|
+
): OAuthProviderPersistence => {
|
|
115
|
+
const identity = getStorageIdentity(storage);
|
|
116
|
+
let byServerId = persistenceByIdentity.get(identity);
|
|
117
|
+
if (!byServerId) {
|
|
118
|
+
byServerId = new Map();
|
|
119
|
+
persistenceByIdentity.set(identity, byServerId);
|
|
120
|
+
}
|
|
121
|
+
|
|
122
|
+
let persistence = byServerId.get(serverId);
|
|
123
|
+
if (!persistence) {
|
|
124
|
+
persistence = {
|
|
125
|
+
cached: null,
|
|
126
|
+
cachePromise: null,
|
|
127
|
+
queue: Promise.resolve(),
|
|
128
|
+
invalidated: false,
|
|
129
|
+
};
|
|
130
|
+
byServerId.set(serverId, persistence);
|
|
131
|
+
}
|
|
132
|
+
return persistence;
|
|
133
|
+
};
|
|
134
|
+
|
|
135
|
+
/**
|
|
136
|
+
* Clears persisted OAuth state after the in-flight load and every queued write
|
|
137
|
+
* for that server have settled, so a discarded provider cannot recreate the
|
|
138
|
+
* record it was mid-save on.
|
|
139
|
+
*/
|
|
140
|
+
export const clearOAuthProviderAuthState = async (
|
|
141
|
+
storage: MCPStorage,
|
|
142
|
+
serverId: string,
|
|
143
|
+
): Promise<void> => {
|
|
144
|
+
const identity = getStorageIdentity(storage);
|
|
145
|
+
const byServerId = persistenceByIdentity.get(identity);
|
|
146
|
+
const persistence = byServerId?.get(serverId);
|
|
147
|
+
if (!byServerId || !persistence) {
|
|
148
|
+
await storage.clearAuthState(serverId);
|
|
149
|
+
return;
|
|
150
|
+
}
|
|
151
|
+
|
|
152
|
+
// Detaching the entry before awaiting keeps a provider built during the clear
|
|
153
|
+
// on a fresh generation instead of inheriting the fenced one.
|
|
154
|
+
persistence.invalidated = true;
|
|
155
|
+
byServerId.delete(serverId);
|
|
156
|
+
if (byServerId.size === 0) persistenceByIdentity.delete(identity);
|
|
157
|
+
|
|
158
|
+
await Promise.allSettled([persistence.cachePromise, persistence.queue]);
|
|
159
|
+
await storage.clearAuthState(serverId);
|
|
160
|
+
};
|
|
161
|
+
|
|
60
162
|
/**
|
|
61
163
|
* Builds an OAuthClientProvider for the MCP SDK, backed by MCPStorage.
|
|
62
164
|
* Token refresh and DCR are handled by the SDK; this provider only mediates
|
|
@@ -66,46 +168,66 @@ export function createOAuthProvider(
|
|
|
66
168
|
opts: CreateOAuthProviderOptions,
|
|
67
169
|
): OAuthClientProvider {
|
|
68
170
|
const { serverId, config, storage, redirectUri, onAuthorizationUrl } = opts;
|
|
171
|
+
const persistence = getPersistence(storage, serverId);
|
|
172
|
+
let pendingState: string | undefined;
|
|
69
173
|
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
if (
|
|
83
|
-
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
|
|
95
|
-
|
|
96
|
-
|
|
97
|
-
|
|
174
|
+
// The cache is shared with every other provider for this (storage, serverId),
|
|
175
|
+
// so a statically configured client stays a read-time overlay owned by this
|
|
176
|
+
// provider. Writing it into the cache would leak this provider's registration
|
|
177
|
+
// to a replacement built for a different, or absent, clientId.
|
|
178
|
+
const staticClientInformation = (():
|
|
179
|
+
| OAuthClientInformationFull
|
|
180
|
+
| undefined => {
|
|
181
|
+
if (!config.clientId) return undefined;
|
|
182
|
+
const ci: OAuthClientInformationFull = {
|
|
183
|
+
client_id: config.clientId,
|
|
184
|
+
redirect_uris: [redirectUri],
|
|
185
|
+
};
|
|
186
|
+
if (config.clientSecret) ci.client_secret = config.clientSecret;
|
|
187
|
+
return ci;
|
|
188
|
+
})();
|
|
189
|
+
|
|
190
|
+
const loadCache = (): Promise<OAuthProviderCache> => {
|
|
191
|
+
if (persistence.cached) return Promise.resolve(persistence.cached);
|
|
192
|
+
if (persistence.cachePromise) return persistence.cachePromise;
|
|
193
|
+
|
|
194
|
+
persistence.cachePromise = storage.loadAuthState(serverId).then(
|
|
195
|
+
(persisted) => {
|
|
196
|
+
const initial: OAuthProviderCache = {};
|
|
197
|
+
if (persisted?.tokens) initial.tokens = persisted.tokens;
|
|
198
|
+
if (persisted?.clientInformation)
|
|
199
|
+
initial.clientInformation = persisted.clientInformation;
|
|
200
|
+
if (persisted?.codeVerifier)
|
|
201
|
+
initial.codeVerifier = persisted.codeVerifier;
|
|
202
|
+
if (persisted?.state) initial.state = persisted.state;
|
|
203
|
+
if (persisted?.discoveryState)
|
|
204
|
+
initial.discoveryState = persisted.discoveryState;
|
|
205
|
+
persistence.cached = initial;
|
|
206
|
+
return initial;
|
|
207
|
+
},
|
|
208
|
+
(error) => {
|
|
209
|
+
persistence.cachePromise = null;
|
|
210
|
+
throw error;
|
|
211
|
+
},
|
|
212
|
+
);
|
|
213
|
+
return persistence.cachePromise;
|
|
98
214
|
};
|
|
99
215
|
|
|
100
|
-
const persist =
|
|
101
|
-
const
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
216
|
+
const persist = () => {
|
|
217
|
+
const task = persistence.queue.then(async () => {
|
|
218
|
+
if (persistence.invalidated) return;
|
|
219
|
+
const c = persistence.cached;
|
|
220
|
+
if (!c) return;
|
|
221
|
+
const next: Parameters<typeof storage.saveAuthState>[1] = {};
|
|
222
|
+
if (c.tokens) next.tokens = c.tokens;
|
|
223
|
+
if (c.clientInformation) next.clientInformation = c.clientInformation;
|
|
224
|
+
if (c.codeVerifier) next.codeVerifier = c.codeVerifier;
|
|
225
|
+
if (c.state) next.state = c.state;
|
|
226
|
+
if (c.discoveryState) next.discoveryState = c.discoveryState;
|
|
227
|
+
await storage.saveAuthState(serverId, next);
|
|
228
|
+
});
|
|
229
|
+
persistence.queue = task.catch(() => {});
|
|
230
|
+
return task;
|
|
109
231
|
};
|
|
110
232
|
|
|
111
233
|
const clientMetadata: OAuthClientMetadata = {
|
|
@@ -131,11 +253,12 @@ export function createOAuthProvider(
|
|
|
131
253
|
typeof crypto !== "undefined" && "randomUUID" in crypto
|
|
132
254
|
? crypto.randomUUID()
|
|
133
255
|
: `${Date.now()}.${Math.random()}`;
|
|
134
|
-
|
|
256
|
+
pendingState = `${encodeServerIdInState(serverId)}.${nonce}`;
|
|
257
|
+
return pendingState;
|
|
135
258
|
},
|
|
136
259
|
async clientInformation() {
|
|
137
260
|
const c = await loadCache();
|
|
138
|
-
return c.clientInformation;
|
|
261
|
+
return staticClientInformation ?? c.clientInformation;
|
|
139
262
|
},
|
|
140
263
|
async saveClientInformation(info) {
|
|
141
264
|
const c = await loadCache();
|
|
@@ -149,6 +272,7 @@ export function createOAuthProvider(
|
|
|
149
272
|
async saveTokens(tokens) {
|
|
150
273
|
const c = await loadCache();
|
|
151
274
|
c.tokens = tokens;
|
|
275
|
+
delete c.state;
|
|
152
276
|
await persist();
|
|
153
277
|
},
|
|
154
278
|
async redirectToAuthorization(url) {
|
|
@@ -157,6 +281,10 @@ export function createOAuthProvider(
|
|
|
157
281
|
async saveCodeVerifier(codeVerifier) {
|
|
158
282
|
const c = await loadCache();
|
|
159
283
|
c.codeVerifier = codeVerifier;
|
|
284
|
+
if (pendingState) {
|
|
285
|
+
c.state = pendingState;
|
|
286
|
+
pendingState = undefined;
|
|
287
|
+
}
|
|
160
288
|
await persist();
|
|
161
289
|
},
|
|
162
290
|
async codeVerifier() {
|
|
@@ -179,7 +307,10 @@ export function createOAuthProvider(
|
|
|
179
307
|
const c = await loadCache();
|
|
180
308
|
if (scope === "all" || scope === "tokens") delete c.tokens;
|
|
181
309
|
if (scope === "all" || scope === "client") delete c.clientInformation;
|
|
182
|
-
if (scope === "all" || scope === "verifier")
|
|
310
|
+
if (scope === "all" || scope === "verifier") {
|
|
311
|
+
delete c.codeVerifier;
|
|
312
|
+
delete c.state;
|
|
313
|
+
}
|
|
183
314
|
if (scope === "all" || scope === "discovery") delete c.discoveryState;
|
|
184
315
|
await persist();
|
|
185
316
|
},
|
package/src/auth/types.ts
CHANGED
|
@@ -8,6 +8,7 @@ export type MCPPersistedAuthState = {
|
|
|
8
8
|
tokens?: OAuthTokens;
|
|
9
9
|
clientInformation?: OAuthClientInformationFull;
|
|
10
10
|
codeVerifier?: string;
|
|
11
|
+
state?: string;
|
|
11
12
|
discoveryState?: OAuthDiscoveryState;
|
|
12
13
|
/** Bearer token (entered at add-form time). */
|
|
13
14
|
token?: string;
|
|
@@ -1,6 +1,12 @@
|
|
|
1
1
|
// @vitest-environment jsdom
|
|
2
2
|
|
|
3
|
-
import {
|
|
3
|
+
import {
|
|
4
|
+
createElement,
|
|
5
|
+
Suspense,
|
|
6
|
+
startTransition,
|
|
7
|
+
type PropsWithChildren,
|
|
8
|
+
} from "react";
|
|
9
|
+
import { act, renderHook, waitFor } from "@testing-library/react";
|
|
4
10
|
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
|
5
11
|
|
|
6
12
|
const mocks = vi.hoisted(() => {
|
|
@@ -98,6 +104,73 @@ describe("useMcpOAuthCallback", () => {
|
|
|
98
104
|
});
|
|
99
105
|
},
|
|
100
106
|
);
|
|
107
|
+
|
|
108
|
+
it("lets the server validate error callback parameters", async () => {
|
|
109
|
+
const authError = new Error("authorization response issuer mismatch");
|
|
110
|
+
const url = `${callbackUrl}&error=access_denied&error_description=untrusted`;
|
|
111
|
+
mocks.completeAuth.mockRejectedValueOnce(authError);
|
|
112
|
+
|
|
113
|
+
const { result } = renderHook(() => useMcpOAuthCallback({ url }));
|
|
114
|
+
|
|
115
|
+
await waitFor(() => expect(result.current.status).toBe("error"));
|
|
116
|
+
expect(mocks.completeAuth).toHaveBeenCalledWith(url);
|
|
117
|
+
expect(result.current.error).toMatchObject({
|
|
118
|
+
message:
|
|
119
|
+
'MCP OAuth callback for server "docs" failed: authorization response issuer mismatch',
|
|
120
|
+
cause: authError,
|
|
121
|
+
});
|
|
122
|
+
});
|
|
123
|
+
|
|
124
|
+
it("keeps callbacks scoped to committed renders", async () => {
|
|
125
|
+
let resolveAuth!: () => void;
|
|
126
|
+
mocks.completeAuth.mockReturnValueOnce(
|
|
127
|
+
new Promise<void>((resolve) => {
|
|
128
|
+
resolveAuth = resolve;
|
|
129
|
+
}),
|
|
130
|
+
);
|
|
131
|
+
const onCompleteA = vi.fn();
|
|
132
|
+
const onCompleteB = vi.fn();
|
|
133
|
+
const interruptedRender = vi.fn();
|
|
134
|
+
const pending = new Promise<never>(() => {});
|
|
135
|
+
let blocked = false;
|
|
136
|
+
const Blocker = () => {
|
|
137
|
+
if (blocked) {
|
|
138
|
+
interruptedRender();
|
|
139
|
+
throw pending;
|
|
140
|
+
}
|
|
141
|
+
return null;
|
|
142
|
+
};
|
|
143
|
+
const Wrapper = ({ children }: PropsWithChildren) =>
|
|
144
|
+
createElement(
|
|
145
|
+
Suspense,
|
|
146
|
+
{ fallback: null },
|
|
147
|
+
children,
|
|
148
|
+
createElement(Blocker),
|
|
149
|
+
);
|
|
150
|
+
|
|
151
|
+
const { rerender } = renderHook(
|
|
152
|
+
({ onComplete }) => useMcpOAuthCallback({ url: callbackUrl, onComplete }),
|
|
153
|
+
{
|
|
154
|
+
initialProps: { onComplete: onCompleteA },
|
|
155
|
+
wrapper: Wrapper,
|
|
156
|
+
},
|
|
157
|
+
);
|
|
158
|
+
await waitFor(() => expect(mocks.completeAuth).toHaveBeenCalledOnce());
|
|
159
|
+
|
|
160
|
+
act(() => {
|
|
161
|
+
blocked = true;
|
|
162
|
+
startTransition(() => rerender({ onComplete: onCompleteB }));
|
|
163
|
+
});
|
|
164
|
+
expect(interruptedRender).toHaveBeenCalled();
|
|
165
|
+
|
|
166
|
+
await act(async () => {
|
|
167
|
+
resolveAuth();
|
|
168
|
+
await Promise.resolve();
|
|
169
|
+
});
|
|
170
|
+
|
|
171
|
+
expect(onCompleteB).not.toHaveBeenCalled();
|
|
172
|
+
expect(onCompleteA).toHaveBeenCalledWith("docs");
|
|
173
|
+
});
|
|
101
174
|
});
|
|
102
175
|
|
|
103
176
|
describe("createMcpOAuthCallbackError", () => {
|
|
@@ -1,4 +1,11 @@
|
|
|
1
|
-
import {
|
|
1
|
+
import {
|
|
2
|
+
type FC,
|
|
3
|
+
type ReactNode,
|
|
4
|
+
useEffect,
|
|
5
|
+
useInsertionEffect,
|
|
6
|
+
useRef,
|
|
7
|
+
useState,
|
|
8
|
+
} from "react";
|
|
2
9
|
import { useAui } from "@assistant-ui/store";
|
|
3
10
|
import { decodeServerIdFromState } from "../auth/createOAuthProvider";
|
|
4
11
|
import { invokeMcpCallback } from "../utils/invokeMcpCallback";
|
|
@@ -44,7 +51,9 @@ export function useMcpOAuthCallback(
|
|
|
44
51
|
// single-use OAuth code is double-redeemed and the second attempt 4xxs.
|
|
45
52
|
const startedRef = useRef<string | null>(null);
|
|
46
53
|
const optsRef = useRef(opts);
|
|
47
|
-
|
|
54
|
+
useInsertionEffect(() => {
|
|
55
|
+
optsRef.current = opts;
|
|
56
|
+
});
|
|
48
57
|
|
|
49
58
|
useEffect(() => {
|
|
50
59
|
const url =
|
|
@@ -59,12 +68,6 @@ export function useMcpOAuthCallback(
|
|
|
59
68
|
const parsed = new URL(url);
|
|
60
69
|
const state = parsed.searchParams.get("state");
|
|
61
70
|
if (state) serverId = decodeServerIdFromState(state);
|
|
62
|
-
const error = parsed.searchParams.get("error");
|
|
63
|
-
if (error) {
|
|
64
|
-
throw new Error(
|
|
65
|
-
parsed.searchParams.get("error_description") ?? error,
|
|
66
|
-
);
|
|
67
|
-
}
|
|
68
71
|
if (!state) throw new Error('missing "state" parameter');
|
|
69
72
|
if (!serverId) {
|
|
70
73
|
throw new Error("state was not created by assistant-ui MCP");
|
|
@@ -9,6 +9,7 @@ import {
|
|
|
9
9
|
import { useAssistantScopeEffect } from "@assistant-ui/store/client";
|
|
10
10
|
import { ModelContext } from "@assistant-ui/core/store";
|
|
11
11
|
import { createMcpId } from "../utils/createMcpId";
|
|
12
|
+
import { clearOAuthProviderAuthState } from "../auth/createOAuthProvider";
|
|
12
13
|
import type { Tool } from "assistant-stream";
|
|
13
14
|
import { McpServerResource } from "./McpServerResource";
|
|
14
15
|
import { McpLocalStorage } from "./storage/McpLocalStorage";
|
|
@@ -303,7 +304,7 @@ const useMcpManagerResource = (
|
|
|
303
304
|
try {
|
|
304
305
|
await lookup.get({ key: id }).remove();
|
|
305
306
|
} catch {
|
|
306
|
-
await storage
|
|
307
|
+
await clearOAuthProviderAuthState(storage, id);
|
|
307
308
|
setCustomServers((prev) => prev.filter((s) => s.id !== id));
|
|
308
309
|
}
|
|
309
310
|
},
|