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:
@@ -0,0 +1,30 @@
|
||||
import type { FallbackEntry } from "../../shared/model-requirements"
|
||||
import type { ModelFallbackStateController } from "./fallback-state-controller"
|
||||
|
||||
export type ModelFallbackControllerAccessor = {
|
||||
register: (controller: ModelFallbackStateController) => void
|
||||
setSessionFallbackChain: (sessionID: string, fallbackChain: FallbackEntry[] | undefined) => void
|
||||
clearSessionFallbackChain: (sessionID: string) => void
|
||||
}
|
||||
|
||||
export function createModelFallbackControllerAccessor(): ModelFallbackControllerAccessor {
|
||||
let controller: ModelFallbackStateController | null = null
|
||||
|
||||
function register(nextController: ModelFallbackStateController): void {
|
||||
controller = nextController
|
||||
}
|
||||
|
||||
function setSessionFallbackChain(sessionID: string, fallbackChain: FallbackEntry[] | undefined): void {
|
||||
controller?.setSessionFallbackChain(sessionID, fallbackChain)
|
||||
}
|
||||
|
||||
function clearSessionFallbackChain(sessionID: string): void {
|
||||
controller?.clearSessionFallbackChain(sessionID)
|
||||
}
|
||||
|
||||
return {
|
||||
register,
|
||||
setSessionFallbackChain,
|
||||
clearSessionFallbackChain,
|
||||
}
|
||||
}
|
||||
@@ -70,22 +70,23 @@ const {
|
||||
setPendingModelFallback,
|
||||
} = await importFreshModelFallbackHookModule()
|
||||
|
||||
type ModelFallbackHook = ReturnType<typeof createModelFallbackHook>
|
||||
|
||||
describe("model fallback hook", () => {
|
||||
let modelFallback: ModelFallbackHook
|
||||
|
||||
beforeEach(() => {
|
||||
modelFallback = createModelFallbackHook()
|
||||
readConnectedProvidersCacheMock.mockReturnValue(null)
|
||||
readProviderModelsCacheMock.mockReturnValue(null)
|
||||
readConnectedProvidersCacheMock.mockClear()
|
||||
readProviderModelsCacheMock.mockClear()
|
||||
selectFallbackProviderMock.mockClear()
|
||||
|
||||
clearPendingModelFallback("ses_model_fallback_main")
|
||||
clearPendingModelFallback("ses_model_fallback_ghcp")
|
||||
clearPendingModelFallback("ses_model_fallback_google")
|
||||
})
|
||||
|
||||
test("applies pending fallback on chat.message by overriding model", async () => {
|
||||
//#given
|
||||
const hook = createModelFallbackHook() as unknown as {
|
||||
const hook = modelFallback as unknown as {
|
||||
"chat.message"?: (
|
||||
input: { sessionID: string },
|
||||
output: { message: Record<string, unknown>; parts: Array<{ type: string; text?: string }> },
|
||||
@@ -93,6 +94,7 @@ describe("model fallback hook", () => {
|
||||
}
|
||||
|
||||
const set = setPendingModelFallback(
|
||||
modelFallback,
|
||||
"ses_model_fallback_main",
|
||||
"Sisyphus - Ultraworker",
|
||||
"anthropic",
|
||||
@@ -123,7 +125,7 @@ describe("model fallback hook", () => {
|
||||
|
||||
test("preserves fallback progression across repeated session.error retries", async () => {
|
||||
//#given
|
||||
const hook = createModelFallbackHook() as unknown as {
|
||||
const hook = modelFallback as unknown as {
|
||||
"chat.message"?: (
|
||||
input: { sessionID: string },
|
||||
output: { message: Record<string, unknown>; parts: Array<{ type: string; text?: string }> },
|
||||
@@ -132,7 +134,7 @@ describe("model fallback hook", () => {
|
||||
const sessionID = "ses_model_fallback_main"
|
||||
|
||||
expect(
|
||||
setPendingModelFallback(sessionID, "Sisyphus - Ultraworker", "anthropic", "claude-opus-4-7-thinking"),
|
||||
setPendingModelFallback(modelFallback, sessionID, "Sisyphus - Ultraworker", "anthropic", "claude-opus-4-7-thinking"),
|
||||
).toBe(true)
|
||||
|
||||
const firstOutput = {
|
||||
@@ -154,7 +156,7 @@ describe("model fallback hook", () => {
|
||||
|
||||
//#when - second error re-arms fallback and should advance to next entry
|
||||
expect(
|
||||
setPendingModelFallback(sessionID, "Sisyphus - Ultraworker", "anthropic", "claude-opus-4-7"),
|
||||
setPendingModelFallback(modelFallback, sessionID, "Sisyphus - Ultraworker", "anthropic", "claude-opus-4-7"),
|
||||
).toBe(true)
|
||||
|
||||
const secondOutput = {
|
||||
@@ -176,16 +178,18 @@ describe("model fallback hook", () => {
|
||||
test("does not re-arm fallback when one is already pending", () => {
|
||||
//#given
|
||||
const sessionID = "ses_model_fallback_pending_guard"
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
//#when
|
||||
const firstSet = setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Sisyphus - Ultraworker",
|
||||
"anthropic",
|
||||
"claude-opus-4-7-thinking",
|
||||
)
|
||||
const secondSet = setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Sisyphus - Ultraworker",
|
||||
"anthropic",
|
||||
@@ -195,28 +199,29 @@ describe("model fallback hook", () => {
|
||||
//#then
|
||||
expect(firstSet).toBe(true)
|
||||
expect(secondSet).toBe(false)
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
})
|
||||
|
||||
test("skips no-op fallback entries that resolve to same provider/model", async () => {
|
||||
//#given
|
||||
const sessionID = "ses_model_fallback_noop_skip"
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
const hook = createModelFallbackHook() as unknown as {
|
||||
const hook = modelFallback as unknown as {
|
||||
"chat.message"?: (
|
||||
input: { sessionID: string },
|
||||
output: { message: Record<string, unknown>; parts: Array<{ type: string; text?: string }> },
|
||||
) => Promise<void>
|
||||
}
|
||||
|
||||
setSessionFallbackChain(sessionID, [
|
||||
setSessionFallbackChain(modelFallback, sessionID, [
|
||||
{ providers: ["anthropic"], model: "claude-opus-4-7" },
|
||||
{ providers: ["opencode"], model: "kimi-k2.5-free" },
|
||||
])
|
||||
|
||||
expect(
|
||||
setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Sisyphus - Ultraworker",
|
||||
"anthropic",
|
||||
@@ -239,28 +244,29 @@ describe("model fallback hook", () => {
|
||||
providerID: "opencode",
|
||||
modelID: "kimi-k2.5-free",
|
||||
})
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
})
|
||||
|
||||
test("skips no-op fallback entries even when variant differs", async () => {
|
||||
//#given
|
||||
const sessionID = "ses_model_fallback_noop_variant_skip"
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
const hook = createModelFallbackHook() as unknown as {
|
||||
const hook = modelFallback as unknown as {
|
||||
"chat.message"?: (
|
||||
input: { sessionID: string },
|
||||
output: { message: Record<string, unknown>; parts: Array<{ type: string; text?: string }> },
|
||||
) => Promise<void>
|
||||
}
|
||||
|
||||
setSessionFallbackChain(sessionID, [
|
||||
setSessionFallbackChain(modelFallback, sessionID, [
|
||||
{ providers: ["quotio"], model: "claude-opus-4-7", variant: "max" },
|
||||
{ providers: ["quotio"], model: "gpt-5.2" },
|
||||
])
|
||||
|
||||
expect(
|
||||
setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Sisyphus - Ultraworker",
|
||||
"quotio",
|
||||
@@ -285,28 +291,29 @@ describe("model fallback hook", () => {
|
||||
modelID: "gpt-5.2",
|
||||
})
|
||||
expect(output.message["variant"]).toBeUndefined()
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
})
|
||||
|
||||
test("uses connected preferred provider when fallback entry providers are disconnected", async () => {
|
||||
//#given
|
||||
const sessionID = "ses_model_fallback_preferred_provider"
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
readConnectedProvidersCacheMock.mockReturnValue(["provider-x"])
|
||||
|
||||
const hook = createModelFallbackHook() as unknown as {
|
||||
const hook = modelFallback as unknown as {
|
||||
"chat.message"?: (
|
||||
input: { sessionID: string },
|
||||
output: { message: Record<string, unknown>; parts: Array<{ type: string; text?: string }> },
|
||||
) => Promise<void>
|
||||
}
|
||||
|
||||
setSessionFallbackChain(sessionID, [
|
||||
setSessionFallbackChain(modelFallback, sessionID, [
|
||||
{ providers: ["provider-y"], model: "fallback-model" },
|
||||
])
|
||||
|
||||
expect(
|
||||
setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Sisyphus - Ultraworker",
|
||||
"provider-x",
|
||||
@@ -329,17 +336,18 @@ describe("model fallback hook", () => {
|
||||
providerID: "provider-x",
|
||||
modelID: "fallback-model",
|
||||
})
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
})
|
||||
|
||||
test("does not fall back to hardcoded agent chain when session explicitly stores no fallback chain [regression #2941]", () => {
|
||||
//#given
|
||||
const sessionID = "ses_model_fallback_explicit_none"
|
||||
clearPendingModelFallback(sessionID)
|
||||
setSessionFallbackChain(sessionID, undefined)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
setSessionFallbackChain(modelFallback, sessionID, undefined)
|
||||
|
||||
//#when
|
||||
const set = setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Sisyphus - Junior",
|
||||
"anthropic",
|
||||
@@ -348,7 +356,7 @@ describe("model fallback hook", () => {
|
||||
|
||||
//#then
|
||||
expect(set).toBe(false)
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
})
|
||||
|
||||
test("shows toast when fallback is applied", async () => {
|
||||
@@ -366,6 +374,7 @@ describe("model fallback hook", () => {
|
||||
}
|
||||
|
||||
const set = setPendingModelFallback(
|
||||
hook,
|
||||
"ses_model_fallback_toast",
|
||||
"Sisyphus - Ultraworker",
|
||||
"anthropic",
|
||||
@@ -392,9 +401,9 @@ describe("model fallback hook", () => {
|
||||
test("transforms model names for github-copilot provider via fallback chain", async () => {
|
||||
//#given
|
||||
const sessionID = "ses_model_fallback_ghcp"
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
const hook = createModelFallbackHook() as unknown as {
|
||||
const hook = modelFallback as unknown as {
|
||||
"chat.message"?: (
|
||||
input: { sessionID: string },
|
||||
output: { message: Record<string, unknown>; parts: Array<{ type: string; text?: string }> },
|
||||
@@ -402,11 +411,12 @@ describe("model fallback hook", () => {
|
||||
}
|
||||
|
||||
// Set a custom fallback chain that routes through github-copilot
|
||||
setSessionFallbackChain(sessionID, [
|
||||
setSessionFallbackChain(modelFallback, sessionID, [
|
||||
{ providers: ["github-copilot"], model: "claude-sonnet-4-6" },
|
||||
])
|
||||
|
||||
const set = setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Atlas - Plan Executor",
|
||||
"github-copilot",
|
||||
@@ -430,15 +440,15 @@ describe("model fallback hook", () => {
|
||||
modelID: "claude-sonnet-4.6",
|
||||
})
|
||||
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
})
|
||||
|
||||
test("preserves canonical google preview model names via fallback chain", async () => {
|
||||
//#given
|
||||
const sessionID = "ses_model_fallback_google"
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
|
||||
const hook = createModelFallbackHook() as unknown as {
|
||||
const hook = modelFallback as unknown as {
|
||||
"chat.message"?: (
|
||||
input: { sessionID: string },
|
||||
output: { message: Record<string, unknown>; parts: Array<{ type: string; text?: string }> },
|
||||
@@ -446,11 +456,12 @@ describe("model fallback hook", () => {
|
||||
}
|
||||
|
||||
// Set a custom fallback chain that routes through google
|
||||
setSessionFallbackChain(sessionID, [
|
||||
setSessionFallbackChain(modelFallback, sessionID, [
|
||||
{ providers: ["google"], model: "gemini-3.1-pro-preview" },
|
||||
])
|
||||
|
||||
const set = setPendingModelFallback(
|
||||
modelFallback,
|
||||
sessionID,
|
||||
"Oracle",
|
||||
"google",
|
||||
@@ -474,7 +485,7 @@ describe("model fallback hook", () => {
|
||||
modelID: "gemini-3.1-pro-preview",
|
||||
})
|
||||
|
||||
clearPendingModelFallback(sessionID)
|
||||
clearPendingModelFallback(modelFallback, sessionID)
|
||||
})
|
||||
})
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ import {
|
||||
createModelFallbackStateController,
|
||||
type ModelFallbackStateController,
|
||||
} from "./fallback-state-controller"
|
||||
import type { ModelFallbackControllerAccessor } from "./controller-accessor"
|
||||
|
||||
type FallbackToast = (input: {
|
||||
title: string
|
||||
@@ -28,26 +29,45 @@ export type ModelFallbackState = {
|
||||
pending: boolean
|
||||
}
|
||||
|
||||
const modelFallbackControllerRef: { current?: ModelFallbackStateController } = {}
|
||||
type ModelFallbackControllerWithState = Pick<
|
||||
ModelFallbackStateController,
|
||||
| "lastToastKey"
|
||||
| "setSessionFallbackChain"
|
||||
| "clearSessionFallbackChain"
|
||||
| "setPendingModelFallback"
|
||||
| "getNextFallback"
|
||||
| "clearPendingModelFallback"
|
||||
| "hasPendingModelFallback"
|
||||
| "getFallbackState"
|
||||
| "reset"
|
||||
>
|
||||
|
||||
function getOrCreateModelFallbackController(): ModelFallbackStateController {
|
||||
if (!modelFallbackControllerRef.current) {
|
||||
createModelFallbackHook()
|
||||
}
|
||||
|
||||
const controller = modelFallbackControllerRef.current
|
||||
if (!controller) {
|
||||
throw new Error("Model fallback controller should be initialized")
|
||||
}
|
||||
return controller
|
||||
export type ModelFallbackHook = ModelFallbackControllerWithState & {
|
||||
"chat.message": (
|
||||
input: ChatMessageInput,
|
||||
output: ChatMessageHandlerOutput,
|
||||
) => Promise<void>
|
||||
}
|
||||
|
||||
export function setSessionFallbackChain(sessionID: string, fallbackChain: FallbackEntry[] | undefined): void {
|
||||
getOrCreateModelFallbackController().setSessionFallbackChain(sessionID, fallbackChain)
|
||||
type ModelFallbackHookArgs = {
|
||||
toast?: FallbackToast
|
||||
onApplied?: FallbackCallback
|
||||
controllerAccessor?: ModelFallbackControllerAccessor
|
||||
}
|
||||
|
||||
export function clearSessionFallbackChain(sessionID: string): void {
|
||||
getOrCreateModelFallbackController().clearSessionFallbackChain(sessionID)
|
||||
export function setSessionFallbackChain(
|
||||
controller: Pick<ModelFallbackStateController, "setSessionFallbackChain">,
|
||||
sessionID: string,
|
||||
fallbackChain: FallbackEntry[] | undefined,
|
||||
): void {
|
||||
controller.setSessionFallbackChain(sessionID, fallbackChain)
|
||||
}
|
||||
|
||||
export function clearSessionFallbackChain(
|
||||
controller: Pick<ModelFallbackStateController, "clearSessionFallbackChain">,
|
||||
sessionID: string,
|
||||
): void {
|
||||
controller.clearSessionFallbackChain(sessionID)
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -55,12 +75,13 @@ export function clearSessionFallbackChain(sessionID: string): void {
|
||||
* Called when a model error is detected in session.error handler.
|
||||
*/
|
||||
export function setPendingModelFallback(
|
||||
controller: Pick<ModelFallbackStateController, "setPendingModelFallback">,
|
||||
sessionID: string,
|
||||
agentName: string,
|
||||
currentProviderID: string,
|
||||
currentModelID: string,
|
||||
): boolean {
|
||||
return getOrCreateModelFallbackController().setPendingModelFallback(
|
||||
return controller.setPendingModelFallback(
|
||||
sessionID,
|
||||
agentName,
|
||||
currentProviderID,
|
||||
@@ -73,54 +94,71 @@ export function setPendingModelFallback(
|
||||
* Increments attemptCount each time called.
|
||||
*/
|
||||
export function getNextFallback(
|
||||
controller: Pick<ModelFallbackStateController, "getNextFallback">,
|
||||
sessionID: string,
|
||||
): { providerID: string; modelID: string; variant?: string } | null {
|
||||
return getOrCreateModelFallbackController().getNextFallback(sessionID)
|
||||
return controller.getNextFallback(sessionID)
|
||||
}
|
||||
|
||||
/**
|
||||
* Clears the pending fallback for a session.
|
||||
* Called after fallback is successfully applied.
|
||||
*/
|
||||
export function clearPendingModelFallback(sessionID: string): void {
|
||||
getOrCreateModelFallbackController().clearPendingModelFallback(sessionID)
|
||||
export function clearPendingModelFallback(
|
||||
controller: Pick<ModelFallbackStateController, "clearPendingModelFallback">,
|
||||
sessionID: string,
|
||||
): void {
|
||||
controller.clearPendingModelFallback(sessionID)
|
||||
}
|
||||
|
||||
/**
|
||||
* Checks if there's a pending fallback for a session.
|
||||
*/
|
||||
export function hasPendingModelFallback(sessionID: string): boolean {
|
||||
return getOrCreateModelFallbackController().hasPendingModelFallback(sessionID)
|
||||
export function hasPendingModelFallback(
|
||||
controller: Pick<ModelFallbackStateController, "hasPendingModelFallback">,
|
||||
sessionID: string,
|
||||
): boolean {
|
||||
return controller.hasPendingModelFallback(sessionID)
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets the current fallback state for a session (for debugging).
|
||||
*/
|
||||
export function getFallbackState(sessionID: string): ModelFallbackState | undefined {
|
||||
return getOrCreateModelFallbackController().getFallbackState(sessionID)
|
||||
export function getFallbackState(
|
||||
controller: Pick<ModelFallbackStateController, "getFallbackState">,
|
||||
sessionID: string,
|
||||
): ModelFallbackState | undefined {
|
||||
return controller.getFallbackState(sessionID)
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates a chat.message hook that applies model fallbacks when pending.
|
||||
*/
|
||||
export function createModelFallbackHook(args?: { toast?: FallbackToast; onApplied?: FallbackCallback }) {
|
||||
if (!modelFallbackControllerRef.current) {
|
||||
const pendingModelFallbacks = new Map<string, ModelFallbackState>()
|
||||
const lastToastKey = new Map<string, string>()
|
||||
const sessionFallbackChains = new Map<string, FallbackEntry[]>()
|
||||
export function createModelFallbackHook(args?: ModelFallbackHookArgs): ModelFallbackHook {
|
||||
const pendingModelFallbacks = new Map<string, ModelFallbackState>()
|
||||
const lastToastKey = new Map<string, string>()
|
||||
const sessionFallbackChains = new Map<string, FallbackEntry[]>()
|
||||
const controller = createModelFallbackStateController({
|
||||
pendingModelFallbacks,
|
||||
lastToastKey,
|
||||
sessionFallbackChains,
|
||||
})
|
||||
|
||||
modelFallbackControllerRef.current = createModelFallbackStateController({
|
||||
pendingModelFallbacks,
|
||||
lastToastKey,
|
||||
sessionFallbackChains,
|
||||
})
|
||||
}
|
||||
args?.controllerAccessor?.register(controller)
|
||||
|
||||
const controller = getOrCreateModelFallbackController()
|
||||
const toast = args?.toast
|
||||
const onApplied = args?.onApplied
|
||||
|
||||
return {
|
||||
lastToastKey: controller.lastToastKey,
|
||||
setSessionFallbackChain: controller.setSessionFallbackChain,
|
||||
clearSessionFallbackChain: controller.clearSessionFallbackChain,
|
||||
setPendingModelFallback: controller.setPendingModelFallback,
|
||||
getNextFallback: controller.getNextFallback,
|
||||
clearPendingModelFallback: controller.clearPendingModelFallback,
|
||||
hasPendingModelFallback: controller.hasPendingModelFallback,
|
||||
getFallbackState: controller.getFallbackState,
|
||||
reset: controller.reset,
|
||||
"chat.message": async (
|
||||
input: ChatMessageInput,
|
||||
output: ChatMessageHandlerOutput,
|
||||
@@ -128,7 +166,7 @@ export function createModelFallbackHook(args?: { toast?: FallbackToast; onApplie
|
||||
const { sessionID } = input
|
||||
if (!sessionID) return
|
||||
|
||||
const fallback = getNextFallback(sessionID)
|
||||
const fallback = getNextFallback(controller, sessionID)
|
||||
if (!fallback) return
|
||||
|
||||
await applyFallbackToChatMessage({
|
||||
@@ -144,9 +182,8 @@ export function createModelFallbackHook(args?: { toast?: FallbackToast; onApplie
|
||||
}
|
||||
|
||||
/**
|
||||
* Resets all module-global state for testing.
|
||||
* Clears pending fallbacks, toast keys, and session chains.
|
||||
* Resets hook-owned state for testing.
|
||||
*/
|
||||
export function _resetForTesting(): void {
|
||||
getOrCreateModelFallbackController().reset()
|
||||
export function _resetForTesting(controller?: Pick<ModelFallbackStateController, "reset">): void {
|
||||
controller?.reset()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
export { createModelFallbackControllerAccessor } from "./controller-accessor"
|
||||
export type { ModelFallbackControllerAccessor } from "./controller-accessor"
|
||||
Reference in New Issue
Block a user