pi-mcp-adapter 2.4.2 → 2.5.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/CHANGELOG.md +17 -0
- package/README.md +4 -1
- package/direct-tools.ts +26 -5
- package/index.ts +2 -2
- package/init.ts +10 -0
- package/mcp-auth-flow.ts +37 -84
- package/mcp-callback-server.ts +1 -1
- package/mcp-oauth-provider.ts +19 -4
- package/package.json +5 -2
- package/proxy-modes.ts +38 -45
- package/sampling-handler.ts +246 -0
- package/server-manager.ts +18 -1
- package/types.ts +8 -0
- package/utils.ts +10 -0
package/CHANGELOG.md
CHANGED
|
@@ -7,6 +7,23 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|
|
7
7
|
|
|
8
8
|
## [Unreleased]
|
|
9
9
|
|
|
10
|
+
## [2.5.1] - 2026-04-24
|
|
11
|
+
|
|
12
|
+
### Fixed
|
|
13
|
+
- Changed OAuth browser callbacks to `http://localhost:<port>/callback` for pre-registered clients such as Slack MCP. Thanks @shenal for PR #53.
|
|
14
|
+
|
|
15
|
+
## [2.5.0] - 2026-04-24
|
|
16
|
+
|
|
17
|
+
### Added
|
|
18
|
+
- Added MCP `sampling/createMessage` support with conservative human approval by default and opt-in `settings.samplingAutoApprove` for non-interactive flows.
|
|
19
|
+
- Added configured Vitest coverage for OAuth provider authorization fallback behavior.
|
|
20
|
+
- Added `test:oauth-provider` for running the root OAuth provider node test with the required TypeScript loader.
|
|
21
|
+
|
|
22
|
+
### Fixed
|
|
23
|
+
- Applied `settings.authRequiredMessage` to proxy and direct-tool auth-required paths, including non-UI `autoAuth` failures.
|
|
24
|
+
- Fixed `/mcp-auth <server>` reporting success for expired stored OAuth tokens without forcing the SDK refresh/re-auth flow.
|
|
25
|
+
- Kept `mcp` search focused on MCP tools and added a direct-call hint when native Pi tools are accidentally routed through the proxy.
|
|
26
|
+
|
|
10
27
|
## [2.4.2] - 2026-04-22
|
|
11
28
|
|
|
12
29
|
### 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:
|
|
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:
|
|
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
|
|
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 =
|
|
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
|
|
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
|
|
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
|
|
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
|
|
90
|
-
|
|
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
|
|
136
|
-
|
|
137
|
-
|
|
138
|
-
|
|
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
|
-
|
|
148
|
-
|
|
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
|
-
|
|
215
|
-
|
|
216
|
-
|
|
217
|
-
|
|
218
|
-
|
|
219
|
-
|
|
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
|
-
|
|
299
|
-
|
|
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
|
package/mcp-callback-server.ts
CHANGED
|
@@ -178,7 +178,7 @@ export async function ensureCallbackServer(options: EnsureCallbackServerOptions
|
|
|
178
178
|
reject(err)
|
|
179
179
|
})
|
|
180
180
|
|
|
181
|
-
candidateServer.listen(candidatePort, "
|
|
181
|
+
candidateServer.listen(candidatePort, "localhost", () => {
|
|
182
182
|
resolve()
|
|
183
183
|
})
|
|
184
184
|
})
|
package/mcp-oauth-provider.ts
CHANGED
|
@@ -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,
|
|
@@ -28,7 +29,7 @@ import {
|
|
|
28
29
|
|
|
29
30
|
// Callback server configuration
|
|
30
31
|
const DEFAULT_OAUTH_CALLBACK_PORT = 19876
|
|
31
|
-
const OAUTH_CALLBACK_PATH = "/
|
|
32
|
+
const OAUTH_CALLBACK_PATH = "/callback"
|
|
32
33
|
|
|
33
34
|
let configuredOAuthCallbackPort = DEFAULT_OAUTH_CALLBACK_PORT
|
|
34
35
|
|
|
@@ -88,7 +89,7 @@ export class McpOAuthProvider implements OAuthClientProvider {
|
|
|
88
89
|
*/
|
|
89
90
|
get redirectUrl(): string | undefined {
|
|
90
91
|
if (this.usesClientCredentials) return undefined
|
|
91
|
-
return `http://
|
|
92
|
+
return `http://localhost:${getOAuthCallbackPort()}${OAUTH_CALLBACK_PATH}`
|
|
92
93
|
}
|
|
93
94
|
|
|
94
95
|
/**
|
|
@@ -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
|
|
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
|
|
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.
|
|
3
|
+
"version": "2.5.1",
|
|
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(
|
|
20
|
-
|
|
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:
|
|
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:
|
|
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 =
|
|
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 =
|
|
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
|
*/
|