refactor(model-fallback): fully encapsulate session state in factory closure
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
This commit is contained in:
@@ -65,13 +65,13 @@ function createChatMessageHandlerHooks(modelFallback: ReturnType<typeof createMo
|
||||
let readConnectedProvidersCacheSpy: { mockRestore: () => void } | undefined
|
||||
let readProviderModelsCacheSpy: { mockRestore: () => void } | undefined
|
||||
|
||||
afterEach(() => {
|
||||
readConnectedProvidersCacheSpy?.mockRestore()
|
||||
readProviderModelsCacheSpy?.mockRestore()
|
||||
readConnectedProvidersCacheSpy = undefined
|
||||
readProviderModelsCacheSpy = undefined
|
||||
_resetForTesting()
|
||||
})
|
||||
afterEach(() => {
|
||||
readConnectedProvidersCacheSpy?.mockRestore()
|
||||
readProviderModelsCacheSpy?.mockRestore()
|
||||
readConnectedProvidersCacheSpy = undefined
|
||||
readProviderModelsCacheSpy = undefined
|
||||
_resetForTesting()
|
||||
})
|
||||
|
||||
describe("createEventHandler - category runtime fallback suppression", () => {
|
||||
test("does not arm retry fallback when category session explicitly stores no fallback chain [regression #2941]", async () => {
|
||||
@@ -83,11 +83,10 @@ describe("createEventHandler - category runtime fallback suppression", () => {
|
||||
readConnectedProvidersCacheSpy = spyOn(connectedProvidersCache, "readConnectedProvidersCache").mockReturnValue(null)
|
||||
readProviderModelsCacheSpy = spyOn(connectedProvidersCache, "readProviderModelsCache").mockReturnValue(null)
|
||||
|
||||
clearPendingModelFallback(sessionID)
|
||||
setSessionAgent(sessionID, "sisyphus-junior")
|
||||
setSessionFallbackChain(sessionID, undefined)
|
||||
|
||||
const modelFallback = createModelFallbackHook()
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
setSessionAgent(sessionID, "sisyphus-junior")
|
||||
setSessionFallbackChain(modelFallback, sessionID, undefined)
|
||||
const eventHandler = createEventHandler({
|
||||
ctx: asEventHandlerContext({
|
||||
directory: "/tmp",
|
||||
|
||||
@@ -142,9 +142,8 @@ describe("createEventHandler - model fallback", () => {
|
||||
//#given
|
||||
const sessionID = "ses_status_retry_fallback"
|
||||
setMainSession(sessionID)
|
||||
clearPendingModelFallback(sessionID)
|
||||
|
||||
const modelFallback = createModelFallbackHook()
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
const { handler, abortCalls, promptCalls } = createHandler({ hooks: { modelFallback } })
|
||||
|
||||
@@ -232,8 +231,8 @@ describe("createEventHandler - model fallback", () => {
|
||||
//#given
|
||||
const sessionID = "ses_status_retry_dedup"
|
||||
setMainSession(sessionID)
|
||||
clearPendingModelFallback(sessionID)
|
||||
const modelFallback = createModelFallbackHook()
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
const { handler, abortCalls, promptCalls } = createHandler({ hooks: { modelFallback } })
|
||||
|
||||
await handler({
|
||||
@@ -293,8 +292,8 @@ describe("createEventHandler - model fallback", () => {
|
||||
//#given
|
||||
const sessionID = "ses_status_retry_runtime_enabled"
|
||||
setMainSession(sessionID)
|
||||
clearPendingModelFallback(sessionID)
|
||||
const modelFallback = createModelFallbackHook()
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
const runtimeFallback = {
|
||||
event: async () => {},
|
||||
"chat.message": async () => {},
|
||||
@@ -346,9 +345,8 @@ describe("createEventHandler - model fallback", () => {
|
||||
//#given
|
||||
const sessionID = "ses_status_retry_user_fallback"
|
||||
setMainSession(sessionID)
|
||||
clearPendingModelFallback(sessionID)
|
||||
|
||||
const modelFallback = createModelFallbackHook()
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
const pluginConfig = {
|
||||
agents: {
|
||||
sisyphus: {
|
||||
@@ -446,9 +444,8 @@ describe("createEventHandler - model fallback", () => {
|
||||
const toastCalls: string[] = []
|
||||
const sessionID = "ses_main_fallback_chain"
|
||||
setMainSession(sessionID)
|
||||
clearPendingModelFallback(sessionID)
|
||||
|
||||
const modelFallback = createModelFallbackHook()
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
setupConnectedProviderCacheMocks()
|
||||
const eventHandler = createEventHandler({
|
||||
|
||||
@@ -761,11 +761,10 @@ describe("createEventHandler - retry dedupe lifecycle", () => {
|
||||
//#given
|
||||
const sessionID = "ses_retry_recovery_rearm"
|
||||
setMainSession(sessionID)
|
||||
clearPendingModelFallback(sessionID)
|
||||
|
||||
const abortCalls: string[] = []
|
||||
const promptCalls: string[] = []
|
||||
const modelFallback = createModelFallbackHook()
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
const eventHandler = createEventHandler({
|
||||
ctx: asEventHandlerContext({
|
||||
|
||||
+22
-9
@@ -15,6 +15,7 @@ import {
|
||||
clearSessionFallbackChain,
|
||||
setSessionFallbackChain,
|
||||
setPendingModelFallback,
|
||||
type ModelFallbackHook,
|
||||
} from "../hooks/model-fallback/hook";
|
||||
import { getRawFallbackModels } from "../hooks/runtime-fallback/fallback-models";
|
||||
import {
|
||||
@@ -111,6 +112,7 @@ function extractProviderModelFromErrorMessage(message: string): { providerID?: s
|
||||
return {};
|
||||
}
|
||||
function applyUserConfiguredFallbackChain(
|
||||
modelFallback: Pick<ModelFallbackHook, "setSessionFallbackChain"> | null | undefined,
|
||||
sessionID: string,
|
||||
agentName: string,
|
||||
currentProviderID: string,
|
||||
@@ -123,7 +125,9 @@ function applyUserConfiguredFallbackChain(
|
||||
const fallbackChain = buildFallbackChainFromModels(rawFallbackModels, currentProviderID);
|
||||
|
||||
if (fallbackChain && fallbackChain.length > 0) {
|
||||
setSessionFallbackChain(sessionID, fallbackChain);
|
||||
if (modelFallback) {
|
||||
setSessionFallbackChain(modelFallback, sessionID, fallbackChain);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -170,6 +174,7 @@ export function createEventHandler(args: {
|
||||
|
||||
const isModelFallbackEnabled =
|
||||
hooks.modelFallback !== null && hooks.modelFallback !== undefined;
|
||||
const modelFallback = hooks.modelFallback;
|
||||
|
||||
// Avoid triggering multiple abort+continue cycles for the same failing assistant message.
|
||||
const lastHandledModelErrorMessageID = new Map<string, string>();
|
||||
@@ -408,8 +413,10 @@ export function createEventHandler(args: {
|
||||
lastHandledModelErrorMessageID.delete(sessionInfo.id);
|
||||
lastHandledRetryStatusKey.delete(sessionInfo.id);
|
||||
lastKnownModelBySession.delete(sessionInfo.id);
|
||||
clearPendingModelFallback(sessionInfo.id);
|
||||
clearSessionFallbackChain(sessionInfo.id);
|
||||
if (modelFallback) {
|
||||
clearPendingModelFallback(modelFallback, sessionInfo.id);
|
||||
clearSessionFallbackChain(modelFallback, sessionInfo.id);
|
||||
}
|
||||
resetMessageCursor(sessionInfo.id);
|
||||
clearBackgroundOutputConsumptionsForParentSession(sessionInfo.id);
|
||||
clearBackgroundOutputConsumptionsForTaskSession(sessionInfo.id);
|
||||
@@ -517,9 +524,11 @@ export function createEventHandler(args: {
|
||||
);
|
||||
const rawModel = (info?.modelID as string | undefined) ?? "claude-opus-4-7";
|
||||
const currentModel = normalizeFallbackModelID(rawModel);
|
||||
applyUserConfiguredFallbackChain(sessionID, agentName, currentProvider, args.pluginConfig);
|
||||
applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig);
|
||||
|
||||
const setFallback = setPendingModelFallback(sessionID, agentName, currentProvider, currentModel);
|
||||
const setFallback = modelFallback
|
||||
? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel)
|
||||
: false;
|
||||
|
||||
if (
|
||||
setFallback &&
|
||||
@@ -580,9 +589,11 @@ export function createEventHandler(args: {
|
||||
const currentProvider = resolveFallbackProviderID(sessionID, parsed.providerID);
|
||||
let currentModel = parsed.modelID ?? lastKnown?.modelID ?? "claude-opus-4-7";
|
||||
currentModel = normalizeFallbackModelID(currentModel);
|
||||
applyUserConfiguredFallbackChain(sessionID, agentName, currentProvider, args.pluginConfig);
|
||||
applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig);
|
||||
|
||||
const setFallback = setPendingModelFallback(sessionID, agentName, currentProvider, currentModel);
|
||||
const setFallback = modelFallback
|
||||
? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel)
|
||||
: false;
|
||||
|
||||
if (
|
||||
setFallback &&
|
||||
@@ -666,9 +677,11 @@ export function createEventHandler(args: {
|
||||
);
|
||||
let currentModel = (props?.modelID as string) || parsed.modelID || "claude-opus-4-7";
|
||||
currentModel = normalizeFallbackModelID(currentModel);
|
||||
applyUserConfiguredFallbackChain(sessionID, agentName, currentProvider, args.pluginConfig);
|
||||
applyUserConfiguredFallbackChain(modelFallback, sessionID, agentName, currentProvider, args.pluginConfig);
|
||||
|
||||
const setFallback = setPendingModelFallback(sessionID, agentName, currentProvider, currentModel);
|
||||
const setFallback = modelFallback
|
||||
? setPendingModelFallback(modelFallback, sessionID, agentName, currentProvider, currentModel)
|
||||
: false;
|
||||
|
||||
if (
|
||||
setFallback &&
|
||||
|
||||
@@ -9,7 +9,6 @@ import { createModelFallbackHook } from "../hooks/model-fallback/hook"
|
||||
import { createRuntimeFallbackHook } from "../hooks/runtime-fallback"
|
||||
import type { RuntimeFallbackPluginInput } from "../hooks/runtime-fallback/types"
|
||||
import { _resetForTesting } from "../features/claude-code-session-state"
|
||||
import { _resetForTesting as _resetModelFallbackForTesting } from "../hooks/model-fallback/hook"
|
||||
import { SessionCategoryRegistry } from "../shared/session-category-registry"
|
||||
import * as connectedProvidersCache from "../shared/connected-providers-cache"
|
||||
|
||||
@@ -369,7 +368,6 @@ function setupConnectedProviderCacheMocks(): void {
|
||||
|
||||
afterEach(() => {
|
||||
_resetForTesting()
|
||||
_resetModelFallbackForTesting()
|
||||
SessionCategoryRegistry.clear()
|
||||
})
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import type { HookName, OhMyOpenCodeConfig } from "../../config"
|
||||
import type { ModelFallbackControllerAccessor } from "../../hooks/model-fallback"
|
||||
import type { PluginContext } from "../types"
|
||||
import type { ModelCacheState } from "../../plugin-state"
|
||||
|
||||
@@ -10,15 +11,17 @@ export function createCoreHooks(args: {
|
||||
ctx: PluginContext
|
||||
pluginConfig: OhMyOpenCodeConfig
|
||||
modelCacheState: ModelCacheState
|
||||
modelFallbackControllerAccessor?: ModelFallbackControllerAccessor
|
||||
isHookEnabled: (hookName: HookName) => boolean
|
||||
safeHookEnabled: boolean
|
||||
}) {
|
||||
const { ctx, pluginConfig, modelCacheState, isHookEnabled, safeHookEnabled } = args
|
||||
const { ctx, pluginConfig, modelCacheState, modelFallbackControllerAccessor, isHookEnabled, safeHookEnabled } = args
|
||||
|
||||
const session = createSessionHooks({
|
||||
ctx,
|
||||
pluginConfig,
|
||||
modelCacheState,
|
||||
modelFallbackControllerAccessor,
|
||||
isHookEnabled,
|
||||
safeHookEnabled,
|
||||
})
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import type { OhMyOpenCodeConfig, HookName } from "../../config"
|
||||
import type { ModelFallbackControllerAccessor } from "../../hooks/model-fallback"
|
||||
import type { ModelCacheState } from "../../plugin-state"
|
||||
import type { PluginContext } from "../types"
|
||||
|
||||
@@ -69,10 +70,11 @@ export function createSessionHooks(args: {
|
||||
ctx: PluginContext
|
||||
pluginConfig: OhMyOpenCodeConfig
|
||||
modelCacheState: ModelCacheState
|
||||
modelFallbackControllerAccessor?: ModelFallbackControllerAccessor
|
||||
isHookEnabled: (hookName: HookName) => boolean
|
||||
safeHookEnabled: boolean
|
||||
}): SessionHooks {
|
||||
const { ctx, pluginConfig, modelCacheState, isHookEnabled, safeHookEnabled } = args
|
||||
const { ctx, pluginConfig, modelCacheState, modelFallbackControllerAccessor, isHookEnabled, safeHookEnabled } = args
|
||||
const safeHook = <T>(hookName: HookName, factory: () => T): T | null =>
|
||||
safeCreateHook(hookName, factory, { enabled: safeHookEnabled })
|
||||
|
||||
@@ -171,6 +173,7 @@ export function createSessionHooks(args: {
|
||||
.catch(() => {})
|
||||
},
|
||||
onApplied: enableFallbackTitle ? updateFallbackTitle : undefined,
|
||||
controllerAccessor: modelFallbackControllerAccessor,
|
||||
}))
|
||||
: null
|
||||
|
||||
|
||||
@@ -144,7 +144,7 @@ export function trimToolsToCap(filteredTools: ToolsRecord, maxTools: number): vo
|
||||
export function createToolRegistry(args: {
|
||||
ctx: PluginContext
|
||||
pluginConfig: OhMyOpenCodeConfig
|
||||
managers: Pick<Managers, "backgroundManager" | "tmuxSessionManager" | "skillMcpManager">
|
||||
managers: Pick<Managers, "backgroundManager" | "tmuxSessionManager" | "skillMcpManager" | "modelFallbackControllerAccessor">
|
||||
skillContext: SkillContext
|
||||
availableCategories: AvailableCategory[]
|
||||
interactiveBashEnabled?: boolean
|
||||
@@ -170,6 +170,7 @@ export function createToolRegistry(args: {
|
||||
pluginConfig.disabled_agents ?? [],
|
||||
pluginConfig.agents,
|
||||
pluginConfig.categories,
|
||||
managers.modelFallbackControllerAccessor,
|
||||
)
|
||||
|
||||
const isMultimodalLookerEnabled = !(pluginConfig.disabled_agents ?? []).some(
|
||||
@@ -191,6 +192,7 @@ export function createToolRegistry(args: {
|
||||
availableSkills: skillContext.availableSkills,
|
||||
sisyphusAgentConfig: pluginConfig.sisyphus_agent,
|
||||
syncPollTimeoutMs: pluginConfig.background_task?.syncPollTimeoutMs,
|
||||
modelFallbackControllerAccessor: managers.modelFallbackControllerAccessor,
|
||||
onSyncSessionCreated: async (event) => {
|
||||
log("[index] onSyncSessionCreated callback", {
|
||||
sessionID: event.sessionID,
|
||||
|
||||
Reference in New Issue
Block a user