test: remove provider and cache mock leak paths

This commit is contained in:
YeonGyu-Kim
2026-04-04 20:06:56 +09:00
parent 963b576b9e
commit 861ce1c161
7 changed files with 88 additions and 77 deletions
@@ -42,7 +42,7 @@ export async function createHttpClient(params: SkillMcpClientConnectionParams):
registerProcessCleanup(state) registerProcessCleanup(state)
const requestInit = await buildHttpRequestInit(config, state.authProviders) const requestInit = await buildHttpRequestInit(config, state.authProviders, state.createOAuthProvider)
const transport = new StreamableHTTPClientTransport(url, { const transport = new StreamableHTTPClientTransport(url, {
requestInit, requestInit,
}) })
+6 -10
View File
@@ -7,7 +7,6 @@ const mockHttpConnect = mock(() => Promise.reject(new Error("Mocked HTTP connect
const mockHttpClose = mock(() => Promise.resolve()) const mockHttpClose = mock(() => Promise.resolve())
let lastTransportInstance: { url?: URL; options?: { requestInit?: RequestInit } } = {} let lastTransportInstance: { url?: URL; options?: { requestInit?: RequestInit } } = {}
// Mock OAuth provider for OAuth integration tests
const mockTokens = mock(() => null as { accessToken: string } | null) const mockTokens = mock(() => null as { accessToken: string } | null)
const mockLogin = mock(() => Promise.resolve({ accessToken: "test-token" }) as Promise<{ accessToken: string } | null>) const mockLogin = mock(() => Promise.resolve({ accessToken: "test-token" }) as Promise<{ accessToken: string } | null>)
@@ -26,14 +25,6 @@ async function importFreshManagerModule(): Promise<typeof import("./manager")> {
}, },
})) }))
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()}`) const module = await import(`./manager?test=${Date.now()}-${Math.random()}`)
mock.restore() mock.restore()
return module return module
@@ -46,7 +37,12 @@ describe("SkillMcpManager", () => {
beforeEach(async () => { beforeEach(async () => {
const { SkillMcpManager } = await importFreshManagerModule() const { SkillMcpManager } = await importFreshManagerModule()
manager = new SkillMcpManager() manager = new SkillMcpManager({
createOAuthProvider: () => ({
tokens: () => mockTokens(),
login: () => mockLogin(),
}),
})
mockHttpConnect.mockClear() mockHttpConnect.mockClear()
mockHttpClose.mockClear() mockHttpClose.mockClear()
mockTokens.mockClear() mockTokens.mockClear()
+25 -13
View File
@@ -1,24 +1,35 @@
import type { Client } from "@modelcontextprotocol/sdk/client/index.js" import type { Client } from "@modelcontextprotocol/sdk/client/index.js"
import type { Prompt, Resource, Tool } from "@modelcontextprotocol/sdk/types.js" import type { Prompt, Resource, Tool } from "@modelcontextprotocol/sdk/types.js"
import type { ClaudeCodeMcpServer } from "../claude-code-mcp-loader/types" import type { ClaudeCodeMcpServer } from "../claude-code-mcp-loader/types"
import { McpOAuthProvider } from "../mcp-oauth/provider"
import { disconnectAll, disconnectSession, forceReconnect } from "./cleanup" import { disconnectAll, disconnectSession, forceReconnect } from "./cleanup"
import { getOrCreateClient, getOrCreateClientWithRetryImpl } from "./connection" import { getOrCreateClient, getOrCreateClientWithRetryImpl } from "./connection"
import { handleStepUpIfNeeded } from "./oauth-handler" import { handleStepUpIfNeeded } from "./oauth-handler"
import type { SkillMcpClientInfo, SkillMcpManagerState, SkillMcpServerContext } from "./types" import type {
OAuthProviderFactory,
SkillMcpClientInfo,
SkillMcpManagerState,
SkillMcpServerContext,
} from "./types"
export class SkillMcpManager { export class SkillMcpManager {
private readonly state: SkillMcpManagerState = { private readonly state: SkillMcpManagerState
clients: new Map(),
pendingConnections: new Map(), constructor(options: { createOAuthProvider?: OAuthProviderFactory } = {}) {
disconnectedSessions: new Map(), this.state = {
authProviders: new Map(), clients: new Map(),
cleanupRegistered: false, pendingConnections: new Map(),
cleanupInterval: null, disconnectedSessions: new Map(),
cleanupHandlers: [], authProviders: new Map(),
idleTimeoutMs: 5 * 60 * 1000, cleanupRegistered: false,
shutdownGeneration: 0, cleanupInterval: null,
inFlightConnections: new Map(), cleanupHandlers: [],
disposed: false, idleTimeoutMs: 5 * 60 * 1000,
shutdownGeneration: 0,
inFlightConnections: new Map(),
disposed: false,
createOAuthProvider: options.createOAuthProvider ?? ((providerOptions) => new McpOAuthProvider(providerOptions)),
}
} }
private getClientKey(info: SkillMcpClientInfo): string { private getClientKey(info: SkillMcpClientInfo): string {
@@ -112,6 +123,7 @@ export class SkillMcpManager {
error: lastError, error: lastError,
config, config,
authProviders: this.state.authProviders, authProviders: this.state.authProviders,
createOAuthProvider: this.state.createOAuthProvider,
}) })
if (stepUpHandled) { if (stepUpHandled) {
await forceReconnect(this.state, this.getClientKey(info)) await forceReconnect(this.state, this.getClientKey(info))
@@ -2,16 +2,18 @@ import type { ClaudeCodeMcpServer } from "../claude-code-mcp-loader/types"
import { McpOAuthProvider } from "../mcp-oauth/provider" import { McpOAuthProvider } from "../mcp-oauth/provider"
import type { OAuthTokenData } from "../mcp-oauth/storage" import type { OAuthTokenData } from "../mcp-oauth/storage"
import { isStepUpRequired, mergeScopes } from "../mcp-oauth/step-up" import { isStepUpRequired, mergeScopes } from "../mcp-oauth/step-up"
import type { OAuthProviderFactory, OAuthProviderLike } from "./types"
export function getOrCreateAuthProvider( export function getOrCreateAuthProvider(
authProviders: Map<string, McpOAuthProvider>, authProviders: Map<string, OAuthProviderLike>,
serverUrl: string, serverUrl: string,
oauth: NonNullable<ClaudeCodeMcpServer["oauth"]> oauth: NonNullable<ClaudeCodeMcpServer["oauth"]>,
): McpOAuthProvider { createOAuthProvider: OAuthProviderFactory = (options) => new McpOAuthProvider(options),
): OAuthProviderLike {
const existing = authProviders.get(serverUrl) const existing = authProviders.get(serverUrl)
if (existing) return existing if (existing) return existing
const provider = new McpOAuthProvider({ const provider = createOAuthProvider({
serverUrl, serverUrl,
clientId: oauth.clientId, clientId: oauth.clientId,
scopes: oauth.scopes, scopes: oauth.scopes,
@@ -27,7 +29,8 @@ function isTokenExpired(tokenData: OAuthTokenData): boolean {
export async function buildHttpRequestInit( export async function buildHttpRequestInit(
config: ClaudeCodeMcpServer, config: ClaudeCodeMcpServer,
authProviders: Map<string, McpOAuthProvider> authProviders: Map<string, OAuthProviderLike>,
createOAuthProvider?: OAuthProviderFactory,
): Promise<RequestInit | undefined> { ): Promise<RequestInit | undefined> {
const headers: Record<string, string> = {} const headers: Record<string, string> = {}
@@ -38,7 +41,7 @@ export async function buildHttpRequestInit(
} }
if (config.oauth && config.url) { 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() let tokenData = provider.tokens()
if (!tokenData || isTokenExpired(tokenData)) { if (!tokenData || isTokenExpired(tokenData)) {
@@ -60,9 +63,10 @@ export async function buildHttpRequestInit(
export async function handleStepUpIfNeeded(params: { export async function handleStepUpIfNeeded(params: {
error: Error error: Error
config: ClaudeCodeMcpServer config: ClaudeCodeMcpServer
authProviders: Map<string, McpOAuthProvider> authProviders: Map<string, OAuthProviderLike>
createOAuthProvider?: OAuthProviderFactory
}): Promise<boolean> { }): Promise<boolean> {
const { error, config, authProviders } = params const { error, config, authProviders, createOAuthProvider } = params
if (!config.oauth || !config.url) { if (!config.oauth || !config.url) {
return false return false
@@ -89,7 +93,7 @@ export async function handleStepUpIfNeeded(params: {
config.oauth.scopes = mergedScopes config.oauth.scopes = mergedScopes
authProviders.delete(config.url) authProviders.delete(config.url)
const provider = getOrCreateAuthProvider(authProviders, config.url, config.oauth) const provider = getOrCreateAuthProvider(authProviders, config.url, config.oauth, createOAuthProvider)
try { try {
await provider.login() await provider.login()
+12
View File
@@ -48,6 +48,17 @@ export interface ProcessCleanupHandler {
listener: () => void listener: () => void
} }
export type OAuthProviderLike = Pick<
McpOAuthProvider,
"tokens" | "login"
>
export type OAuthProviderFactory = (options: {
serverUrl: string
clientId?: string
scopes?: string[]
}) => OAuthProviderLike
export interface SkillMcpManagerState { export interface SkillMcpManagerState {
clients: Map<string, ManagedClient> clients: Map<string, ManagedClient>
pendingConnections: Map<string, Promise<Client>> pendingConnections: Map<string, Promise<Client>>
@@ -60,6 +71,7 @@ export interface SkillMcpManagerState {
shutdownGeneration: number shutdownGeneration: number
inFlightConnections: Map<string, number> inFlightConnections: Map<string, number>
disposed: boolean disposed: boolean
createOAuthProvider: OAuthProviderFactory
} }
export interface SkillMcpClientConnectionParams { export interface SkillMcpClientConnectionParams {
+17 -23
View File
@@ -1,30 +1,17 @@
import type { ModelCapabilitiesSnapshot } from "./model-capabilities" import type { ModelCapabilitiesSnapshot } from "./model-capabilities"
import { afterAll, describe, expect, test, mock } from "bun:test" import { afterEach, describe, expect, test, spyOn } from "bun:test"
import * as connectedProvidersCache from "./connected-providers-cache"
afterAll(() => { import { getModelCapabilities, getBundledModelCapabilitiesSnapshot } from "./model-capabilities"
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 { AGENT_MODEL_REQUIREMENTS, CATEGORY_MODEL_REQUIREMENTS } from "./model-requirements" import { AGENT_MODEL_REQUIREMENTS, CATEGORY_MODEL_REQUIREMENTS } from "./model-requirements"
describe("getModelCapabilities", () => { describe("getModelCapabilities", () => {
let findProviderModelMetadataSpy: ReturnType<typeof spyOn> | undefined
afterEach(() => {
findProviderModelMetadataSpy?.mockRestore()
findProviderModelMetadataSpy = undefined
})
const bundledSnapshot: ModelCapabilitiesSnapshot = { const bundledSnapshot: ModelCapabilitiesSnapshot = {
generatedAt: "2026-03-25T00:00:00.000Z", generatedAt: "2026-03-25T00:00:00.000Z",
sourceUrl: "https://models.dev/api.json", sourceUrl: "https://models.dev/api.json",
@@ -76,6 +63,7 @@ describe("getModelCapabilities", () => {
} }
test("uses runtime metadata before snapshot data", () => { test("uses runtime metadata before snapshot data", () => {
findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined)
const result = getModelCapabilities({ const result = getModelCapabilities({
providerID: "anthropic", providerID: "anthropic",
modelID: "claude-opus-4-6", modelID: "claude-opus-4-6",
@@ -107,6 +95,7 @@ describe("getModelCapabilities", () => {
}) })
test("reads structured runtime capabilities from the SDK v2 shape", () => { test("reads structured runtime capabilities from the SDK v2 shape", () => {
findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined)
const result = getModelCapabilities({ const result = getModelCapabilities({
providerID: "openai", providerID: "openai",
modelID: "gpt-5.4", modelID: "gpt-5.4",
@@ -147,6 +136,7 @@ describe("getModelCapabilities", () => {
}) })
test("respects root-level thinking flags when providers do not nest them under capabilities", () => { test("respects root-level thinking flags when providers do not nest them under capabilities", () => {
findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined)
const result = getModelCapabilities({ const result = getModelCapabilities({
providerID: "custom-proxy", providerID: "custom-proxy",
modelID: "gpt-5.4", modelID: "gpt-5.4",
@@ -166,6 +156,7 @@ describe("getModelCapabilities", () => {
}) })
test("accepts runtime variant arrays without corrupting them into numeric keys", () => { test("accepts runtime variant arrays without corrupting them into numeric keys", () => {
findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined)
const result = getModelCapabilities({ const result = getModelCapabilities({
providerID: "openai", providerID: "openai",
modelID: "gpt-5.4", modelID: "gpt-5.4",
@@ -179,6 +170,7 @@ describe("getModelCapabilities", () => {
}) })
test("normalizes the legacy Claude Opus thinking alias before snapshot lookup", () => { test("normalizes the legacy Claude Opus thinking alias before snapshot lookup", () => {
findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined)
const result = getModelCapabilities({ const result = getModelCapabilities({
providerID: "anthropic", providerID: "anthropic",
modelID: "claude-opus-4-6-thinking", modelID: "claude-opus-4-6-thinking",
@@ -203,6 +195,7 @@ describe("getModelCapabilities", () => {
}) })
test("maps local gemini aliases to canonical models.dev entries", () => { test("maps local gemini aliases to canonical models.dev entries", () => {
findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined)
const result = getModelCapabilities({ const result = getModelCapabilities({
providerID: "google", providerID: "google",
modelID: "gemini-3.1-pro-high", modelID: "gemini-3.1-pro-high",
@@ -227,6 +220,7 @@ describe("getModelCapabilities", () => {
}) })
test("prefers runtime models.dev cache over bundled snapshot", () => { test("prefers runtime models.dev cache over bundled snapshot", () => {
findProviderModelMetadataSpy = spyOn(connectedProvidersCache, "findProviderModelMetadata").mockReturnValue(undefined)
const runtimeSnapshot: ModelCapabilitiesSnapshot = { const runtimeSnapshot: ModelCapabilitiesSnapshot = {
...bundledSnapshot, ...bundledSnapshot,
models: { models: {
+14 -21
View File
@@ -1,26 +1,19 @@
declare const require: (name: string) => any 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) let readConnectedProvidersCacheSpy: ReturnType<typeof spyOn> | undefined
const { shouldRetryError, selectFallbackProvider } = await import("./model-error-classifier")
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()
describe("model-error-classifier", () => { describe("model-error-classifier", () => {
beforeEach(() => { beforeEach(() => {
readConnectedProvidersCacheMock.mockReturnValue(null) readConnectedProvidersCacheSpy?.mockRestore()
readConnectedProvidersCacheMock.mockClear() readConnectedProvidersCacheSpy = spyOn(connectedProvidersCache, "readConnectedProvidersCache").mockReturnValue(null)
})
afterEach(() => {
readConnectedProvidersCacheSpy?.mockRestore()
readConnectedProvidersCacheSpy = undefined
}) })
test("treats overloaded retry messages as retryable", () => { test("treats overloaded retry messages as retryable", () => {
@@ -50,7 +43,7 @@ describe("model-error-classifier", () => {
test("selectFallbackProvider prefers first connected provider in preference order", () => { test("selectFallbackProvider prefers first connected provider in preference order", () => {
//#given //#given
readConnectedProvidersCacheMock.mockReturnValue(["anthropic", "nvidia"]) readConnectedProvidersCacheSpy?.mockReturnValue(["anthropic", "nvidia"])
//#when //#when
const provider = selectFallbackProvider(["anthropic", "nvidia"], "nvidia") 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", () => { test("selectFallbackProvider falls back to next connected provider when first is disconnected", () => {
//#given //#given
readConnectedProvidersCacheMock.mockReturnValue(["nvidia"]) readConnectedProvidersCacheSpy?.mockReturnValue(["nvidia"])
//#when //#when
const provider = selectFallbackProvider(["anthropic", "nvidia"]) const provider = selectFallbackProvider(["anthropic", "nvidia"])
@@ -82,7 +75,7 @@ describe("model-error-classifier", () => {
test("selectFallbackProvider uses connected preferred provider when fallback providers are unavailable", () => { test("selectFallbackProvider uses connected preferred provider when fallback providers are unavailable", () => {
//#given //#given
readConnectedProvidersCacheMock.mockReturnValue(["provider-x"]) readConnectedProvidersCacheSpy?.mockReturnValue(["provider-x"])
//#when //#when
const provider = selectFallbackProvider(["provider-y"], "provider-x") const provider = selectFallbackProvider(["provider-y"], "provider-x")