diff --git a/packages/prompts-core/src/index.ts b/packages/prompts-core/src/index.ts index f12a83c8b..d5b68ec06 100644 --- a/packages/prompts-core/src/index.ts +++ b/packages/prompts-core/src/index.ts @@ -6,3 +6,5 @@ export type { RuntimeInjection, VariantTable, } from "./types" +export { resolveVariant } from "./variant-resolver" +export type { ResolveVariantInput } from "./variant-resolver" diff --git a/packages/prompts-core/src/variant-resolver.test.ts b/packages/prompts-core/src/variant-resolver.test.ts new file mode 100644 index 000000000..303e3d028 --- /dev/null +++ b/packages/prompts-core/src/variant-resolver.test.ts @@ -0,0 +1,50 @@ +import { describe, expect, test } from "bun:test" +import type { PromptSource, VariantTable } from "./types" +import { resolveVariant } from "./variant-resolver" + +const promptSource = (baseDir: string): PromptSource => ({ baseDir }) + +const variants = { + planner: promptSource("/prompts/planner"), + gpt: promptSource("/prompts/gpt"), + gemini: promptSource("/prompts/gemini"), + kimi: promptSource("/prompts/kimi"), + glm: promptSource("/prompts/glm"), + default: promptSource("/prompts/default"), +} satisfies VariantTable + +describe("resolveVariant", () => { + test("#given Claude Opus 4.7 model #then resolves default variant", () => { + expect(resolveVariant({ modelID: "claude-opus-4-7", variants })).toBe("default") + }) + + test("#given GPT model #then resolves gpt variant", () => { + expect(resolveVariant({ modelID: "gpt-5-5", variants })).toBe("gpt") + }) + + test("#given Gemini model #then resolves gemini variant", () => { + expect(resolveVariant({ modelID: "gemini-3-1-pro", variants })).toBe("gemini") + }) + + test("#given Kimi K2 model #then resolves kimi variant", () => { + expect(resolveVariant({ modelID: "kimi-k2-6", variants })).toBe("kimi") + }) + + test("#given GLM model #then resolves glm variant", () => { + expect(resolveVariant({ modelID: "glm-5-1", variants })).toBe("glm") + }) + + test("#given Prometheus agent #then planner overrides model variant", () => { + expect(resolveVariant({ agentName: "prometheus", modelID: "gpt-5-5", variants })).toBe( + "planner" + ) + }) + + test("#given unknown model #then falls back to default variant", () => { + expect(resolveVariant({ modelID: "claude-haiku-4-5", variants })).toBe("default") + }) + + test("#given empty variants table #then throws TypeError", () => { + expect(() => resolveVariant({ modelID: "gpt-5-5", variants: {} })).toThrow(TypeError) + }) +}) diff --git a/packages/prompts-core/src/variant-resolver.ts b/packages/prompts-core/src/variant-resolver.ts new file mode 100644 index 000000000..f2f5e1d81 --- /dev/null +++ b/packages/prompts-core/src/variant-resolver.ts @@ -0,0 +1,59 @@ +import { + isClaudeOpus47Model, + isGeminiModel, + isGlmModel, + isGptModel, + isKimiK2Model, + isMiniMaxModel, +} from "@oh-my-opencode/model-core" +import type { VariantTable } from "./types" + +export type ResolveVariantInput = { + readonly modelID?: string + readonly agentName?: string + readonly variants: VariantTable +} + +const PLANNER_AGENT_NAMES = new Set(["prometheus"]) + +const MODEL_MATCHERS: Record boolean> = { + gpt: isGptModel, + gemini: isGeminiModel, + kimi: isKimiK2Model, + glm: isGlmModel, + "opus-4-7": isClaudeOpus47Model, + minimax: isMiniMaxModel, +} + +export function resolveVariant(input: ResolveVariantInput): string { + const variantNames = Object.keys(input.variants) + if (variantNames.length === 0) { + throw new TypeError("resolveVariant requires at least one prompt variant") + } + + if (isPlannerAgent(input.agentName) && variantNames.includes("planner")) { + return "planner" + } + + if (input.modelID !== undefined) { + for (const variantName of variantNames) { + if (matchesModelVariant(variantName, input.modelID)) return variantName + } + } + + if (variantNames.includes("default")) return "default" + + const firstVariant = variantNames[0] + if (firstVariant !== undefined) return firstVariant + + throw new TypeError("resolveVariant requires at least one prompt variant") +} + +function isPlannerAgent(agentName: string | undefined): boolean { + return agentName !== undefined && PLANNER_AGENT_NAMES.has(agentName.toLowerCase()) +} + +function matchesModelVariant(variantName: string, modelID: string): boolean { + const matcher = MODEL_MATCHERS[variantName] + return matcher?.(modelID) ?? false +}