Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 32 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
name: CI

on:
pull_request:
push:
branches: [main]

permissions:
contents: read

jobs:
test:
runs-on: ubuntu-latest
timeout-minutes: 60

steps:
- uses: actions/checkout@v4
- uses: actions/setup-node@v4
with:
node-version: 22
cache: npm

- run: npm ci
- name: Build interactive visualizer example
run: |
npm install --prefix examples/interactive-visualizer --no-package-lock --no-audit --no-fund
npm run --prefix examples/interactive-visualizer build
- run: npx tsc --noEmit
- run: npm test
- run: npm run test:oauth
- run: npm run test:conformance
- run: npm pack --dry-run
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
node_modules/
.pi/
conformance/results/
*.log
.DS_Store

Expand Down
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]

### Added
- Migrated the MCP client and interactive visualizer to the exact-pinned MCP SDK v2 beta.5 packages, with automatic protocol negotiation and client conformance coverage. Thanks Matt Carey (@mattzcarey) for PR #210.
- Added disabled MCP server definitions plus `/mcp disable` and `/mcp enable` project-local overrides that preserve visibility while preventing execution. Thanks Ömer Ulusoy (@ulusoyomer) for PR #61.
- Added argument completions for `/mcp` subcommands and reconnect/logout server names. Thanks @sting8k for PR #8.
- Surfaced MCP connection failure reasons from bounded stdio diagnostics in status output and the `/mcp` panel, with a shortcut to copy the selected failure. Thanks @parkuman for PR #197.
Expand Down
11 changes: 4 additions & 7 deletions OAUTH.md
Original file line number Diff line number Diff line change
Expand Up @@ -323,18 +323,15 @@ The OAuth implementation uses the following modules:

## SDK Integration

The implementation uses these SDK exports:
The implementation uses these MCP SDK v2 exports:

```typescript
import {
auth,
UnauthorizedError,
OAuthClientProvider,
} from "@modelcontextprotocol/sdk/client/auth.js"

import {
StreamableHTTPClientTransport,
} from "@modelcontextprotocol/sdk/client/streamableHttp.js"
UnauthorizedError,
type OAuthClientProvider,
} from "@modelcontextprotocol/client"
```

