From d1fc6629c246d11e984ccf8a0b2415cf0f6f98eb Mon Sep 17 00:00:00 2001 From: YeonGyu-Kim Date: Wed, 11 Mar 2026 20:10:34 +0900 Subject: [PATCH] fix(skill-mcp-manager): remove process listeners on disconnect and guard connection races H7: Process 'exit'/'SIGINT' listeners registered per-session were never removed when all sessions disconnected, accumulating handlers. - Add unregisterProcessCleanup() called in disconnectAll() H8: Race condition where disconnectSession() during pending connection left orphan clients in state.clients. - Add disconnectedSessions Set to track mid-flight disconnects - Check disconnect marker after connection resolves, close if stale - Clear marker on reconnection for same session Tests: 6 pass (3 disconnect + 3 race) --- src/features/skill-mcp-manager/cleanup.ts | 11 ++ .../skill-mcp-manager/connection-race.test.ts | 165 ++++++++++++++++++ src/features/skill-mcp-manager/connection.ts | 17 +- .../disconnect-cleanup.test.ts | 108 ++++++++++++ src/features/skill-mcp-manager/manager.ts | 1 + src/features/skill-mcp-manager/types.ts | 1 + 6 files changed, 301 insertions(+), 2 deletions(-) create mode 100644 src/features/skill-mcp-manager/connection-race.test.ts create mode 100644 src/features/skill-mcp-manager/disconnect-cleanup.test.ts diff --git a/src/features/skill-mcp-manager/cleanup.ts b/src/features/skill-mcp-manager/cleanup.ts index 805c8c506..bbfe68baf 100644 --- a/src/features/skill-mcp-manager/cleanup.ts +++ b/src/features/skill-mcp-manager/cleanup.ts @@ -24,6 +24,7 @@ export function registerProcessCleanup(state: SkillMcpManagerState): void { } state.clients.clear() state.pendingConnections.clear() + state.disconnectedSessions.clear() } // Note: Node's 'exit' event is synchronous-only, so we rely on signal handlers for async cleanup. @@ -81,10 +82,12 @@ async function cleanupIdleClients(state: SkillMcpManagerState): Promise { if (state.clients.size === 0) { stopCleanupTimer(state) + unregisterProcessCleanup(state) } } export async function disconnectSession(state: SkillMcpManagerState, sessionID: string): Promise { + state.disconnectedSessions.add(sessionID) const keysToRemove: string[] = [] for (const [key, managed] of state.clients.entries()) { @@ -96,12 +99,19 @@ export async function disconnectSession(state: SkillMcpManagerState, sessionID: } } + for (const key of state.pendingConnections.keys()) { + if (key.startsWith(`${sessionID}:`)) { + keysToRemove.push(key) + } + } + for (const key of keysToRemove) { state.pendingConnections.delete(key) } if (state.clients.size === 0) { stopCleanupTimer(state) + unregisterProcessCleanup(state) } } @@ -112,6 +122,7 @@ export async function disconnectAll(state: SkillMcpManagerState): Promise const clients = Array.from(state.clients.values()) state.clients.clear() state.pendingConnections.clear() + state.disconnectedSessions.clear() state.authProviders.clear() for (const managed of clients) { diff --git a/src/features/skill-mcp-manager/connection-race.test.ts b/src/features/skill-mcp-manager/connection-race.test.ts new file mode 100644 index 000000000..f08750c27 --- /dev/null +++ b/src/features/skill-mcp-manager/connection-race.test.ts @@ -0,0 +1,165 @@ +import { afterEach, beforeEach, describe, expect, it, mock } from "bun:test" +import type { ClaudeCodeMcpServer } from "../claude-code-mcp-loader/types" +import type { SkillMcpClientInfo, SkillMcpManagerState } from "./types" + +type Deferred = { + promise: Promise + resolve: (value: TValue) => void + reject: (error: Error) => void +} + +const pendingConnects: Deferred[] = [] +const trackedStates: SkillMcpManagerState[] = [] +const createdClients: MockClient[] = [] +const createdTransports: MockStdioClientTransport[] = [] + +class MockClient { + readonly close = mock(async () => {}) + + constructor( + _clientInfo: { name: string; version: string }, + _options: { capabilities: Record } + ) { + createdClients.push(this) + } + + async connect(_transport: MockStdioClientTransport): Promise { + const pendingConnect = pendingConnects.shift() + if (pendingConnect) { + await pendingConnect.promise + } + } +} + +class MockStdioClientTransport { + readonly close = mock(async () => {}) + + constructor(_options: { command: string; args?: string[]; env?: Record; stderr?: string }) { + createdTransports.push(this) + } +} + +mock.module("@modelcontextprotocol/sdk/client/index.js", () => ({ + Client: MockClient, +})) + +mock.module("@modelcontextprotocol/sdk/client/stdio.js", () => ({ + StdioClientTransport: MockStdioClientTransport, +})) + +const { disconnectAll, disconnectSession } = await import("./cleanup") +const { getOrCreateClient } = await import("./connection") + +function createDeferred(): Deferred { + let resolvePromise: ((value: TValue) => void) | null = null + let rejectPromise: ((error: Error) => void) | null = null + const promise = new Promise((resolve, reject) => { + resolvePromise = resolve + rejectPromise = reject + }) + + if (!resolvePromise || !rejectPromise) { + throw new Error("Failed to create deferred promise") + } + + return { + promise, + resolve: resolvePromise, + reject: rejectPromise, + } +} + +function createState(): SkillMcpManagerState { + const state: SkillMcpManagerState = { + clients: new Map(), + pendingConnections: new Map(), + disconnectedSessions: new Set(), + authProviders: new Map(), + cleanupRegistered: false, + cleanupInterval: null, + cleanupHandlers: [], + idleTimeoutMs: 5 * 60 * 1000, + } + + trackedStates.push(state) + return state +} + +function createClientInfo(sessionID: string): SkillMcpClientInfo { + return { + serverName: "race-server", + skillName: "race-skill", + sessionID, + } +} + +function createClientKey(info: SkillMcpClientInfo): string { + return `${info.sessionID}:${info.skillName}:${info.serverName}` +} + +const stdioConfig: ClaudeCodeMcpServer = { + command: "mock-mcp-server", +} + +beforeEach(() => { + pendingConnects.length = 0 + createdClients.length = 0 + createdTransports.length = 0 +}) + +afterEach(async () => { + for (const state of trackedStates) { + await disconnectAll(state) + } + + trackedStates.length = 0 + pendingConnects.length = 0 + createdClients.length = 0 + createdTransports.length = 0 +}) + +describe("getOrCreateClient disconnect race", () => { + it("#given pending connection for session A #when disconnectSession(A) is called before connection completes #then completed client is not added to state.clients", async () => { + const state = createState() + const info = createClientInfo("session-a") + const clientKey = createClientKey(info) + const pendingConnect = createDeferred() + pendingConnects.push(pendingConnect) + + const clientPromise = getOrCreateClient({ state, clientKey, info, config: stdioConfig }) + expect(state.pendingConnections.has(clientKey)).toBe(true) + + await disconnectSession(state, info.sessionID) + pendingConnect.resolve(undefined) + + await expect(clientPromise).rejects.toThrow(/disconnected during MCP connection setup/) + expect(state.clients.has(clientKey)).toBe(false) + expect(state.pendingConnections.has(clientKey)).toBe(false) + expect(state.disconnectedSessions.has(info.sessionID)).toBe(true) + expect(createdClients).toHaveLength(1) + expect(createdClients[0]?.close).toHaveBeenCalledTimes(1) + expect(createdTransports[0]?.close).toHaveBeenCalledTimes(1) + }) + + it("#given session A in disconnectedSessions #when new connection is requested for session A #then session A is removed from disconnectedSessions and connection proceeds normally", async () => { + const state = createState() + const info = createClientInfo("session-a") + const clientKey = createClientKey(info) + state.disconnectedSessions.add(info.sessionID) + + const client = await getOrCreateClient({ state, clientKey, info, config: stdioConfig }) + + expect(state.disconnectedSessions.has(info.sessionID)).toBe(false) + expect(state.clients.get(clientKey)?.client).toBe(client) + expect(createdClients[0]?.close).not.toHaveBeenCalled() + }) + + it("#given no pending connections #when disconnectSession is called #then no errors occur and the session is added to disconnectedSessions", async () => { + const state = createState() + + await expect(disconnectSession(state, "session-a")).resolves.toBeUndefined() + expect(state.disconnectedSessions.has("session-a")).toBe(true) + expect(state.pendingConnections.size).toBe(0) + expect(state.clients.size).toBe(0) + }) +}) diff --git a/src/features/skill-mcp-manager/connection.ts b/src/features/skill-mcp-manager/connection.ts index 1a3559d06..e35215706 100644 --- a/src/features/skill-mcp-manager/connection.ts +++ b/src/features/skill-mcp-manager/connection.ts @@ -14,6 +14,7 @@ export async function getOrCreateClient(params: { config: ClaudeCodeMcpServer }): Promise { const { state, clientKey, info, config } = params + state.disconnectedSessions.delete(info.sessionID) const existing = state.clients.get(clientKey) if (existing) { @@ -28,14 +29,26 @@ export async function getOrCreateClient(params: { } const expandedConfig = expandEnvVarsInObject(config) - const connectionPromise = createClient({ state, clientKey, info, config: expandedConfig }) + const connectionPromise = (async () => { + const client = await createClient({ state, clientKey, info, config: expandedConfig }) + + if (state.disconnectedSessions.has(info.sessionID)) { + await forceReconnect(state, clientKey) + throw new Error(`Session "${info.sessionID}" disconnected during MCP connection setup.`) + } + + return client + })() + state.pendingConnections.set(clientKey, connectionPromise) try { const client = await connectionPromise return client } finally { - state.pendingConnections.delete(clientKey) + if (state.pendingConnections.get(clientKey) === connectionPromise) { + state.pendingConnections.delete(clientKey) + } } } diff --git a/src/features/skill-mcp-manager/disconnect-cleanup.test.ts b/src/features/skill-mcp-manager/disconnect-cleanup.test.ts new file mode 100644 index 000000000..e4cfa1499 --- /dev/null +++ b/src/features/skill-mcp-manager/disconnect-cleanup.test.ts @@ -0,0 +1,108 @@ +import { Client } from "@modelcontextprotocol/sdk/client/index.js" +import { StreamableHTTPClientTransport } from "@modelcontextprotocol/sdk/client/streamableHttp.js" +import { afterEach, describe, expect, it } from "bun:test" +import { disconnectSession, registerProcessCleanup, unregisterProcessCleanup } from "./cleanup" +import type { ManagedClient, SkillMcpManagerState } from "./types" + +const trackedStates: SkillMcpManagerState[] = [] + +afterEach(() => { + for (const state of trackedStates) { + unregisterProcessCleanup(state) + } + + trackedStates.length = 0 +}) + +function createState(): SkillMcpManagerState { + const state: SkillMcpManagerState = { + clients: new Map(), + pendingConnections: new Map(), + disconnectedSessions: new Set(), + authProviders: new Map(), + cleanupRegistered: false, + cleanupInterval: null, + cleanupHandlers: [], + idleTimeoutMs: 5 * 60 * 1000, + } + + trackedStates.push(state) + return state +} + +function createManagedClient(skillName: string): ManagedClient { + return { + client: new Client( + { name: `test-${skillName}`, version: "1.0.0" }, + { capabilities: {} } + ), + transport: new StreamableHTTPClientTransport(new URL("https://example.com/mcp")), + skillName, + lastUsedAt: Date.now(), + connectionType: "http", + } +} + +describe("disconnectSession cleanup registration", () => { + it("#given state with 1 client and cleanup registered #when disconnectSession removes last client #then process cleanup handlers are unregistered", async () => { + // given + const state = createState() + const signalIntCountBeforeRegister = process.listenerCount("SIGINT") + const signalTermCountBeforeRegister = process.listenerCount("SIGTERM") + + state.clients.set("session-1:skill-1:server-1", createManagedClient("skill-1")) + registerProcessCleanup(state) + + // when + await disconnectSession(state, "session-1") + + // then + expect(state.cleanupRegistered).toBe(false) + expect(state.cleanupHandlers).toEqual([]) + expect(process.listenerCount("SIGINT")).toBe(signalIntCountBeforeRegister) + expect(process.listenerCount("SIGTERM")).toBe(signalTermCountBeforeRegister) + }) + + it("#given state with 2 clients in different sessions #when disconnectSession removes one session #then process cleanup handlers remain registered", async () => { + // given + const state = createState() + const signalIntCountBeforeRegister = process.listenerCount("SIGINT") + const signalTermCountBeforeRegister = process.listenerCount("SIGTERM") + + state.clients.set("session-1:skill-1:server-1", createManagedClient("skill-1")) + state.clients.set("session-2:skill-2:server-2", createManagedClient("skill-2")) + registerProcessCleanup(state) + + // when + await disconnectSession(state, "session-1") + + // then + expect(state.clients.has("session-2:skill-2:server-2")).toBe(true) + expect(state.cleanupRegistered).toBe(true) + expect(state.cleanupHandlers).toHaveLength(2) + expect(process.listenerCount("SIGINT")).toBe(signalIntCountBeforeRegister + 1) + expect(process.listenerCount("SIGTERM")).toBe(signalTermCountBeforeRegister + 1) + }) + + it("#given state with 2 clients in different sessions #when both sessions disconnected #then process cleanup handlers are unregistered", async () => { + // given + const state = createState() + const signalIntCountBeforeRegister = process.listenerCount("SIGINT") + const signalTermCountBeforeRegister = process.listenerCount("SIGTERM") + + state.clients.set("session-1:skill-1:server-1", createManagedClient("skill-1")) + state.clients.set("session-2:skill-2:server-2", createManagedClient("skill-2")) + registerProcessCleanup(state) + + // when + await disconnectSession(state, "session-1") + await disconnectSession(state, "session-2") + + // then + expect(state.clients.size).toBe(0) + expect(state.cleanupRegistered).toBe(false) + expect(state.cleanupHandlers).toEqual([]) + expect(process.listenerCount("SIGINT")).toBe(signalIntCountBeforeRegister) + expect(process.listenerCount("SIGTERM")).toBe(signalTermCountBeforeRegister) + }) +}) diff --git a/src/features/skill-mcp-manager/manager.ts b/src/features/skill-mcp-manager/manager.ts index 71d3cf781..150609efc 100644 --- a/src/features/skill-mcp-manager/manager.ts +++ b/src/features/skill-mcp-manager/manager.ts @@ -10,6 +10,7 @@ export class SkillMcpManager { private readonly state: SkillMcpManagerState = { clients: new Map(), pendingConnections: new Map(), + disconnectedSessions: new Set(), authProviders: new Map(), cleanupRegistered: false, cleanupInterval: null, diff --git a/src/features/skill-mcp-manager/types.ts b/src/features/skill-mcp-manager/types.ts index b7a9f46ca..7023749cb 100644 --- a/src/features/skill-mcp-manager/types.ts +++ b/src/features/skill-mcp-manager/types.ts @@ -51,6 +51,7 @@ export interface ProcessCleanupHandler { export interface SkillMcpManagerState { clients: Map pendingConnections: Map> + disconnectedSessions: Set authProviders: Map cleanupRegistered: boolean cleanupInterval: ReturnType | null