From 29e1136813e87cb4378b924fbd2a7d25096d14eb Mon Sep 17 00:00:00 2001 From: YeonGyu-Kim Date: Wed, 11 Mar 2026 15:15:50 +0900 Subject: [PATCH] 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 --- src/plugin/chat-message.ts | 9 +- src/plugin/ultrawork-model-override.ts | 136 ++++++++----- .../ultrawork-variant-availability.test.ts | 186 ++++++++++++++++++ src/plugin/ultrawork-variant-availability.ts | 51 +++++ 4 files changed, 327 insertions(+), 55 deletions(-) create mode 100644 src/plugin/ultrawork-variant-availability.test.ts create mode 100644 src/plugin/ultrawork-variant-availability.ts diff --git a/src/plugin/chat-message.ts b/src/plugin/chat-message.ts index 96555d021..750eaf667 100644 --- a/src/plugin/chat-message.ts +++ b/src/plugin/chat-message.ts @@ -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, + ) } } diff --git a/src/plugin/ultrawork-model-override.ts b/src/plugin/ultrawork-model-override.ts index 736926bf6..980de1752 100644 --- a/src/plugin/ultrawork-model-override.ts +++ b/src/plugin/ultrawork-model-override.ts @@ -4,6 +4,7 @@ import { getSessionAgent } from "../features/claude-code-session-state" import { log } from "../shared" import { getAgentConfigKey } from "../shared/agent-display-names" import { scheduleDeferredModelOverride } from "./ultrawork-db-model-override" +import { resolveValidUltraworkVariant } from "./ultrawork-variant-availability" const CODE_BLOCK = /```[\s\S]*?```/g const INLINE_CODE = /`[^`]+`/g @@ -15,7 +16,7 @@ export function detectUltrawork(text: string): boolean { } 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 = { @@ -36,22 +37,26 @@ export type UltraworkOverrideResult = { variant?: string } -function isSameModel( - current: unknown, - target: { providerID: string; modelID: string }, -): boolean { - if (typeof current !== "object" || current === null) return false - const currentRecord = current as Record - return ( - currentRecord["providerID"] === target.providerID - && currentRecord["modelID"] === target.modelID - ) +type ModelDescriptor = { + providerID: string + modelID: string +} + +function isSameModel(current: unknown, target: ModelDescriptor): boolean { + if (typeof current !== "object" || current === null) return false + const currentRecord = current as Record + 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 + 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( pluginConfig: OhMyOpenCodeConfig, inputAgentName: string | undefined, @@ -76,9 +81,7 @@ export function resolveUltraworkOverride( if (!ultraworkConfig?.model && !ultraworkConfig?.variant) return null if (!ultraworkConfig.model) { - return { - variant: ultraworkConfig.variant, - } + return { variant: ultraworkConfig.variant } } const modelParts = ultraworkConfig.model.split("/") @@ -91,37 +94,20 @@ export function resolveUltraworkOverride( } } -/** - * Applies ultrawork model override using a deferred DB update strategy. - * - * Instead of directly mutating output.message.model (which would cause the TUI - * bottom bar to show the override model), this schedules a queueMicrotask that - * updates the message model directly in SQLite AFTER Session.updateMessage() - * saves the original model, but BEFORE loop() reads it for the API call. - * - * Result: API call uses opus, TUI bottom bar stays on sonnet. - */ -export function applyUltraworkModelOverrideOnMessage( - pluginConfig: OhMyOpenCodeConfig, - inputAgentName: string | undefined, - output: { - message: Record - 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 +function applyResolvedUltraworkOverride(args: { + override: UltraworkOverrideResult + validatedVariant: string | undefined + output: { message: Record } + inputAgentName: string | undefined + tui: unknown +}): void { + const { override, validatedVariant, output, inputAgentName, tui } = args + if (validatedVariant) { + output.message["variant"] = validatedVariant + output.message["thinking"] = validatedVariant } - if (!override.providerID || !override.modelID) { - return - } + if (!override.providerID || !override.modelID) return const targetModel = { providerID: override.providerID, modelID: override.modelID } 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") output.message.model = targetModel return - } 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"), ) - scheduleDeferredModelOverride( - messageId, - targetModel, - override.variant, - ) + scheduleDeferredModelOverride(messageId, targetModel, validatedVariant) log(`[ultrawork-model-override] ${fromModel} -> ${override.modelID} (deferred DB)`, { agent: agentConfigKey, @@ -156,6 +137,53 @@ export function applyUltraworkModelOverrideOnMessage( showToast( tui, "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 + parts: Array<{ type: string; text?: string; [key: string]: unknown }> + }, + tui: unknown, + sessionID?: string, + client?: unknown, +): void | Promise { + 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 }) + }) +} diff --git a/src/plugin/ultrawork-variant-availability.test.ts b/src/plugin/ultrawork-variant-availability.test.ts new file mode 100644 index 000000000..1fb8a0910 --- /dev/null +++ b/src/plugin/ultrawork-variant-availability.test.ts @@ -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>) { + 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>) { + 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[0] + + const output = { + message: { + id: "msg_123", + model: { providerID: "anthropic", modelID: "claude-sonnet-4-6" }, + } as Record, + 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[0] + + const output = { + message: { + model: { providerID: "anthropic", modelID: "claude-sonnet-4-6" }, + } as Record, + 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() + }) +}) diff --git a/src/plugin/ultrawork-variant-availability.ts b/src/plugin/ultrawork-variant-availability.ts new file mode 100644 index 000000000..b1ce97a87 --- /dev/null +++ b/src/plugin/ultrawork-variant-availability.ts @@ -0,0 +1,51 @@ +import { normalizeSDKResponse } from "../shared" + +type ModelDescriptor = { + providerID: string + modelID: string +} + +type ProviderListClient = { + provider?: { + list?: () => Promise + } +} + +type ProviderModelMetadata = { + variants?: Record +} + +type ProviderListEntry = { + id?: string + models?: Record +} + +type ProviderListData = { + all?: ProviderListEntry[] +} + +export async function resolveValidUltraworkVariant( + client: unknown, + model: ModelDescriptor | undefined, + variant: string | undefined, +): Promise { + 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(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 +}