Guard ultrawork variant overrides with SDK metadata

Ultrawork now checks provider SDK metadata before forcing a variant, so unsupported variants are skipped instead of being written into the message state.

Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
This commit is contained in:
YeonGyu-Kim
2026-03-11 15:15:50 +09:00
parent 3ba4ada04c
commit 29e1136813
4 changed files with 327 additions and 55 deletions
+8 -1
View File
@@ -158,6 +158,13 @@ export function createChatMessageHandler(args: {
} }
} }
applyUltraworkModelOverrideOnMessage(pluginConfig, input.agent, output, pluginContext.client.tui, input.sessionID) await applyUltraworkModelOverrideOnMessage(
pluginConfig,
input.agent,
output,
pluginContext.client.tui,
input.sessionID,
pluginContext.client,
)
} }
} }
+82 -54
View File
@@ -4,6 +4,7 @@ import { getSessionAgent } from "../features/claude-code-session-state"
import { log } from "../shared" import { log } from "../shared"
import { getAgentConfigKey } from "../shared/agent-display-names" import { getAgentConfigKey } from "../shared/agent-display-names"
import { scheduleDeferredModelOverride } from "./ultrawork-db-model-override" import { scheduleDeferredModelOverride } from "./ultrawork-db-model-override"
import { resolveValidUltraworkVariant } from "./ultrawork-variant-availability"
const CODE_BLOCK = /```[\s\S]*?```/g const CODE_BLOCK = /```[\s\S]*?```/g
const INLINE_CODE = /`[^`]+`/g const INLINE_CODE = /`[^`]+`/g
@@ -15,7 +16,7 @@ export function detectUltrawork(text: string): boolean {
} }
function extractPromptText(parts: Array<{ type: string; text?: string }>): string { function extractPromptText(parts: Array<{ type: string; text?: string }>): string {
return parts.filter((p) => p.type === "text").map((p) => p.text || "").join("") return parts.filter((part) => part.type === "text").map((part) => part.text || "").join("")
} }
type ToastFn = { type ToastFn = {
@@ -36,22 +37,26 @@ export type UltraworkOverrideResult = {
variant?: string variant?: string
} }
function isSameModel( type ModelDescriptor = {
current: unknown, providerID: string
target: { providerID: string; modelID: string }, modelID: string
): boolean { }
if (typeof current !== "object" || current === null) return false
const currentRecord = current as Record<string, unknown> function isSameModel(current: unknown, target: ModelDescriptor): boolean {
return ( if (typeof current !== "object" || current === null) return false
currentRecord["providerID"] === target.providerID const currentRecord = current as Record<string, unknown>
&& currentRecord["modelID"] === target.modelID return currentRecord["providerID"] === target.providerID && currentRecord["modelID"] === target.modelID
) }
function getMessageModel(current: unknown): ModelDescriptor | undefined {
if (typeof current !== "object" || current === null) return undefined
const currentRecord = current as Record<string, unknown>
const providerID = currentRecord["providerID"]
const modelID = currentRecord["modelID"]
if (typeof providerID !== "string" || typeof modelID !== "string") return undefined
return { providerID, modelID }
} }
/**
* Resolves the ultrawork model override config for the given agent and prompt text.
* Returns null if no override should be applied.
*/
export function resolveUltraworkOverride( export function resolveUltraworkOverride(
pluginConfig: OhMyOpenCodeConfig, pluginConfig: OhMyOpenCodeConfig,
inputAgentName: string | undefined, inputAgentName: string | undefined,
@@ -76,9 +81,7 @@ export function resolveUltraworkOverride(
if (!ultraworkConfig?.model && !ultraworkConfig?.variant) return null if (!ultraworkConfig?.model && !ultraworkConfig?.variant) return null
if (!ultraworkConfig.model) { if (!ultraworkConfig.model) {
return { return { variant: ultraworkConfig.variant }
variant: ultraworkConfig.variant,
}
} }
const modelParts = ultraworkConfig.model.split("/") const modelParts = ultraworkConfig.model.split("/")
@@ -91,37 +94,20 @@ export function resolveUltraworkOverride(
} }
} }
/** function applyResolvedUltraworkOverride(args: {
* Applies ultrawork model override using a deferred DB update strategy. override: UltraworkOverrideResult
* validatedVariant: string | undefined
* Instead of directly mutating output.message.model (which would cause the TUI output: { message: Record<string, unknown> }
* bottom bar to show the override model), this schedules a queueMicrotask that inputAgentName: string | undefined
* updates the message model directly in SQLite AFTER Session.updateMessage() tui: unknown
* saves the original model, but BEFORE loop() reads it for the API call. }): void {
* const { override, validatedVariant, output, inputAgentName, tui } = args
* Result: API call uses opus, TUI bottom bar stays on sonnet. if (validatedVariant) {
*/ output.message["variant"] = validatedVariant
export function applyUltraworkModelOverrideOnMessage( output.message["thinking"] = validatedVariant
pluginConfig: OhMyOpenCodeConfig,
inputAgentName: string | undefined,
output: {
message: Record<string, unknown>
parts: Array<{ type: string; text?: string; [key: string]: unknown }>
},
tui: unknown,
sessionID?: string,
): void {
const override = resolveUltraworkOverride(pluginConfig, inputAgentName, output, sessionID)
if (!override) return
if (override.variant) {
output.message["variant"] = override.variant
output.message["thinking"] = override.variant
} }
if (!override.providerID || !override.modelID) { if (!override.providerID || !override.modelID) return
return
}
const targetModel = { providerID: override.providerID, modelID: override.modelID } const targetModel = { providerID: override.providerID, modelID: override.modelID }
if (isSameModel(output.message.model, targetModel)) { if (isSameModel(output.message.model, targetModel)) {
@@ -134,7 +120,6 @@ export function applyUltraworkModelOverrideOnMessage(
log("[ultrawork-model-override] No message ID found, falling back to direct mutation") log("[ultrawork-model-override] No message ID found, falling back to direct mutation")
output.message.model = targetModel output.message.model = targetModel
return return
} }
const fromModel = (output.message.model as { modelID?: string } | undefined)?.modelID ?? "unknown" const fromModel = (output.message.model as { modelID?: string } | undefined)?.modelID ?? "unknown"
@@ -143,11 +128,7 @@ export function applyUltraworkModelOverrideOnMessage(
(typeof output.message["agent"] === "string" ? (output.message["agent"] as string) : "unknown"), (typeof output.message["agent"] === "string" ? (output.message["agent"] as string) : "unknown"),
) )
scheduleDeferredModelOverride( scheduleDeferredModelOverride(messageId, targetModel, validatedVariant)
messageId,
targetModel,
override.variant,
)
log(`[ultrawork-model-override] ${fromModel} -> ${override.modelID} (deferred DB)`, { log(`[ultrawork-model-override] ${fromModel} -> ${override.modelID} (deferred DB)`, {
agent: agentConfigKey, agent: agentConfigKey,
@@ -156,6 +137,53 @@ export function applyUltraworkModelOverrideOnMessage(
showToast( showToast(
tui, tui,
"Ultrawork Model Override", "Ultrawork Model Override",
`${fromModel} \u2192 ${override.modelID}. Maximum precision engaged.`, `${fromModel} ${override.modelID}. Maximum precision engaged.`,
) )
} }
export function applyUltraworkModelOverrideOnMessage(
pluginConfig: OhMyOpenCodeConfig,
inputAgentName: string | undefined,
output: {
message: Record<string, unknown>
parts: Array<{ type: string; text?: string; [key: string]: unknown }>
},
tui: unknown,
sessionID?: string,
client?: unknown,
): void | Promise<void> {
const override = resolveUltraworkOverride(pluginConfig, inputAgentName, output, sessionID)
if (!override) return
const currentModel = getMessageModel(output.message.model)
const variantTargetModel = override.providerID && override.modelID
? { providerID: override.providerID, modelID: override.modelID }
: currentModel
if (!client || typeof (client as { provider?: { list?: unknown } }).provider?.list !== "function") {
applyResolvedUltraworkOverride({ override, validatedVariant: override.variant, output, inputAgentName, tui })
return
}
return resolveValidUltraworkVariant(client, variantTargetModel, override.variant)
.then((validatedVariant) => {
if (override.variant && !validatedVariant) {
log("[ultrawork-model-override] Skip invalid ultrawork variant override", {
variant: override.variant,
providerID: variantTargetModel?.providerID,
modelID: variantTargetModel?.modelID,
})
}
applyResolvedUltraworkOverride({ override, validatedVariant, output, inputAgentName, tui })
})
.catch((error) => {
log("[ultrawork-model-override] Failed to validate ultrawork variant via SDK", {
variant: override.variant,
error: String(error),
providerID: variantTargetModel?.providerID,
modelID: variantTargetModel?.modelID,
})
applyResolvedUltraworkOverride({ override, validatedVariant: undefined, output, inputAgentName, tui })
})
}
@@ -0,0 +1,186 @@
import { describe, expect, spyOn, test } from "bun:test"
import * as dbOverrideModule from "./ultrawork-db-model-override"
import { applyUltraworkModelOverrideOnMessage } from "./ultrawork-model-override"
import { resolveValidUltraworkVariant } from "./ultrawork-variant-availability"
describe("resolveValidUltraworkVariant", () => {
function createClient(models: Record<string, Record<string, unknown>>) {
return {
provider: {
list: async () => ({
data: {
all: Object.entries(models).map(([providerID, providerModels]) => ({
id: providerID,
models: providerModels,
})),
},
}),
},
}
}
test("#given provider sdk metadata #when variant exists #then returns variant", async () => {
// given
const client = createClient({
anthropic: {
"claude-opus-4-6": {
variants: {
max: {},
high: {},
},
},
},
})
// when
const result = await resolveValidUltraworkVariant(
client,
{ providerID: "anthropic", modelID: "claude-opus-4-6" },
"max",
)
// then
expect(result).toBe("max")
})
test("#given provider sdk metadata #when variant does not exist #then returns undefined", async () => {
// given
const client = createClient({
anthropic: {
"claude-opus-4-6": {
variants: {
high: {},
},
},
},
})
// when
const result = await resolveValidUltraworkVariant(
client,
{ providerID: "anthropic", modelID: "claude-opus-4-6" },
"max",
)
// then
expect(result).toBeUndefined()
})
})
describe("applyUltraworkModelOverrideOnMessage variant guard", () => {
function createClient(models: Record<string, Record<string, unknown>>) {
return {
provider: {
list: async () => ({
data: {
all: Object.entries(models).map(([providerID, providerModels]) => ({
id: providerID,
models: providerModels,
})),
},
}),
},
}
}
test("#given ultrawork variant missing from target model #when override applies #then skips forced variant change", async () => {
// given
const client = createClient({
anthropic: {
"claude-opus-4-6": {
variants: {
high: {},
},
},
},
})
const dbOverrideSpy = spyOn(dbOverrideModule, "scheduleDeferredModelOverride").mockImplementation(() => {})
const config = {
agents: {
sisyphus: {
ultrawork: {
model: "anthropic/claude-opus-4-6",
variant: "max",
},
},
},
} as Parameters<typeof applyUltraworkModelOverrideOnMessage>[0]
const output = {
message: {
id: "msg_123",
model: { providerID: "anthropic", modelID: "claude-sonnet-4-6" },
} as Record<string, unknown>,
parts: [{ type: "text", text: "ultrawork do something" }],
}
// when
await applyUltraworkModelOverrideOnMessage(
config,
"sisyphus",
output,
{ showToast: async () => {} },
undefined,
client,
)
// then
expect(output.message["variant"]).toBeUndefined()
expect(output.message["thinking"]).toBeUndefined()
expect(dbOverrideSpy).toHaveBeenCalledWith(
"msg_123",
{ providerID: "anthropic", modelID: "claude-opus-4-6" },
undefined,
)
dbOverrideSpy.mockRestore()
})
test("#given variant only ultrawork config without valid current model variant #when override applies #then skips override entirely", async () => {
// given
const client = createClient({
anthropic: {
"claude-sonnet-4-6": {
variants: {
high: {},
},
},
},
})
const dbOverrideSpy = spyOn(dbOverrideModule, "scheduleDeferredModelOverride").mockImplementation(() => {})
const config = {
agents: {
sisyphus: {
ultrawork: {
variant: "max",
},
},
},
} as Parameters<typeof applyUltraworkModelOverrideOnMessage>[0]
const output = {
message: {
model: { providerID: "anthropic", modelID: "claude-sonnet-4-6" },
} as Record<string, unknown>,
parts: [{ type: "text", text: "ultrawork do something" }],
}
// when
await applyUltraworkModelOverrideOnMessage(
config,
"sisyphus",
output,
{ showToast: async () => {} },
undefined,
client,
)
// then
expect(output.message["variant"]).toBeUndefined()
expect(output.message["thinking"]).toBeUndefined()
expect(dbOverrideSpy).not.toHaveBeenCalled()
expect(output.message.model).toEqual({ providerID: "anthropic", modelID: "claude-sonnet-4-6" })
dbOverrideSpy.mockRestore()
})
})
@@ -0,0 +1,51 @@
import { normalizeSDKResponse } from "../shared"
type ModelDescriptor = {
providerID: string
modelID: string
}
type ProviderListClient = {
provider?: {
list?: () => Promise<unknown>
}
}
type ProviderModelMetadata = {
variants?: Record<string, unknown>
}
type ProviderListEntry = {
id?: string
models?: Record<string, ProviderModelMetadata>
}
type ProviderListData = {
all?: ProviderListEntry[]
}
export async function resolveValidUltraworkVariant(
client: unknown,
model: ModelDescriptor | undefined,
variant: string | undefined,
): Promise<string | undefined> {
if (!model || !variant) {
return undefined
}
const providerList = (client as ProviderListClient | null | undefined)?.provider?.list
if (typeof providerList !== "function") {
return undefined
}
const response = await providerList()
const data = normalizeSDKResponse<ProviderListData>(response, {})
const providerEntry = data.all?.find((entry) => entry.id === model.providerID)
const variants = providerEntry?.models?.[model.modelID]?.variants
if (!variants) {
return undefined
}
return Object.hasOwn(variants, variant) ? variant : undefined
}