From 861ce1c1610d3c2fd6cd908dba7461752c964550 Mon Sep 17 00:00:00 2001 From: YeonGyu-Kim Date: Sat, 4 Apr 2026 20:06:56 +0900 Subject: [PATCH] test: remove provider and cache mock leak paths --- src/features/skill-mcp-manager/http-client.ts | 2 +- .../skill-mcp-manager/manager.test.ts | 16 +++----- src/features/skill-mcp-manager/manager.ts | 38 ++++++++++++------ .../skill-mcp-manager/oauth-handler.ts | 22 +++++----- src/features/skill-mcp-manager/types.ts | 12 ++++++ src/shared/model-capabilities.test.ts | 40 ++++++++----------- src/shared/model-error-classifier.test.ts | 35 +++++++--------- 7 files changed, 88 insertions(+), 77 deletions(-) diff --git a/src/features/skill-mcp-manager/http-client.ts b/src/features/skill-mcp-manager/http-client.ts index d3f00f292..74bc43598 100644 --- a/src/features/skill-mcp-manager/http-client.ts +++ b/src/features/skill-mcp-manager/http-client.ts @@ -42,7 +42,7 @@ export async function createHttpClient(params: SkillMcpClientConnectionParams): registerProcessCleanup(state) - const requestInit = await buildHttpRequestInit(config, state.authProviders) + const requestInit = await buildHttpRequestInit(config, state.authProviders, state.createOAuthProvider) const transport = new StreamableHTTPClientTransport(url, { requestInit, }) diff --git a/src/features/skill-mcp-manager/manager.test.ts b/src/features/skill-mcp-manager/manager.test.ts index 541ab18f8..ae648222f 100644 --- a/src/features/skill-mcp-manager/manager.test.ts +++ b/src/features/skill-mcp-manager/manager.test.ts @@ -7,7 +7,6 @@ const mockHttpConnect = mock(() => Promise.reject(new Error("Mocked HTTP connect const mockHttpClose = mock(() => Promise.resolve()) let lastTransportInstance: { url?: URL; options?: { requestInit?: RequestInit } } = {} -// Mock OAuth provider for OAuth integration tests const mockTokens = mock(() => null as { accessToken: string } | null) const mockLogin = mock(() => Promise.resolve({ accessToken: "test-token" }) as Promise<{ accessToken: string } | null>) @@ -26,14 +25,6 @@ async function importFreshManagerModule(): Promise { }, })) - mock.module("../mcp-oauth/provider", () => ({ - McpOAuthProvider: class MockMcpOAuthProvider { - tokens = mockTokens - login = mockLogin - constructor(_opts: unknown) {} - }, - })) - const module = await import(`./manager?test=${Date.now()}-${Math.random()}`) mock.restore() return module @@ -46,7 +37,12 @@ describe("SkillMcpManager", () => { beforeEach(async () => { const { SkillMcpManager } = await importFreshManagerModule() - manager = new SkillMcpManager() + manager = new SkillMcpManager({ + createOAuthProvider: () => ({ + tokens: () => mockTokens(), + login: () => mockLogin(), + }), + }) mockHttpConnect.mockClear() mockHttpClose.mockClear() mockTokens.mockClear() diff --git a/src/features/skill-mcp-manager/manager.ts b/src/features/skill-mcp-manager/manager.ts index 00980e987..473d5f390 100644 --- a/src/features/skill-mcp-manager/manager.ts +++ b/src/features/skill-mcp-manager/manager.ts @@ -1,24 +1,35 @@ import type { Client } from "@modelcontextprotocol/sdk/client/index.js" import type { Prompt, Resource, Tool } from "@modelcontextprotocol/sdk/types.js" import type { ClaudeCodeMcpServer } from "../claude-code-mcp-loader/types" +import { McpOAuthProvider } from "../mcp-oauth/provider" import { disconnectAll, disconnectSession, forceReconnect } from "./cleanup" import { getOrCreateClient, getOrCreateClientWithRetryImpl } from "./connection" import { handleStepUpIfNeeded } from "./oauth-handler" -import type { SkillMcpClientInfo, SkillMcpManagerState, SkillMcpServerContext } from "./types" +import type { + OAuthProviderFactory, + SkillMcpClientInfo, + SkillMcpManagerState, + SkillMcpServerContext, +} from "./types" export class SkillMcpManager { - private readonly state: SkillMcpManagerState = { - clients: new Map(), - pendingConnections: new Map(), - disconnectedSessions: new Map(), - authProviders: new Map(), - cleanupRegistered: false, - cleanupInterval: null, - cleanupHandlers: [], - idleTimeoutMs: 5 * 60 * 1000, - shutdownGeneration: 0, - inFlightConnections: new Map(), - disposed: false, + private readonly state: SkillMcpManagerState + + constructor(options: { createOAuthProvider?: OAuthProviderFactory } = {}) { + this.state = { + clients: new Map(), + pendingConnections: new Map(), + disconnectedSessions: new Map(), + authProviders: new Map(), + cleanupRegistered: false, + cleanupInterval: null, + cleanupHandlers: [], + idleTimeoutMs: 5 * 60 * 1000, + shutdownGeneration: 0, + inFlightConnections: new Map(), + disposed: false, + createOAuthProvider: options.createOAuthProvider ?? ((providerOptions) => new McpOAuthProvider(providerOptions)), + } } private getClientKey(info: SkillMcpClientInfo): string { @@ -112,6 +123,7 @@ export class SkillMcpManager { error: lastError, config, authProviders: this.state.authProviders, + createOAuthProvider: this.state.createOAuthProvider, }) if (stepUpHandled) { await forceReconnect(this.state, this.getClientKey(info)) diff --git a/src/features/skill-mcp-manager/oauth-handler.ts b/src/features/skill-mcp-manager/oauth-handler.ts index 66e12b3e6..c09c845f2 100644 --- a/src/features/skill-mcp-manager/oauth-handler.ts +++ b/src/features/skill-mcp-manager/oauth-handler.ts @@ -2,16 +2,18 @@ import type { ClaudeCodeMcpServer } from "../claude-code-mcp-loader/types" import { McpOAuthProvider } from "../mcp-oauth/provider" import type { OAuthTokenData } from "../mcp-oauth/storage" import { isStepUpRequired, mergeScopes } from "../mcp-oauth/step-up" +import type { OAuthProviderFactory, OAuthProviderLike } from "./types" export function getOrCreateAuthProvider( - authProviders: Map, + authProviders: Map, serverUrl: string, - oauth: NonNullable -): McpOAuthProvider { + oauth: NonNullable, + createOAuthProvider: OAuthProviderFactory = (options) => new McpOAuthProvider(options), +): OAuthProviderLike { const existing = authProviders.get(serverUrl) if (existing) return existing - const provider = new McpOAuthProvider({ + const provider = createOAuthProvider({ serverUrl, clientId: oauth.clientId, scopes: oauth.scopes, @@ -27,7 +29,8 @@ function isTokenExpired(tokenData: OAuthTokenData): boolean { export async function buildHttpRequestInit( config: ClaudeCodeMcpServer, - authProviders: Map + authProviders: Map, + createOAuthProvider?: OAuthProviderFactory, ): Promise { const headers: Record = {} @@ -38,7 +41,7 @@ export async function buildHttpRequestInit( } if (config.oauth && config.url) { - const provider = getOrCreateAuthProvider(authProviders, config.url, config.oauth) + const provider = getOrCreateAuthProvider(authProviders, config.url, config.oauth, createOAuthProvider) let tokenData = provider.tokens() if (!tokenData || isTokenExpired(tokenData)) { @@ -60,9 +63,10 @@ export async function buildHttpRequestInit( export async function handleStepUpIfNeeded(params: { error: Error config: ClaudeCodeMcpServer - authProviders: Map + authProviders: Map + createOAuthProvider?: OAuthProviderFactory }): Promise { - const { error, config, authProviders } = params + const { error, config, authProviders, createOAuthProvider } = params if (!config.oauth || !config.url) { return false @@ -89,7 +93,7 @@ export async function handleStepUpIfNeeded(params: { config.oauth.scopes = mergedScopes authProviders.delete(config.url) - const provider = getOrCreateAuthProvider(authProviders, config.url, config.oauth) + const provider = getOrCreateAuthProvider(authProviders, config.url, config.oauth, createOAuthProvider) try { await provider.login() diff --git a/src/features/skill-mcp-manager/types.ts b/src/features/skill-mcp-manager/types.ts index 1fb704a69..cbaa29204 100644 --- a/src/features/skill-mcp-manager/types.ts +++ b/src/features/skill-mcp-manager/types.ts @@ -48,6 +48,17 @@ export interface ProcessCleanupHandler { listener: () => void } +export type OAuthProviderLike = Pick< + McpOAuthProvider, + "tokens" | "login" +> + +export type OAuthProviderFactory = (options: { + serverUrl: string + clientId?: string + scopes?: string[] +}) => OAuthProviderLike + export interface SkillMcpManagerState { clients: Map pendingConnections: Map> @@ -60,6 +71,7 @@ export interface SkillMcpManagerState { shutdownGeneration: number inFlightConnections: Map disposed: boolean + createOAuthProvider: OAuthProviderFactory } export interface SkillMcpClientConnectionParams { diff --git a/src/shared/model-capabilities.test.ts b/src/shared/model-capabilities.test.ts index 711ee7906..736657af0 100644 --- a/src/shared/model-capabilities.test.ts +++ b/src/shared/model-capabilities.test.ts @@ -1,30 +1,17 @@ import type { ModelCapabilitiesSnapshot } from "./model-capabilities" -import { afterAll, describe, expect, test, mock } from "bun:test" - -afterAll(() => { - mock.restore() -}) - -async function importFreshModelCapabilitiesModule() { - // Mock connected-providers-cache to prevent local disk cache from polluting test results. - // Without this, findProviderModelMetadata reads real cached model metadata (e.g., from opencode serve) - // which causes the "prefers runtime models.dev cache" test to get different values than expected. - mock.module("./connected-providers-cache", () => ({ - findProviderModelMetadata: () => undefined, - readConnectedProvidersCache: () => null, - hasConnectedProvidersCache: () => false, - hasProviderModelsCache: () => false, - })) - - const module = await import(`./model-capabilities?test=${Date.now()}-${Math.random()}`) - mock.restore() - return module -} - -const { getModelCapabilities, getBundledModelCapabilitiesSnapshot } = await importFreshModelCapabilitiesModule() +import { afterEach, describe, expect, test, spyOn } from "bun:test" +import * as connectedProvidersCache from "./connected-providers-cache" +import { getModelCapabilities, getBundledModelCapabilitiesSnapshot } from "./model-capabilities" import { AGENT_MODEL_REQUIREMENTS, CATEGORY_MODEL_REQUIREMENTS } from "./model-requirements" describe("getModelCapabilities", () => { + let findProviderModelMetadataSpy: ReturnType | undefined + + afterEach(() => { + findProviderModelMetadataSpy?.mockRestore() + findProviderModelMetadataSpy = undefined + }) + const bundledSnapshot: ModelCapabilitiesSnapshot = { generatedAt: "2026-03-25T00:00:00.000Z", sourceUrl: "https://models.dev/api.json", @@ -76,6 +63,7 @@ describe("getModelCapabilities", () => { } test("uses runtime metadata before snapshot data", () => { + findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined) const result = getModelCapabilities({ providerID: "anthropic", modelID: "claude-opus-4-6", @@ -107,6 +95,7 @@ describe("getModelCapabilities", () => { }) test("reads structured runtime capabilities from the SDK v2 shape", () => { + findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined) const result = getModelCapabilities({ providerID: "openai", modelID: "gpt-5.4", @@ -147,6 +136,7 @@ describe("getModelCapabilities", () => { }) test("respects root-level thinking flags when providers do not nest them under capabilities", () => { + findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined) const result = getModelCapabilities({ providerID: "custom-proxy", modelID: "gpt-5.4", @@ -166,6 +156,7 @@ describe("getModelCapabilities", () => { }) test("accepts runtime variant arrays without corrupting them into numeric keys", () => { + findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined) const result = getModelCapabilities({ providerID: "openai", modelID: "gpt-5.4", @@ -179,6 +170,7 @@ describe("getModelCapabilities", () => { }) test("normalizes the legacy Claude Opus thinking alias before snapshot lookup", () => { + findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined) const result = getModelCapabilities({ providerID: "anthropic", modelID: "claude-opus-4-6-thinking", @@ -203,6 +195,7 @@ describe("getModelCapabilities", () => { }) test("maps local gemini aliases to canonical models.dev entries", () => { + findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined) const result = getModelCapabilities({ providerID: "google", modelID: "gemini-3.1-pro-high", @@ -227,6 +220,7 @@ describe("getModelCapabilities", () => { }) test("prefers runtime models.dev cache over bundled snapshot", () => { + findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined) const runtimeSnapshot: ModelCapabilitiesSnapshot = { ...bundledSnapshot, models: { diff --git a/src/shared/model-error-classifier.test.ts b/src/shared/model-error-classifier.test.ts index 60c2572a9..812b21e07 100644 --- a/src/shared/model-error-classifier.test.ts +++ b/src/shared/model-error-classifier.test.ts @@ -1,26 +1,19 @@ declare const require: (name: string) => any -const { describe, expect, test, beforeEach, mock, afterAll } = require("bun:test") +const { describe, expect, test, beforeEach, afterEach, mock, spyOn } = require("bun:test") +import * as connectedProvidersCache from "./connected-providers-cache" -const readConnectedProvidersCacheMock = mock(() => null) - -async function importFreshModelErrorClassifierModule() { - mock.module("./connected-providers-cache", () => ({ - readConnectedProvidersCache: readConnectedProvidersCacheMock, - })) - - const module = await import(`./model-error-classifier?test=${Date.now()}-${Math.random()}`) - mock.restore() - return module -} - -afterAll(() => { mock.restore() }) - -const { shouldRetryError, selectFallbackProvider } = await importFreshModelErrorClassifierModule() +let readConnectedProvidersCacheSpy: ReturnType | undefined +const { shouldRetryError, selectFallbackProvider } = await import("./model-error-classifier") describe("model-error-classifier", () => { beforeEach(() => { - readConnectedProvidersCacheMock.mockReturnValue(null) - readConnectedProvidersCacheMock.mockClear() + readConnectedProvidersCacheSpy?.mockRestore() + readConnectedProvidersCacheSpy = spyOn(connectedProvidersCache, "readConnectedProvidersCache").mockReturnValue(null) + }) + + afterEach(() => { + readConnectedProvidersCacheSpy?.mockRestore() + readConnectedProvidersCacheSpy = undefined }) test("treats overloaded retry messages as retryable", () => { @@ -50,7 +43,7 @@ describe("model-error-classifier", () => { test("selectFallbackProvider prefers first connected provider in preference order", () => { //#given - readConnectedProvidersCacheMock.mockReturnValue(["anthropic", "nvidia"]) + readConnectedProvidersCacheSpy?.mockReturnValue(["anthropic", "nvidia"]) //#when const provider = selectFallbackProvider(["anthropic", "nvidia"], "nvidia") @@ -61,7 +54,7 @@ describe("model-error-classifier", () => { test("selectFallbackProvider falls back to next connected provider when first is disconnected", () => { //#given - readConnectedProvidersCacheMock.mockReturnValue(["nvidia"]) + readConnectedProvidersCacheSpy?.mockReturnValue(["nvidia"]) //#when const provider = selectFallbackProvider(["anthropic", "nvidia"]) @@ -82,7 +75,7 @@ describe("model-error-classifier", () => { test("selectFallbackProvider uses connected preferred provider when fallback providers are unavailable", () => { //#given - readConnectedProvidersCacheMock.mockReturnValue(["provider-x"]) + readConnectedProvidersCacheSpy?.mockReturnValue(["provider-x"]) //#when const provider = selectFallbackProvider(["provider-y"], "provider-x")