pi-mcp-adapter 2.4.2 → 2.5.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 CHANGED
@@ -7,6 +7,18 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
7
7
 
8
8
  ## [Unreleased]
9
9
 
10
+ ## [2.5.0] - 2026-04-24
11
+
12
+ ### Added
13
+ - Added MCP `sampling/createMessage` support with conservative human approval by default and opt-in `settings.samplingAutoApprove` for non-interactive flows.
14
+ - Added configured Vitest coverage for OAuth provider authorization fallback behavior.
15
+ - Added `test:oauth-provider` for running the root OAuth provider node test with the required TypeScript loader.
16
+
17
+ ### Fixed
18
+ - Applied `settings.authRequiredMessage` to proxy and direct-tool auth-required paths, including non-UI `autoAuth` failures.
19
+ - Fixed `/mcp-auth <server>` reporting success for expired stored OAuth tokens without forcing the SDK refresh/re-auth flow.
20
+ - Kept `mcp` search focused on MCP tools and added a direct-call hint when native Pi tools are accidentally routed through the proxy.
21
+
10
22
  ## [2.4.2] - 2026-04-22
11
23
 
12
24
  ### Fixed
package/README.md CHANGED
@@ -158,12 +158,14 @@ Pi-specific files are the write targets for imported or shared global servers wh
158
158
  | `directTools` | Global default for all servers (default: false). Per-server overrides this. |
159
159
  | `disableProxyTool` | Hide the `mcp` proxy tool once configured direct tools are fully available from cache. |
160
160
  | `autoAuth` | Auto-run OAuth on `connect`/tool calls when a server needs auth, then retry once (default: false). |
161
+ | `sampling` | Allow MCP servers to request LLM sampling through Pi's current/default model (default: true when UI approval is available). |
162
+ | `samplingAutoApprove` | Skip sampling confirmation prompts. Required for sampling in non-UI sessions (default: false). |
161
163
 
162
164
  Per-server `idleTimeout` overrides the global setting.
163
165
 
164
166
  ### Direct Tools
165
167
 
166
- By default, all MCP tools are accessed through the single `mcp` proxy tool. This keeps context small but means the LLM has to discover tools via 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.
168
+ 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.
167
169
 
168
170
  Per-server:
169
171
 
@@ -360,3 +362,4 @@ If `settings.autoAuth` is `true`, `mcp({ connect: ... })`, `mcp({ tool: ... })`,
360
362
  ## Limitations
361
363
 
362
364
  - Cross-session server sharing not yet implemented (each Pi session runs its own server processes)
365
+ - MCP sampling support is text-only; context inclusion, tools, stop sequences, audio, and image content are rejected with explicit errors.
package/direct-tools.ts CHANGED
@@ -10,6 +10,7 @@ import { maybeStartUiSession, type UiSessionRuntime } from "./ui-session.js";
10
10
  import { formatToolName, isToolExcluded } from "./types.js";
11
11
  import { resourceNameToToolName } from "./resource-tools.js";
12
12
  import { authenticate, supportsOAuth } from "./mcp-auth-flow.js";
13
+ import { formatAuthRequiredMessage } from "./utils.js";
13
14
 
14
15
  const BUILTIN_NAMES = new Set(["read", "bash", "edit", "write", "grep", "find", "ls", "mcp"]);
15
16
 
@@ -18,6 +19,22 @@ type DirectAutoAuthResult =
18
19
  | { status: "success" }
19
20
  | { status: "failed"; message: string };
20
21
 
