test(ci): remove suite-order mock coupling
This commit is contained in:
@@ -1,11 +1,13 @@
|
||||
import { afterAll, beforeEach, describe, expect, mock, test } from "bun:test"
|
||||
import { tryFallbackRetry, type FallbackRetryHandlerDeps } from "./fallback-retry-handler"
|
||||
import type { FallbackEntry } from "../../shared/model-requirements"
|
||||
|
||||
const sharedLogMock = mock(() => {})
|
||||
const readConnectedProvidersCacheMock = mock(() => null)
|
||||
const readProviderModelsCacheMock = mock((): { connected: string[] } | null => null)
|
||||
const shouldRetryErrorMock = mock(() => true)
|
||||
const getNextFallbackMock = mock((chain: Array<{ model: string }>, attempt: number) => chain[attempt])
|
||||
const hasMoreFallbacksMock = mock((chain: Array<{ model: string }>, attempt: number) => attempt < chain.length)
|
||||
const getNextFallbackMock = mock((chain: FallbackEntry[], attempt: number) => chain[attempt])
|
||||
const hasMoreFallbacksMock = mock((chain: FallbackEntry[], attempt: number) => attempt < chain.length)
|
||||
const selectFallbackProviderMock = mock((providers: string[]) => providers[0])
|
||||
const transformModelForProviderMock = mock((_provider: string, model: string) => model)
|
||||
|
||||
@@ -13,41 +15,17 @@ import type { BackgroundTask } from "./types"
|
||||
import type { ConcurrencyManager } from "./concurrency"
|
||||
import type { OpencodeClient, QueueItem } from "./constants"
|
||||
|
||||
async function importFreshFallbackRetryHandlerModule() {
|
||||
mock.module("../../shared/logger", () => ({
|
||||
log: sharedLogMock,
|
||||
}))
|
||||
|
||||
mock.module("../../shared/connected-providers-cache", () => ({
|
||||
readConnectedProvidersCache: readConnectedProvidersCacheMock,
|
||||
readProviderModelsCache: readProviderModelsCacheMock,
|
||||
}))
|
||||
|
||||
mock.module("../../shared/model-error-classifier", () => ({
|
||||
shouldRetryError: shouldRetryErrorMock,
|
||||
getNextFallback: getNextFallbackMock,
|
||||
hasMoreFallbacks: hasMoreFallbacksMock,
|
||||
selectFallbackProvider: selectFallbackProviderMock,
|
||||
}))
|
||||
|
||||
mock.module("../../shared/provider-model-id-transform", () => ({
|
||||
transformModelForProvider: transformModelForProviderMock,
|
||||
}))
|
||||
|
||||
const retryHandlerModule = await import(`./fallback-retry-handler?test=${Date.now()}-${Math.random()}`)
|
||||
mock.restore()
|
||||
|
||||
return {
|
||||
tryFallbackRetry: retryHandlerModule.tryFallbackRetry,
|
||||
shouldRetryError: shouldRetryErrorMock,
|
||||
selectFallbackProvider: selectFallbackProviderMock,
|
||||
readProviderModelsCache: readProviderModelsCacheMock,
|
||||
}
|
||||
const retryHandlerDeps: Partial<FallbackRetryHandlerDeps> = {
|
||||
log: sharedLogMock,
|
||||
readConnectedProvidersCache: readConnectedProvidersCacheMock,
|
||||
readProviderModelsCache: readProviderModelsCacheMock,
|
||||
shouldRetryError: shouldRetryErrorMock,
|
||||
getNextFallback: getNextFallbackMock,
|
||||
hasMoreFallbacks: hasMoreFallbacksMock,
|
||||
selectFallbackProvider: selectFallbackProviderMock,
|
||||
transformModelForProvider: transformModelForProviderMock,
|
||||
}
|
||||
|
||||
const { tryFallbackRetry, shouldRetryError, selectFallbackProvider, readProviderModelsCache } =
|
||||
await importFreshFallbackRetryHandlerModule()
|
||||
|
||||
function createDeferredPromise(): {
|
||||
promise: Promise<void>
|
||||
resolve: () => void
|
||||
@@ -124,6 +102,7 @@ function createDefaultArgs(taskOverrides: Partial<BackgroundTask> = {}) {
|
||||
idleDeferralTimers,
|
||||
queuesByKey,
|
||||
processKey: processKeyFn,
|
||||
deps: retryHandlerDeps,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -133,9 +112,13 @@ describe("tryFallbackRetry", () => {
|
||||
})
|
||||
|
||||
beforeEach(() => {
|
||||
shouldRetryError.mockImplementation(() => true)
|
||||
selectFallbackProvider.mockImplementation((providers: string[]) => providers[0])
|
||||
readProviderModelsCache.mockReturnValue(null)
|
||||
shouldRetryErrorMock.mockImplementation(() => true)
|
||||
selectFallbackProviderMock.mockImplementation((providers: string[]) => providers[0])
|
||||
readProviderModelsCacheMock.mockReturnValue(null)
|
||||
readConnectedProvidersCacheMock.mockReturnValue(null)
|
||||
getNextFallbackMock.mockImplementation((chain: FallbackEntry[], attempt: number) => chain[attempt])
|
||||
hasMoreFallbacksMock.mockImplementation((chain: FallbackEntry[], attempt: number) => attempt < chain.length)
|
||||
transformModelForProviderMock.mockImplementation((_provider: string, model: string) => model)
|
||||
})
|
||||
|
||||
describe("#given retryable error with fallback chain", () => {
|
||||
@@ -332,7 +315,7 @@ describe("tryFallbackRetry", () => {
|
||||
|
||||
describe("#given non-retryable error", () => {
|
||||
test("returns false when shouldRetryError returns false", async () => {
|
||||
shouldRetryError.mockImplementation(() => false)
|
||||
shouldRetryErrorMock.mockImplementation(() => false)
|
||||
const args = createDefaultArgs()
|
||||
|
||||
const result = await tryFallbackRetry(args)
|
||||
@@ -433,8 +416,8 @@ describe("tryFallbackRetry", () => {
|
||||
|
||||
describe("#given disconnected fallback providers with connected preferred provider", () => {
|
||||
test("keeps fallback entry and selects connected preferred provider", async () => {
|
||||
readProviderModelsCache.mockReturnValueOnce({ connected: ["provider-a"] })
|
||||
selectFallbackProvider.mockImplementationOnce(
|
||||
readProviderModelsCacheMock.mockReturnValueOnce({ connected: ["provider-a"] })
|
||||
selectFallbackProviderMock.mockImplementationOnce(
|
||||
(_providers: string[], preferredProviderID?: string) => preferredProviderID ?? "provider-b",
|
||||
)
|
||||
|
||||
|
||||
@@ -17,6 +17,28 @@ function canonicalizeModelID(modelID: string): string {
|
||||
return modelID.toLowerCase().replace(/\./g, "-")
|
||||
}
|
||||
|
||||
export type FallbackRetryHandlerDeps = {
|
||||
log: typeof log
|
||||
readProviderModelsCache: typeof readProviderModelsCache
|
||||
readConnectedProvidersCache: typeof readConnectedProvidersCache
|
||||
shouldRetryError: typeof shouldRetryError
|
||||
getNextFallback: typeof getNextFallback
|
||||
hasMoreFallbacks: typeof hasMoreFallbacks
|
||||
selectFallbackProvider: typeof selectFallbackProvider
|
||||
transformModelForProvider: typeof transformModelForProvider
|
||||
}
|
||||
|
||||
const defaultFallbackRetryHandlerDeps: FallbackRetryHandlerDeps = {
|
||||
log,
|
||||
readProviderModelsCache,
|
||||
readConnectedProvidersCache,
|
||||
shouldRetryError,
|
||||
getNextFallback,
|
||||
hasMoreFallbacks,
|
||||
selectFallbackProvider,
|
||||
transformModelForProvider,
|
||||
}
|
||||
|
||||
export async function tryFallbackRetry(args: {
|
||||
task: BackgroundTask
|
||||
errorInfo: { name?: string; message?: string }
|
||||
@@ -34,20 +56,22 @@ export async function tryFallbackRetry(args: {
|
||||
failedError?: string
|
||||
nextModel: string
|
||||
}) => void
|
||||
deps?: Partial<FallbackRetryHandlerDeps>
|
||||
}): Promise<boolean> {
|
||||
const { task, errorInfo, source, concurrencyManager, client, idleDeferralTimers, queuesByKey, processKey, onRetrying } = args
|
||||
const deps = { ...defaultFallbackRetryHandlerDeps, ...args.deps }
|
||||
const fallbackChain = task.fallbackChain
|
||||
const canRetry =
|
||||
shouldRetryError(errorInfo) &&
|
||||
deps.shouldRetryError(errorInfo) &&
|
||||
fallbackChain &&
|
||||
fallbackChain.length > 0 &&
|
||||
hasMoreFallbacks(fallbackChain, task.attemptCount ?? 0)
|
||||
deps.hasMoreFallbacks(fallbackChain, task.attemptCount ?? 0)
|
||||
|
||||
if (!canRetry) return false
|
||||
|
||||
const attemptCount = task.attemptCount ?? 0
|
||||
const providerModelsCache = readProviderModelsCache()
|
||||
const connectedProviders = providerModelsCache?.connected ?? readConnectedProvidersCache()
|
||||
const providerModelsCache = deps.readProviderModelsCache()
|
||||
const connectedProviders = providerModelsCache?.connected ?? deps.readConnectedProvidersCache()
|
||||
const connectedSet = connectedProviders ? new Set(connectedProviders.map(p => p.toLowerCase())) : null
|
||||
const preferredProvider = task.model?.providerID?.toLowerCase()
|
||||
|
||||
@@ -63,11 +87,11 @@ export async function tryFallbackRetry(args: {
|
||||
let nextFallback: FallbackEntry | undefined
|
||||
let nextProviderID: string | undefined
|
||||
while (fallbackChain && selectedAttemptCount < fallbackChain.length) {
|
||||
const candidate = getNextFallback(fallbackChain, selectedAttemptCount)
|
||||
const candidate = deps.getNextFallback(fallbackChain, selectedAttemptCount)
|
||||
if (!candidate) break
|
||||
selectedAttemptCount++
|
||||
if (!isReachable(candidate)) {
|
||||
log("[background-agent] Skipping unreachable fallback:", {
|
||||
deps.log("[background-agent] Skipping unreachable fallback:", {
|
||||
taskId: task.id,
|
||||
source,
|
||||
model: candidate.model,
|
||||
@@ -75,17 +99,17 @@ export async function tryFallbackRetry(args: {
|
||||
})
|
||||
continue
|
||||
}
|
||||
const candidateProviderID = selectFallbackProvider(
|
||||
const candidateProviderID = deps.selectFallbackProvider(
|
||||
candidate.providers,
|
||||
task.model?.providerID,
|
||||
)
|
||||
const candidateModelID = transformModelForProvider(candidateProviderID, candidate.model)
|
||||
const candidateModelID = deps.transformModelForProvider(candidateProviderID, candidate.model)
|
||||
const isNoOpFallback =
|
||||
!!task.model &&
|
||||
candidateProviderID.toLowerCase() === task.model.providerID.toLowerCase() &&
|
||||
canonicalizeModelID(candidateModelID) === canonicalizeModelID(task.model.modelID)
|
||||
if (isNoOpFallback) {
|
||||
log("[background-agent] Skipping no-op fallback:", {
|
||||
deps.log("[background-agent] Skipping no-op fallback:", {
|
||||
taskId: task.id,
|
||||
source,
|
||||
model: candidate.model,
|
||||
@@ -99,12 +123,12 @@ export async function tryFallbackRetry(args: {
|
||||
}
|
||||
if (!nextFallback) return false
|
||||
|
||||
const providerID = nextProviderID ?? selectFallbackProvider(
|
||||
const providerID = nextProviderID ?? deps.selectFallbackProvider(
|
||||
nextFallback.providers,
|
||||
task.model?.providerID,
|
||||
)
|
||||
|
||||
log("[background-agent] Retryable error, attempting fallback:", {
|
||||
deps.log("[background-agent] Retryable error, attempting fallback:", {
|
||||
taskId: task.id,
|
||||
source,
|
||||
errorName: errorInfo.name,
|
||||
@@ -127,7 +151,7 @@ export async function tryFallbackRetry(args: {
|
||||
const previousSessionID = task.sessionId
|
||||
const previousModel = task.model
|
||||
|
||||
const transformedModelId = transformModelForProvider(providerID, nextFallback.model)
|
||||
const transformedModelId = deps.transformModelForProvider(providerID, nextFallback.model)
|
||||
const nextModel = {
|
||||
providerID,
|
||||
modelID: transformedModelId,
|
||||
|
||||
@@ -16,16 +16,6 @@ import { promptAsyncAfterSessionIdle } from "../../shared/prompt-async-gate"
|
||||
import { initTaskToastManager, _resetTaskToastManagerForTesting } from "../task-toast-manager/manager"
|
||||
import { _resetForTesting as resetProcessCleanupState } from "./process-cleanup"
|
||||
|
||||
mock.module("../../shared/connected-providers-cache", () => ({
|
||||
readConnectedProvidersCache: () => null,
|
||||
readProviderModelsCache: () => null,
|
||||
hasConnectedProvidersCache: () => false,
|
||||
hasProviderModelsCache: () => false,
|
||||
writeProviderModelsCache: () => {},
|
||||
updateConnectedProvidersCache: () => {},
|
||||
}))
|
||||
mock.restore()
|
||||
|
||||
|
||||
const TASK_TTL_MS = 30 * 60 * 1000
|
||||
type PendingParentWakeForTest = {
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
/// <reference types="bun-types" />
|
||||
|
||||
import { afterEach, describe, expect, mock, test } from "bun:test"
|
||||
import { afterEach, beforeEach, describe, expect, mock, test } from "bun:test"
|
||||
|
||||
import { TeamModeConfigSchema } from "../../../config/schema/team-mode"
|
||||
import type { BackgroundManager } from "../../background-agent/manager"
|
||||
@@ -14,6 +14,10 @@ import {
|
||||
} from "./session-cleanup"
|
||||
|
||||
describe("session team cleanup", () => {
|
||||
beforeEach(() => {
|
||||
clearSessionTeamRunCleanupRegistry()
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
clearSessionTeamRunCleanupRegistry()
|
||||
mock.restore()
|
||||
|
||||
Reference in New Issue
Block a user