The `McpOAuthProvider` class implements `OAuthClientProvider` and is passed to `StreamableHTTPClientTransport`:
Expand Down
2 changes: 0 additions & 2 deletions __tests__/abort-signal.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,6 @@ describe("AbortSignal propagation", () => {
expect(result.details.error).toBe("aborted");
expect(callTool).toHaveBeenCalledWith(
{ name: "slow", arguments: {}, _meta: undefined },
undefined,
{ signal: controller.signal },
);
expect(state.manager.decrementInFlight).toHaveBeenCalledWith("demo");
Expand All @@ -91,7 +90,6 @@ describe("AbortSignal propagation", () => {
expect(result.details.error).toBe("aborted");
expect(callTool).toHaveBeenCalledWith(
{ name: "slow", arguments: {}, _meta: undefined },
undefined,
{ signal: controller.signal },
);
expect(state.manager.decrementInFlight).toHaveBeenCalledWith("demo");
Expand Down
4 changes: 1 addition & 3 deletions __tests__/direct-tools-auto-auth.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,6 @@ describe("direct tools auto auth", () => {
arguments: { q: "hello" },
_meta: undefined,
},
undefined,
{ timeout: 4321 },
);
expect(result.content[0].text).toContain("ok");
Expand Down Expand Up @@ -157,7 +156,6 @@ describe("direct tools auto auth", () => {
expect(state.manager.getRequestOptions).toHaveBeenCalledWith("demo", controller.signal);
expect(connection.client.callTool).toHaveBeenCalledWith(
{ name: "search", arguments: {}, _meta: undefined },
undefined,
requestOptions,
);
expect(result.details).toMatchObject({ error: "aborted", server: "demo" });
Expand Down Expand Up @@ -246,7 +244,7 @@ describe("direct tools auto auth", () => {
});

it("runs URL elicitations returned by a URL-required tool error", async () => {
const { UrlElicitationRequiredError } = await import("@modelcontextprotocol/sdk/types.js");
const { UrlElicitationRequiredError } = await import("@modelcontextprotocol/client");
const { createDirectToolExecutor } = await import("../direct-tools.ts");
const error = new UrlElicitationRequiredError([{
mode: "url",
Expand Down
2 changes: 1 addition & 1 deletion __tests__/elicitation-handler.test.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
import type { ElicitRequest } from "@modelcontextprotocol/sdk/types.js";
import type { ElicitRequest } from "@modelcontextprotocol/client";

const mocks = vi.hoisted(() => ({
open: vi.fn(async () => undefined),
Expand Down
2 changes: 1 addition & 1 deletion __tests__/elicitation-sdk-integration.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ describe("elicitation with the real MCP SDK", () => {
const connection = manager.getConnection("real")!;

await expect(connection.client.callTool({ name: "url", arguments: {} })).rejects.toThrow(
/does not support URL-mode elicitation requests/,
/does not support (?:URL-mode|url) elicitation/,
);
expect(mocks.open).not.toHaveBeenCalled();
});
Expand Down
9 changes: 4 additions & 5 deletions __tests__/fixtures/delayed-mcp-server.mjs
Original file line number Diff line number Diff line change
@@ -1,8 +1,7 @@
import { rename, writeFile } from "node:fs/promises";
import { join } from "node:path";
import { Server } from "@modelcontextprotocol/sdk/server/index.js";
import { StdioServerTransport } from "@modelcontextprotocol/sdk/server/stdio.js";
import { ListResourcesRequestSchema, ListToolsRequestSchema } from "@modelcontextprotocol/sdk/types.js";
import { Server } from "@modelcontextprotocol/server";
import { StdioServerTransport } from "@modelcontextprotocol/server/stdio";

const pidPath = process.env.MCP_RELOAD_PID_DIR ? join(process.env.MCP_RELOAD_PID_DIR, `${process.pid}.json`) : undefined;
const identity = { pid: process.pid, toolName: "reload_identity" };
Expand All @@ -19,11 +18,11 @@ const server = new Server(
{ name: "delayed-reload-fixture", version: "1.0.0" },
{ capabilities: { tools: {}, resources: {} } },
);
server.setRequestHandler(ListToolsRequestSchema, async () => {
server.setRequestHandler("tools/list", async () => {
await new Promise(resolve => setTimeout(resolve, 100));
return {
tools: [{ name: identity.toolName, description: "reload identity", inputSchema: { type: "object", properties: {} } }],
};
});
server.setRequestHandler(ListResourcesRequestSchema, async () => ({ resources: [] }));
server.setRequestHandler("resources/list", async () => ({ resources: [] }));
await server.connect(new StdioServerTransport());
54 changes: 20 additions & 34 deletions __tests__/fixtures/elicitation-server.mjs
Original file line number Diff line number Diff line change
@@ -1,13 +1,5 @@
import { Server } from "@modelcontextprotocol/sdk/server/index.js";
import { StdioServerTransport } from "@modelcontextprotocol/sdk/server/stdio.js";
import {
CallToolRequestSchema,
ElicitResultSchema,
ListResourcesRequestSchema,
ListToolsRequestSchema,
ReadResourceRequestSchema,
UrlElicitationRequiredError,
} from "@modelcontextprotocol/sdk/types.js";
import { Server, UrlElicitationRequiredError } from "@modelcontextprotocol/server";
import { StdioServerTransport } from "@modelcontextprotocol/server/stdio";

const server = new Server(
{ name: "elicitation-integration-server", version: "1.0.0" },
Expand All @@ -23,7 +15,7 @@ function urlRequiredError() {
}]);
}

server.setRequestHandler(ListToolsRequestSchema, async () => ({
server.setRequestHandler("tools/list", async () => ({
tools: [
{ name: "capabilities", inputSchema: { type: "object", properties: {} } },
{ name: "form", inputSchema: { type: "object", properties: {} } },
Expand All @@ -32,21 +24,21 @@ server.setRequestHandler(ListToolsRequestSchema, async () => ({
],
}));

server.setRequestHandler(ListResourcesRequestSchema, async () => ({
server.setRequestHandler("resources/list", async () => ({
resources: [
{ name: "URL-required resource", uri: "test://url-required" },
{ name: "URL-required UI resource", uri: "ui://url-required" },
],
}));

server.setRequestHandler(ReadResourceRequestSchema, async request => {
server.setRequestHandler("resources/read", async request => {
if (request.params.uri === "test://url-required" || request.params.uri === "ui://url-required") {
throw urlRequiredError();
}
return { contents: [] };
});

server.setRequestHandler(CallToolRequestSchema, async request => {
server.setRequestHandler("tools/call", async request => {
if (request.params.name === "capabilities") {
return {
content: [{ type: "text", text: JSON.stringify(server.getClientCapabilities()?.elicitation ?? null) }],
Expand All @@ -56,31 +48,25 @@ server.setRequestHandler(CallToolRequestSchema, async request => {
if (request.params.name === "url-required") throw urlRequiredError();

if (request.params.name === "form") {
const result = await server.request({
method: "elicitation/create",
params: {
mode: "form",
message: "Provide a name",
requestedSchema: {
type: "object",
properties: { name: { type: "string", minLength: 1 } },
required: ["name"],
},
const result = await server.elicitInput({
mode: "form",
message: "Provide a name",
requestedSchema: {
type: "object",
properties: { name: { type: "string", minLength: 1 } },
required: ["name"],
},
}, ElicitResultSchema);
});
return { content: [{ type: "text", text: JSON.stringify(result) }] };
}

if (request.params.name === "url") {
const result = await server.request({
method: "elicitation/create",
params: {
mode: "url",
message: "Connect your account",
elicitationId: "requested-1",
url: "https://example.com/authorize",
},
}, ElicitResultSchema);
const result = await server.elicitInput({
mode: "url",
message: "Connect your account",
elicitationId: "requested-1",
url: "https://example.com/authorize",
});
if (result.action === "accept") {
for (const elicitationId of ["unknown", "requested-1", "requested-1"]) {
await server.notification({
Expand Down
7 changes: 4 additions & 3 deletions __tests__/mcp-auth-flow-client-credentials.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,8 @@ const mocks = vi.hoisted(() => ({

class MockUnauthorizedError extends Error {}

vi.mock("@modelcontextprotocol/sdk/client/auth.js", () => ({
vi.mock("@modelcontextprotocol/client", async (importOriginal) => ({
...(await importOriginal<Record<string, unknown>>()),
auth: mocks.sdkAuth,
extractWWWAuthenticateParams: (response: Response) => {
const header = response.headers.get("www-authenticate") ?? "";
Expand Down Expand Up @@ -327,7 +328,7 @@ describe("mcp-auth-flow explicit auth", () => {

it("preserves stored dynamic client info when tokens exist", async () => {
mocks.sdkAuth.mockImplementationOnce(async (provider) => {
expect(await provider.clientInformation()).toEqual({ client_id: "stored-client", client_secret: "stored-secret" });
expect(await provider.clientInformation()).toEqual({ client_id: "stored-client", client_secret: "stored-secret", redirect_uris: ["http://localhost:19876/callback"] });
await provider.redirectToAuthorization(new URL("https://auth.example.com/authorize"));
return "REDIRECT";
});
Expand Down Expand Up @@ -474,7 +475,7 @@ describe("mcp-auth-flow explicit auth", () => {

it("refreshes expired tokens even when cached dynamic redirect URIs are stale", async () => {
mocks.sdkAuth.mockImplementationOnce(async (provider) => {
expect(await provider.clientInformation()).toEqual({ client_id: "refresh-client", client_secret: "refresh-secret" });
expect(await provider.clientInformation()).toEqual({ client_id: "refresh-client", client_secret: "refresh-secret", redirect_uris: ["http://localhost:19876/callback"] });
await provider.saveTokens({
access_token: "new-access",
token_type: "Bearer",
Expand Down
83 changes: 82 additions & 1 deletion __tests__/mcp-oauth-provider.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ import { afterEach, beforeEach, describe, expect, it } from "vitest";
import { mkdtempSync, rmSync } from "node:fs";
import { join } from "node:path";
import { tmpdir } from "node:os";
import { UnauthorizedError } from "@modelcontextprotocol/sdk/client/auth.js";
import { UnauthorizedError } from "@modelcontextprotocol/client";
import { McpOAuthProvider } from "../mcp-oauth-provider.ts";
import { saveAuthEntry } from "../mcp-auth.ts";

Expand Down Expand Up @@ -158,6 +158,87 @@ describe("McpOAuthProvider addClientAuthentication", () => {
expect([...params.entries()]).toEqual([["grant_type", "authorization_code"]]);
expect([...headers.entries()]).toEqual([]);
});

it("does not persist a pre-registered issuer stub after deactivation", async () => {
const provider = new McpOAuthProvider(
"inactive-client-info",
serverUrl,
{ clientId: "my-client", clientSecret: "my-secret" },
{ onRedirect: async () => {} },
);
provider.deactivate();

await expect(provider.saveClientInformation({
client_id: "my-client",
issuer: "https://auth.example.com",
})).rejects.toThrow("OAuth flow is no longer active");

const { getAuthForUrl } = await import("../mcp-auth.ts");
expect(getAuthForUrl("inactive-client-info", serverUrl)).toBeUndefined();
});
});

describe("McpOAuthProvider discovery state", () => {
const originalOAuthDir = process.env.MCP_OAUTH_DIR;
const serverUrl = "https://api.example.com/mcp";
let authDir: string;

beforeEach(() => {
authDir = mkdtempSync(join(tmpdir(), "pi-mcp-oauth-discovery-"));
process.env.MCP_OAUTH_DIR = authDir;
});

afterEach(() => {
rmSync(authDir, { recursive: true, force: true });
if (originalOAuthDir === undefined) {
delete process.env.MCP_OAUTH_DIR;
} else {
process.env.MCP_OAUTH_DIR = originalOAuthDir;
}
});

it("round-trips callback-leg discovery state and invalidates it independently", async () => {
const provider = new McpOAuthProvider(
"discovery-state",
serverUrl,
{},
{ onRedirect: async () => {} },
);
const discoveryState = {
authorizationServerUrl: "https://auth.example.com",
resourceMetadataUrl: "https://api.example.com/.well-known/oauth-protected-resource/mcp",
authorizationServerMetadata: {
issuer: "https://auth.example.com",
authorization_endpoint: "https://auth.example.com/authorize",
token_endpoint: "https://auth.example.com/token",
response_types_supported: ["code"],
},
};

await provider.saveDiscoveryState(discoveryState);
expect(await provider.discoveryState()).toEqual(discoveryState);

const otherRuntimeProvider = new McpOAuthProvider(
"discovery-state",
serverUrl,
{},
{ onRedirect: async () => {} },
);
expect(await otherRuntimeProvider.discoveryState()).toBeUndefined();

await provider.saveTokens({
access_token: "access-token",
token_type: "Bearer",
issuer: "https://auth.example.com",
});
expect(await provider.discoveryState()).toBeUndefined();

await provider.saveDiscoveryState(discoveryState);
await provider.invalidateCredentials("discovery");

expect(await provider.discoveryState()).toBeUndefined();
expect((await provider.tokens())?.access_token).toBe("access-token");
});
});

describe("McpOAuthProvider authorization fallback", () => {
Expand Down
Loading