22
+ function getDirectAuthRequiredMessage(
23
+ state: McpExtensionState,
24
+ serverName: string,
25
+ defaultMessage = `MCP server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} first.`,
26
+ ): string {
27
+ return formatAuthRequiredMessage(state.config, serverName, defaultMessage);
28
+ }
29
+
30
+ function getDirectAuthFailedMessage(state: McpExtensionState, serverName: string, message: string): string {
31
+ const customGuidance = state.config.settings?.authRequiredMessage;
32
+ if (customGuidance) {
33
+ return `OAuth authentication failed for "${serverName}": ${message}. ${getDirectAuthRequiredMessage(state, serverName)}`;
34
+ }
35
+ return `OAuth authentication failed for "${serverName}": ${message}. Run /mcp-auth ${serverName} first.`;
36
+ }
37
+
21
38
  async function attemptDirectAutoAuth(
22
39
  state: McpExtensionState,
23
40
  serverName: string,
@@ -35,7 +52,11 @@ async function attemptDirectAutoAuth(
35
52
  if (!state.ui && grantType !== "client_credentials") {
36
53
  return {
37
54
  status: "failed",
38
- message: `MCP server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} in an interactive session.`,
55
+ message: getDirectAuthRequiredMessage(
56
+ state,
57
+ serverName,
58
+ `MCP server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} in an interactive session.`,
59
+ ),
39
60
  };
40
61
  }
41
62
 
@@ -46,7 +67,7 @@ async function attemptDirectAutoAuth(
46
67
  const message = error instanceof Error ? error.message : String(error);
47
68
  return {
48
69
  status: "failed",
49
- message: `OAuth authentication failed for "${serverName}": ${message}. Run /mcp-auth ${serverName} first.`,
70
+ message: getDirectAuthFailedMessage(state, serverName, message),
50
71
  };
51
72
  }
52
73
  }
@@ -187,7 +208,7 @@ export function buildProxyDescription(
187
208
  directSpecs: DirectToolSpec[],
188
209
  ): string {
189
210
  const prefix = config.settings?.toolPrefix ?? "server";
190
- let desc = `MCP gateway - connect to MCP servers and call their tools.\n`;
211
+ let desc = `MCP gateway - connect to MCP servers and call their tools. Non-MCP Pi tools should be called directly, not through mcp.\n`;
191
212
 
192
213
  const directByServer = new Map<string, number>();
193
214
  for (const spec of directSpecs) {
@@ -229,7 +250,7 @@ export function buildProxyDescription(
229
250
  desc += `\nUsage:\n`;
230
251
  desc += ` mcp({ }) → Show server status\n`;
231
252
  desc += ` mcp({ server: "name" }) → List tools from server\n`;
232
- desc += ` mcp({ search: "query" }) → Search for tools (MCP + pi, space-separated words OR'd)\n`;
253
+ desc += ` mcp({ search: "query" }) → Search MCP tools by name/description\n`;
233
254
  desc += ` mcp({ describe: "tool_name" }) → Show tool details and parameters\n`;
234
255
  desc += ` mcp({ connect: "server-name" }) → Connect to a server and refresh metadata\n`;
235
256
  desc += ` mcp({ tool: "name", args: '{"key": "value"}' }) → Call a tool (args is JSON string)\n`;
@@ -290,7 +311,7 @@ export function createDirectToolExecutor(
290
311
  if (!connected) {
291
312
  const authConnection = state.manager.getConnection(spec.serverName);
292
313
  if (authConnection?.status === "needs-auth") {
293
- const message = `MCP server "${spec.serverName}" requires OAuth authentication. Run /mcp-auth ${spec.serverName} first.`;
314
+ const message = getDirectAuthRequiredMessage(state, spec.serverName);
294
315
  return {
295
316
  content: [{ type: "text" as const, text: message }],
296
317
  details: { error: "auth_required", server: spec.serverName, message, autoAuthAttempted },
package/index.ts CHANGED
@@ -296,7 +296,7 @@ export default function mcpAdapter(pi: ExtensionAPI) {
296
296
  return executeUiMessages(state);
297
297
  }
298
298
  if (params.tool) {
299
- return executeCall(state, params.tool, parsedArgs, params.server);
299
+ return executeCall(state, params.tool, parsedArgs, params.server, getPiTools);
300
300
  }
301
301
  if (params.connect) {
302
302
  return executeConnect(state, params.connect);
@@ -305,7 +305,7 @@ export default function mcpAdapter(pi: ExtensionAPI) {
305
305
  return executeDescribe(state, params.describe);
306
306
  }
307
307
  if (params.search) {
308
- return executeSearch(state, params.search, params.regex, params.server, params.includeSchemas, getPiTools);
308
+ return executeSearch(state, params.search, params.regex, params.server, params.includeSchemas);
309
309
  }
310
310
  if (params.server) {
311
311
  return executeList(state, params.server);
package/init.ts CHANGED
@@ -33,6 +33,16 @@ export async function initializeMcp(
33
33
  const config = loadMcpConfig(configPath);
34
34
 
35
35
  const manager = new McpServerManager();
36
+ const samplingAutoApprove = config.settings?.samplingAutoApprove === true;
37
+ if (config.settings?.sampling !== false && (ctx.hasUI || samplingAutoApprove)) {
38
+ manager.setSamplingConfig({
39
+ autoApprove: samplingAutoApprove,
40
+ ui: ctx.hasUI ? ctx.ui : undefined,
41
+ modelRegistry: ctx.modelRegistry,
42
+ getCurrentModel: () => ctx.model,
43
+ getSignal: () => ctx.signal,
44
+ });
45
+ }
36
46
  const lifecycle = new McpLifecycleManager(manager);
37
47
  const toolMetadata = new Map<string, ToolMetadata[]>();
38
48
  const failureTracker = new Map<string, number>();
package/mcp-auth-flow.ts CHANGED
@@ -2,14 +2,13 @@
2
2
  * MCP Auth Flow
3
3
  *
4
4
  * High-level OAuth flow management using the MCP SDK's built-in auth functions.
5
- * Follows the OpenCode pattern: let the SDK handle discovery internally via transport.
6
5
  */
7
6
 
8
7
  import {
8
+ auth as runSdkAuth,
9
9
  UnauthorizedError,
10
10
  } from "@modelcontextprotocol/sdk/client/auth.js"
11
11
  import { StreamableHTTPClientTransport } from "@modelcontextprotocol/sdk/client/streamableHttp.js"
12
- import { Client } from "@modelcontextprotocol/sdk/client/index.js"
13
12
  import open from "open"
14
13
  import { McpOAuthProvider, type McpOAuthConfig } from "./mcp-oauth-provider.js"
15
14
  import {
@@ -66,18 +65,13 @@ function extractOAuthConfig(definition: ServerEntry): McpOAuthConfig {
66
65
 
67
66
  /**
68
67
  * Start OAuth authentication flow for a server.
69
- * Returns the authorization URL that should be opened in a browser.
70
- *
71
- * This follows the OpenCode pattern:
72
- * 1. Create transport with auth provider
73
- * 2. Try to connect - SDK handles discovery internally
74
- * 3. If UnauthorizedError, capture the auth URL from onRedirect
68
+ * Returns the authorization URL when browser authorization is required.
75
69
  */
76
70
  export async function startAuth(
77
71
  serverName: string,
78
72
  serverUrl: string,
79
73
  definition?: ServerEntry
80
- ): Promise<{ authorizationUrl: string; transport: StreamableHTTPClientTransport }> {
74
+ ): Promise<{ authorizationUrl: string }> {
81
75
  const config = definition ? extractOAuthConfig(definition) : {}
82
76
 
83
77
  if (config.grantType === "client_credentials") {
@@ -86,33 +80,20 @@ export async function startAuth(
86
80
  throw new Error("Browser redirect is not used for client_credentials flow")
87
81
  },
88
82
  })
89
- const transport = new StreamableHTTPClientTransport(new URL(serverUrl), {
90
- authProvider,
91
- })
92
- const client = new Client({
93
- name: "pi-mcp",
94
- version: "3.0.0",
95
- })
96
-
97
- try {
98
- await client.connect(transport)
99
- return { authorizationUrl: "", transport }
100
- } finally {
101
- await client.close().catch(() => {})
102
- await transport.close().catch(() => {})
83
+ const result = await runSdkAuth(authProvider, { serverUrl })
84
+ if (result !== "AUTHORIZED") {
85
+ throw new UnauthorizedError("Failed to authorize")
103
86
  }
87
+ return { authorizationUrl: "" }
104
88
  }
105
89
 
106
90
  // Start the callback server.
107
91
  // Pre-registered OAuth clients require an exact redirect URI, so enforce strict port binding.
108
92
  await ensureCallbackServer({ strictPort: Boolean(config.clientId) })
109
93
 
110
- // Generate and store OAuth state BEFORE creating the provider
111
- // The SDK will call provider.state() to read this value
112
94
  const oauthState = generateState()
113
95
  await updateOAuthState(serverName, oauthState)
114
96
 
115
- // Create the auth provider
116
97
  let capturedUrl: URL | undefined
117
98
  const authProvider = new McpOAuthProvider(serverName, serverUrl, config, {
118
99
  onRedirect: async (url) => {
@@ -120,32 +101,22 @@ export async function startAuth(
120
101
  },
121
102
  })
122
103
 
123
- // Create transport with auth provider
124
- // The SDK handles OAuth discovery internally when connecting
125
- const transport = new StreamableHTTPClientTransport(new URL(serverUrl), {
126
- authProvider,
127
- })
128
- const client = new Client({
129
- name: "pi-mcp",
130
- version: "3.0.0",
131
- })
132
-
133
- // Try to connect - this triggers the OAuth flow
134
104
  try {
135
- await client.connect(transport)
136
- // If we get here, we're already authenticated
137
- await client.close().catch(() => {})
138
- await transport.close().catch(() => {})
139
- return { authorizationUrl: "", transport }
140
- } catch (error) {
141
- if (error instanceof UnauthorizedError && capturedUrl) {
142
- await client.close().catch(() => {})
143
- // Store transport for later finishAuth
144
- pendingTransports.set(serverName, transport)
145
- return { authorizationUrl: capturedUrl.toString(), transport }
105
+ const result = await runSdkAuth(authProvider, { serverUrl })
106
+ if (result === "AUTHORIZED") {
107
+ await clearOAuthState(serverName)
108
+ return { authorizationUrl: "" }
146
109
  }
147
- await client.close().catch(() => {})
148
- await transport.close().catch(() => {})
110
+ if (!capturedUrl) {
111
+ throw new UnauthorizedError("OAuth authorization URL was not provided")
112
+ }
113
+ pendingTransports.set(
114
+ serverName,
115
+ new StreamableHTTPClientTransport(new URL(serverUrl), { authProvider }),
116
+ )
117
+ return { authorizationUrl: capturedUrl.toString() }
118
+ } catch (error) {
119
+ await clearOAuthState(serverName)
149
120
  throw error
150
121
  }
151
122
  }
@@ -208,19 +179,19 @@ export async function authenticate(
208
179
  // Register the callback BEFORE opening the browser
209
180
  const callbackPromise = waitForCallback(oauthState)
210
181
 
211
- // Open browser
212
- console.log(`MCP Auth: Opening browser for ${serverName}`)
213
182
  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
- }
183
+ // Open browser
184
+ console.log(`MCP Auth: Opening browser for ${serverName}`)
185
+ try {
186
+ await open(authorizationUrl)
187
+ } catch (error) {
188
+ console.warn(`MCP Auth: Failed to open browser for ${serverName}`, { error })
189
+ throw new Error(
190
+ `Could not open browser. Please open this URL manually: ${authorizationUrl}`,
191
+ { cause: error },
192
+ )
193
+ }
222
194
 
223
- try {
224
195
  // Wait for callback
225
196
  const code = await callbackPromise
226
197
 
@@ -236,6 +207,7 @@ export async function authenticate(
236
207
  return await completeAuth(serverName, code)
237
208
  } catch (error) {
238
209
  cancelPendingCallback(oauthState)
210
+ await clearOAuthState(serverName)
239
211
  const pendingTransport = pendingTransports.get(serverName)
240
212
  if (pendingTransport) {
241
213
  pendingTransports.delete(serverName)
@@ -295,31 +267,12 @@ export async function getValidToken(
295
267
  return null
296
268
  }
297
269
 
298
- // Try to get tokens to find the token endpoint
299
- const existingTokens = await authProvider.tokens()
300
- if (!existingTokens) {
301
- return null
302
- }
303
-
304
- // Create transport to trigger refresh
305
- const transport = new StreamableHTTPClientTransport(new URL(serverUrl), {
306
- authProvider,
307
- })
308
-
309
- // Try to connect - SDK will attempt token refresh internally
310
- const client = new Client({ name: "pi-mcp", version: "3.0.0" })
311
- try {
312
- await client.connect(transport)
313
- // Get refreshed tokens
314
- const refreshed = await getAuthForUrl(serverName, serverUrl)
315
- return refreshed?.tokens ?? null
316
- } catch (error) {
317
- console.error(`MCP Auth: Token refresh failed for ${serverName}`, { error })
270
+ const result = await runSdkAuth(authProvider, { serverUrl })
271
+ if (result !== "AUTHORIZED") {
318
272
  return null
319
- } finally {
320
- await client.close().catch(() => {})
321
- await transport.close().catch(() => {})
322
273
  }
274
+ const refreshed = await getAuthForUrl(serverName, serverUrl)
275
+ return refreshed?.tokens ?? null
323
276
  } catch (error) {
324
277
  console.error(`MCP Auth: Token refresh failed for ${serverName}`, { error })
325
278
  return null
@@ -6,6 +6,7 @@
6
6
  */
7
7
 
8
8
  import type { OAuthClientProvider } from "@modelcontextprotocol/sdk/client/auth.js"
9
+ import { UnauthorizedError } from "@modelcontextprotocol/sdk/client/auth.js"
9
10
  import type {
10
11
  OAuthClientMetadata,
11
12
  OAuthTokens,
@@ -200,11 +201,23 @@ export class McpOAuthProvider implements OAuthClientProvider {
200
201
  /**
201
202
  * Redirect the user to the authorization URL.
202
203
  * This opens the browser for the user to authenticate.
204
+ *
205
+ * Throws UnauthorizedError when called outside of a user-initiated flow
206
+ * (no oauthState saved by startAuth). That path is reached when the SDK
207
+ * falls through from a failed refresh into a fresh authorization_code
208
+ * flow, which library hosts cannot complete in-process.
203
209
  */
204
210
  async redirectToAuthorization(authorizationUrl: URL): Promise<void> {
205
211
  if (this.usesClientCredentials) {
206
212
  throw new Error("redirectToAuthorization is not used for client_credentials flow")
207
213
  }
214
+ // No saved oauthState means we're on the post-refresh authorize fallback.
215
+ const entry = await getAuthEntry(this.serverName)
216
+ if (!entry?.oauthState) {
217
+ throw new UnauthorizedError(
218
+ `Re-authentication required for MCP server: ${this.serverName}`,
219
+ )
220
+ }
208
221
  // URL is passed to callback, not logged (may contain sensitive params)
209
222
  await this.callbacks.onRedirect(authorizationUrl)
210
223
  }
@@ -240,7 +253,7 @@ export class McpOAuthProvider implements OAuthClientProvider {
240
253
 
241
254
  /**
242
255
  * Get the stored OAuth state parameter.
243
- * @throws Error if no state is stored
256
+ * @throws UnauthorizedError if no flow is in progress (see redirectToAuthorization)
244
257
  */
245
258
  async state(): Promise<string> {
246
259
  if (this.usesClientCredentials) {
@@ -248,7 +261,9 @@ export class McpOAuthProvider implements OAuthClientProvider {
248
261
  }
249
262
  const entry = await getAuthEntry(this.serverName)
250
263
  if (!entry?.oauthState) {
251
- throw new Error(`No OAuth state saved for MCP server: ${this.serverName}`)
264
+ throw new UnauthorizedError(
265
+ `Re-authentication required for MCP server: ${this.serverName}`,
266
+ )
252
267
  }
253
268
  return entry.oauthState
254
269
  }
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "pi-mcp-adapter",
3
- "version": "2.4.2",
3
+ "version": "2.5.0",
4
4
  "description": "MCP (Model Context Protocol) adapter extension for Pi coding agent",
5
5
  "type": "module",
6
6
  "license": "MIT",
@@ -11,7 +11,8 @@
11
11
  "scripts": {
12
12
  "test": "vitest run",
13
13
  "test:watch": "vitest",
14
- "test:coverage": "vitest run --coverage"
14
+ "test:coverage": "vitest run --coverage",
15
+ "test:oauth-provider": "node --import tsx --test mcp-oauth-provider.test.ts"
15
16
  },
16
17
  "repository": {
17
18
  "type": "git",
@@ -51,6 +52,7 @@
51
52
  "ui-stream-types.ts",
52
53
  "config.ts",
53
54
  "server-manager.ts",
55
+ "sampling-handler.ts",
54
56
  "tool-registrar.ts",
55
57
  "resource-tools.ts",
56
58
  "lifecycle.ts",
@@ -77,6 +79,7 @@
77
79
  "dependencies": {
78
80
  "@modelcontextprotocol/ext-apps": "^1.2.2",
79
81
  "@modelcontextprotocol/sdk": "^1.25.1",
82
+ "@mariozechner/pi-ai": "^0.70.2",
80
83
  "typebox": "^1.1.24",
81
84
  "open": "^10.2.0",
82
85
  "zod": "^3.25.0 || ^4.0.0"
package/proxy-modes.ts CHANGED
@@ -6,7 +6,7 @@ import { lazyConnect, updateServerMetadata, updateMetadataCache, getFailureAgeSe
6
6
  import { buildToolMetadata, getToolNames, findToolByName, formatSchema } from "./tool-metadata.js";
7
7
  import { transformMcpContent } from "./tool-registrar.js";
8
8
  import { maybeStartUiSession, type UiSessionRuntime } from "./ui-session.js";
9
- import { truncateAtWord } from "./utils.js";
9
+ import { formatAuthRequiredMessage, truncateAtWord } from "./utils.js";
10
10
  import { authenticate, supportsOAuth } from "./mcp-auth-flow.js";
11
11
 
12
12
  type ProxyToolResult = AgentToolResult<Record<string, unknown>>;
@@ -16,8 +16,20 @@ type AutoAuthResult =
16
16
  | { status: "success" }
17
17
  | { status: "failed"; message: string };
18
18
 
19
- function getAuthRequiredMessage(serverName: string): string {
20
- return `Server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} first.`;
19
+ function getAuthRequiredMessage(
20
+ state: McpExtensionState,
21
+ serverName: string,
22
+ defaultMessage = `Server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} first.`,
23
+ ): string {
24
+ return formatAuthRequiredMessage(state.config, serverName, defaultMessage);
25
+ }
26
+
27
+ function getAuthFailedMessage(state: McpExtensionState, serverName: string, message: string): string {
28
+ const customGuidance = state.config.settings?.authRequiredMessage;
29
+ if (customGuidance) {
30
+ return `OAuth authentication failed for "${serverName}": ${message}. ${getAuthRequiredMessage(state, serverName)}`;
31
+ }
32
+ return `OAuth authentication failed for "${serverName}": ${message}. Run /mcp-auth ${serverName} first.`;
21
33
  }
22
34
 
23
35
  async function attemptAutoAuth(
@@ -37,7 +49,11 @@ async function attemptAutoAuth(
37
49
  if (!state.ui && grantType !== "client_credentials") {
38
50
  return {
39
51
  status: "failed",
40
- message: `Server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} in an interactive session.`,
52
+ message: getAuthRequiredMessage(
53
+ state,
54
+ serverName,
55
+ `Server "${serverName}" requires OAuth authentication. Run /mcp-auth ${serverName} in an interactive session.`,
56
+ ),
41
57
  };
42
58
  }
43
59
 
@@ -48,7 +64,7 @@ async function attemptAutoAuth(
48
64
  const message = error instanceof Error ? error.message : String(error);
49
65
  return {
50
66
  status: "failed",
51
- message: `OAuth authentication failed for "${serverName}": ${message}. Run /mcp-auth ${serverName} first.`,
67
+ message: getAuthFailedMessage(state, serverName, message),
52
68
  };
53
69
  }
54
70
  }
@@ -234,7 +250,6 @@ export function executeSearch(
234
250
  regex?: boolean,
235
251
  server?: string,
236
252
  includeSchemas?: boolean,
237
- getPiTools?: () => ToolInfo[]
238
253
  ): ProxyToolResult {
239
254
  const showSchemas = includeSchemas !== false;
240
255
 
@@ -262,21 +277,6 @@ export function executeSearch(
262
277
  };
263
278
  }
264
279
 
265
- const piMatches: Array<{ name: string; description: string }> = [];
266
- if (!server && getPiTools) {
267
- const piTools = getPiTools();
268
- for (const tool of piTools) {
269
- if (tool.name === "mcp") continue;
270
-
271
- if (pattern.test(tool.name) || pattern.test(tool.description ?? "")) {
272
- piMatches.push({
273
- name: tool.name,
274
- description: tool.description ?? "",
275
- });
276
- }
277
- }
278
- }
279
-
280
280
  for (const [serverName, metadata] of state.toolMetadata.entries()) {
281
281
  if (server && serverName !== server) continue;
282
282
  for (const tool of metadata) {
@@ -289,7 +289,7 @@ export function executeSearch(
289
289
  }
290
290
  }
291
291
 
292
- const totalCount = piMatches.length + matches.length;
292
+ const totalCount = matches.length;
293
293
 
294
294
  if (totalCount === 0) {
295
295
  const msg = server
@@ -303,21 +303,6 @@ export function executeSearch(
303
303
 
304
304
  let text = `Found ${totalCount} tool${totalCount === 1 ? "" : "s"} matching "${query}":\n\n`;
305
305
 
306
- for (const match of piMatches) {
307
- if (showSchemas) {
308
- text += `[pi tool] ${match.name}\n`;
309
- text += ` ${match.description || "(no description)"}\n`;
310
- text += ` No parameters (call directly).\n`;
311
- text += "\n";
312
- } else {
313
- text += `[pi tool] ${match.name}`;
314
- if (match.description) {
315
- text += ` - ${truncateAtWord(match.description, 50)}`;
316
- }
317
- text += "\n";
318
- }
319
- }
320
-
321
306
  for (const match of matches) {
322
307
  if (showSchemas) {
323
308
  text += `${match.tool.name}\n`;
@@ -341,10 +326,7 @@ export function executeSearch(
341
326
  content: [{ type: "text" as const, text: text.trim() }],
342
327
  details: {
343
328
  mode: "search",
344
- matches: [
345
- ...piMatches.map(m => ({ server: "pi", tool: m.name })),
346
- ...matches.map(m => ({ server: m.server, tool: m.tool.name })),
347
- ],
329
+ matches: matches.map(m => ({ server: m.server, tool: m.tool.name })),
348
330
  count: totalCount,
349
331
  query,
350
332
  },
@@ -433,7 +415,7 @@ export async function executeConnect(state: McpExtensionState, serverName: strin
433
415
  connection = await state.manager.connect(serverName, definition);
434
416
  }
435
417
  if (connection.status === "needs-auth") {
436
- const message = getAuthRequiredMessage(serverName);
418
+ const message = getAuthRequiredMessage(state, serverName);
437
419
  return {
438
420
  content: [{ type: "text" as const, text: message }],
439
421
  details: { mode: "connect", error: "auth_required", server: serverName, message },
@@ -463,6 +445,7 @@ export async function executeCall(
463
445
  toolName: string,
464
446
  args?: Record<string, unknown>,
465
447
  serverOverride?: string,
448
+ getPiTools?: () => ToolInfo[],
466
449
  ): Promise<ProxyToolResult> {
467
450
  let serverName: string | undefined = serverOverride;
468
451
  let toolMeta: ToolMetadata | undefined;
@@ -522,7 +505,7 @@ export async function executeCall(
522
505
  }
523
506
 
524
507
  if (!toolMeta && state.manager.getConnection(serverName)?.status === "needs-auth") {
525
- const message = getAuthRequiredMessage(serverName);
508
+ const message = getAuthRequiredMessage(state, serverName);
526
509
  return {
527
510
  content: [{ type: "text" as const, text: message }],
528
511
  details: { mode: "call", error: "auth_required", server: serverName, message },
@@ -583,6 +566,16 @@ export async function executeCall(
583
566
  }
584
567
 
585
568
  if (!serverName || !toolMeta) {
569
+ const nativeTool = !serverOverride
570
+ ? getPiTools?.().find((tool) => tool.name === toolName && tool.name !== "mcp")
571
+ : undefined;
572
+ if (nativeTool) {
573
+ return {
574
+ content: [{ type: "text" as const, text: `"${toolName}" is a native Pi tool. Call ${toolName} directly instead of using mcp({ tool: "${toolName}" }).` }],
575
+ details: { mode: "call", error: "native_tool", requestedTool: toolName },
576
+ };
577
+ }
578
+
586
579
  const hintServer = serverName ?? prefixMatchedServer;
587
580
  const available = hintServer ? getToolNames(state, hintServer) : [];
588
581
  let msg = `Tool "${toolName}" not found.`;
@@ -616,7 +609,7 @@ export async function executeCall(
616
609
  }
617
610
 
618
611
  if (connection?.status === "needs-auth") {
619
- const message = getAuthRequiredMessage(serverName);
612
+ const message = getAuthRequiredMessage(state, serverName);
620
613
  return {
621
614
  content: [{ type: "text" as const, text: message }],
622
615
  details: { mode: "call", error: "auth_required", server: serverName, message },
@@ -662,7 +655,7 @@ export async function executeCall(
662
655
  }
663
656
 
664
657
  if (connection.status === "needs-auth") {
665
- const message = getAuthRequiredMessage(serverName);
658
+ const message = getAuthRequiredMessage(state, serverName);
666
659
  return {
667
660
  content: [{ type: "text" as const, text: message }],
668
661
  details: { mode: "call", error: "auth_required", server: serverName, message },
@@ -0,0 +1,246 @@
1
+ import { complete, type AssistantMessage, type Message, type Model, type TextContent } from "@mariozechner/pi-ai";
2
+ import { truncateAtWord } from "./utils.js";
3
+ import type { ExtensionUIContext, ModelRegistry } from "@mariozechner/pi-coding-agent";
4
+ import type { Client } from "@modelcontextprotocol/sdk/client/index.js";
5
+ import {
6
+ CreateMessageRequestSchema,
7
+ type CreateMessageRequest,
8
+ type CreateMessageResult,
9
+ type SamplingMessage,
10
+ type SamplingMessageContentBlock,
11
+ } from "@modelcontextprotocol/sdk/types.js";
12
+
13
+ export interface SamplingHandlerOptions {
14
+ serverName: string;
15
+ autoApprove: boolean;
16
+ ui?: ExtensionUIContext;
17
+ modelRegistry: ModelRegistry;
18
+ getCurrentModel: () => Model<any> | undefined;
19
+ getSignal: () => AbortSignal | undefined;
20
+ }
21
+
22
+ export type ServerSamplingConfig = Omit<SamplingHandlerOptions, "serverName">;
23
+
24
+ export function registerSamplingHandler(client: Client, options: SamplingHandlerOptions): void {
25
+ client.setRequestHandler(CreateMessageRequestSchema, (request) => {
26
+ return handleSamplingRequest(options, request as CreateMessageRequest);
27
+ });
28
+ }
29
+
30
+ export async function handleSamplingRequest(
31
+ options: SamplingHandlerOptions,
32
+ request: CreateMessageRequest,
33
+ ): Promise<CreateMessageResult> {
34
+ const params = request.params;
35
+
36
+ if ("task" in params && params.task) {
37
+ throw new Error("MCP sampling tasks are not supported");
38
+ }
39
+ if (params.includeContext && params.includeContext !== "none") {
40
+ throw new Error("MCP sampling context inclusion is not supported");
41
+ }
42
+ if (params.tools?.length) {
43
+ throw new Error("MCP sampling tool use is not supported");
44
+ }
45
+ if (params.toolChoice) {
46
+ throw new Error("MCP sampling tool choice is not supported");
47
+ }
48
+ if (params.stopSequences?.length) {
49
+ throw new Error("MCP sampling stop sequences are not supported");
50
+ }
51
+
52
+ const messages = params.messages.map(convertSamplingMessage);
53
+ const { model, apiKey, headers } = await resolveSamplingModel(options);
54
+ await confirmSampling(
55
+ options,
56
+ "Approve MCP sampling request",
57
+ formatRequestApproval(options.serverName, `${model.provider}/${model.id}`, params.systemPrompt, messages),
58
+ );
59
+
60
+ const result = await complete(
61
+ model,
62
+ {
63
+ systemPrompt: params.systemPrompt,
64
+ messages,
65
+ },
66
+ {
67
+ apiKey,
68
+ headers,
69
+ maxTokens: params.maxTokens,
70
+ temperature: params.temperature,
71
+ metadata: params.metadata as Record<string, unknown> | undefined,
72
+ signal: options.getSignal(),
73
+ },
74
+ );
75
+
76
+ const converted = convertAssistantResult(result);
77
+ await confirmSampling(
78
+ options,
79
+ "Return MCP sampling response",
80
+ formatResponseApproval(options.serverName, converted),
81
+ );
82
+ return converted;
83
+ }
84
+
85
+ function formatRequestApproval(
86
+ serverName: string,
87
+ modelName: string,
88
+ systemPrompt: string | undefined,
89
+ messages: Message[],
90
+ ): string {
91
+ const lines = [`${serverName} wants to sample ${messages.length} message${messages.length === 1 ? "" : "s"} with ${modelName}.`];
92
+ if (systemPrompt) {
93
+ lines.push(`System: ${truncateAtWord(systemPrompt, 400)}`);
94
+ }
95
+ for (const [index, message] of messages.entries()) {
96
+ lines.push(`${index + 1}. ${message.role}: ${truncateAtWord(messageText(message), 400)}`);
97
+ }
98
+ return lines.join("\n\n");
99
+ }
100
+
101
+ function formatResponseApproval(serverName: string, response: CreateMessageResult): string {
102
+ const text = response.content.type === "text" ? response.content.text : `[${response.content.type} content]`;
103
+ return `${serverName} will receive this response from ${response.model}:\n\n${truncateAtWord(text, 1000)}`;
104
+ }
105
+
106
+ function messageText(message: Message): string {
107
+ if (typeof message.content === "string") return message.content;
108
+ return message.content.map((block) => {
109
+ if (block.type === "text") return block.text;
110
+ if (block.type === "image") return `[image: ${block.mimeType}]`;
111
+ if (block.type === "thinking") return "[thinking]";
112
+ if (block.type === "toolCall") return `[tool call: ${block.name}]`;
113
+ return "[content]";
114
+ }).join("\n");
115
+ }
116
+
117
+ async function resolveSamplingModel(options: SamplingHandlerOptions): Promise<{
118
+ model: Model<any>;
119
+ apiKey?: string;
120
+ headers?: Record<string, string>;
121
+ }> {
122
+ const candidates: Model<any>[] = [];
123
+ const currentModel = options.getCurrentModel();
124
+ if (currentModel) candidates.push(currentModel);
125
+
126
+ for (const model of options.modelRegistry.getAvailable()) {
127
+ if (!candidates.some((candidate) => candidate.provider === model.provider && candidate.id === model.id)) {
128
+ candidates.push(model);
129
+ }
130
+ }
131
+
132
+ const errors: string[] = [];
133
+ for (const model of candidates) {
134
+ const auth = await options.modelRegistry.getApiKeyAndHeaders(model);
135
+ if (auth.ok) {
136
+ return { model, apiKey: auth.apiKey, headers: auth.headers };
137
+ }
138
+ errors.push(`${model.provider}/${model.id}: ${auth.error}`);
139
+ }
140
+
141
+ if (errors.length > 0) {
142
+ throw new Error(`No configured auth for MCP sampling model. ${errors.join("; ")}`);
143
+ }
144
+ throw new Error("No Pi model is available for MCP sampling");
145
+ }
146
+
147
+ async function confirmSampling(options: SamplingHandlerOptions, title: string, message: string): Promise<void> {
148
+ if (options.autoApprove) return;
149
+ if (!options.ui) {
150
+ throw new Error("MCP sampling requires interactive approval. Set settings.samplingAutoApprove to true to allow it without UI.");
151
+ }
152
+ const approved = await options.ui.confirm(title, message);
153
+ if (!approved) {
154
+ throw new Error("MCP sampling request was declined");
155
+ }
156
+ }
157
+
158
+ function convertSamplingMessage(message: SamplingMessage): Message {
159
+ const blocks = Array.isArray(message.content) ? message.content : [message.content];
160
+ if (message.role === "user") {
161
+ return {
162
+ role: "user",
163
+ content: blocks.map(convertUserContent),
164
+ timestamp: Date.now(),
165
+ };
166
+ }
167
+
168
+ return {
169
+ role: "assistant",
170
+ content: blocks.map(convertAssistantContent),
171
+ api: "mcp-sampling",
172
+ provider: "mcp",
173
+ model: "sampling-request",
174
+ usage: zeroUsage(),
175
+ stopReason: "stop",
176
+ timestamp: Date.now(),
177
+ };
178
+ }
179
+
180
+ function convertUserContent(block: SamplingMessageContentBlock): TextContent {
181
+ if (block.type === "text") {
182
+ return { type: "text", text: block.text };
183
+ }
184
+ throw new Error(`MCP sampling ${block.type} content is not supported`);
185
+ }
186
+
187
+ function convertAssistantContent(block: SamplingMessageContentBlock): TextContent {
188
+ if (block.type === "text") {
189
+ return { type: "text", text: block.text };
190
+ }
191
+ throw new Error(`MCP sampling assistant ${block.type} content is not supported`);
192
+ }
193
+
194
+ function convertAssistantResult(message: AssistantMessage): CreateMessageResult {
195
+ if (message.stopReason === "error") {
196
+ throw new Error(message.errorMessage ?? "MCP sampling model call failed");
197
+ }
198
+ if (message.stopReason === "aborted") {
199
+ throw new Error(message.errorMessage ?? "MCP sampling model call was aborted");
200
+ }
201
+
202
+ const text = message.content
203
+ .map((block) => {
204
+ if (block.type === "text") return block.text;
205
+ if (block.type === "thinking") return undefined;
206
+ throw new Error(`MCP sampling result ${block.type} content is not supported`);
207
+ })
208
+ .filter((value): value is string => value !== undefined)
209
+ .join("\n\n")
210
+ .trim();
211
+
212
+ if (!text) {
213
+ throw new Error("MCP sampling result did not contain text content");
214
+ }
215
+
216
+ return {
217
+ role: "assistant",
218
+ content: { type: "text", text },
219
+ model: `${message.provider}/${message.model}`,
220
+ stopReason: mapStopReason(message.stopReason),
221
+ };
222
+ }
223
+
224
+ function mapStopReason(reason: AssistantMessage["stopReason"]): CreateMessageResult["stopReason"] {
225
+ if (reason === "stop") return "endTurn";
226
+ if (reason === "length") return "maxTokens";
227
+ if (reason === "toolUse") return "toolUse";
228
+ return reason;
229
+ }
230
+
231
+ function zeroUsage(): AssistantMessage["usage"] {
232
+ return {
233
+ input: 0,
234
+ output: 0,
235
+ cacheRead: 0,
236
+ cacheWrite: 0,
237
+ totalTokens: 0,
238
+ cost: {
239
+ input: 0,
240
+ output: 0,
241
+ cacheRead: 0,
242
+ cacheWrite: 0,
243
+ total: 0,
244
+ },
245
+ };
246
+ }
package/server-manager.ts CHANGED
@@ -16,6 +16,7 @@ import { resolveNpxBinary } from "./npx-resolver.js";
16
16
  import { logger } from "./logger.js";
17
17
  import { McpOAuthProvider } from "./mcp-oauth-provider.js";
18
18
  import { supportsOAuth } from "./mcp-auth-flow.js";
19
+ import { registerSamplingHandler, type ServerSamplingConfig } from "./sampling-handler.js";
19
20
 
20
21
  interface ServerConnection {
21
22
  client: Client;
@@ -34,6 +35,11 @@ export class McpServerManager {
34
35
  private connections = new Map<string, ServerConnection>();
35
36
  private connectPromises = new Map<string, Promise<ServerConnection>>();
36
37
  private uiStreamListeners = new Map<string, UiStreamListener>();
38
+ private samplingConfig: ServerSamplingConfig | undefined;
39
+
40
+ setSamplingConfig(config: ServerSamplingConfig | undefined): void {
41
+ this.samplingConfig = config;
42
+ }
37
43
 
38
44
  async connect(name: string, definition: ServerDefinition): Promise<ServerConnection> {
39
45
  // Dedupe concurrent connection attempts
@@ -64,7 +70,7 @@ export class McpServerManager {
64
70
  name: string,
65
71
  definition: ServerDefinition
66
72
  ): Promise<ServerConnection> {
67
- const client = new Client({ name: `pi-mcp-${name}`, version: "1.0.0" });
73
+ const client = this.createClient(name);
68
74
 
69
75
  let transport: Transport;
70
76
 
@@ -141,6 +147,17 @@ export class McpServerManager {
141
147
  }
142
148
  }
143
149
 
150
+ private createClient(serverName: string): Client {
151
+ const client = new Client(
152
+ { name: `pi-mcp-${serverName}`, version: "1.0.0" },
153
+ this.samplingConfig ? { capabilities: { sampling: {} } } : undefined,
154
+ );
155
+ if (this.samplingConfig) {
156
+ registerSamplingHandler(client, { ...this.samplingConfig, serverName });
157
+ }
158
+ return client;
159
+ }
160
+
144
161
  private async createHttpTransport(
145
162
  definition: ServerDefinition,
146
163
  serverName: string
package/types.ts CHANGED
@@ -318,6 +318,14 @@ export interface McpSettings {
318
318
  directTools?: boolean;
319
319
  disableProxyTool?: boolean;
320
320
  autoAuth?: boolean;
321
+ sampling?: boolean;
322
+ samplingAutoApprove?: boolean;
323
+ /**
324
+ * Message returned in tool results when a server needs (re-)authentication.
325
+ * "${server}" is substituted with the server name. Defaults to a TUI
326
+ * instruction when unset.
327
+ */
328
+ authRequiredMessage?: string;
321
329
  }
322
330
 
323
331
  // Root config
package/utils.ts CHANGED
@@ -1,5 +1,6 @@
1
1
  import type { ExtensionAPI } from "@mariozechner/pi-coding-agent";
2
2
  import { platform } from "node:os";
3
+ import type { McpConfig } from "./types.js";
3
4
 
4
5
  async function execOpen(pi: ExtensionAPI, target: string, browser?: string) {
5
6
  const os = platform();
@@ -70,6 +71,15 @@ export function truncateAtWord(text: string, target: number): string {
70
71
  return truncated + "...";
71
72
  }
72
73
 
74
+ export function formatAuthRequiredMessage(
75
+ config: Pick<McpConfig, "settings">,
76
+ serverName: string,
77
+ defaultMessage: string,
78
+ ): string {
79
+ const template = config.settings?.authRequiredMessage;
80
+ return template ? template.replaceAll("${server}", serverName) : defaultMessage;
81
+ }
82
+
73
83
  /**
74
84
  * Extract the adapter-owned UI stream mode from tool metadata.
75
85
  */