pi-mcp-adapter 2.3.5 → 2.4.1

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/direct-tools.ts CHANGED
@@ -7,11 +7,50 @@ import { isServerCacheValid } from "./metadata-cache.js";
7
7
  import { formatSchema } from "./tool-metadata.js";
8
8
  import { transformMcpContent } from "./tool-registrar.js";
9
9
  import { maybeStartUiSession, type UiSessionRuntime } from "./ui-session.js";
10
- import { formatToolName } from "./types.js";
10
+ import { formatToolName, isToolExcluded } from "./types.js";
11
11
  import { resourceNameToToolName } from "./resource-tools.js";
12
+ import { authenticate, supportsOAuth } from "./mcp-auth-flow.js";
12
13
 
13
14
  const BUILTIN_NAMES = new Set(["read", "bash", "edit", "write", "grep", "find", "ls", "mcp"]);
14
15
 
16
+ type DirectAutoAuthResult =
17
+ | { status: "skipped" }
18
+ | { status: "success" }
19
+ | { status: "failed"; message: string };
20
+
21
+ async function attemptDirectAutoAuth(
22
+ state: McpExtensionState,
23
+ serverName: string,
24
+ ): Promise<DirectAutoAuthResult> {
25
+ if (state.config.settings?.autoAuth !== true) {
26
+ return { status: "skipped" };
27
+ }
28
+
29
+ const definition = state.config.mcpServers[serverName];
30
+ if (!definition || !supportsOAuth(definition) || !definition.url) {
31
+ return { status: "skipped" };
32
+ }
33
+
34
+ const grantType = definition.oauth?.grantType ?? "authorization_code";
35
+ if (!state.ui && grantType !== "client_credentials") {
36
+ return {
37
+ status: "failed",
38
+ message: `MCP server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} in an interactive session.`,
39
+ };
40
+ }
41
+
42
+ try {
43
+ await authenticate(serverName, definition.url, definition);
44
+ return { status: "success" };
45
+ } catch (error) {
46
+ const message = error instanceof Error ? error.message : String(error);
47
+ return {
48
+ status: "failed",
49
+ message: `OAuth authentication failed for "${serverName}": ${message}. Run /mcp-auth ${serverName} first.`,
50
+ };
51
+ }
52
+ }
53
+
15
54
  export function resolveDirectTools(
16
55
  config: McpConfig,
17
56
  cache: MetadataCache | null,
@@ -68,6 +107,7 @@ export function resolveDirectTools(
68
107
 
69
108
  for (const tool of serverCache.tools ?? []) {
70
109
  if (toolFilter !== true && !toolFilter.includes(tool.name)) continue;
110
+ if (isToolExcluded(tool.name, serverName, prefix, definition.excludeTools)) continue;
71
111
  const prefixedName = formatToolName(tool.name, serverName, prefix);
72
112
  if (BUILTIN_NAMES.has(prefixedName)) {
73
113
  console.warn(`MCP: skipping direct tool "${prefixedName}" (collides with builtin)`);
@@ -93,6 +133,7 @@ export function resolveDirectTools(
93
133
  for (const resource of serverCache.resources ?? []) {
94
134
  const baseName = `get_${resourceNameToToolName(resource.name)}`;
95
135
  if (toolFilter !== true && !toolFilter.includes(baseName)) continue;
136
+ if (isToolExcluded(baseName, serverName, prefix, definition.excludeTools)) continue;
96
137
  const prefixedName = formatToolName(baseName, serverName, prefix);
97
138
  if (BUILTIN_NAMES.has(prefixedName)) {
98
139
  console.warn(`MCP: skipping direct resource tool "${prefixedName}" (collides with builtin)`);
@@ -117,11 +158,35 @@ export function resolveDirectTools(
117
158
  return specs;
118
159
  }
119
160
 
161
+ export function getMissingConfiguredDirectToolServers(
162
+ config: McpConfig,
163
+ cache: MetadataCache | null,
164
+ ): string[] {
165
+ const missing: string[] = [];
166
+ const globalDirect = config.settings?.directTools;
167
+
168
+ for (const [serverName, definition] of Object.entries(config.mcpServers)) {
169
+ const hasDirectTools = definition.directTools !== undefined
170
+ ? !!definition.directTools
171
+ : !!globalDirect;
172
+
173
+ if (!hasDirectTools) continue;
174
+
175
+ const serverCache = cache?.servers?.[serverName];
176
+ if (!serverCache || !isServerCacheValid(serverCache, definition)) {
177
+ missing.push(serverName);
178
+ }
179
+ }
180
+
181
+ return missing;
182
+ }
183
+
120
184
  export function buildProxyDescription(
121
185
  config: McpConfig,
122
186
  cache: MetadataCache | null,
123
187
  directSpecs: DirectToolSpec[],
124
188
  ): string {
189
+ const prefix = config.settings?.toolPrefix ?? "server";
125
190
  let desc = `MCP gateway - connect to MCP servers and call their tools.\n`;
126
191
 
127
192
  const directByServer = new Map<string, number>();
@@ -139,8 +204,15 @@ export function buildProxyDescription(
139
204
  for (const serverName of Object.keys(config.mcpServers)) {
140
205
  const entry = cache?.servers?.[serverName];
141
206
  const definition = config.mcpServers[serverName];
142
- const toolCount = entry?.tools?.length ?? 0;
143
- const resourceCount = definition?.exposeResources !== false ? (entry?.resources?.length ?? 0) : 0;
207
+ const toolCount = (entry?.tools ?? []).filter(
208
+ (tool) => !isToolExcluded(tool.name, serverName, prefix, definition.excludeTools),
209
+ ).length;
210
+ const resourceCount = definition?.exposeResources !== false
211
+ ? (entry?.resources ?? []).filter((resource) => {
212
+ const baseName = `get_${resourceNameToToolName(resource.name)}`;
213
+ return !isToolExcluded(baseName, serverName, prefix, definition.excludeTools);
214
+ }).length
215
+ : 0;
144
216
  const totalItems = toolCount + resourceCount;
145
217
  if (totalItems === 0) continue;
146
218
  const directCount = directByServer.get(serverName) ?? 0;
@@ -196,13 +268,32 @@ export function createDirectToolExecutor(
196
268
  };
197
269
  }
198
270
 
199
- const connected = await lazyConnect(state, spec.serverName);
271
+ let connected = await lazyConnect(state, spec.serverName);
272
+ let autoAuthAttempted = false;
273
+
274
+ if (!connected && state.manager.getConnection(spec.serverName)?.status === "needs-auth") {
275
+ autoAuthAttempted = true;
276
+ const autoAuth = await attemptDirectAutoAuth(state, spec.serverName);
277
+ if (autoAuth.status === "failed") {
278
+ return {
279
+ content: [{ type: "text" as const, text: autoAuth.message }],
280
+ details: { error: "auth_required", server: spec.serverName, message: autoAuth.message },
281
+ };
282
+ }
283
+ if (autoAuth.status === "success") {
284
+ await state.manager.close(spec.serverName);
285
+ state.failureTracker.delete(spec.serverName);
286
+ connected = await lazyConnect(state, spec.serverName);
287
+ }
288
+ }
289
+
200
290
  if (!connected) {
201
291
  const authConnection = state.manager.getConnection(spec.serverName);
202
292
  if (authConnection?.status === "needs-auth") {
293
+ const message = `MCP server "${spec.serverName}" requires OAuth authentication. Run /mcp-auth ${spec.serverName} first.`;
203
294
  return {
204
- content: [{ type: "text" as const, text: `MCP server "${spec.serverName}" requires OAuth authentication. Run /mcp-auth ${spec.serverName} first.` }],
205
- details: { error: "auth_required", server: spec.serverName },
295
+ content: [{ type: "text" as const, text: message }],
296
+ details: { error: "auth_required", server: spec.serverName, message, autoAuthAttempted },
206
297
  };
207
298
  }
208
299
  const failedAgo = getFailureAgeSeconds(state, spec.serverName);
package/index.ts CHANGED
@@ -1,9 +1,9 @@
1
1
  import type { ExtensionAPI, ToolInfo } from "@mariozechner/pi-coding-agent";
2
2
  import type { McpExtensionState } from "./state.js";
3
3
  import { Type } from "@sinclair/typebox";
4
- import { showStatus, showTools, reconnectServers, authenticateServer, openMcpPanel } from "./commands.js";
4
+ import { showStatus, showTools, reconnectServers, authenticateServer, openMcpPanel, openMcpSetup } from "./commands.js";
5
5
  import { loadMcpConfig } from "./config.js";
6
- import { buildProxyDescription, createDirectToolExecutor, resolveDirectTools } from "./direct-tools.js";
6
+ import { buildProxyDescription, createDirectToolExecutor, getMissingConfiguredDirectToolServers, resolveDirectTools } from "./direct-tools.js";
7
7
  import { flushMetadataCache, initializeMcp, updateStatusBar } from "./init.js";
8
8
  import { loadMetadataCache } from "./metadata-cache.js";
9
9
  import { executeCall, executeConnect, executeDescribe, executeList, executeSearch, executeStatus, executeUiMessages } from "./proxy-modes.js";
@@ -59,6 +59,11 @@ export default function mcpAdapter(pi: ExtensionAPI) {
59
59
  prefix,
60
60
  envRaw?.split(",").map(s => s.trim()).filter(Boolean),
61
61
  );
62
+ const missingConfiguredDirectToolServers = getMissingConfiguredDirectToolServers(earlyConfig, earlyCache);
63
+ const shouldRegisterProxyTool =
64
+ earlyConfig.settings?.disableProxyTool !== true
65
+ || directSpecs.length === 0
66
+ || missingConfiguredDirectToolServers.length > 0;
62
67
 
63
68
  for (const spec of directSpecs) {
64
69
  pi.registerTool({
@@ -173,11 +178,23 @@ export default function mcpAdapter(pi: ExtensionAPI) {
173
178
  case "tools":
174
179
  await showTools(state, ctx);
175
180
  break;
181
+ case "setup": {
182
+ const result = await openMcpSetup(state, pi, ctx, earlyConfigPath, "setup");
183
+ if (result?.configChanged) {
184
+ await ctx.reload();
185
+ return;
186
+ }
187
+ break;
188
+ }
176
189
  case "status":
177
190
  case "":
178
191
  default:
179
192
  if (ctx.hasUI) {
180
- await openMcpPanel(state, pi, ctx, earlyConfigPath);
193
+ const result = await openMcpPanel(state, pi, ctx, earlyConfigPath);
194
+ if (result?.configChanged) {
195
+ await ctx.reload();
196
+ return;
197
+ }
181
198
  } else {
182
199
  await showStatus(state, ctx);
183
200
  }
@@ -213,86 +230,88 @@ export default function mcpAdapter(pi: ExtensionAPI) {
213
230
  },
214
231
  });
215
232
 
216
- pi.registerTool({
217
- name: "mcp",
218
- label: "MCP",
219
- description: buildProxyDescription(earlyConfig, earlyCache, directSpecs),
220
- promptSnippet: "MCP gateway - connect to MCP servers and call their tools",
221
- parameters: Type.Object({
222
- tool: Type.Optional(Type.String({ description: "Tool name to call (e.g., 'xcodebuild_list_sims')" })),
223
- args: Type.Optional(Type.String({ description: "Arguments as JSON string (e.g., '{\"key\": \"value\"}')" })),
224
- connect: Type.Optional(Type.String({ description: "Server name to connect (lazy connect + metadata refresh)" })),
225
- describe: Type.Optional(Type.String({ description: "Tool name to describe (shows parameters)" })),
226
- search: Type.Optional(Type.String({ description: "Search tools by name/description" })),
227
- regex: Type.Optional(Type.Boolean({ description: "Treat search as regex (default: substring match)" })),
228
- includeSchemas: Type.Optional(Type.Boolean({ description: "Include parameter schemas in search results (default: true)" })),
229
- server: Type.Optional(Type.String({ description: "Filter to specific server (also disambiguates tool calls)" })),
230
- action: Type.Optional(Type.String({ description: "Action: 'ui-messages' to retrieve prompts/intents from UI sessions" })),
231
- }),
232
- async execute(_toolCallId, params: {
233
- tool?: string;
234
- args?: string;
235
- connect?: string;
236
- describe?: string;
237
- search?: string;
238
- regex?: boolean;
239
- includeSchemas?: boolean;
240
- server?: string;
241
- action?: string;
242
- }, _signal, _onUpdate, _ctx) {
243
- let parsedArgs: Record<string, unknown> | undefined;
244
- if (params.args) {
245
- try {
246
- parsedArgs = JSON.parse(params.args);
247
- if (typeof parsedArgs !== "object" || parsedArgs === null || Array.isArray(parsedArgs)) {
248
- const gotType = Array.isArray(parsedArgs) ? "array" : parsedArgs === null ? "null" : typeof parsedArgs;
249
- throw new Error(`Invalid args: expected a JSON object, got ${gotType}`);
250
- }
251
- } catch (error) {
252
- if (error instanceof SyntaxError) {
253
- throw new Error(`Invalid args JSON: ${error.message}`, { cause: error });
233
+ if (shouldRegisterProxyTool) {
234
+ pi.registerTool({
235
+ name: "mcp",
236
+ label: "MCP",
237
+ description: buildProxyDescription(earlyConfig, earlyCache, directSpecs),
238
+ promptSnippet: "MCP gateway - connect to MCP servers and call their tools",
239
+ parameters: Type.Object({
240
+ tool: Type.Optional(Type.String({ description: "Tool name to call (e.g., 'xcodebuild_list_sims')" })),
241
+ args: Type.Optional(Type.String({ description: "Arguments as JSON string (e.g., '{\"key\": \"value\"}')" })),
242
+ connect: Type.Optional(Type.String({ description: "Server name to connect (lazy connect + metadata refresh)" })),
243
+ describe: Type.Optional(Type.String({ description: "Tool name to describe (shows parameters)" })),
244
+ search: Type.Optional(Type.String({ description: "Search tools by name/description" })),
245
+ regex: Type.Optional(Type.Boolean({ description: "Treat search as regex (default: substring match)" })),
246
+ includeSchemas: Type.Optional(Type.Boolean({ description: "Include parameter schemas in search results (default: true)" })),
247
+ server: Type.Optional(Type.String({ description: "Filter to specific server (also disambiguates tool calls)" })),
248
+ action: Type.Optional(Type.String({ description: "Action: 'ui-messages' to retrieve prompts/intents from UI sessions" })),
249
+ }),
250
+ async execute(_toolCallId, params: {
251
+ tool?: string;
252
+ args?: string;
253
+ connect?: string;
254
+ describe?: string;
255
+ search?: string;
256
+ regex?: boolean;
257
+ includeSchemas?: boolean;
258
+ server?: string;
259
+ action?: string;
260
+ }, _signal, _onUpdate, _ctx) {
261
+ let parsedArgs: Record<string, unknown> | undefined;
262
+ if (params.args) {
263
+ try {
264
+ parsedArgs = JSON.parse(params.args);
265
+ if (typeof parsedArgs !== "object" || parsedArgs === null || Array.isArray(parsedArgs)) {
266
+ const gotType = Array.isArray(parsedArgs) ? "array" : parsedArgs === null ? "null" : typeof parsedArgs;
267
+ throw new Error(`Invalid args: expected a JSON object, got ${gotType}`);
268
+ }
269
+ } catch (error) {
270
+ if (error instanceof SyntaxError) {
271
+ throw new Error(`Invalid args JSON: ${error.message}`, { cause: error });
272
+ }
273
+ throw error;
254
274
  }
255
- throw error;
256
275
  }
257
- }
258
276
 
259
- if (!state && initPromise) {
260
- try {
261
- state = await initPromise;
262
- } catch (error) {
263
- const message = error instanceof Error ? error.message : String(error);
277
+ if (!state && initPromise) {
278
+ try {
279
+ state = await initPromise;
280
+ } catch (error) {
281
+ const message = error instanceof Error ? error.message : String(error);
282
+ return {
283
+ content: [{ type: "text" as const, text: `MCP initialization failed: ${message}` }],
284
+ details: { error: "init_failed", message },
285
+ };
286
+ }
287
+ }
288
+ if (!state) {
264
289
  return {
265
- content: [{ type: "text" as const, text: `MCP initialization failed: ${message}` }],
266
- details: { error: "init_failed", message },
290
+ content: [{ type: "text" as const, text: "MCP not initialized" }],
291
+ details: { error: "not_initialized" },
267
292
  };
268
293
  }
269
- }
270
- if (!state) {
271
- return {
272
- content: [{ type: "text" as const, text: "MCP not initialized" }],
273
- details: { error: "not_initialized" },
274
- };
275
- }
276
294
 
277
- if (params.action === "ui-messages") {
278
- return executeUiMessages(state);
279
- }
280
- if (params.tool) {
281
- return executeCall(state, params.tool, parsedArgs, params.server);
282
- }
283
- if (params.connect) {
284
- return executeConnect(state, params.connect);
285
- }
286
- if (params.describe) {
287
- return executeDescribe(state, params.describe);
288
- }
289
- if (params.search) {
290
- return executeSearch(state, params.search, params.regex, params.server, params.includeSchemas, getPiTools);
291
- }
292
- if (params.server) {
293
- return executeList(state, params.server);
294
- }
295
- return executeStatus(state);
296
- },
297
- });
295
+ if (params.action === "ui-messages") {
296
+ return executeUiMessages(state);
297
+ }
298
+ if (params.tool) {
299
+ return executeCall(state, params.tool, parsedArgs, params.server);
300
+ }
301
+ if (params.connect) {
302
+ return executeConnect(state, params.connect);
303
+ }
304
+ if (params.describe) {
305
+ return executeDescribe(state, params.describe);
306
+ }
307
+ if (params.search) {
308
+ return executeSearch(state, params.search, params.regex, params.server, params.includeSchemas, getPiTools);
309
+ }
310
+ if (params.server) {
311
+ return executeList(state, params.server);
312
+ }
313
+ return executeStatus(state);
314
+ },
315
+ });
316
+ }
298
317
  }
package/init.ts CHANGED
@@ -21,6 +21,7 @@ import { buildToolMetadata, totalToolCount } from "./tool-metadata.js";
21
21
  import { UiResourceHandler } from "./ui-resource-handler.js";
22
22
  import { openUrl, parallelLimit } from "./utils.js";
23
23
  import { logger } from "./logger.js";
24
+ import { getMissingConfiguredDirectToolServers } from "./direct-tools.js";
24
25
 
25
26
  const FAILURE_BACKOFF_MS = 60 * 1000;
26
27
 
@@ -89,7 +90,7 @@ export async function initializeMcp(
89
90
  }
90
91
 
91
92
  if (cache?.servers?.[name] && isServerCacheValid(cache.servers[name], definition)) {
92
- const metadata = reconstructToolMetadata(name, cache.servers[name], prefix, definition.exposeResources);
93
+ const metadata = reconstructToolMetadata(name, cache.servers[name], prefix, definition);
93
94
  toolMetadata.set(name, metadata);
94
95
  }
95
96
  }
@@ -151,18 +152,8 @@ export async function initializeMcp(
151
152
 
152
153
  const envDirect = process.env.MCP_DIRECT_TOOLS;
153
154
  if (envDirect !== "__none__") {
154
- const missingCacheServers: string[] = [];
155
155
  const currentCache = loadMetadataCache();
156
- for (const [name, definition] of serverEntries) {
157
- const hasDirect = definition.directTools !== undefined
158
- ? !!definition.directTools
159
- : !!config.settings?.directTools;
160
- if (!hasDirect) continue;
161
- const entry = currentCache?.servers?.[name];
162
- if (!entry || !isServerCacheValid(entry, definition)) {
163
- missingCacheServers.push(name);
164
- }
165
- }
156
+ const missingCacheServers = getMissingConfiguredDirectToolServers(config, currentCache);
166
157
 
167
158
  if (missingCacheServers.length > 0) {
168
159
  const bootstrapResults = await parallelLimit(
package/mcp-auth-flow.ts CHANGED
@@ -36,6 +36,9 @@ export type AuthStatus = "authenticated" | "expired" | "not_authenticated"
36
36
  // Track pending transports for auth completion
37
37
  const pendingTransports = new Map<string, StreamableHTTPClientTransport>()
38
38
 
39
+ // Deduplicate concurrent authenticate() calls per server.
40
+ const pendingAuthentications = new Map<string, Promise<AuthStatus>>()
41
+
39
42
  /**
40
43
  * Generate a cryptographically secure random state parameter.
41
44
  */
@@ -182,56 +185,74 @@ export async function authenticate(
182
185
  serverUrl: string,
183
186
  definition?: ServerEntry,
184
187
  ): Promise<AuthStatus> {
185
- // Start auth flow
186
- const { authorizationUrl } = await startAuth(serverName, serverUrl, definition)
187
-
188
- // If no auth URL needed, already authenticated
189
- if (!authorizationUrl) {
190
- return "authenticated"
188
+ const inFlight = pendingAuthentications.get(serverName)
189
+ if (inFlight) {
190
+ return inFlight
191
191
  }
192
192
 
193
- // Get the state that was already generated and stored in startAuth()
194
- const oauthState = await getOAuthState(serverName)
195
- if (!oauthState) {
196
- throw new Error("OAuth state not found - this should not happen")
197
- }
193
+ const operation = (async (): Promise<AuthStatus> => {
194
+ // Start auth flow
195
+ const { authorizationUrl } = await startAuth(serverName, serverUrl, definition)
198
196
 
199
- // Register the callback BEFORE opening the browser
200
- const callbackPromise = waitForCallback(oauthState)
197
+ // If no auth URL needed, already authenticated
198
+ if (!authorizationUrl) {
199
+ return "authenticated"
200
+ }
201
201
 
202
- // Open browser
203
- console.log(`MCP Auth: Opening browser for ${serverName}`)
204
- try {
205
- await open(authorizationUrl)
206
- } catch (error) {
207
- console.warn(`MCP Auth: Failed to open browser for ${serverName}`, { error })
208
- throw new Error(
209
- `Could not open browser. Please open this URL manually: ${authorizationUrl}`
210
- )
211
- }
202
+ // Get the state that was already generated and stored in startAuth()
203
+ const oauthState = await getOAuthState(serverName)
204
+ if (!oauthState) {
205
+ throw new Error("OAuth state not found - this should not happen")
206
+ }
212
207
 
213
- try {
214
- // Wait for callback
215
- const code = await callbackPromise
208
+ // Register the callback BEFORE opening the browser
209
+ const callbackPromise = waitForCallback(oauthState)
210
+
211
+ // Open browser
212
+ console.log(`MCP Auth: Opening browser for ${serverName}`)
213
+ try {
214
+ await open(authorizationUrl)
215
+ } catch (error) {
216
+ console.warn(`MCP Auth: Failed to open browser for ${serverName}`, { error })
217
+ throw new Error(
218
+ `Could not open browser. Please open this URL manually: ${authorizationUrl}`,
219
+ { cause: error },
220
+ )
221
+ }
216
222
 
217
- // Validate state
218
- const storedState = await getOAuthState(serverName)
219
- if (storedState !== oauthState) {
223
+ try {
224
+ // Wait for callback
225
+ const code = await callbackPromise
226
+
227
+ // Validate state
228
+ const storedState = await getOAuthState(serverName)
229
+ if (storedState !== oauthState) {
230
+ await clearOAuthState(serverName)
231
+ throw new Error("OAuth state mismatch - potential CSRF attack")
232
+ }
220
233
  await clearOAuthState(serverName)
221
- throw new Error("OAuth state mismatch - potential CSRF attack")
234
+
235
+ // Complete the auth
236
+ return await completeAuth(serverName, code)
237
+ } catch (error) {
238
+ cancelPendingCallback(oauthState)
239
+ const pendingTransport = pendingTransports.get(serverName)
240
+ if (pendingTransport) {
241
+ pendingTransports.delete(serverName)
242
+ await pendingTransport.close().catch(() => {})
243
+ }
244
+ throw error
222
245
  }
223
- await clearOAuthState(serverName)
246
+ })()
224
247
 
225
- // Complete the auth
226
- return await completeAuth(serverName, code)
227
- } catch (error) {
228
- cancelPendingCallback(oauthState)
229
- const pendingTransport = pendingTransports.get(serverName)
230
- if (pendingTransport) {
231
- pendingTransports.delete(serverName)
232
- await pendingTransport.close().catch(() => {})
248
+ pendingAuthentications.set(serverName, operation)
249
+
250
+ try {
251
+ return await operation
252
+ } finally {
253
+ if (pendingAuthentications.get(serverName) === operation) {
254
+ pendingAuthentications.delete(serverName)
233
255
  }
234
- throw error
235
256
  }
236
257
  }
237
258