pi-mcp-adapter 2.7.0 → 2.9.0
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/CHANGELOG.md +18 -0
- package/README.md +14 -0
- package/elicitation-handler.ts +353 -0
- package/init.ts +10 -0
- package/mcp-auth-flow.ts +123 -23
- package/mcp-auth.ts +1 -0
- package/mcp-callback-server.ts +150 -62
- package/mcp-oauth-provider.ts +28 -8
- package/package.json +2 -1
- package/server-manager.ts +27 -9
- package/types.ts +8 -0
package/CHANGELOG.md
CHANGED
|
@@ -7,6 +7,24 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|
|
7
7
|
|
|
8
8
|
## [Unreleased]
|
|
9
9
|
|
|
10
|
+
## [2.9.0] - 2026-06-04
|
|
11
|
+
|
|
12
|
+
### Added
|
|
13
|
+
- Added MCP elicitation support with Pi form prompts and browser-opening URL requests.
|
|
14
|
+
|
|
15
|
+
### Fixed
|
|
16
|
+
- Rejected non-http/https MCP URL elicitations before prompting or opening a browser.
|
|
17
|
+
- Preserved empty string form values for MCP string elicitations unless schema constraints reject them.
|
|
18
|
+
|
|
19
|
+
## [2.8.0] - 2026-05-25
|
|
20
|
+
|
|
21
|
+
### Added
|
|
22
|
+
- Added per-server OAuth `redirectUri`, `clientName`, and `clientUri` overrides for pre-registered callbacks and dynamic client metadata.
|
|
23
|
+
|
|
24
|
+
### Fixed
|
|
25
|
+
- Avoided OAuth callback port exhaustion by starting the callback server lazily and using OS-assigned ports for dynamic OAuth flows.
|
|
26
|
+
- Re-register dynamic OAuth clients before browser auth when cached redirect URI metadata is missing or no longer matches the active callback URI.
|
|
27
|
+
|
|
10
28
|
## [2.7.0] - 2026-05-22
|
|
11
29
|
|
|
12
30
|
### Added
|
package/README.md
CHANGED
|
@@ -126,6 +126,12 @@ Pi-specific files are the write targets for imported or shared global servers wh
|
|
|
126
126
|
| `headers` | HTTP headers; supports `${VAR}` and `$env:VAR` interpolation |
|
|
127
127
|
| `auth` | `"bearer"` or `"oauth"` |
|
|
128
128
|
| `oauth.grantType` | `"authorization_code"` (default) or `"client_credentials"` for non-interactive machine auth |
|
|
129
|
+
| `oauth.clientId` | Pre-registered OAuth client ID; dynamic registration is used when omitted |
|
|
130
|
+
| `oauth.clientSecret` | OAuth client secret for confidential clients |
|
|
131
|
+
| `oauth.scope` | Requested OAuth scopes |
|
|
132
|
+
| `oauth.redirectUri` | Exact localhost redirect URI for browser OAuth, including port and path, for providers that pre-register callbacks |
|
|
133
|
+
| `oauth.clientName` | Client display name advertised during dynamic registration |
|
|
134
|
+
| `oauth.clientUri` | Client homepage URI advertised during dynamic registration |
|
|
129
135
|
| `bearerToken` / `bearerTokenEnv` | Token or env var name; `bearerToken` supports `${VAR}` and `$env:VAR` interpolation |
|
|
130
136
|
| `lifecycle` | `"lazy"` (default), `"eager"`, or `"keep-alive"` |
|
|
131
137
|
| `idleTimeout` | Minutes before idle disconnect (overrides global) |
|
|
@@ -134,6 +140,8 @@ Pi-specific files are the write targets for imported or shared global servers wh
|
|
|
134
140
|
| `excludeTools` | `string[]` of tool names to hide (matches original names like `get_screenshot` and prefixed names like `figma_get_screenshot`) |
|
|
135
141
|
| `debug` | Show server stderr (default: false) |
|
|
136
142
|
|
|
143
|
+
For pre-registered browser OAuth clients, set `oauth.redirectUri` to the exact callback registered with the provider, for example `"http://localhost:3118/callback"`. Dynamic clients normally omit it and use a lazy OS-assigned localhost callback port.
|
|
144
|
+
|
|
137
145
|
### Lifecycle Modes
|
|
138
146
|
|
|
139
147
|
- **`lazy`** (default) — Don't connect at startup. Connect on first tool call. Disconnect after idle timeout. Cached metadata keeps search/list working without connections.
|
|
@@ -161,9 +169,15 @@ Pi-specific files are the write targets for imported or shared global servers wh
|
|
|
161
169
|
| `autoAuth` | Auto-run OAuth on `connect`/tool calls when a server needs auth, then retry once (default: false). |
|
|
162
170
|
| `sampling` | Allow MCP servers to sample through Pi models, honoring `modelPreferences.hints` before current/default fallback (default: true when UI approval is available). |
|
|
163
171
|
| `samplingAutoApprove` | Skip sampling confirmation prompts. Required for sampling in non-UI sessions (default: false). |
|
|
172
|
+
| `elicitation` | Allow MCP servers to request user input through Pi UI forms/URL prompts (default: true when Pi UI form support is available). |
|
|
173
|
+
| `elicitationAutoOpenUrls` | Automatically open URL elicitations without prompting first (default: false). |
|
|
164
174
|
|
|
165
175
|
Per-server `idleTimeout` overrides the global setting.
|
|
166
176
|
|
|
177
|
+
### MCP Elicitation
|
|
178
|
+
|
|
179
|
+
When Pi exposes UI, the adapter advertises MCP elicitation support. Form elicitations are rendered with `ctx.ui.form()` and map Pi actions to MCP actions: submit → `accept`, secondary → `decline`, cancel → `cancel`. URL elicitations prompt before opening a browser unless `elicitationAutoOpenUrls` is enabled.
|
|
180
|
+
|
|
167
181
|
### Direct Tools
|
|
168
182
|
|
|
169
183
|
By default, all MCP tools are accessed through the single `mcp` proxy tool. This keeps context small but means the LLM has to discover MCP tools via proxy search. If you want specific tools to show up directly in the agent's tool list — alongside `read`, `bash`, `edit`, etc. — add `directTools` to your config.
|
|
@@ -0,0 +1,353 @@
|
|
|
1
|
+
import type { ExtensionUIContext } from "@earendil-works/pi-coding-agent";
|
|
2
|
+
import type { Client } from "@modelcontextprotocol/sdk/client/index.js";
|
|
3
|
+
import {
|
|
4
|
+
ElicitRequestSchema,
|
|
5
|
+
type ElicitRequest,
|
|
6
|
+
type ElicitRequestFormParams,
|
|
7
|
+
type ElicitRequestURLParams,
|
|
8
|
+
type ElicitResult,
|
|
9
|
+
} from "@modelcontextprotocol/sdk/types.js";
|
|
10
|
+
import open from "open";
|
|
11
|
+
|
|
12
|
+
export type ExtensionUIFormValue = string | number | boolean | string[] | undefined;
|
|
13
|
+
|
|
14
|
+
export interface ExtensionUIFormSelectOption {
|
|
15
|
+
value: string;
|
|
16
|
+
label?: string;
|
|
17
|
+
description?: string;
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
export type ExtensionUIFormField =
|
|
21
|
+
| {
|
|
22
|
+
type: "text";
|
|
23
|
+
name: string;
|
|
24
|
+
label: string;
|
|
25
|
+
description?: string;
|
|
26
|
+
placeholder?: string;
|
|
27
|
+
required?: boolean;
|
|
28
|
+
defaultValue?: string;
|
|
29
|
+
minLength?: number;
|
|
30
|
+
maxLength?: number;
|
|
31
|
+
pattern?: string;
|
|
32
|
+
}
|
|
33
|
+
| {
|
|
34
|
+
type: "number" | "integer";
|
|
35
|
+
name: string;
|
|
36
|
+
label: string;
|
|
37
|
+
description?: string;
|
|
38
|
+
required?: boolean;
|
|
39
|
+
defaultValue?: number;
|
|
40
|
+
minimum?: number;
|
|
41
|
+
maximum?: number;
|
|
42
|
+
}
|
|
43
|
+
| {
|
|
44
|
+
type: "boolean";
|
|
45
|
+
name: string;
|
|
46
|
+
label: string;
|
|
47
|
+
description?: string;
|
|
48
|
+
defaultValue?: boolean;
|
|
49
|
+
}
|
|
50
|
+
| {
|
|
51
|
+
type: "select";
|
|
52
|
+
name: string;
|
|
53
|
+
label: string;
|
|
54
|
+
description?: string;
|
|
55
|
+
required?: boolean;
|
|
56
|
+
options: ExtensionUIFormSelectOption[];
|
|
57
|
+
defaultValue?: string;
|
|
58
|
+
}
|
|
59
|
+
| {
|
|
60
|
+
type: "multiSelect";
|
|
61
|
+
name: string;
|
|
62
|
+
label: string;
|
|
63
|
+
description?: string;
|
|
64
|
+
required?: boolean;
|
|
65
|
+
options: ExtensionUIFormSelectOption[];
|
|
66
|
+
defaultValue?: string[];
|
|
67
|
+
};
|
|
68
|
+
|
|
69
|
+
export interface ExtensionUIFormRequest {
|
|
70
|
+
title: string;
|
|
71
|
+
message?: string;
|
|
72
|
+
fields: ExtensionUIFormField[];
|
|
73
|
+
submitLabel?: string;
|
|
74
|
+
secondaryLabel?: string;
|
|
75
|
+
cancelLabel?: string;
|
|
76
|
+
}
|
|
77
|
+
|
|
78
|
+
export type ExtensionUIFormResult =
|
|
79
|
+
| { action: "submit"; values: Record<string, ExtensionUIFormValue> }
|
|
80
|
+
| { action: "secondary" }
|
|
81
|
+
| { action: "cancel" };
|
|
82
|
+
|
|
83
|
+
export interface ElicitationUIContext extends ExtensionUIContext {
|
|
84
|
+
form(request: ExtensionUIFormRequest): Promise<ExtensionUIFormResult>;
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
export interface ElicitationHandlerOptions {
|
|
88
|
+
serverName: string;
|
|
89
|
+
ui: ElicitationUIContext;
|
|
90
|
+
autoOpenUrls: boolean;
|
|
91
|
+
}
|
|
92
|
+
|
|
93
|
+
export type ServerElicitationConfig = Omit<ElicitationHandlerOptions, "serverName">;
|
|
94
|
+
|
|
95
|
+
export function registerElicitationHandler(client: Client, options: ElicitationHandlerOptions): void {
|
|
96
|
+
client.setRequestHandler(ElicitRequestSchema, (request) => {
|
|
97
|
+
return handleElicitationRequest(options, request as ElicitRequest);
|
|
98
|
+
});
|
|
99
|
+
}
|
|
100
|
+
|
|
101
|
+
export async function handleElicitationRequest(
|
|
102
|
+
options: ElicitationHandlerOptions,
|
|
103
|
+
request: ElicitRequest,
|
|
104
|
+
): Promise<ElicitResult> {
|
|
105
|
+
const params = request.params;
|
|
106
|
+
if (params.mode === "url") {
|
|
107
|
+
return handleUrlElicitation(options, params);
|
|
108
|
+
}
|
|
109
|
+
return handleFormElicitation(options, params);
|
|
110
|
+
}
|
|
111
|
+
|
|
112
|
+
export async function handleFormElicitation(
|
|
113
|
+
options: ElicitationHandlerOptions,
|
|
114
|
+
params: ElicitRequestFormParams,
|
|
115
|
+
): Promise<ElicitResult> {
|
|
116
|
+
const form = convertMcpSchemaToPiForm(options.serverName, params);
|
|
117
|
+
const result = await options.ui.form(form);
|
|
118
|
+
if (result.action !== "submit") {
|
|
119
|
+
return convertPiFormResultToMcpResult(result);
|
|
120
|
+
}
|
|
121
|
+
return {
|
|
122
|
+
action: "accept",
|
|
123
|
+
content: coerceAndValidateFormValues(params, result.values),
|
|
124
|
+
};
|
|
125
|
+
}
|
|
126
|
+
|
|
127
|
+
export async function handleUrlElicitation(
|
|
128
|
+
options: ElicitationHandlerOptions,
|
|
129
|
+
params: ElicitRequestURLParams,
|
|
130
|
+
): Promise<ElicitResult> {
|
|
131
|
+
const browserUrl = getBrowserElicitationUrl(params.url);
|
|
132
|
+
if (!options.autoOpenUrls) {
|
|
133
|
+
const result = await options.ui.form({
|
|
134
|
+
title: "MCP Browser Request",
|
|
135
|
+
message: [
|
|
136
|
+
`Server: ${options.serverName}`,
|
|
137
|
+
"",
|
|
138
|
+
params.message,
|
|
139
|
+
"",
|
|
140
|
+
`Domain: ${browserUrl.host}`,
|
|
141
|
+
`URL: ${browserUrl.toString()}`,
|
|
142
|
+
"",
|
|
143
|
+
"Open this URL in your browser?",
|
|
144
|
+
].join("\n"),
|
|
145
|
+
fields: [],
|
|
146
|
+
submitLabel: "Open",
|
|
147
|
+
secondaryLabel: "Decline",
|
|
148
|
+
cancelLabel: "Cancel",
|
|
149
|
+
});
|
|
150
|
+
if (result.action === "secondary") return { action: "decline" };
|
|
151
|
+
if (result.action === "cancel") return { action: "cancel" };
|
|
152
|
+
}
|
|
153
|
+
|
|
154
|
+
await open(browserUrl.toString());
|
|
155
|
+
options.ui.notify("Opened browser for MCP elicitation.", "info");
|
|
156
|
+
return { action: "accept" };
|
|
157
|
+
}
|
|
158
|
+
|
|
159
|
+
export function convertMcpSchemaToPiForm(
|
|
160
|
+
serverName: string,
|
|
161
|
+
params: ElicitRequestFormParams,
|
|
162
|
+
): ExtensionUIFormRequest {
|
|
163
|
+
const required = new Set(params.requestedSchema.required ?? []);
|
|
164
|
+
return {
|
|
165
|
+
title: "MCP Input Request",
|
|
166
|
+
message: `Server: ${serverName}\n\n${params.message}`,
|
|
167
|
+
submitLabel: "Submit",
|
|
168
|
+
secondaryLabel: "Decline",
|
|
169
|
+
cancelLabel: "Cancel",
|
|
170
|
+
fields: Object.entries(params.requestedSchema.properties).map(([name, schema]): ExtensionUIFormField => {
|
|
171
|
+
const label = schema.title ?? humanizeName(name);
|
|
172
|
+
const base = {
|
|
173
|
+
name,
|
|
174
|
+
label,
|
|
175
|
+
description: schema.description,
|
|
176
|
+
required: required.has(name),
|
|
177
|
+
};
|
|
178
|
+
|
|
179
|
+
if (schema.type === "string" && "oneOf" in schema && Array.isArray(schema.oneOf)) {
|
|
180
|
+
return omitUndefined({
|
|
181
|
+
...base,
|
|
182
|
+
type: "select" as const,
|
|
183
|
+
options: schema.oneOf.map((option) => ({ value: option.const, label: option.title })),
|
|
184
|
+
defaultValue: schema.default,
|
|
185
|
+
});
|
|
186
|
+
}
|
|
187
|
+
|
|
188
|
+
if (schema.type === "string" && "enum" in schema && Array.isArray(schema.enum)) {
|
|
189
|
+
const enumNames = "enumNames" in schema && Array.isArray(schema.enumNames) ? schema.enumNames : undefined;
|
|
190
|
+
return omitUndefined({
|
|
191
|
+
...base,
|
|
192
|
+
type: "select" as const,
|
|
193
|
+
options: schema.enum.map((value, index) => omitUndefined({ value, label: enumNames?.[index] })),
|
|
194
|
+
defaultValue: schema.default,
|
|
195
|
+
});
|
|
196
|
+
}
|
|
197
|
+
|
|
198
|
+
if (schema.type === "array") {
|
|
199
|
+
return omitUndefined({
|
|
200
|
+
...base,
|
|
201
|
+
type: "multiSelect" as const,
|
|
202
|
+
options: extractMultiSelectOptions(schema),
|
|
203
|
+
defaultValue: schema.default,
|
|
204
|
+
});
|
|
205
|
+
}
|
|
206
|
+
|
|
207
|
+
if (schema.type === "number" || schema.type === "integer") {
|
|
208
|
+
return omitUndefined({
|
|
209
|
+
...base,
|
|
210
|
+
type: schema.type,
|
|
211
|
+
defaultValue: schema.default,
|
|
212
|
+
minimum: schema.minimum,
|
|
213
|
+
maximum: schema.maximum,
|
|
214
|
+
});
|
|
215
|
+
}
|
|
216
|
+
|
|
217
|
+
if (schema.type === "boolean") {
|
|
218
|
+
return omitUndefined({
|
|
219
|
+
type: "boolean" as const,
|
|
220
|
+
name,
|
|
221
|
+
label,
|
|
222
|
+
description: schema.description,
|
|
223
|
+
defaultValue: schema.default,
|
|
224
|
+
});
|
|
225
|
+
}
|
|
226
|
+
|
|
227
|
+
const stringSchema = schema as { default?: string; minLength?: number; maxLength?: number };
|
|
228
|
+
return omitUndefined({
|
|
229
|
+
...base,
|
|
230
|
+
type: "text" as const,
|
|
231
|
+
defaultValue: stringSchema.default,
|
|
232
|
+
minLength: stringSchema.minLength,
|
|
233
|
+
maxLength: stringSchema.maxLength,
|
|
234
|
+
});
|
|
235
|
+
}),
|
|
236
|
+
};
|
|
237
|
+
}
|
|
238
|
+
|
|
239
|
+
export function convertPiFormResultToMcpResult(result: ExtensionUIFormResult): ElicitResult {
|
|
240
|
+
if (result.action === "secondary") return { action: "decline" };
|
|
241
|
+
if (result.action === "cancel") return { action: "cancel" };
|
|
242
|
+
return { action: "accept", content: stripUndefined(result.values) as ElicitResult["content"] };
|
|
243
|
+
}
|
|
244
|
+
|
|
245
|
+
export function coerceAndValidateFormValues(
|
|
246
|
+
params: ElicitRequestFormParams,
|
|
247
|
+
values: Record<string, ExtensionUIFormValue>,
|
|
248
|
+
): Record<string, string | number | boolean | string[]> {
|
|
249
|
+
const output: Record<string, string | number | boolean | string[]> = {};
|
|
250
|
+
const required = new Set(params.requestedSchema.required ?? []);
|
|
251
|
+
|
|
252
|
+
for (const [name, schema] of Object.entries(params.requestedSchema.properties)) {
|
|
253
|
+
const raw = values[name] ?? schema.default;
|
|
254
|
+
if (raw === undefined || (raw === "" && schema.type !== "string")) {
|
|
255
|
+
if (required.has(name)) throw new Error(`Missing required elicitation field: ${name}`);
|
|
256
|
+
continue;
|
|
257
|
+
}
|
|
258
|
+
|
|
259
|
+
if (schema.type === "string") {
|
|
260
|
+
const stringSchema = schema as { minLength?: number; maxLength?: number };
|
|
261
|
+
const value = String(raw);
|
|
262
|
+
if (stringSchema.minLength !== undefined && value.length < stringSchema.minLength) {
|
|
263
|
+
throw new Error(`Elicitation field ${name} is shorter than minimum length ${stringSchema.minLength}`);
|
|
264
|
+
}
|
|
265
|
+
if (stringSchema.maxLength !== undefined && value.length > stringSchema.maxLength) {
|
|
266
|
+
throw new Error(`Elicitation field ${name} is longer than maximum length ${stringSchema.maxLength}`);
|
|
267
|
+
}
|
|
268
|
+
if ("enum" in schema && Array.isArray(schema.enum) && !schema.enum.includes(value)) {
|
|
269
|
+
throw new Error(`Elicitation field ${name} is not an allowed value`);
|
|
270
|
+
}
|
|
271
|
+
if ("oneOf" in schema && Array.isArray(schema.oneOf) && !schema.oneOf.some((option) => option.const === value)) {
|
|
272
|
+
throw new Error(`Elicitation field ${name} is not an allowed value`);
|
|
273
|
+
}
|
|
274
|
+
output[name] = value;
|
|
275
|
+
continue;
|
|
276
|
+
}
|
|
277
|
+
|
|
278
|
+
if (schema.type === "number" || schema.type === "integer") {
|
|
279
|
+
const value = typeof raw === "number" ? raw : Number(raw);
|
|
280
|
+
if (!Number.isFinite(value)) throw new Error(`Elicitation field ${name} must be a number`);
|
|
281
|
+
if (schema.type === "integer" && !Number.isInteger(value)) throw new Error(`Elicitation field ${name} must be an integer`);
|
|
282
|
+
if (schema.minimum !== undefined && value < schema.minimum) {
|
|
283
|
+
throw new Error(`Elicitation field ${name} is below minimum ${schema.minimum}`);
|
|
284
|
+
}
|
|
285
|
+
if (schema.maximum !== undefined && value > schema.maximum) {
|
|
286
|
+
throw new Error(`Elicitation field ${name} is above maximum ${schema.maximum}`);
|
|
287
|
+
}
|
|
288
|
+
output[name] = value;
|
|
289
|
+
continue;
|
|
290
|
+
}
|
|
291
|
+
|
|
292
|
+
if (schema.type === "boolean") {
|
|
293
|
+
output[name] = typeof raw === "boolean" ? raw : raw === "true";
|
|
294
|
+
continue;
|
|
295
|
+
}
|
|
296
|
+
|
|
297
|
+
if (schema.type === "array") {
|
|
298
|
+
if (!Array.isArray(raw)) throw new Error(`Elicitation field ${name} must be a list`);
|
|
299
|
+
const allowed = new Set(extractMultiSelectOptions(schema).map((option) => option.value));
|
|
300
|
+
const value = raw.map(String);
|
|
301
|
+
if (schema.minItems !== undefined && value.length < schema.minItems) {
|
|
302
|
+
throw new Error(`Elicitation field ${name} has fewer than ${schema.minItems} selections`);
|
|
303
|
+
}
|
|
304
|
+
if (schema.maxItems !== undefined && value.length > schema.maxItems) {
|
|
305
|
+
throw new Error(`Elicitation field ${name} has more than ${schema.maxItems} selections`);
|
|
306
|
+
}
|
|
307
|
+
for (const item of value) {
|
|
308
|
+
if (!allowed.has(item)) throw new Error(`Elicitation field ${name} contains an invalid selection`);
|
|
309
|
+
}
|
|
310
|
+
output[name] = value;
|
|
311
|
+
}
|
|
312
|
+
}
|
|
313
|
+
|
|
314
|
+
return output;
|
|
315
|
+
}
|
|
316
|
+
|
|
317
|
+
function extractMultiSelectOptions(schema: Extract<ElicitRequestFormParams["requestedSchema"]["properties"][string], { type: "array" }>): ExtensionUIFormSelectOption[] {
|
|
318
|
+
const items = schema.items as { enum?: string[]; anyOf?: Array<{ const: string; title: string }> };
|
|
319
|
+
if (Array.isArray(items.anyOf)) {
|
|
320
|
+
return items.anyOf.map((option) => ({ value: option.const, label: option.title }));
|
|
321
|
+
}
|
|
322
|
+
return (items.enum ?? []).map((value) => ({ value }));
|
|
323
|
+
}
|
|
324
|
+
|
|
325
|
+
function humanizeName(name: string): string {
|
|
326
|
+
return name
|
|
327
|
+
.replace(/[_-]+/g, " ")
|
|
328
|
+
.replace(/([a-z0-9])([A-Z])/g, "$1 $2")
|
|
329
|
+
.replace(/^./, (char) => char.toUpperCase());
|
|
330
|
+
}
|
|
331
|
+
|
|
332
|
+
function getBrowserElicitationUrl(url: string): URL {
|
|
333
|
+
const parsed = new URL(url);
|
|
334
|
+
if (parsed.protocol !== "http:" && parsed.protocol !== "https:") {
|
|
335
|
+
throw new Error(`MCP URL elicitation only supports http/https URLs: ${parsed.protocol}`);
|
|
336
|
+
}
|
|
337
|
+
return parsed;
|
|
338
|
+
}
|
|
339
|
+
|
|
340
|
+
function stripUndefined(values: Record<string, ExtensionUIFormValue>): Record<string, string | number | boolean | string[]> {
|
|
341
|
+
const output: Record<string, string | number | boolean | string[]> = {};
|
|
342
|
+
for (const [key, value] of Object.entries(values)) {
|
|
343
|
+
if (value !== undefined) output[key] = value;
|
|
344
|
+
}
|
|
345
|
+
return output;
|
|
346
|
+
}
|
|
347
|
+
|
|
348
|
+
function omitUndefined<T extends Record<string, unknown>>(value: T): T {
|
|
349
|
+
for (const key of Object.keys(value)) {
|
|
350
|
+
if (value[key] === undefined) delete value[key];
|
|
351
|
+
}
|
|
352
|
+
return value;
|
|
353
|
+
}
|
package/init.ts
CHANGED
|
@@ -43,6 +43,16 @@ export async function initializeMcp(
|
|
|
43
43
|
getSignal: () => ctx.signal,
|
|
44
44
|
});
|
|
45
45
|
}
|
|
46
|
+
const elicitationEnabled =
|
|
47
|
+
config.settings?.elicitation !== false &&
|
|
48
|
+
ctx.hasUI &&
|
|
49
|
+
typeof (ctx.ui as { form?: unknown }).form === "function";
|
|
50
|
+
if (elicitationEnabled) {
|
|
51
|
+
manager.setElicitationConfig({
|
|
52
|
+
ui: ctx.ui as any,
|
|
53
|
+
autoOpenUrls: config.settings?.elicitationAutoOpenUrls === true,
|
|
54
|
+
});
|
|
55
|
+
}
|
|
46
56
|
const lifecycle = new McpLifecycleManager(manager);
|
|
47
57
|
const toolMetadata = new Map<string, ToolMetadata[]>();
|
|
48
58
|
const failureTracker = new Map<string, number>();
|
package/mcp-auth-flow.ts
CHANGED
|
@@ -16,6 +16,7 @@ import {
|
|
|
16
16
|
waitForCallback,
|
|
17
17
|
cancelPendingCallback,
|
|
18
18
|
stopCallbackServer,
|
|
19
|
+
releaseCallbackServer,
|
|
19
20
|
} from "./mcp-callback-server.ts"
|
|
20
21
|
import {
|
|
21
22
|
getAuthForUrl,
|
|
@@ -23,6 +24,7 @@ import {
|
|
|
23
24
|
hasStoredTokens,
|
|
24
25
|
clearAllCredentials,
|
|
25
26
|
clearClientInfo,
|
|
27
|
+
clearTokens,
|
|
26
28
|
clearCodeVerifier,
|
|
27
29
|
updateOAuthState,
|
|
28
30
|
getOAuthState,
|
|
@@ -52,17 +54,82 @@ function generateState(): string {
|
|
|
52
54
|
/**
|
|
53
55
|
* Extract OAuth configuration from a ServerEntry.
|
|
54
56
|
*/
|
|
55
|
-
function extractOAuthConfig(definition: ServerEntry): McpOAuthConfig {
|
|
56
|
-
// If oauth is explicitly false, return empty config
|
|
57
|
+
export function extractOAuthConfig(definition: ServerEntry): McpOAuthConfig {
|
|
57
58
|
if (definition.oauth === false) {
|
|
58
59
|
return {}
|
|
59
60
|
}
|
|
60
|
-
|
|
61
|
-
|
|
62
|
-
|
|
63
|
-
|
|
64
|
-
|
|
61
|
+
|
|
62
|
+
const config: McpOAuthConfig = {}
|
|
63
|
+
if (definition.oauth?.grantType !== undefined) config.grantType = definition.oauth.grantType
|
|
64
|
+
if (definition.oauth?.clientId !== undefined) config.clientId = definition.oauth.clientId
|
|
65
|
+
if (definition.oauth?.clientSecret !== undefined) config.clientSecret = definition.oauth.clientSecret
|
|
66
|
+
if (definition.oauth?.scope !== undefined) config.scope = definition.oauth.scope
|
|
67
|
+
if (definition.oauth?.redirectUri !== undefined) {
|
|
68
|
+
if (typeof definition.oauth.redirectUri !== "string") {
|
|
69
|
+
throw new Error("OAuth redirectUri must be a string")
|
|
70
|
+
}
|
|
71
|
+
const redirectUri = definition.oauth.redirectUri.trim()
|
|
72
|
+
if (!redirectUri) {
|
|
73
|
+
throw new Error("OAuth redirectUri must not be empty")
|
|
74
|
+
}
|
|
75
|
+
config.redirectUri = redirectUri
|
|
65
76
|
}
|
|
77
|
+
if (definition.oauth?.clientName !== undefined) {
|
|
78
|
+
if (typeof definition.oauth.clientName !== "string") {
|
|
79
|
+
throw new Error("OAuth clientName must be a string")
|
|
80
|
+
}
|
|
81
|
+
const clientName = definition.oauth.clientName.trim()
|
|
82
|
+
if (!clientName) {
|
|
83
|
+
throw new Error("OAuth clientName must not be empty")
|
|
84
|
+
}
|
|
85
|
+
config.clientName = clientName
|
|
86
|
+
}
|
|
87
|
+
if (definition.oauth?.clientUri !== undefined) {
|
|
88
|
+
if (typeof definition.oauth.clientUri !== "string") {
|
|
89
|
+
throw new Error("OAuth clientUri must be a string")
|
|
90
|
+
}
|
|
91
|
+
const clientUri = definition.oauth.clientUri.trim()
|
|
92
|
+
if (!clientUri) {
|
|
93
|
+
throw new Error("OAuth clientUri must not be empty")
|
|
94
|
+
}
|
|
95
|
+
config.clientUri = clientUri
|
|
96
|
+
}
|
|
97
|
+
return config
|
|
98
|
+
}
|
|
99
|
+
|
|
100
|
+
function parseOAuthRedirectUri(redirectUri: string): { port: number; callbackHost: string; callbackPath: string } {
|
|
101
|
+
let url: URL
|
|
102
|
+
try {
|
|
103
|
+
url = new URL(redirectUri)
|
|
104
|
+
} catch (error) {
|
|
105
|
+
throw new Error(`Invalid OAuth redirectUri: ${redirectUri}`, { cause: error })
|
|
106
|
+
}
|
|
107
|
+
|
|
108
|
+
const hostname = url.hostname.toLowerCase()
|
|
109
|
+
const isLocalhost = hostname === "localhost" || hostname === "127.0.0.1" || hostname === "[::1]" || hostname === "::1"
|
|
110
|
+
if (url.protocol !== "http:" || !isLocalhost) {
|
|
111
|
+
throw new Error("OAuth redirectUri must be an http:// localhost or loopback URI")
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
if (url.username || url.password) {
|
|
115
|
+
throw new Error("OAuth redirectUri must not include username or password")
|
|
116
|
+
}
|
|
117
|
+
|
|
118
|
+
if (url.hash) {
|
|
119
|
+
throw new Error("OAuth redirectUri must not include a fragment")
|
|
120
|
+
}
|
|
121
|
+
|
|
122
|
+
if (!url.port) {
|
|
123
|
+
throw new Error("OAuth redirectUri must include an explicit numeric port")
|
|
124
|
+
}
|
|
125
|
+
|
|
126
|
+
const port = Number.parseInt(url.port, 10)
|
|
127
|
+
if (!Number.isInteger(port) || port <= 0 || port > 65535) {
|
|
128
|
+
throw new Error("OAuth redirectUri must include an explicit numeric port")
|
|
129
|
+
}
|
|
130
|
+
|
|
131
|
+
const callbackHost = hostname === "[::1]" ? "::1" : hostname
|
|
132
|
+
return { port, callbackHost, callbackPath: url.pathname }
|
|
66
133
|
}
|
|
67
134
|
|
|
68
135
|
/**
|
|
@@ -76,14 +143,14 @@ export async function startAuth(
|
|
|
76
143
|
): Promise<{ authorizationUrl: string }> {
|
|
77
144
|
const config = definition ? extractOAuthConfig(definition) : {}
|
|
78
145
|
|
|
79
|
-
const storedAuth = await getAuthForUrl(serverName, serverUrl)
|
|
80
|
-
if (storedAuth?.clientInfo && !storedAuth.tokens && !config.clientId) {
|
|
81
|
-
clearClientInfo(serverName)
|
|
82
|
-
clearCodeVerifier(serverName)
|
|
83
|
-
await clearOAuthState(serverName)
|
|
84
|
-
}
|
|
85
|
-
|
|
86
146
|
if (config.grantType === "client_credentials") {
|
|
147
|
+
const storedAuth = await getAuthForUrl(serverName, serverUrl)
|
|
148
|
+
if (storedAuth?.clientInfo && !storedAuth.tokens && !config.clientId) {
|
|
149
|
+
clearClientInfo(serverName)
|
|
150
|
+
clearCodeVerifier(serverName)
|
|
151
|
+
await clearOAuthState(serverName)
|
|
152
|
+
}
|
|
153
|
+
|
|
87
154
|
const authProvider = new McpOAuthProvider(serverName, serverUrl, config, {
|
|
88
155
|
onRedirect: async () => {
|
|
89
156
|
throw new Error("Browser redirect is not used for client_credentials flow")
|
|
@@ -96,12 +163,20 @@ export async function startAuth(
|
|
|
96
163
|
return { authorizationUrl: "" }
|
|
97
164
|
}
|
|
98
165
|
|
|
99
|
-
|
|
100
|
-
// Pre-registered OAuth clients require an exact redirect URI, so enforce strict port binding.
|
|
101
|
-
await ensureCallbackServer({ strictPort: Boolean(config.clientId) })
|
|
102
|
-
|
|
166
|
+
const redirectCallback = config.redirectUri !== undefined ? parseOAuthRedirectUri(config.redirectUri) : undefined
|
|
103
167
|
const oauthState = generateState()
|
|
104
|
-
|
|
168
|
+
|
|
169
|
+
try {
|
|
170
|
+
await ensureCallbackServer({
|
|
171
|
+
strictPort: Boolean(config.clientId) || config.redirectUri !== undefined,
|
|
172
|
+
oauthState,
|
|
173
|
+
reserveState: true,
|
|
174
|
+
...(redirectCallback ? { port: redirectCallback.port, callbackHost: redirectCallback.callbackHost, callbackPath: redirectCallback.callbackPath } : {}),
|
|
175
|
+
})
|
|
176
|
+
} catch (error) {
|
|
177
|
+
await clearOAuthState(serverName)
|
|
178
|
+
throw error
|
|
179
|
+
}
|
|
105
180
|
|
|
106
181
|
let capturedUrl: URL | undefined
|
|
107
182
|
const authProvider = new McpOAuthProvider(serverName, serverUrl, config, {
|
|
@@ -111,8 +186,28 @@ export async function startAuth(
|
|
|
111
186
|
})
|
|
112
187
|
|
|
113
188
|
try {
|
|
189
|
+
const storedAuth = await getAuthForUrl(serverName, serverUrl)
|
|
190
|
+
if (storedAuth?.clientInfo && !config.clientId) {
|
|
191
|
+
if (!storedAuth.tokens) {
|
|
192
|
+
clearClientInfo(serverName)
|
|
193
|
+
clearCodeVerifier(serverName)
|
|
194
|
+
await clearOAuthState(serverName)
|
|
195
|
+
} else {
|
|
196
|
+
const redirectUris = storedAuth.clientInfo.redirectUris
|
|
197
|
+
if (!Array.isArray(redirectUris) || !redirectUris.includes(authProvider.redirectUrl ?? "")) {
|
|
198
|
+
clearClientInfo(serverName)
|
|
199
|
+
clearTokens(serverName)
|
|
200
|
+
clearCodeVerifier(serverName)
|
|
201
|
+
await clearOAuthState(serverName)
|
|
202
|
+
}
|
|
203
|
+
}
|
|
204
|
+
}
|
|
205
|
+
|
|
206
|
+
await updateOAuthState(serverName, oauthState, serverUrl)
|
|
207
|
+
|
|
114
208
|
const result = await runSdkAuth(authProvider, { serverUrl })
|
|
115
209
|
if (result === "AUTHORIZED") {
|
|
210
|
+
releaseCallbackServer(oauthState)
|
|
116
211
|
await clearOAuthState(serverName)
|
|
117
212
|
return { authorizationUrl: "" }
|
|
118
213
|
}
|
|
@@ -125,6 +220,7 @@ export async function startAuth(
|
|
|
125
220
|
)
|
|
126
221
|
return { authorizationUrl: capturedUrl.toString() }
|
|
127
222
|
} catch (error) {
|
|
223
|
+
releaseCallbackServer(oauthState)
|
|
128
224
|
await clearOAuthState(serverName)
|
|
129
225
|
throw error
|
|
130
226
|
}
|
|
@@ -142,11 +238,17 @@ export async function completeAuth(
|
|
|
142
238
|
throw new Error(`No pending OAuth flow for server: ${serverName}`)
|
|
143
239
|
}
|
|
144
240
|
|
|
241
|
+
const oauthState = await getOAuthState(serverName)
|
|
242
|
+
|
|
145
243
|
try {
|
|
146
244
|
// Complete the auth using the transport's finishAuth method
|
|
147
245
|
await transport.finishAuth(authorizationCode)
|
|
148
246
|
return "authenticated"
|
|
149
247
|
} finally {
|
|
248
|
+
if (oauthState) {
|
|
249
|
+
releaseCallbackServer(oauthState)
|
|
250
|
+
}
|
|
251
|
+
await clearOAuthState(serverName)
|
|
150
252
|
pendingTransports.delete(serverName)
|
|
151
253
|
await transport.close().catch(() => {})
|
|
152
254
|
}
|
|
@@ -347,11 +449,9 @@ export function supportsOAuth(definition: ServerEntry): boolean {
|
|
|
347
449
|
|
|
348
450
|
/**
|
|
349
451
|
* Initialize the OAuth system on startup.
|
|
350
|
-
*
|
|
452
|
+
* OAuth callback binding is lazy and starts from startAuth() only.
|
|
351
453
|
*/
|
|
352
|
-
export async function initializeOAuth(): Promise<void> {
|
|
353
|
-
await ensureCallbackServer()
|
|
354
|
-
}
|
|
454
|
+
export async function initializeOAuth(): Promise<void> {}
|
|
355
455
|
|
|
356
456
|
/**
|
|
357
457
|
* Shutdown the OAuth system.
|
package/mcp-auth.ts
CHANGED
package/mcp-callback-server.ts
CHANGED
|
@@ -7,9 +7,11 @@
|
|
|
7
7
|
|
|
8
8
|
import { createServer, type Server, type IncomingMessage, type ServerResponse } from "http"
|
|
9
9
|
import {
|
|
10
|
-
|
|
10
|
+
DEFAULT_OAUTH_CALLBACK_PATH,
|
|
11
11
|
getConfiguredOAuthCallbackPort,
|
|
12
|
+
getOAuthCallbackPath,
|
|
12
13
|
getOAuthCallbackPort,
|
|
14
|
+
setOAuthCallbackPath,
|
|
13
15
|
setOAuthCallbackPort,
|
|
14
16
|
} from "./mcp-oauth-provider.ts"
|
|
15
17
|
|
|
@@ -34,6 +36,15 @@ const HTML_SUCCESS = `<!DOCTYPE html>
|
|
|
34
36
|
</body>
|
|
35
37
|
</html>`
|
|
36
38
|
|
|
39
|
+
function escapeHtml(value: string): string {
|
|
40
|
+
return value
|
|
41
|
+
.replace(/&/g, "&")
|
|
42
|
+
.replace(/</g, "<")
|
|
43
|
+
.replace(/>/g, ">")
|
|
44
|
+
.replace(/"/g, """)
|
|
45
|
+
.replace(/'/g, "'")
|
|
46
|
+
}
|
|
47
|
+
|
|
37
48
|
const HTML_ERROR = (error: string) => `<!DOCTYPE html>
|
|
38
49
|
<html>
|
|
39
50
|
<head>
|
|
@@ -50,7 +61,7 @@ const HTML_ERROR = (error: string) => `<!DOCTYPE html>
|
|
|
50
61
|
<div class="container">
|
|
51
62
|
<h1>Authorization Failed</h1>
|
|
52
63
|
<p>An error occurred during authorization.</p>
|
|
53
|
-
<div class="error">${error}</div>
|
|
64
|
+
<div class="error">${escapeHtml(error)}</div>
|
|
54
65
|
</div>
|
|
55
66
|
</body>
|
|
56
67
|
</html>`
|
|
@@ -64,17 +75,25 @@ interface PendingAuth {
|
|
|
64
75
|
|
|
65
76
|
/** Server singleton state */
|
|
66
77
|
let server: Server | undefined
|
|
78
|
+
let bindingPromise: Promise<void> | undefined
|
|
67
79
|
const pendingAuths = new Map<string, PendingAuth>()
|
|
80
|
+
const reservedAuthStates = new Set<string>()
|
|
68
81
|
|
|
69
82
|
/** Timeout for callback completion (5 minutes) */
|
|
70
83
|
const CALLBACK_TIMEOUT_MS = 5 * 60 * 1000
|
|
71
84
|
|
|
72
|
-
const MAX_PORT_SCAN_ATTEMPTS = 25
|
|
73
|
-
|
|
74
85
|
interface EnsureCallbackServerOptions {
|
|
75
86
|
strictPort?: boolean
|
|
87
|
+
port?: number
|
|
88
|
+
callbackHost?: string
|
|
89
|
+
callbackPath?: string
|
|
90
|
+
oauthState?: string
|
|
91
|
+
reserveState?: boolean
|
|
76
92
|
}
|
|
77
93
|
|
|
94
|
+
const DEFAULT_OAUTH_CALLBACK_HOST = "localhost"
|
|
95
|
+
let callbackServerHost = DEFAULT_OAUTH_CALLBACK_HOST
|
|
96
|
+
|
|
78
97
|
/**
|
|
79
98
|
* Handle incoming HTTP requests to the callback server.
|
|
80
99
|
*/
|
|
@@ -82,7 +101,7 @@ function handleRequest(req: IncomingMessage, res: ServerResponse): void {
|
|
|
82
101
|
const url = new URL(req.url || "/", `http://${req.headers.host}`)
|
|
83
102
|
|
|
84
103
|
// Only handle the callback path
|
|
85
|
-
if (url.pathname !==
|
|
104
|
+
if (url.pathname !== getOAuthCallbackPath()) {
|
|
86
105
|
res.writeHead(404, { "Content-Type": "text/plain" })
|
|
87
106
|
res.end("Not found")
|
|
88
107
|
return
|
|
@@ -101,15 +120,25 @@ function handleRequest(req: IncomingMessage, res: ServerResponse): void {
|
|
|
101
120
|
return
|
|
102
121
|
}
|
|
103
122
|
|
|
104
|
-
|
|
123
|
+
const pending = pendingAuths.get(state)
|
|
124
|
+
const isReserved = reservedAuthStates.has(state)
|
|
125
|
+
|
|
126
|
+
// Handle OAuth errors only for a state that belongs to an active flow.
|
|
105
127
|
if (error) {
|
|
128
|
+
if (!pending && !isReserved) {
|
|
129
|
+
const errorMsg = "Invalid or expired state parameter - potential CSRF attack"
|
|
130
|
+
res.writeHead(400, { "Content-Type": "text/html" })
|
|
131
|
+
res.end(HTML_ERROR(errorMsg))
|
|
132
|
+
return
|
|
133
|
+
}
|
|
134
|
+
|
|
106
135
|
const errorMsg = errorDescription || error
|
|
107
136
|
// Send HTTP response first before rejecting promise
|
|
108
137
|
res.writeHead(200, { "Content-Type": "text/html" })
|
|
109
138
|
res.end(HTML_ERROR(errorMsg))
|
|
139
|
+
reservedAuthStates.delete(state)
|
|
110
140
|
// Reject promise after response is sent (defer to allow test to attach handler)
|
|
111
|
-
if (
|
|
112
|
-
const pending = pendingAuths.get(state)!
|
|
141
|
+
if (pending) {
|
|
113
142
|
clearTimeout(pending.timeout)
|
|
114
143
|
pendingAuths.delete(state)
|
|
115
144
|
setTimeout(() => pending.reject(new Error(errorMsg)), 0)
|
|
@@ -117,22 +146,20 @@ function handleRequest(req: IncomingMessage, res: ServerResponse): void {
|
|
|
117
146
|
return
|
|
118
147
|
}
|
|
119
148
|
|
|
120
|
-
// Require authorization code
|
|
121
|
-
if (!code) {
|
|
122
|
-
res.writeHead(400, { "Content-Type": "text/html" })
|
|
123
|
-
res.end(HTML_ERROR("No authorization code provided"))
|
|
124
|
-
return
|
|
125
|
-
}
|
|
126
|
-
|
|
127
149
|
// Validate state parameter
|
|
128
|
-
if (!
|
|
150
|
+
if (!pending) {
|
|
129
151
|
const errorMsg = "Invalid or expired state parameter - potential CSRF attack"
|
|
130
152
|
res.writeHead(400, { "Content-Type": "text/html" })
|
|
131
153
|
res.end(HTML_ERROR(errorMsg))
|
|
132
154
|
return
|
|
133
155
|
}
|
|
134
156
|
|
|
135
|
-
|
|
157
|
+
// Require authorization code
|
|
158
|
+
if (!code) {
|
|
159
|
+
res.writeHead(400, { "Content-Type": "text/html" })
|
|
160
|
+
res.end(HTML_ERROR("No authorization code provided"))
|
|
161
|
+
return
|
|
162
|
+
}
|
|
136
163
|
|
|
137
164
|
// Clear timeout and resolve the pending promise
|
|
138
165
|
clearTimeout(pending.timeout)
|
|
@@ -146,72 +173,128 @@ function handleRequest(req: IncomingMessage, res: ServerResponse): void {
|
|
|
146
173
|
/**
|
|
147
174
|
* Ensure the callback server is running.
|
|
148
175
|
* If strictPort is true, requires binding on the configured callback port.
|
|
149
|
-
* If strictPort is false,
|
|
176
|
+
* If strictPort is false, asks the OS for an available local port.
|
|
150
177
|
*/
|
|
151
178
|
export async function ensureCallbackServer(options: EnsureCallbackServerOptions = {}): Promise<void> {
|
|
152
|
-
|
|
153
|
-
|
|
179
|
+
while (bindingPromise) {
|
|
180
|
+
await bindingPromise
|
|
181
|
+
}
|
|
154
182
|
|
|
155
|
-
|
|
156
|
-
|
|
183
|
+
const operation = ensureCallbackServerLocked(options)
|
|
184
|
+
bindingPromise = operation
|
|
185
|
+
try {
|
|
186
|
+
await operation
|
|
187
|
+
} finally {
|
|
188
|
+
if (bindingPromise === operation) {
|
|
189
|
+
bindingPromise = undefined
|
|
190
|
+
}
|
|
191
|
+
}
|
|
192
|
+
}
|
|
157
193
|
|
|
158
|
-
|
|
194
|
+
async function ensureCallbackServerLocked(options: EnsureCallbackServerOptions = {}): Promise<void> {
|
|
195
|
+
const requiredPort = options.port ?? getConfiguredOAuthCallbackPort()
|
|
196
|
+
const strictPort = options.strictPort === true
|
|
197
|
+
const requestedHost = options.callbackHost ?? DEFAULT_OAUTH_CALLBACK_HOST
|
|
198
|
+
const rawRequestedPath = options.callbackPath ?? DEFAULT_OAUTH_CALLBACK_PATH
|
|
199
|
+
const requestedPath = rawRequestedPath.startsWith("/") ? rawRequestedPath : `/${rawRequestedPath}`
|
|
200
|
+
if (options.reserveState && !options.oauthState) {
|
|
201
|
+
throw new Error("OAuth callback reservation requires an oauthState")
|
|
202
|
+
}
|
|
203
|
+
let reservedState: string | undefined
|
|
204
|
+
|
|
205
|
+
const previousServer = server
|
|
206
|
+
const needsStrictRebind = Boolean(previousServer && strictPort && getOAuthCallbackPort() !== requiredPort)
|
|
207
|
+
const needsHostSwitch = Boolean(previousServer && callbackServerHost !== requestedHost)
|
|
208
|
+
const needsPathSwitch = Boolean(previousServer && getOAuthCallbackPath() !== requestedPath)
|
|
209
|
+
|
|
210
|
+
if (previousServer) {
|
|
211
|
+
if (!needsStrictRebind && !needsHostSwitch) {
|
|
212
|
+
if (needsPathSwitch) {
|
|
213
|
+
if (pendingAuths.size > 0 || reservedAuthStates.size > 0) {
|
|
214
|
+
throw new Error(
|
|
215
|
+
`OAuth callback server is using path ${getOAuthCallbackPath()}, but callback path ${requestedPath} is required and cannot be switched while authorizations are pending`
|
|
216
|
+
)
|
|
217
|
+
}
|
|
218
|
+
setOAuthCallbackPath(requestedPath)
|
|
219
|
+
}
|
|
220
|
+
if (options.reserveState && options.oauthState) {
|
|
221
|
+
reservedAuthStates.add(options.oauthState)
|
|
222
|
+
reservedState = options.oauthState
|
|
223
|
+
}
|
|
224
|
+
return
|
|
225
|
+
}
|
|
226
|
+
|
|
227
|
+
if (pendingAuths.size > 0 || reservedAuthStates.size > 0) {
|
|
159
228
|
throw new Error(
|
|
160
|
-
`OAuth callback server is running on
|
|
229
|
+
`OAuth callback server is running on ${callbackServerHost}:${getOAuthCallbackPort()}, but strict callback endpoint ${requestedHost}:${requiredPort} is required and cannot be switched while authorizations are pending`
|
|
161
230
|
)
|
|
162
231
|
}
|
|
163
|
-
|
|
164
|
-
await stopCallbackServer()
|
|
165
232
|
}
|
|
166
233
|
|
|
167
|
-
const
|
|
168
|
-
const
|
|
169
|
-
let lastError: Error | undefined
|
|
234
|
+
const candidateServer = createServer(handleRequest)
|
|
235
|
+
const listenPort = strictPort ? requiredPort : 0
|
|
170
236
|
|
|
171
|
-
|
|
172
|
-
|
|
173
|
-
|
|
174
|
-
|
|
175
|
-
|
|
176
|
-
await new Promise<void>((resolve, reject) => {
|
|
177
|
-
candidateServer.once("error", (err) => {
|
|
178
|
-
reject(err)
|
|
179
|
-
})
|
|
237
|
+
try {
|
|
238
|
+
await new Promise<void>((resolve, reject) => {
|
|
239
|
+
candidateServer.once("error", (err) => {
|
|
240
|
+
reject(err)
|
|
241
|
+
})
|
|
180
242
|
|
|
181
|
-
|
|
182
|
-
|
|
183
|
-
})
|
|
243
|
+
candidateServer.listen(listenPort, requestedHost, () => {
|
|
244
|
+
resolve()
|
|
184
245
|
})
|
|
246
|
+
})
|
|
185
247
|
|
|
186
|
-
|
|
187
|
-
|
|
188
|
-
|
|
189
|
-
|
|
190
|
-
|
|
191
|
-
|
|
248
|
+
if (strictPort) {
|
|
249
|
+
setOAuthCallbackPort(requiredPort)
|
|
250
|
+
} else {
|
|
251
|
+
const address = candidateServer.address()
|
|
252
|
+
if (!address || typeof address === "string" || typeof address.port !== "number") {
|
|
253
|
+
throw new Error("OAuth callback server did not report an assigned port")
|
|
254
|
+
}
|
|
255
|
+
setOAuthCallbackPort(address.port)
|
|
256
|
+
}
|
|
257
|
+
|
|
258
|
+
if (previousServer && (needsStrictRebind || needsHostSwitch)) {
|
|
192
259
|
await new Promise<void>((resolve) => {
|
|
193
|
-
|
|
260
|
+
previousServer.close(() => resolve())
|
|
194
261
|
})
|
|
262
|
+
}
|
|
195
263
|
|
|
196
|
-
|
|
197
|
-
|
|
198
|
-
|
|
264
|
+
callbackServerHost = requestedHost
|
|
265
|
+
setOAuthCallbackPath(requestedPath)
|
|
266
|
+
server = candidateServer
|
|
267
|
+
if (options.reserveState && options.oauthState) {
|
|
268
|
+
reservedAuthStates.add(options.oauthState)
|
|
269
|
+
reservedState = options.oauthState
|
|
270
|
+
}
|
|
271
|
+
server.unref()
|
|
272
|
+
} catch (error) {
|
|
273
|
+
if (reservedState) {
|
|
274
|
+
reservedAuthStates.delete(reservedState)
|
|
275
|
+
}
|
|
276
|
+
const nodeError = error as NodeJS.ErrnoException
|
|
277
|
+
await new Promise<void>((resolve) => {
|
|
278
|
+
candidateServer.close(() => resolve())
|
|
279
|
+
})
|
|
199
280
|
|
|
200
|
-
|
|
281
|
+
if (strictPort && nodeError.code === "EADDRINUSE") {
|
|
282
|
+
throw new Error(
|
|
283
|
+
`OAuth callback port ${requiredPort} is already in use. Pre-registered OAuth clients require an exact redirect URI; set MCP_OAUTH_CALLBACK_PORT to your registered port or free port ${requiredPort}`,
|
|
284
|
+
{ cause: error }
|
|
285
|
+
)
|
|
201
286
|
}
|
|
202
|
-
}
|
|
203
287
|
|
|
204
|
-
|
|
205
|
-
throw new Error(
|
|
206
|
-
`OAuth callback port ${preferredPort} is already in use. Pre-registered OAuth clients require an exact redirect URI; set MCP_OAUTH_CALLBACK_PORT to your registered port or free port ${preferredPort}`,
|
|
207
|
-
{ cause: lastError }
|
|
208
|
-
)
|
|
288
|
+
throw error
|
|
209
289
|
}
|
|
290
|
+
}
|
|
291
|
+
|
|
292
|
+
export function reserveCallbackServer(oauthState: string): void {
|
|
293
|
+
reservedAuthStates.add(oauthState)
|
|
294
|
+
}
|
|
210
295
|
|
|
211
|
-
|
|
212
|
-
|
|
213
|
-
{ cause: lastError }
|
|
214
|
-
)
|
|
296
|
+
export function releaseCallbackServer(oauthState: string): void {
|
|
297
|
+
reservedAuthStates.delete(oauthState)
|
|
215
298
|
}
|
|
216
299
|
|
|
217
300
|
/**
|
|
@@ -219,6 +302,7 @@ export async function ensureCallbackServer(options: EnsureCallbackServerOptions
|
|
|
219
302
|
* Returns a promise that resolves with the authorization code.
|
|
220
303
|
*/
|
|
221
304
|
export function waitForCallback(oauthState: string): Promise<string> {
|
|
305
|
+
reservedAuthStates.delete(oauthState)
|
|
222
306
|
return new Promise((resolve, reject) => {
|
|
223
307
|
const timeout = setTimeout(() => {
|
|
224
308
|
if (pendingAuths.has(oauthState)) {
|
|
@@ -235,6 +319,7 @@ export function waitForCallback(oauthState: string): Promise<string> {
|
|
|
235
319
|
* Cancel a pending authorization by state.
|
|
236
320
|
*/
|
|
237
321
|
export function cancelPendingCallback(oauthState: string): void {
|
|
322
|
+
reservedAuthStates.delete(oauthState)
|
|
238
323
|
const pending = pendingAuths.get(oauthState)
|
|
239
324
|
if (pending) {
|
|
240
325
|
clearTimeout(pending.timeout)
|
|
@@ -257,10 +342,13 @@ export async function stopCallbackServer(): Promise<void> {
|
|
|
257
342
|
}
|
|
258
343
|
|
|
259
344
|
setOAuthCallbackPort(getConfiguredOAuthCallbackPort())
|
|
345
|
+
callbackServerHost = DEFAULT_OAUTH_CALLBACK_HOST
|
|
346
|
+
setOAuthCallbackPath(DEFAULT_OAUTH_CALLBACK_PATH)
|
|
260
347
|
|
|
261
348
|
// Reject all pending auths (defer to allow any pending operations to complete)
|
|
262
349
|
const pendingList = Array.from(pendingAuths.entries())
|
|
263
350
|
pendingAuths.clear()
|
|
351
|
+
reservedAuthStates.clear()
|
|
264
352
|
setTimeout(() => {
|
|
265
353
|
for (const [, pending] of pendingList) {
|
|
266
354
|
clearTimeout(pending.timeout)
|
package/mcp-oauth-provider.ts
CHANGED
|
@@ -28,7 +28,7 @@ import {
|
|
|
28
28
|
|
|
29
29
|
// Callback server configuration
|
|
30
30
|
const DEFAULT_OAUTH_CALLBACK_PORT = 19876
|
|
31
|
-
const
|
|
31
|
+
const DEFAULT_OAUTH_CALLBACK_PATH = "/callback"
|
|
32
32
|
|
|
33
33
|
let configuredOAuthCallbackPort = DEFAULT_OAUTH_CALLBACK_PORT
|
|
34
34
|
|
|
@@ -40,6 +40,7 @@ if (process.env.MCP_OAUTH_CALLBACK_PORT) {
|
|
|
40
40
|
}
|
|
41
41
|
|
|
42
42
|
let oauthCallbackPort = configuredOAuthCallbackPort
|
|
43
|
+
let oauthCallbackPath = DEFAULT_OAUTH_CALLBACK_PATH
|
|
43
44
|
|
|
44
45
|
export function getConfiguredOAuthCallbackPort(): number {
|
|
45
46
|
return configuredOAuthCallbackPort
|
|
@@ -53,12 +54,23 @@ export function setOAuthCallbackPort(port: number): void {
|
|
|
53
54
|
oauthCallbackPort = port
|
|
54
55
|
}
|
|
55
56
|
|
|
57
|
+
export function getOAuthCallbackPath(): string {
|
|
58
|
+
return oauthCallbackPath
|
|
59
|
+
}
|
|
60
|
+
|
|
61
|
+
export function setOAuthCallbackPath(path: string): void {
|
|
62
|
+
oauthCallbackPath = path.startsWith("/") ? path : `/${path}`
|
|
63
|
+
}
|
|
64
|
+
|
|
56
65
|
/** Configuration options for OAuth */
|
|
57
66
|
export interface McpOAuthConfig {
|
|
58
67
|
grantType?: "authorization_code" | "client_credentials"
|
|
59
68
|
clientId?: string
|
|
60
69
|
clientSecret?: string
|
|
61
70
|
scope?: string
|
|
71
|
+
redirectUri?: string
|
|
72
|
+
clientName?: string
|
|
73
|
+
clientUri?: string
|
|
62
74
|
}
|
|
63
75
|
|
|
64
76
|
/** Callbacks for OAuth flow interactions */
|
|
@@ -71,12 +83,18 @@ export interface McpOAuthCallbacks {
|
|
|
71
83
|
* Implements the OAuthClientProvider interface from the MCP SDK.
|
|
72
84
|
*/
|
|
73
85
|
export class McpOAuthProvider implements OAuthClientProvider {
|
|
86
|
+
private readonly redirectUrlSnapshot: string | undefined
|
|
87
|
+
|
|
74
88
|
constructor(
|
|
75
89
|
private serverName: string,
|
|
76
90
|
private serverUrl: string,
|
|
77
91
|
private config: McpOAuthConfig,
|
|
78
92
|
private callbacks: McpOAuthCallbacks,
|
|
79
|
-
) {
|
|
93
|
+
) {
|
|
94
|
+
this.redirectUrlSnapshot = config.grantType === "client_credentials"
|
|
95
|
+
? undefined
|
|
96
|
+
: config.redirectUri ?? `http://localhost:${getOAuthCallbackPort()}${getOAuthCallbackPath()}`
|
|
97
|
+
}
|
|
80
98
|
|
|
81
99
|
private get usesClientCredentials(): boolean {
|
|
82
100
|
return this.config.grantType === "client_credentials"
|
|
@@ -87,8 +105,7 @@ export class McpOAuthProvider implements OAuthClientProvider {
|
|
|
87
105
|
* This must match the redirect_uri in client metadata.
|
|
88
106
|
*/
|
|
89
107
|
get redirectUrl(): string | undefined {
|
|
90
|
-
|
|
91
|
-
return `http://localhost:${getOAuthCallbackPort()}${OAUTH_CALLBACK_PATH}`
|
|
108
|
+
return this.redirectUrlSnapshot
|
|
92
109
|
}
|
|
93
110
|
|
|
94
111
|
/**
|
|
@@ -98,7 +115,8 @@ export class McpOAuthProvider implements OAuthClientProvider {
|
|
|
98
115
|
get clientMetadata(): OAuthClientMetadata {
|
|
99
116
|
if (this.usesClientCredentials) {
|
|
100
117
|
return {
|
|
101
|
-
client_name: "Pi Coding Agent",
|
|
118
|
+
client_name: this.config.clientName ?? "Pi Coding Agent",
|
|
119
|
+
client_uri: this.config.clientUri ?? "https://github.com/nicobailon/pi-mcp-adapter",
|
|
102
120
|
redirect_uris: [],
|
|
103
121
|
grant_types: ["client_credentials"],
|
|
104
122
|
token_endpoint_auth_method: this.config.clientSecret ? "client_secret_post" : "none",
|
|
@@ -112,8 +130,8 @@ export class McpOAuthProvider implements OAuthClientProvider {
|
|
|
112
130
|
|
|
113
131
|
return {
|
|
114
132
|
redirect_uris: [redirectUrl],
|
|
115
|
-
client_name: "Pi Coding Agent",
|
|
116
|
-
client_uri: "https://github.com/nicobailon/pi-mcp-adapter",
|
|
133
|
+
client_name: this.config.clientName ?? "Pi Coding Agent",
|
|
134
|
+
client_uri: this.config.clientUri ?? "https://github.com/nicobailon/pi-mcp-adapter",
|
|
117
135
|
grant_types: ["authorization_code", "refresh_token"],
|
|
118
136
|
response_types: ["code"],
|
|
119
137
|
token_endpoint_auth_method: this.config.clientSecret ? "client_secret_post" : "none",
|
|
@@ -155,11 +173,13 @@ export class McpOAuthProvider implements OAuthClientProvider {
|
|
|
155
173
|
* Save client information from dynamic registration.
|
|
156
174
|
*/
|
|
157
175
|
async saveClientInformation(info: OAuthClientInformationFull): Promise<void> {
|
|
176
|
+
const redirectUris = info.redirect_uris ?? (this.redirectUrl ? [this.redirectUrl] : undefined)
|
|
158
177
|
const clientInfo: StoredClientInfo = {
|
|
159
178
|
clientId: info.client_id,
|
|
160
179
|
clientSecret: info.client_secret,
|
|
161
180
|
clientIdIssuedAt: info.client_id_issued_at,
|
|
162
181
|
clientSecretExpiresAt: info.client_secret_expires_at,
|
|
182
|
+
redirectUris,
|
|
163
183
|
}
|
|
164
184
|
updateClientInfo(this.serverName, clientInfo, this.serverUrl)
|
|
165
185
|
}
|
|
@@ -299,4 +319,4 @@ export class McpOAuthProvider implements OAuthClientProvider {
|
|
|
299
319
|
}
|
|
300
320
|
}
|
|
301
321
|
|
|
302
|
-
export { DEFAULT_OAUTH_CALLBACK_PORT,
|
|
322
|
+
export { DEFAULT_OAUTH_CALLBACK_PORT, DEFAULT_OAUTH_CALLBACK_PATH }
|
package/package.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "pi-mcp-adapter",
|
|
3
|
-
"version": "2.
|
|
3
|
+
"version": "2.9.0",
|
|
4
4
|
"description": "MCP (Model Context Protocol) adapter extension for Pi coding agent",
|
|
5
5
|
"type": "module",
|
|
6
6
|
"license": "MIT",
|
|
@@ -54,6 +54,7 @@
|
|
|
54
54
|
"config.ts",
|
|
55
55
|
"server-manager.ts",
|
|
56
56
|
"sampling-handler.ts",
|
|
57
|
+
"elicitation-handler.ts",
|
|
57
58
|
"tool-registrar.ts",
|
|
58
59
|
"tool-result-renderer.ts",
|
|
59
60
|
"resource-tools.ts",
|
package/server-manager.ts
CHANGED
|
@@ -15,8 +15,9 @@ import { serverStreamResultPatchNotificationSchema } from "./types.ts";
|
|
|
15
15
|
import { resolveNpxBinary } from "./npx-resolver.ts";
|
|
16
16
|
import { logger } from "./logger.ts";
|
|
17
17
|
import { McpOAuthProvider } from "./mcp-oauth-provider.ts";
|
|
18
|
-
import { supportsOAuth } from "./mcp-auth-flow.ts";
|
|
18
|
+
import { extractOAuthConfig, supportsOAuth } from "./mcp-auth-flow.ts";
|
|
19
19
|
import { registerSamplingHandler, type ServerSamplingConfig } from "./sampling-handler.ts";
|
|
20
|
+
import { registerElicitationHandler, type ServerElicitationConfig } from "./elicitation-handler.ts";
|
|
20
21
|
import { interpolateEnvRecord, resolveBearerToken, resolveConfigPath } from "./utils.ts";
|
|
21
22
|
|
|
22
23
|
interface ServerConnection {
|
|
@@ -37,10 +38,15 @@ export class McpServerManager {
|
|
|
37
38
|
private connectPromises = new Map<string, Promise<ServerConnection>>();
|
|
38
39
|
private uiStreamListeners = new Map<string, UiStreamListener>();
|
|
39
40
|
private samplingConfig: ServerSamplingConfig | undefined;
|
|
41
|
+
private elicitationConfig: ServerElicitationConfig | undefined;
|
|
40
42
|
|
|
41
43
|
setSamplingConfig(config: ServerSamplingConfig | undefined): void {
|
|
42
44
|
this.samplingConfig = config;
|
|
43
45
|
}
|
|
46
|
+
|
|
47
|
+
setElicitationConfig(config: ServerElicitationConfig | undefined): void {
|
|
48
|
+
this.elicitationConfig = config;
|
|
49
|
+
}
|
|
44
50
|
|
|
45
51
|
async connect(name: string, definition: ServerDefinition): Promise<ServerConnection> {
|
|
46
52
|
// Dedupe concurrent connection attempts
|
|
@@ -148,14 +154,32 @@ export class McpServerManager {
|
|
|
148
154
|
}
|
|
149
155
|
}
|
|
150
156
|
|
|
157
|
+
private buildClientCapabilities() {
|
|
158
|
+
return {
|
|
159
|
+
...(this.samplingConfig ? { sampling: {} } : {}),
|
|
160
|
+
...(this.elicitationConfig
|
|
161
|
+
? {
|
|
162
|
+
elicitation: {
|
|
163
|
+
form: { applyDefaults: true },
|
|
164
|
+
url: {},
|
|
165
|
+
},
|
|
166
|
+
}
|
|
167
|
+
: {}),
|
|
168
|
+
};
|
|
169
|
+
}
|
|
170
|
+
|
|
151
171
|
private createClient(serverName: string): Client {
|
|
172
|
+
const capabilities = this.buildClientCapabilities();
|
|
152
173
|
const client = new Client(
|
|
153
174
|
{ name: `pi-mcp-${serverName}`, version: "1.0.0" },
|
|
154
|
-
|
|
175
|
+
Object.keys(capabilities).length > 0 ? { capabilities } : undefined,
|
|
155
176
|
);
|
|
156
177
|
if (this.samplingConfig) {
|
|
157
178
|
registerSamplingHandler(client, { ...this.samplingConfig, serverName });
|
|
158
179
|
}
|
|
180
|
+
if (this.elicitationConfig) {
|
|
181
|
+
registerElicitationHandler(client, { ...this.elicitationConfig, serverName });
|
|
182
|
+
}
|
|
159
183
|
return client;
|
|
160
184
|
}
|
|
161
185
|
|
|
@@ -182,13 +206,7 @@ export class McpServerManager {
|
|
|
182
206
|
// For OAuth servers, create an auth provider
|
|
183
207
|
let authProvider: McpOAuthProvider | undefined;
|
|
184
208
|
if (supportsOAuth(definition)) {
|
|
185
|
-
|
|
186
|
-
const oauthConfig = definition.oauth === false ? {} : {
|
|
187
|
-
grantType: definition.oauth?.grantType,
|
|
188
|
-
clientId: definition.oauth?.clientId,
|
|
189
|
-
clientSecret: definition.oauth?.clientSecret,
|
|
190
|
-
scope: definition.oauth?.scope,
|
|
191
|
-
};
|
|
209
|
+
const oauthConfig = extractOAuthConfig(definition);
|
|
192
210
|
authProvider = new McpOAuthProvider(
|
|
193
211
|
serverName,
|
|
194
212
|
definition.url!,
|
package/types.ts
CHANGED
|
@@ -272,6 +272,12 @@ export interface OAuthConfig {
|
|
|
272
272
|
clientSecret?: string;
|
|
273
273
|
/** Requested OAuth scopes */
|
|
274
274
|
scope?: string;
|
|
275
|
+
/** Exact authorization-code redirect URI for pre-registered clients */
|
|
276
|
+
redirectUri?: string;
|
|
277
|
+
/** Client display name for dynamic registration */
|
|
278
|
+
clientName?: string;
|
|
279
|
+
/** Client homepage URI for dynamic registration */
|
|
280
|
+
clientUri?: string;
|
|
275
281
|
}
|
|
276
282
|
|
|
277
283
|
// Server configuration
|
|
@@ -320,6 +326,8 @@ export interface McpSettings {
|
|
|
320
326
|
autoAuth?: boolean;
|
|
321
327
|
sampling?: boolean;
|
|
322
328
|
samplingAutoApprove?: boolean;
|
|
329
|
+
elicitation?: boolean;
|
|
330
|
+
elicitationAutoOpenUrls?: boolean;
|
|
323
331
|
/**
|
|
324
332
|
* Message returned in tool results when a server needs (re-)authentication.
|
|
325
333
|
* "${server}" is substituted with the server name. Defaults to a TUI
|