diff --git a/src/shared/connected-providers-cache.test.ts b/src/shared/connected-providers-cache.test.ts index 0a21e22fe..a170c90b9 100644 --- a/src/shared/connected-providers-cache.test.ts +++ b/src/shared/connected-providers-cache.test.ts @@ -1,27 +1,47 @@ -import { describe, test, expect, beforeEach, afterEach, spyOn } from "bun:test" -import { existsSync, mkdirSync, rmSync } from "fs" -import { join } from "path" -import * as dataPath from "./data-path" -import { updateConnectedProvidersCache, readProviderModelsCache } from "./connected-providers-cache" +/// -const TEST_CACHE_DIR = join(import.meta.dir, "__test-cache__") +import { beforeAll, beforeEach, afterEach, describe, expect, mock, test } from "bun:test" + +import { existsSync, mkdtempSync, rmSync } from "node:fs" +import { tmpdir } from "node:os" +import { join } from "node:path" +import * as dataPath from "./data-path" + +let testCacheDir = "" +let moduleImportCounter = 0 + +const getOmoOpenCodeCacheDirMock = mock(() => testCacheDir) + +let updateConnectedProvidersCache: typeof import("./connected-providers-cache").updateConnectedProvidersCache +let readProviderModelsCache: typeof import("./connected-providers-cache").readProviderModelsCache describe("updateConnectedProvidersCache", () => { - let cacheDirSpy: ReturnType + beforeAll(() => { + mock.restore() + }) - beforeEach(() => { - cacheDirSpy = spyOn(dataPath, "getOmoOpenCodeCacheDir").mockReturnValue(TEST_CACHE_DIR) - if (existsSync(TEST_CACHE_DIR)) { - rmSync(TEST_CACHE_DIR, { recursive: true }) + beforeEach(async () => { + mock.restore() + const realCacheDir = join(dataPath.getCacheDir(), "oh-my-opencode") + if (existsSync(realCacheDir)) { + rmSync(realCacheDir, { recursive: true, force: true }) } - mkdirSync(TEST_CACHE_DIR, { recursive: true }) + + testCacheDir = mkdtempSync(join(tmpdir(), "connected-providers-cache-test-")) + getOmoOpenCodeCacheDirMock.mockClear() + mock.module("./data-path", () => ({ + getOmoOpenCodeCacheDir: getOmoOpenCodeCacheDirMock, + })) + moduleImportCounter += 1 + ;({ updateConnectedProvidersCache, readProviderModelsCache } = await import(`./connected-providers-cache?test=${moduleImportCounter}`)) }) afterEach(() => { - cacheDirSpy.mockRestore() - if (existsSync(TEST_CACHE_DIR)) { - rmSync(TEST_CACHE_DIR, { recursive: true }) + mock.restore() + if (existsSync(testCacheDir)) { + rmSync(testCacheDir, { recursive: true, force: true }) } + testCacheDir = "" }) test("extracts models from provider.list().all response", async () => { diff --git a/src/shared/model-availability.test.ts b/src/shared/model-availability.test.ts index be1ec8332..7be5e363c 100644 --- a/src/shared/model-availability.test.ts +++ b/src/shared/model-availability.test.ts @@ -1,8 +1,9 @@ declare const require: (name: string) => any -const { describe, it, expect, beforeEach, afterEach, beforeAll } = require("bun:test") -import { mkdtempSync, writeFileSync, rmSync } from "fs" +const { describe, it, expect, beforeEach, afterEach, beforeAll, spyOn } = require("bun:test") +import { mkdtempSync, writeFileSync, rmSync, existsSync, readFileSync } from "fs" import { tmpdir } from "os" import { join } from "path" +import * as connectedProvidersCache from "./connected-providers-cache" let __resetModelCache: () => void let fetchAvailableModels: (client?: unknown, options?: { connectedProviders?: string[] | null }) => Promise> @@ -33,25 +34,27 @@ beforeAll(async () => { }) describe("fetchAvailableModels", () => { - let tempDir: string + let tempDir: string let originalXdgCache: string | undefined + let providerModelsCacheSpy: { mockRestore(): void } | undefined - - beforeEach(() => { - __resetModelCache() - tempDir = mkdtempSync(join(tmpdir(), "opencode-test-")) + beforeEach(() => { + __resetModelCache() + tempDir = mkdtempSync(join(tmpdir(), "opencode-test-")) originalXdgCache = process.env.XDG_CACHE_HOME process.env.XDG_CACHE_HOME = tempDir - }) + providerModelsCacheSpy = spyOn(connectedProvidersCache, "readProviderModelsCache").mockReturnValue(null) + }) - afterEach(() => { - if (originalXdgCache !== undefined) { + afterEach(() => { + providerModelsCacheSpy?.mockRestore() + if (originalXdgCache !== undefined) { process.env.XDG_CACHE_HOME = originalXdgCache } else { delete process.env.XDG_CACHE_HOME } - rmSync(tempDir, { recursive: true, force: true }) - }) + rmSync(tempDir, { recursive: true, force: true }) + }) function writeModelsCache(data: Record) { const cacheDir = join(tempDir, "opencode") @@ -485,15 +488,18 @@ describe("getConnectedProviders", () => { describe("fetchAvailableModels with connected providers filtering", () => { let tempDir: string let originalXdgCache: string | undefined + let providerModelsCacheSpy: { mockRestore(): void } | undefined beforeEach(() => { __resetModelCache() tempDir = mkdtempSync(join(tmpdir(), "opencode-test-")) originalXdgCache = process.env.XDG_CACHE_HOME process.env.XDG_CACHE_HOME = tempDir + providerModelsCacheSpy = spyOn(connectedProvidersCache, "readProviderModelsCache").mockReturnValue(null) }) afterEach(() => { + providerModelsCacheSpy?.mockRestore() if (originalXdgCache !== undefined) { process.env.XDG_CACHE_HOME = originalXdgCache } else { @@ -652,15 +658,24 @@ describe("fetchAvailableModels with connected providers filtering", () => { describe("fetchAvailableModels with provider-models cache (whitelist-filtered)", () => { let tempDir: string let originalXdgCache: string | undefined + let providerModelsCacheSpy: { mockRestore(): void } | undefined beforeEach(() => { __resetModelCache() tempDir = mkdtempSync(join(tmpdir(), "opencode-test-")) originalXdgCache = process.env.XDG_CACHE_HOME process.env.XDG_CACHE_HOME = tempDir + providerModelsCacheSpy = spyOn(connectedProvidersCache, "readProviderModelsCache").mockImplementation(() => { + const cacheFile = join(tempDir, "oh-my-opencode", "provider-models.json") + if (!existsSync(cacheFile)) { + return null + } + return JSON.parse(readFileSync(cacheFile, "utf-8")) + }) }) afterEach(() => { + providerModelsCacheSpy?.mockRestore() if (originalXdgCache !== undefined) { process.env.XDG_CACHE_HOME = originalXdgCache } else { @@ -878,21 +893,23 @@ describe("isModelAvailable", () => { describe("fallback model availability", () => { let tempDir: string - let originalXdgCache: string | undefined + let connectedProvidersCacheSpy: { mockRestore(): void } | undefined beforeEach(() => { // given tempDir = mkdtempSync(join(tmpdir(), "opencode-test-")) - originalXdgCache = process.env.XDG_CACHE_HOME - process.env.XDG_CACHE_HOME = tempDir + connectedProvidersCacheSpy = spyOn(connectedProvidersCache, "readConnectedProvidersCache").mockImplementation(() => { + const cacheFile = join(tempDir, "oh-my-opencode", "connected-providers.json") + if (!existsSync(cacheFile)) { + return null + } + const cache = JSON.parse(readFileSync(cacheFile, "utf-8")) as { connected?: string[] } + return Array.isArray(cache.connected) ? cache.connected : null + }) }) afterEach(() => { - if (originalXdgCache !== undefined) { - process.env.XDG_CACHE_HOME = originalXdgCache - } else { - delete process.env.XDG_CACHE_HOME - } + connectedProvidersCacheSpy?.mockRestore() rmSync(tempDir, { recursive: true, force: true }) }) diff --git a/src/shared/model-error-classifier.test.ts b/src/shared/model-error-classifier.test.ts index 17470199c..d359c26d3 100644 --- a/src/shared/model-error-classifier.test.ts +++ b/src/shared/model-error-classifier.test.ts @@ -1,28 +1,18 @@ -import { describe, expect, test, beforeEach, afterEach, spyOn } from "bun:test" +declare const require: (name: string) => any +const { describe, expect, test, beforeEach, mock } = require("bun:test") + +const readConnectedProvidersCacheMock = mock(() => null) + +mock.module("./connected-providers-cache", () => ({ + readConnectedProvidersCache: readConnectedProvidersCacheMock, +})) -import { mkdirSync, rmSync, writeFileSync, existsSync } from "node:fs" -import { join } from "node:path" -import * as dataPath from "./data-path" import { shouldRetryError, selectFallbackProvider } from "./model-error-classifier" -const TEST_CACHE_DIR = join(import.meta.dir, "__test-cache__") - describe("model-error-classifier", () => { - let cacheDirSpy: ReturnType - beforeEach(() => { - cacheDirSpy = spyOn(dataPath, "getOmoOpenCodeCacheDir").mockReturnValue(TEST_CACHE_DIR) - if (existsSync(TEST_CACHE_DIR)) { - rmSync(TEST_CACHE_DIR, { recursive: true }) - } - mkdirSync(TEST_CACHE_DIR, { recursive: true }) - }) - - afterEach(() => { - cacheDirSpy.mockRestore() - if (existsSync(TEST_CACHE_DIR)) { - rmSync(TEST_CACHE_DIR, { recursive: true }) - } + readConnectedProvidersCacheMock.mockReturnValue(null) + readConnectedProvidersCacheMock.mockClear() }) test("treats overloaded retry messages as retryable", () => { @@ -52,10 +42,7 @@ describe("model-error-classifier", () => { test("selectFallbackProvider prefers first connected provider in preference order", () => { //#given - writeFileSync( - join(TEST_CACHE_DIR, "connected-providers.json"), - JSON.stringify({ connected: ["anthropic", "nvidia"], updatedAt: new Date().toISOString() }, null, 2), - ) + readConnectedProvidersCacheMock.mockReturnValue(["anthropic", "nvidia"]) //#when const provider = selectFallbackProvider(["anthropic", "nvidia"], "nvidia") @@ -66,10 +53,7 @@ describe("model-error-classifier", () => { test("selectFallbackProvider falls back to next connected provider when first is disconnected", () => { //#given - writeFileSync( - join(TEST_CACHE_DIR, "connected-providers.json"), - JSON.stringify({ connected: ["nvidia"], updatedAt: new Date().toISOString() }, null, 2), - ) + readConnectedProvidersCacheMock.mockReturnValue(["nvidia"]) //#when const provider = selectFallbackProvider(["anthropic", "nvidia"]) @@ -90,10 +74,7 @@ describe("model-error-classifier", () => { test("selectFallbackProvider uses connected preferred provider when fallback providers are unavailable", () => { //#given - writeFileSync( - join(TEST_CACHE_DIR, "connected-providers.json"), - JSON.stringify({ connected: ["provider-x"], updatedAt: new Date().toISOString() }, null, 2), - ) + readConnectedProvidersCacheMock.mockReturnValue(["provider-x"]) //#when const provider = selectFallbackProvider(["provider-y"], "provider-x")