diff --git a/README.md b/README.md index d70019483..588b2a5e7 100644 --- a/README.md +++ b/README.md @@ -868,10 +868,19 @@ providers: auth: none discovery: type: llama.cpp + +equivalence: + overrides: + zenmux/codex: gpt-5.3-codex + p-codex/codex: gpt-5.3-codex + exclude: + - demo/codex-preview ``` **Supported APIs:** `openai-completions`, `openai-responses`, `openai-codex-responses`, `azure-openai-responses`, `anthropic-messages`, `google-generative-ai`, `google-vertex` +Canonical ids are official upstream model ids such as `claude-sonnet-4-6` or `gpt-5.3-codex`. Use `equivalence.overrides` to map custom provider variants into those canonical groups while keeping explicit `provider/model` selection available. + ### Settings File Global settings are stored in: @@ -889,13 +898,17 @@ theme: enabledModels: - "anthropic/*" - - "*gpt*" + - "gpt-5.3-codex" - "gemini-2.5-pro:high" modelRoles: - default: anthropic/claude-sonnet-4-20250514 - plan: anthropic/claude-opus-4-1:high - smol: anthropic/claude-sonnet-4-20250514 + default: claude-sonnet-4-6 + plan: claude-opus-4-6:high + smol: anthropic/claude-sonnet-4-6 +modelProviderOrder: + - github-copilot + - zenmux + - openai defaultThinkingLevel: high retry: @@ -955,6 +968,8 @@ task: merge: patch # patch | branch ``` +`modelRoles` may use either canonical ids or explicit `provider/model` selectors. `modelProviderOrder` decides which provider backs a canonical model when multiple equivalent variants are available. + Legacy migration notes: - `settings.json` → `config.yml` diff --git a/docs/models.md b/docs/models.md index a83afd4d6..2393aaa8a 100644 --- a/docs/models.md +++ b/docs/models.md @@ -29,10 +29,20 @@ Legacy behavior still present: providers: : # provider-level config +equivalence: + overrides: + /: + exclude: + - / ``` `provider-id` is the canonical provider key used across selection and auth lookup. +`equivalence` is optional and configures canonical model grouping on top of concrete provider models: + +- `overrides` maps an exact concrete selector (`provider/modelId`) to an official upstream canonical id +- `exclude` opts a concrete selector out of canonical grouping + ## Provider-level fields ```yaml @@ -134,6 +144,71 @@ ModelRegistry pipeline (on refresh): - otherwise append 6. Apply runtime-discovered models (currently Ollama and LM Studio), then re-apply model overrides. +## Canonical model equivalence and coalescing + +The registry keeps every concrete provider model and then builds a canonical layer above them. + +Canonical ids are official upstream ids only, for example: + +- `claude-opus-4-6` +- `claude-haiku-4-5` +- `gpt-5.3-codex` +### `models.yml` equivalence config + +Example: + +```yaml +providers: + zenmux: + baseUrl: https://api.zenmux.example/v1 + apiKey: ZENMUX_API_KEY + api: openai-codex-responses + models: + - id: codex + name: Zenmux Codex + reasoning: true + input: [text] + cost: + input: 0 + output: 0 + cacheRead: 0 + cacheWrite: 0 + contextWindow: 200000 + maxTokens: 32768 + +equivalence: + overrides: + zenmux/codex: gpt-5.3-codex + p-codex/codex: gpt-5.3-codex + exclude: + - demo/codex-preview +``` + +Build order for canonical grouping: + +1. exact user override from `equivalence.overrides` +2. bundled official-id matches from built-in model metadata +3. conservative heuristic normalization for gateway/provider variants +4. fallback to the concrete model's own id + +Current heuristics are intentionally narrow: + +- embedded upstream prefixes can be stripped when present, for example `anthropic/...` or `openai/...` +- dotted and dashed version variants can normalize only when they map to an existing official id, for example `4.6 -> 4-6` +- ambiguous families or versions are not merged without a bundled match or explicit override + +### Canonical resolution behavior + +When multiple concrete variants share a canonical id, resolution uses: + +1. availability and auth +2. `config.yml` `modelProviderOrder` +3. existing registry/provider order if `modelProviderOrder` is unset + +Disabled or unauthenticated providers are skipped. + +Session state and transcripts continue to record the concrete provider/model that actually executed the turn. + Provider defaults vs per-model overrides: - Provider `headers` are baseline. @@ -244,6 +319,7 @@ So a model can exist in registry but not be selectable until auth is available. `model-resolver.ts` supports: - exact `provider/modelId` +- exact canonical model id - exact model id (provider inferred) - fuzzy/substring matching - glob scope patterns in `--models` (e.g. `openai/*`, `*sonnet*`) @@ -251,6 +327,13 @@ So a model can exist in registry but not be selectable until auth is available. `--provider` is legacy; `--model` is preferred. +Resolution precedence for exact selectors: + +1. exact `provider/modelId` bypasses coalescing +2. exact canonical id resolves through the canonical index +3. exact bare concrete id still works +4. fuzzy and glob matching run after the exact paths + ### Initial model selection priority `findInitialModel(...)` uses this order: @@ -275,9 +358,32 @@ Related settings: - `modelRoles` (record) - `enabledModels` (scoped pattern list) +- `modelProviderOrder` (global canonical-provider precedence) - `providers.kimiApiFormat` (`openai` or `anthropic` request format) - `providers.openaiWebsockets` (`auto|off|on` websocket preference for OpenAI Codex transport) +`modelRoles` may store either: + +- `provider/modelId` to pin a concrete provider variant +- a canonical id such as `gpt-5.3-codex` to allow provider coalescing + +For `enabledModels` and CLI `--models`: + +- exact canonical ids expand to all concrete variants in that canonical group +- explicit `provider/modelId` entries stay exact +- globs and fuzzy matches still operate on concrete models + +## `/model` and `--list-models` + +Both surfaces keep provider-prefixed models visible and selectable. + +They now also expose canonical/coalesced models: + +- `/model` includes a canonical view alongside provider tabs +- `--list-models` prints a canonical section plus the concrete provider rows + +Selecting a canonical entry stores the canonical selector. Selecting a provider row stores the explicit `provider/modelId`. + ## Context promotion (model-level fallback chains) Context promotion is an overflow recovery mechanism for small-context variants (for example `*-spark`) that automatically promotes to a larger-context sibling when the API rejects a request with a context length error. diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index d389c4f8f..ecff2b268 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,16 +1,27 @@ # Changelog ## [Unreleased] - ### Added - Added rendering of usage report entries for accounts with no usage limits, including account label and optional plan type with a `-- no limits` indicator - Updated account label resolution to fall back to email or accountId so unlabeled unlimited-plan accounts display a meaningful name +- Added canonical model equivalence and provider coalescing across `models.yml`, `enabledModels`, `--models`, `/model`, and `--list-models` +- Added `equivalence` overrides/exclusions to `models.yml` and `modelProviderOrder` to `config.yml` for global canonical-provider preference ### Changed +- Updated interactive and CLI model listings/selectors to work with canonical model ids while resolving them to concrete provider variants for actual execution +- Updated role assignment persistence so selected model settings now store the selector used by users, including thinking-level suffixes, while runtime continues to run against the resolved concrete provider model +- Updated model scope resolution to expand exact canonical model ids into all matching provider variants when filtering supported model sets - Changed the agent to avoid giving time estimates or task-duration predictions in user responses, focusing on required work instead - Changed generated code guidance to avoid speculative abstractions and extra compatibility scaffolding, favoring direct implementations that match current needs +- Changed model role resolution so roles can store either canonical model ids or explicit `provider/model` selectors while sessions continue to record the concrete model actually used + +### Fixed + +- Fixed model resolution for commit message generation, title generation, memory consolidation, and image inspection when role strings use canonical ids instead of raw provider/model values +- Fixed default-model updates so previously configured thinking levels were preserved when reassigning a role +- Fixed model scope and selection handling in CLI/session startup paths that previously failed to resolve aliases consistently across features ## [14.0.5] - 2026-04-11 ### Added diff --git a/packages/coding-agent/src/cli/list-models.ts b/packages/coding-agent/src/cli/list-models.ts index 163be17cf..4a1af2f80 100644 --- a/packages/coding-agent/src/cli/list-models.ts +++ b/packages/coding-agent/src/cli/list-models.ts @@ -6,6 +6,45 @@ import { formatNumber } from "@oh-my-pi/pi-utils"; import type { ModelRegistry } from "../config/model-registry"; import { fuzzyFilter } from "../utils/fuzzy"; +interface ProviderRow { + provider: string; + model: string; + context: string; + maxOut: string; + thinking: string; + images: string; +} + +interface CanonicalRow { + canonical: string; + selected: string; + variants: string; + context: string; + maxOut: string; +} + +function writeLine(line = ""): void { + process.stdout.write(`${line}\n`); +} + +function renderTable>(rows: T[], headers: T): void { + const widths = Object.fromEntries( + Object.keys(headers).map(key => [key, Math.max(headers[key]!.length, ...rows.map(row => row[key]!.length))]), + ) as Record; + + const headerLine = Object.keys(headers) + .map(key => headers[key as keyof T]!.padEnd(widths[key as keyof T])) + .join(" "); + writeLine(headerLine); + + for (const row of rows) { + const line = Object.keys(headers) + .map(key => row[key as keyof T]!.padEnd(widths[key as keyof T])) + .join(" "); + writeLine(line); + } +} + /** * List available models, optionally filtered by search pattern */ @@ -13,77 +52,77 @@ export async function listModels(modelRegistry: ModelRegistry, searchPattern?: s const models = modelRegistry.getAvailable(); if (models.length === 0) { - console.log("No models available. Set API keys in environment variables."); + writeLine("No models available. Set API keys in environment variables."); return; } - // Apply fuzzy filter if search pattern provided let filteredModels: Model[] = models; if (searchPattern) { - filteredModels = fuzzyFilter(models, searchPattern, m => `${m.provider} ${m.id}`); + filteredModels = fuzzyFilter(models, searchPattern, model => `${model.provider} ${model.id}`); } - if (filteredModels.length === 0) { - console.log(`No models matching "${searchPattern}"`); + const filteredCanonical = modelRegistry + .getCanonicalModels({ availableOnly: true, candidates: filteredModels }) + .map(record => { + const selected = modelRegistry.resolveCanonicalModel(record.id, { + availableOnly: true, + candidates: filteredModels, + }); + if (!selected) return undefined; + return { + canonical: record.id, + selected: `${selected.provider}/${selected.id}`, + variants: String(record.variants.length), + context: formatNumber(selected.contextWindow), + maxOut: formatNumber(selected.maxTokens), + } satisfies CanonicalRow; + }) + .filter((row): row is CanonicalRow => row !== undefined) + .sort((left, right) => left.canonical.localeCompare(right.canonical)); + + if (filteredModels.length === 0 && filteredCanonical.length === 0) { + writeLine(`No models matching "${searchPattern}"`); return; } - // Sort by provider, then by model id - filteredModels.sort((a, b) => { - const providerCmp = a.provider.localeCompare(b.provider); + filteredModels.sort((left, right) => { + const providerCmp = left.provider.localeCompare(right.provider); if (providerCmp !== 0) return providerCmp; - return a.id.localeCompare(b.id); + return left.id.localeCompare(right.id); }); - // Calculate column widths - const rows = filteredModels.map(m => ({ - provider: m.provider, - model: m.id, - context: formatNumber(m.contextWindow), - maxOut: formatNumber(m.maxTokens), - thinking: m.thinking ? getSupportedEfforts(m).join(",") : m.reasoning ? "yes" : "-", - images: m.input.includes("image") ? "yes" : "no", - })); + const providerRows = filteredModels.map(model => ({ + provider: model.provider, + model: model.id, + context: formatNumber(model.contextWindow), + maxOut: formatNumber(model.maxTokens), + thinking: model.thinking ? getSupportedEfforts(model).join(",") : model.reasoning ? "yes" : "-", + images: model.input.includes("image") ? "yes" : "no", + })) satisfies ProviderRow[]; - const headers = { - provider: "provider", - model: "model", - context: "context", - maxOut: "max-out", - thinking: "thinking", - images: "images", - }; + if (filteredCanonical.length > 0) { + writeLine("Canonical models"); + renderTable(filteredCanonical, { + canonical: "canonical", + selected: "selected", + variants: "variants", + context: "context", + maxOut: "max-out", + }); + if (providerRows.length > 0) { + writeLine(); + } + } - const widths = { - provider: Math.max(headers.provider.length, ...rows.map(r => r.provider.length)), - model: Math.max(headers.model.length, ...rows.map(r => r.model.length)), - context: Math.max(headers.context.length, ...rows.map(r => r.context.length)), - maxOut: Math.max(headers.maxOut.length, ...rows.map(r => r.maxOut.length)), - thinking: Math.max(headers.thinking.length, ...rows.map(r => r.thinking.length)), - images: Math.max(headers.images.length, ...rows.map(r => r.images.length)), - }; - - // Print header - const headerLine = [ - headers.provider.padEnd(widths.provider), - headers.model.padEnd(widths.model), - headers.context.padEnd(widths.context), - headers.maxOut.padEnd(widths.maxOut), - headers.thinking.padEnd(widths.thinking), - headers.images.padEnd(widths.images), - ].join(" "); - console.log(headerLine); - - // Print rows - for (const row of rows) { - const line = [ - row.provider.padEnd(widths.provider), - row.model.padEnd(widths.model), - row.context.padEnd(widths.context), - row.maxOut.padEnd(widths.maxOut), - row.thinking.padEnd(widths.thinking), - row.images.padEnd(widths.images), - ].join(" "); - console.log(line); + if (providerRows.length > 0) { + writeLine("Provider models"); + renderTable(providerRows, { + provider: "provider", + model: "model", + context: "context", + maxOut: "max-out", + thinking: "thinking", + images: "images", + }); } } diff --git a/packages/coding-agent/src/commit/model-selection.ts b/packages/coding-agent/src/commit/model-selection.ts index 2416256b2..3045dfdf6 100644 --- a/packages/coding-agent/src/commit/model-selection.ts +++ b/packages/coding-agent/src/commit/model-selection.ts @@ -1,7 +1,12 @@ import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core"; import type { Api, Model } from "@oh-my-pi/pi-ai"; import { MODEL_ROLE_IDS } from "../config/model-registry"; -import { parseModelPattern, resolveModelRoleValue, resolveRoleSelection } from "../config/model-resolver"; +import { + type ModelLookupRegistry, + parseModelPattern, + resolveModelRoleValue, + resolveRoleSelection, +} from "../config/model-resolver"; import type { Settings } from "../config/settings"; import MODEL_PRIO from "../priority.json" with { type: "json" }; @@ -11,19 +16,20 @@ export interface ResolvedCommitModel { thinkingLevel?: ThinkingLevel; } +type CommitModelRegistry = ModelLookupRegistry & { + getApiKey: (model: Model) => Promise; +}; + export async function resolvePrimaryModel( override: string | undefined, settings: Settings, - modelRegistry: { - getAvailable: () => Model[]; - getApiKey: (model: Model) => Promise; - }, + modelRegistry: CommitModelRegistry, ): Promise { const available = modelRegistry.getAvailable(); const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() }; const resolved = override - ? resolveModelRoleValue(override, available, { settings, matchPreferences }) - : resolveRoleSelection(["commit", "smol", ...MODEL_ROLE_IDS], settings, available); + ? resolveModelRoleValue(override, available, { settings, matchPreferences, modelRegistry }) + : resolveRoleSelection(["commit", "smol", ...MODEL_ROLE_IDS], settings, available, modelRegistry); const model = resolved?.model; if (!model) { throw new Error("No model available for commit generation"); @@ -37,15 +43,12 @@ export async function resolvePrimaryModel( export async function resolveSmolModel( settings: Settings, - modelRegistry: { - getAvailable: () => Model[]; - getApiKey: (model: Model) => Promise; - }, + modelRegistry: CommitModelRegistry, fallbackModel: Model, fallbackApiKey: string, ): Promise { const available = modelRegistry.getAvailable(); - const resolvedSmol = resolveRoleSelection(["smol"], settings, available); + const resolvedSmol = resolveRoleSelection(["smol"], settings, available, modelRegistry); if (resolvedSmol?.model) { const apiKey = await modelRegistry.getApiKey(resolvedSmol.model); if (apiKey) return { model: resolvedSmol.model, apiKey, thinkingLevel: resolvedSmol.thinkingLevel }; @@ -53,7 +56,7 @@ export async function resolveSmolModel( const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() }; for (const pattern of MODEL_PRIO.smol) { - const candidate = parseModelPattern(pattern, available, matchPreferences).model; + const candidate = parseModelPattern(pattern, available, matchPreferences, { modelRegistry }).model; if (!candidate) continue; const apiKey = await modelRegistry.getApiKey(candidate); if (apiKey) return { model: candidate, apiKey }; diff --git a/packages/coding-agent/src/config/model-equivalence.ts b/packages/coding-agent/src/config/model-equivalence.ts new file mode 100644 index 000000000..12c932d34 --- /dev/null +++ b/packages/coding-agent/src/config/model-equivalence.ts @@ -0,0 +1,674 @@ +import { type Api, getBundledModels, getBundledProviders, type Model } from "@oh-my-pi/pi-ai"; + +export type CanonicalModelSource = "override" | "bundled" | "heuristic" | "fallback"; + +export interface ModelEquivalenceConfig { + overrides?: Record; + exclude?: string[]; +} + +export interface CanonicalModelVariant { + canonicalId: string; + selector: string; + model: Model; + source: CanonicalModelSource; +} + +export interface CanonicalModelRecord { + id: string; + name: string; + variants: CanonicalModelVariant[]; +} + +export interface CanonicalModelIndex { + records: CanonicalModelRecord[]; + byId: Map; + bySelector: Map; +} + +interface CanonicalReferenceData { + references: Map>; + officialIds: Set; +} + +interface CompiledEquivalenceConfig { + overrides: Map; + exclude: Set; +} + +interface ResolvedCanonicalModel { + id: string; + source: CanonicalModelSource; +} + +const TRAILING_CANONICAL_MARKERS = [ + "thinking", + "customtools", + "high", + "low", + "medium", + "minimal", + "xhigh", + "free", + "exacto", + "original", + "optimized", + "nvfp4", + "fp8", + "fp4", + "bf16", + "int8", + "int4", +] as const; +const WRAPPER_PREFIXES = ["duo-chat-"] as const; +const FAMILY_EXTRACTION_PATTERNS = [ + /(?:^|[/:._-])((?:claude|gemini|gpt|grok|glm|qwen|minimax|kimi|deepseek|llama|gemma|nova|mistral|ministral|pixtral|codestral|devstral|magistral|ernie|doubao|seed|aion|olmo|molmo|nemotron|palmyra|command|codex|coder|o[1345])[-a-z0-9.]+)(?::|$)/i, + /(?:^|[/:._-])((?:claude|gemini|gpt|grok|glm|qwen|minimax|kimi|deepseek|llama|gemma|nova|mistral|ministral|pixtral|codestral|devstral|magistral|ernie|doubao|seed|aion|olmo|molmo|nemotron|palmyra|command|codex|coder|o[1345])[-a-z0-9.]+(?:[-_/][a-z0-9.]+)*)(?::|$)/i, +] as const; + +function shouldReplaceReference(existing: Model | undefined, candidate: Model): boolean { + if (!existing) return true; + if (candidate.contextWindow !== existing.contextWindow) { + return candidate.contextWindow > existing.contextWindow; + } + if (candidate.maxTokens !== existing.maxTokens) { + return candidate.maxTokens > existing.maxTokens; + } + return existing.provider !== "openai" && candidate.provider === "openai"; +} + +function createCanonicalReferenceData(): CanonicalReferenceData { + const references = new Map>(); + for (const provider of getBundledProviders()) { + for (const model of getBundledModels(provider as Parameters[0])) { + const candidate = model as Model; + const existing = references.get(candidate.id); + if (shouldReplaceReference(existing, candidate)) { + references.set(candidate.id, candidate); + } + } + } + return { + references, + officialIds: new Set(references.keys()), + }; +} + +function normalizeSelectorKey(selector: string): string { + return selector.trim().toLowerCase(); +} + +function normalizeCanonicalIdKey(canonicalId: string): string { + return canonicalId.trim().toLowerCase(); +} + +export function formatCanonicalVariantSelector(model: Model): string { + return `${model.provider}/${model.id}`; +} + +function buildOverrideMap(overrides: Record | undefined): Map { + const result = new Map(); + if (!overrides) { + return result; + } + for (const [selector, canonicalId] of Object.entries(overrides)) { + const normalizedSelector = normalizeSelectorKey(selector); + const normalizedCanonicalId = canonicalId.trim(); + if (!normalizedSelector || !normalizedCanonicalId) { + continue; + } + result.set(normalizedSelector, normalizedCanonicalId); + } + return result; +} + +function buildExclusionSet(exclusions: readonly string[] | undefined): Set { + const result = new Set(); + for (const selector of exclusions ?? []) { + const normalized = normalizeSelectorKey(selector); + if (normalized) { + result.add(normalized); + } + } + return result; +} + +function compileEquivalenceConfig(config: ModelEquivalenceConfig | undefined): CompiledEquivalenceConfig { + return { + overrides: buildOverrideMap(config?.overrides), + exclude: buildExclusionSet(config?.exclude), + }; +} + +function addCanonicalCandidate(candidates: Set, candidate: string): void { + const normalized = candidate.trim(); + if (normalized) { + candidates.add(normalized); + } +} + +function stripTrailingMarker(candidate: string): string | undefined { + for (const marker of TRAILING_CANONICAL_MARKERS) { + for (const separator of ["-", ":"] as const) { + const suffix = `${separator}${marker}`; + if (candidate.toLowerCase().endsWith(suffix)) { + return candidate.slice(0, -suffix.length); + } + } + } + return undefined; +} + +function lowercaseCandidate(candidate: string): string | undefined { + const lowercased = candidate.toLowerCase(); + return lowercased !== candidate ? lowercased : undefined; +} + +function stripSyntheticPrefix(candidate: string): string | undefined { + const stripped = candidate.replace(/^hf:/i, ""); + return stripped !== candidate ? stripped : undefined; +} + +function stripLatestSuffix(candidate: string): string | undefined { + const stripped = candidate.replace(/-latest$/i, ""); + return stripped !== candidate ? stripped : undefined; +} + +function stripLegacyGlmTurboSuffix(candidate: string): string | undefined { + const stripped = candidate.replace(/^(glm-4(?:\.\d+)?v?)-turbo$/i, "$1"); + return stripped !== candidate ? stripped : undefined; +} + +function reorderAnthropicFamily(candidate: string): string | undefined { + const match = /^claude-(\d+(?:[.-]\d+)+)-(opus|sonnet|haiku)$/i.exec(candidate); + if (!match) { + return undefined; + } + const [, version, family] = match; + return `claude-${family.toLowerCase()}-${version}`; +} + +function stripProviderVersionSuffix(candidate: string): string | undefined { + const stripped = candidate.replace(/-v\d+(?::\d+)?$/i, ""); + return stripped !== candidate ? stripped : undefined; +} + +function stripDateSuffix(candidate: string): string | undefined { + const stripped = candidate.replace(/-\d{8}$/i, ""); + return stripped !== candidate ? stripped : undefined; +} + +function insertAttachedFamilyVersionSeparator(candidate: string): string | undefined { + const inserted = candidate.replace( + /(^|[/:._-])((?:claude|gemini|gpt|grok|glm|qwen|minimax|kimi|deepseek|llama|gemma|nova|mistral|ministral|pixtral|codestral|devstral|magistral|ernie|doubao|seed|aion|olmo|molmo|nemotron|palmyra|command|codex|coder))(\d+(?:[.-]\d+)*)(?=$|[-_/.:a-z])/gi, + "$1$2-$3", + ); + return inserted !== candidate ? inserted : undefined; +} + +function toggleSeriesMinorVersionSeparators(candidate: string): string[] { + const toggled = new Set(); + const dotToDash = candidate.replace(/(^|[/:._-])([a-z])(\d)\.(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3-$4"); + if (dotToDash !== candidate) { + toggled.add(dotToDash); + } + const dashToDot = candidate.replace(/(^|[/:._-])([a-z])(\d)-(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3.$4"); + if (dashToDot !== candidate) { + toggled.add(dashToDot); + } + return [...toggled]; +} + +function expandCompactSeriesMinorVersions(candidate: string): string[] { + const expanded = new Set(); + const compactToDash = candidate.replace(/(^|[/:._-])([a-z])(\d)(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3-$4"); + if (compactToDash !== candidate) { + expanded.add(compactToDash); + } + const compactToDot = candidate.replace(/(^|[/:._-])([a-z])(\d)(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3.$4"); + if (compactToDot !== candidate) { + expanded.add(compactToDot); + } + return [...expanded]; +} + +function getQualifiedNamespaceSuffixes(candidate: string): string[] { + const results = new Set(); + for (let index = 1; index < candidate.length; index += 1) { + if (!/[/:.]/.test(candidate[index - 1]!)) { + continue; + } + const suffix = candidate.slice(index); + if (suffix.length < 4) { + continue; + } + if (!/[a-z]/i.test(suffix) || !/\d/.test(suffix)) { + continue; + } + addCanonicalCandidate(results, suffix); + } + return [...results]; +} + +function extractUpstreamFamilyCandidate(candidate: string): string | undefined { + for (const pattern of FAMILY_EXTRACTION_PATTERNS) { + const match = pattern.exec(candidate); + if (match?.[1]) { + return match[1]; + } + } + return undefined; +} + +function getCandidatePenalty(candidate: string): number { + let penalty = 0; + if (candidate.includes("/")) { + penalty += 100; + } + if (candidate.includes(":")) { + penalty += 40; + } + if (/-\d{8}$/i.test(candidate)) { + penalty += 25; + } + if (/-v\d+(?::\d+)?$/i.test(candidate)) { + penalty += 25; + } + if (stripTrailingMarker(candidate)) { + penalty += 20; + } + if (/[A-Z]/.test(candidate)) { + penalty += 10; + } + if (/^claude-\d/i.test(candidate)) { + penalty += 20; + } + if (/^claude-(?:opus|sonnet|haiku)-\d{2}(?=$|[-_a-z])/i.test(candidate)) { + penalty += 10; + } + if (/(?:^|[/:._-])[a-z]\d-\d(?=$|[-_/.:a-z])/i.test(candidate)) { + penalty += 6; + } + if (/(?:^|[-_/])\d-\d(?=$|[-_a-z])/.test(candidate) && !/^claude-(?:opus|sonnet|haiku)-\d-\d/i.test(candidate)) { + penalty += 4; + } + penalty += candidate.length * 0.01; + return penalty; +} + +function compareCandidatePreference(left: string, right: string): number { + const penaltyDiff = getCandidatePenalty(left) - getCandidatePenalty(right); + if (penaltyDiff !== 0) { + return penaltyDiff; + } + if (left.length !== right.length) { + return left.length - right.length; + } + return left.localeCompare(right); +} + +function selectBestOfficialCandidate(candidates: readonly string[]): string | undefined { + if (candidates.length === 0) { + return undefined; + } + const ranked = [...new Set(candidates)].sort(compareCandidatePreference); + return ranked[0]; +} + +function getWrapperCanonicalCandidates(candidate: string): string[] { + const results = new Set(); + for (const prefix of WRAPPER_PREFIXES) { + if (!candidate.toLowerCase().startsWith(prefix)) { + continue; + } + const stripped = candidate.slice(prefix.length); + addCanonicalCandidate(results, stripped); + if (/^(opus|sonnet|haiku)-/i.test(stripped)) { + addCanonicalCandidate(results, `claude-${stripped}`); + } + } + return [...results]; +} + +function getAnthropicAliasOfficial(candidate: string, officialIds: Set): string | undefined { + const reordered = reorderAnthropicFamily(candidate); + if (!reordered) { + return undefined; + } + const candidates = [reordered, ...toggleShortVersionSeparators(reordered)].filter(officialId => + officialIds.has(officialId), + ); + return selectBestOfficialCandidate(candidates); +} + +function compareVersionSegments(left: readonly number[], right: readonly number[]): number { + const maxLength = Math.max(left.length, right.length); + for (let index = 0; index < maxLength; index += 1) { + const diff = (left[index] ?? Number.NEGATIVE_INFINITY) - (right[index] ?? Number.NEGATIVE_INFINITY); + if (diff !== 0) { + return diff; + } + } + return 0; +} + +function parseClaudeFamilyVersionSegments(candidate: string, prefix: string): number[] { + const normalizedCandidate = candidate.toLowerCase(); + const normalizedPrefix = prefix.toLowerCase(); + if (!normalizedCandidate.startsWith(`${normalizedPrefix}-`)) { + return []; + } + const rawSuffix = normalizedCandidate.slice(normalizedPrefix.length + 1); + if (!rawSuffix) { + return []; + } + const versionSegments: number[] = []; + for (const token of rawSuffix.split("-")) { + if (!token) { + break; + } + if (/^\d{8}$/.test(token)) { + break; + } + if (/^\d{2}$/.test(token)) { + versionSegments.push(Number(token[0]), Number(token[1])); + continue; + } + if (/^\d+(?:\.\d+)*$/.test(token)) { + versionSegments.push(...token.split(".").map(part => Number(part))); + continue; + } + break; + } + return versionSegments; +} + +function getClaudeFamilyAliasOfficial(candidate: string, officialIds: Set): string | undefined { + const match = /^(?:anthropic\/)?(claude(?:-\d(?:[.-]\d+)?)?-(?:haiku|opus|sonnet))(?:-latest)?$/i.exec(candidate); + if (!match?.[1]) { + return undefined; + } + const familyPrefix = match[1].toLowerCase(); + const familyMatches = [...officialIds].filter(officialId => { + const normalizedOfficialId = officialId.toLowerCase(); + return normalizedOfficialId.startsWith(`${familyPrefix}-`) || normalizedOfficialId === familyPrefix; + }); + if (familyMatches.length === 0) { + return undefined; + } + return [...familyMatches].sort((left, right) => { + const versionDiff = compareVersionSegments( + parseClaudeFamilyVersionSegments(right, familyPrefix), + parseClaudeFamilyVersionSegments(left, familyPrefix), + ); + if (versionDiff !== 0) { + return versionDiff; + } + const leftHasDate = /-\d{8}(?:$|-)/i.test(left); + const rightHasDate = /-\d{8}(?:$|-)/i.test(right); + if (leftHasDate !== rightHasDate) { + return leftHasDate ? 1 : -1; + } + const leftHasMarker = stripTrailingMarker(left) !== undefined; + const rightHasMarker = stripTrailingMarker(right) !== undefined; + if (leftHasMarker !== rightHasMarker) { + return leftHasMarker ? 1 : -1; + } + return compareCandidatePreference(left, right); + })[0]; +} + +function toggleShortVersionSeparators(candidate: string): string[] { + const toggled = new Set(); + const dotToDash = candidate.replace(/(^|[-_/])(\d{1,2})\.(\d{1,2})(?=$|[-_a-z])/gi, "$1$2-$3"); + if (dotToDash !== candidate) { + toggled.add(dotToDash); + } + const dashToDot = candidate.replace(/(^|[-_/])(\d{1,2})-(\d{1,2})(?=$|[-_a-z])/gi, "$1$2.$3"); + if (dashToDot !== candidate) { + toggled.add(dashToDot); + } + return [...toggled]; +} + +function expandCompactMinorVersions(candidate: string): string[] { + const expanded = new Set(); + const compactToDash = candidate.replace(/(^|[-_/])(\d)(\d)(?=$|[-_a-z])/g, "$1$2-$3"); + if (compactToDash !== candidate) { + expanded.add(compactToDash); + } + const compactToDot = candidate.replace(/(^|[-_/])(\d)(\d)(?=$|[-_a-z])/g, "$1$2.$3"); + if (compactToDot !== candidate) { + expanded.add(compactToDot); + } + return [...expanded]; +} + +function getHeuristicCanonicalCandidates(modelId: string): string[] { + const candidates = new Set(); + const queue = [modelId]; + const visited = new Set(); + + while (queue.length > 0) { + const candidate = queue.shift(); + if (!candidate) { + continue; + } + const normalized = candidate.trim(); + if (!normalized || visited.has(normalized)) { + continue; + } + visited.add(normalized); + addCanonicalCandidate(candidates, normalized); + + const lowercased = lowercaseCandidate(normalized); + if (lowercased) { + queue.push(lowercased); + } + + const pathSegments = normalized.split("/"); + for (let index = 1; index < pathSegments.length; index += 1) { + queue.push(pathSegments.slice(index).join("/")); + } + + for (const suffix of getQualifiedNamespaceSuffixes(normalized)) { + queue.push(suffix); + } + + for (const toggled of toggleShortVersionSeparators(normalized)) { + queue.push(toggled); + } + + const attachedFamilyVersion = insertAttachedFamilyVersionSeparator(normalized); + if (attachedFamilyVersion) { + queue.push(attachedFamilyVersion); + } + + for (const toggledSeriesVersion of toggleSeriesMinorVersionSeparators(normalized)) { + queue.push(toggledSeriesVersion); + } + + for (const expandedVersion of expandCompactMinorVersions(normalized)) { + queue.push(expandedVersion); + } + + for (const expandedSeriesVersion of expandCompactSeriesMinorVersions(normalized)) { + queue.push(expandedSeriesVersion); + } + + for (const wrapperCandidate of getWrapperCanonicalCandidates(normalized)) { + queue.push(wrapperCandidate); + } + + const strippedSyntheticPrefix = stripSyntheticPrefix(normalized); + if (strippedSyntheticPrefix) { + queue.push(strippedSyntheticPrefix); + } + + const strippedLatest = stripLatestSuffix(normalized); + if (strippedLatest) { + queue.push(strippedLatest); + } + + const strippedLegacyGlmTurbo = stripLegacyGlmTurboSuffix(normalized); + if (strippedLegacyGlmTurbo) { + queue.push(strippedLegacyGlmTurbo); + } + + const extractedFamily = extractUpstreamFamilyCandidate(normalized); + if (extractedFamily) { + queue.push(extractedFamily); + } + + const strippedProviderVersion = stripProviderVersionSuffix(normalized); + if (strippedProviderVersion) { + queue.push(strippedProviderVersion); + } + + const strippedDate = stripDateSuffix(normalized); + if (strippedDate) { + queue.push(strippedDate); + } + + const strippedMarker = stripTrailingMarker(normalized); + if (strippedMarker) { + queue.push(strippedMarker); + } + + const reorderedAnthropic = reorderAnthropicFamily(normalized); + if (reorderedAnthropic) { + queue.push(reorderedAnthropic); + } + } + + return [...candidates]; +} + +function getPreferredFallbackCanonicalCandidate(modelId: string, candidates: readonly string[]): string | undefined { + if (!/[/:.]/.test(modelId)) { + return undefined; + } + const cleanCandidates = candidates.filter(candidate => { + if (!candidate || candidate === modelId) { + return false; + } + if (candidate.includes("/") || candidate.includes(":")) { + return false; + } + if (candidate.toLowerCase() !== candidate) { + return false; + } + const extractedFamily = extractUpstreamFamilyCandidate(candidate); + return extractedFamily?.toLowerCase() === candidate; + }); + return selectBestOfficialCandidate(cleanCandidates); +} + +function resolveCanonicalIdForModel( + model: Model, + equivalence: CompiledEquivalenceConfig, + referenceData: CanonicalReferenceData, +): ResolvedCanonicalModel { + const selector = formatCanonicalVariantSelector(model); + const normalizedSelector = normalizeSelectorKey(selector); + + if (equivalence.overrides.has(normalizedSelector)) { + return { id: equivalence.overrides.get(normalizedSelector)!, source: "override" }; + } + + if (equivalence.exclude.has(normalizedSelector)) { + return { id: model.id, source: "fallback" }; + } + + const anthropicAlias = getAnthropicAliasOfficial(model.id, referenceData.officialIds); + if (anthropicAlias) { + return { id: anthropicAlias, source: anthropicAlias === model.id ? "bundled" : "heuristic" }; + } + + const claudeFamilyAlias = getClaudeFamilyAliasOfficial(model.id, referenceData.officialIds); + if (claudeFamilyAlias) { + return { id: claudeFamilyAlias, source: claudeFamilyAlias === model.id ? "bundled" : "heuristic" }; + } + + const heuristicCandidates = getHeuristicCanonicalCandidates(model.id); + const officialMatches = heuristicCandidates.filter(candidate => referenceData.officialIds.has(candidate)); + const preferredFallback = getPreferredFallbackCanonicalCandidate(model.id, heuristicCandidates); + const match = selectBestOfficialCandidate(officialMatches); + if (match) { + if ( + preferredFallback && + (match.includes("/") || match.includes(":")) && + compareCandidatePreference(preferredFallback, match) < 0 + ) { + return { id: preferredFallback, source: "heuristic" }; + } + return { id: match, source: match === model.id ? "bundled" : "heuristic" }; + } + + if (preferredFallback) { + return { id: preferredFallback, source: "heuristic" }; + } + + return { id: model.id, source: "fallback" }; +} + +function getCanonicalRecordName( + record: CanonicalModelRecord | undefined, + canonicalId: string, + variant: CanonicalModelVariant, + referenceData: CanonicalReferenceData, +): string { + if (record) { + return record.name; + } + return referenceData.references.get(canonicalId)?.name ?? variant.model.name ?? canonicalId; +} + +function compareCanonicalRecords(left: CanonicalModelRecord, right: CanonicalModelRecord): number { + return left.id.localeCompare(right.id); +} + +function compareCanonicalVariants(left: CanonicalModelVariant, right: CanonicalModelVariant): number { + const leftSelector = left.selector; + const rightSelector = right.selector; + return leftSelector.localeCompare(rightSelector); +} + +export function buildCanonicalModelIndex( + models: readonly Model[], + equivalence?: ModelEquivalenceConfig, +): CanonicalModelIndex { + const referenceData = createCanonicalReferenceData(); + const compiledEquivalence = compileEquivalenceConfig(equivalence); + const byId = new Map(); + const bySelector = new Map(); + + for (const model of models) { + const canonical = resolveCanonicalIdForModel(model, compiledEquivalence, referenceData); + const selector = formatCanonicalVariantSelector(model); + const variant: CanonicalModelVariant = { + canonicalId: canonical.id, + selector, + model, + source: canonical.source, + }; + const canonicalKey = normalizeCanonicalIdKey(canonical.id); + const existing = byId.get(canonicalKey); + const nextRecord: CanonicalModelRecord = existing ?? { + id: canonical.id, + name: getCanonicalRecordName(existing, canonical.id, variant, referenceData), + variants: [], + }; + nextRecord.name = getCanonicalRecordName(existing, canonical.id, variant, referenceData); + nextRecord.variants.push(variant); + byId.set(canonicalKey, nextRecord); + bySelector.set(normalizeSelectorKey(selector), canonical.id); + } + + const records = [...byId.values()].sort(compareCanonicalRecords); + for (const record of records) { + record.variants.sort(compareCanonicalVariants); + } + + return { records, byId, bySelector }; +} diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 19b776d5f..ff2bc24a6 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -31,8 +31,18 @@ import { type ConfigError, ConfigFile } from "../config"; import { parseModelString } from "../config/model-resolver"; import { isValidThemeColor, type ThemeColor } from "../modes/theme/theme"; import type { AuthStorage, OAuthCredential } from "../session/auth-storage"; +import { + buildCanonicalModelIndex, + type CanonicalModelIndex, + type CanonicalModelRecord, + type CanonicalModelVariant, + formatCanonicalVariantSelector, + type ModelEquivalenceConfig, +} from "./model-equivalence"; import { type Settings, settings } from "./settings"; +export type { CanonicalModelIndex, CanonicalModelRecord, CanonicalModelVariant, ModelEquivalenceConfig }; + export const kNoAuth = "N/A"; export function isAuthenticated(apiKey: string | undefined | null): apiKey is string { @@ -263,8 +273,14 @@ const ProviderConfigSchema = Type.Object({ modelOverrides: Type.Optional(Type.Record(Type.String(), ModelOverrideSchema)), }); +const EquivalenceConfigSchema = Type.Object({ + overrides: Type.Optional(Type.Record(Type.String(), Type.String({ minLength: 1 }))), + exclude: Type.Optional(Type.Array(Type.String({ minLength: 1 }))), +}); + const ModelsConfigSchema = Type.Object({ - providers: Type.Record(Type.String(), ProviderConfigSchema), + providers: Type.Optional(Type.Record(Type.String(), ProviderConfigSchema)), + equivalence: Type.Optional(EquivalenceConfigSchema), }); type ModelsConfig = Static; @@ -356,7 +372,7 @@ function validateProviderConfiguration( export const ModelsConfigFile = new ConfigFile("models", ModelsConfigSchema).withValidation( "models", config => { - for (const [providerName, providerConfig] of Object.entries(config.providers)) { + for (const [providerName, providerConfig] of Object.entries(config.providers ?? {})) { validateProviderConfiguration( providerName, { @@ -405,6 +421,11 @@ export interface ProviderDiscoveryState { error?: string; } +export interface CanonicalModelQueryOptions { + availableOnly?: boolean; + candidates?: readonly Model[]; +} + /** Result of loading custom models from models.json */ interface CustomModelsResult { models?: CustomModelOverlay[]; @@ -413,6 +434,7 @@ interface CustomModelsResult { keylessProviders?: Set; discoverableProviders?: DiscoveryProviderConfig[]; configuredProviders?: Set; + equivalence?: ModelEquivalenceConfig; error?: ConfigError; found: boolean; } @@ -739,17 +761,27 @@ function getDisabledProviderIdsFromSettings(): Set { } } +function getConfiguredProviderOrderFromSettings(): string[] { + try { + return settings.get("modelProviderOrder"); + } catch { + return []; + } +} + /** * Model registry - loads and manages models, resolves API keys via AuthStorage. */ export class ModelRegistry { #models: Model[] = []; + #canonicalIndex: CanonicalModelIndex = { records: [], byId: new Map(), bySelector: new Map() }; #customProviderApiKeys: Map = new Map(); #keylessProviders: Set = new Set(); #discoverableProviders: DiscoveryProviderConfig[] = []; #customModelOverlays: CustomModelOverlay[] = []; #providerOverrides: Map = new Map(); #modelOverrides: Map> = new Map(); + #equivalenceConfig: ModelEquivalenceConfig | undefined; #configError: ConfigError | undefined = undefined; #modelsConfigFile: ConfigFile; #registeredProviderSources: Set = new Set(); @@ -824,6 +856,7 @@ export class ModelRegistry { this.#discoverableProviders = []; this.#providerOverrides.clear(); this.#modelOverrides.clear(); + this.#equivalenceConfig = undefined; this.#configError = undefined; this.#providerDiscoveryStates.clear(); this.#loadModels(); @@ -845,6 +878,7 @@ export class ModelRegistry { keylessProviders = new Set(), discoverableProviders = [], configuredProviders = new Set(), + equivalence, error: configError, } = this.#loadCustomModels(); this.#configError = configError; @@ -853,6 +887,7 @@ export class ModelRegistry { this.#customModelOverlays = customModels; this.#providerOverrides = overrides; this.#modelOverrides = modelOverrides; + this.#equivalenceConfig = equivalence; this.#addImplicitDiscoverableProviders(configuredProviders); const builtInModels = this.#applyHardcodedModelPolicies(this.#loadBuiltInModels(overrides)); @@ -861,6 +896,7 @@ export class ModelRegistry { const combined = this.#mergeCustomModels(resolvedDefaults, this.#customModelOverlays); this.#models = this.#applyModelOverrides(combined, this.#modelOverrides); + this.#rebuildCanonicalIndex(); } /** Load built-in models, applying provider-level overrides only. @@ -1045,9 +1081,10 @@ export class ModelRegistry { const allModelOverrides = new Map>(); const keylessProviders = new Set(); const discoverableProviders: DiscoveryProviderConfig[] = []; - const configuredProviders = new Set(Object.keys(value.providers)); + const providerEntries = Object.entries(value.providers ?? {}); + const configuredProviders = new Set(Object.keys(value.providers ?? {})); - for (const [providerName, providerConfig] of Object.entries(value.providers)) { + for (const [providerName, providerConfig] of providerEntries) { // Always set overrides when baseUrl/headers/apiKey/compat are present if (providerConfig.baseUrl || providerConfig.headers || providerConfig.apiKey || providerConfig.compat) { overrides.set(providerName, { @@ -1097,6 +1134,7 @@ export class ModelRegistry { keylessProviders, discoverableProviders, configuredProviders, + equivalence: value.equivalence, found: true, }; } @@ -1147,6 +1185,7 @@ export class ModelRegistry { const resolved = this.#mergeResolvedModels(this.#models, discoveredModels); const combined = this.#mergeCustomModels(resolved, this.#customModelOverlays); this.#models = this.#applyModelOverrides(combined, this.#modelOverrides); + this.#rebuildCanonicalIndex(); } async #discoverProviderModels( @@ -1653,10 +1692,14 @@ export class ModelRegistry { }); } + #rebuildCanonicalIndex(): void { + this.#canonicalIndex = buildCanonicalModelIndex(this.#models, this.#equivalenceConfig); + } + #parseModels(config: ModelsConfig): CustomModelOverlay[] { const models: CustomModelOverlay[] = []; - for (const [providerName, providerConfig] of Object.entries(config.providers)) { + for (const [providerName, providerConfig] of Object.entries(config.providers ?? {})) { const modelDefs = providerConfig.models ?? []; if (modelDefs.length === 0) continue; // Override-only, no custom models if (providerConfig.apiKey) { @@ -1688,17 +1731,139 @@ export class ModelRegistry { return this.#models; } + #isModelAvailable(model: Model): boolean { + const disabledProviders = getDisabledProviderIdsFromSettings(); + return ( + !disabledProviders.has(model.provider) && + (this.#keylessProviders.has(model.provider) || this.authStorage.hasAuth(model.provider)) + ); + } + + #filterCanonicalVariants( + record: CanonicalModelRecord, + options: CanonicalModelQueryOptions | undefined, + ): CanonicalModelVariant[] { + const candidateKeys = options?.candidates + ? new Set(options.candidates.map(candidate => formatCanonicalVariantSelector(candidate))) + : undefined; + return record.variants.filter(variant => { + if (candidateKeys && !candidateKeys.has(variant.selector)) { + return false; + } + if (options?.availableOnly && !this.#isModelAvailable(variant.model)) { + return false; + } + return true; + }); + } + + #providerRank(models: readonly Model[]): Map { + const configuredProviders = getConfiguredProviderOrderFromSettings(); + const result = new Map(); + let nextRank = 0; + for (const provider of configuredProviders) { + const normalized = provider.trim().toLowerCase(); + if (!normalized || result.has(normalized)) { + continue; + } + result.set(normalized, nextRank); + nextRank += 1; + } + for (const model of models) { + const normalized = model.provider.toLowerCase(); + if (result.has(normalized)) { + continue; + } + result.set(normalized, nextRank); + nextRank += 1; + } + return result; + } + + #resolveCanonicalVariant( + variants: readonly CanonicalModelVariant[], + allCandidates: readonly Model[], + ): CanonicalModelVariant | undefined { + if (variants.length === 0) { + return undefined; + } + const providerRank = this.#providerRank(allCandidates); + const modelOrder = new Map(); + for (let index = 0; index < allCandidates.length; index += 1) { + modelOrder.set(formatCanonicalVariantSelector(allCandidates[index]!), index); + } + const sourceRank: Record = { + override: 1, + bundled: 1, + heuristic: 2, + fallback: 3, + }; + return [...variants].sort((left, right) => { + const leftProviderRank = providerRank.get(left.model.provider.toLowerCase()) ?? Number.MAX_SAFE_INTEGER; + const rightProviderRank = providerRank.get(right.model.provider.toLowerCase()) ?? Number.MAX_SAFE_INTEGER; + if (leftProviderRank !== rightProviderRank) { + return leftProviderRank - rightProviderRank; + } + const leftExact = left.model.id === left.canonicalId ? 0 : 1; + const rightExact = right.model.id === right.canonicalId ? 0 : 1; + if (leftExact !== rightExact) { + return leftExact - rightExact; + } + if (sourceRank[left.source] !== sourceRank[right.source]) { + return sourceRank[left.source] - sourceRank[right.source]; + } + if (left.model.id.length !== right.model.id.length) { + return left.model.id.length - right.model.id.length; + } + const leftOrder = modelOrder.get(left.selector) ?? Number.MAX_SAFE_INTEGER; + const rightOrder = modelOrder.get(right.selector) ?? Number.MAX_SAFE_INTEGER; + return leftOrder - rightOrder; + })[0]; + } + + getCanonicalModels(options?: CanonicalModelQueryOptions): CanonicalModelRecord[] { + const records: CanonicalModelRecord[] = []; + for (const record of this.#canonicalIndex.records) { + const variants = this.#filterCanonicalVariants(record, options); + if (variants.length === 0) { + continue; + } + records.push({ + id: record.id, + name: record.name, + variants, + }); + } + return records; + } + + getCanonicalVariants(canonicalId: string, options?: CanonicalModelQueryOptions): CanonicalModelVariant[] { + const record = this.#canonicalIndex.byId.get(canonicalId.trim().toLowerCase()); + if (!record) { + return []; + } + return this.#filterCanonicalVariants(record, options); + } + + resolveCanonicalModel(canonicalId: string, options?: CanonicalModelQueryOptions): Model | undefined { + const variants = this.getCanonicalVariants(canonicalId, options); + if (variants.length === 0) { + return undefined; + } + const candidates = options?.candidates ?? (options?.availableOnly ? this.getAvailable() : this.getAll()); + return this.#resolveCanonicalVariant(variants, candidates)?.model; + } + + getCanonicalId(model: Model): string | undefined { + return this.#canonicalIndex.bySelector.get(formatCanonicalVariantSelector(model).toLowerCase()); + } + /** * Get only models that have auth configured. * This is a fast check that doesn't refresh OAuth tokens. */ getAvailable(): Model[] { - const disabledProviders = getDisabledProviderIdsFromSettings(); - return this.#models.filter( - m => - !disabledProviders.has(m.provider) && - (this.#keylessProviders.has(m.provider) || this.authStorage.hasAuth(m.provider)), - ); + return this.#models.filter(model => this.#isModelAvailable(model)); } getDiscoverableProviders(): string[] { @@ -1853,11 +2018,13 @@ export class ModelRegistry { const credential = this.authStorage.getOAuthCredential(providerName); if (credential) { this.#models = config.oauth.modifyModels(nextModels, credential); + this.#rebuildCanonicalIndex(); return; } } this.#models = nextModels; + this.#rebuildCanonicalIndex(); return; } @@ -1870,6 +2037,7 @@ export class ModelRegistry { headers: config.headers ? { ...m.headers, ...config.headers } : m.headers, }; }); + this.#rebuildCanonicalIndex(); } } diff --git a/packages/coding-agent/src/config/model-resolver.ts b/packages/coding-agent/src/config/model-resolver.ts index 522ae2d03..ee4545430 100644 --- a/packages/coding-agent/src/config/model-resolver.ts +++ b/packages/coding-agent/src/config/model-resolver.ts @@ -58,6 +58,10 @@ export function formatModelString(model: Model): string { return `${model.provider}/${model.id}`; } +export function formatModelSelectorValue(selector: string, thinkingLevel: ThinkingLevel | undefined): string { + return thinkingLevel && thinkingLevel !== ThinkingLevel.Inherit ? `${selector}:${thinkingLevel}` : selector; +} + export interface ModelMatchPreferences { /** Most-recently-used model keys (provider/modelId) to prefer when ambiguous. */ usageOrder?: string[]; @@ -65,6 +69,14 @@ export interface ModelMatchPreferences { deprioritizeProviders?: string[]; } +export type CanonicalModelRegistry = Partial< + Pick +>; +export type ModelLookupRegistry = Pick & Partial; +type CliModelRegistry = Pick & Partial; +type InitialModelRegistry = Pick; +type RestorableModelRegistry = Pick; + interface ModelPreferenceContext { modelUsageRank: Map; providerUsageRank: Map; @@ -142,9 +154,8 @@ function isAlias(id: string): boolean { } /** - * Find an exact model reference match. - * Supports either a bare model id or a canonical provider/modelId reference. - * When matching by bare id, ambiguous matches across providers are rejected. + * Find an exact explicit provider/model match. + * Bare model ids are handled separately so canonical ids can coalesce variants. */ export function findExactModelReferenceMatch( modelReference: string, @@ -155,18 +166,6 @@ export function findExactModelReferenceMatch( return undefined; } - const normalizedReference = trimmedReference.toLowerCase(); - - const canonicalMatches = availableModels.filter( - model => `${model.provider}/${model.id}`.toLowerCase() === normalizedReference, - ); - if (canonicalMatches.length === 1) { - return canonicalMatches[0]; - } - if (canonicalMatches.length > 1) { - return undefined; - } - const slashIndex = trimmedReference.indexOf("/"); if (slashIndex !== -1) { const provider = trimmedReference.substring(0, slashIndex).trim(); @@ -185,9 +184,25 @@ export function findExactModelReferenceMatch( } } } + return undefined; +} - const idMatches = availableModels.filter(model => model.id.toLowerCase() === normalizedReference); - return idMatches.length === 1 ? idMatches[0] : undefined; +function findExactCanonicalModelMatch( + modelReference: string, + availableModels: Model[], + modelRegistry: CanonicalModelRegistry | undefined, +): Model | undefined { + if (!modelRegistry) { + return undefined; + } + const trimmedReference = modelReference.trim(); + if (!trimmedReference || trimmedReference.includes("/")) { + return undefined; + } + return modelRegistry.resolveCanonicalModel?.(trimmedReference, { + availableOnly: false, + candidates: availableModels, + }); } /** @@ -198,13 +213,20 @@ function tryMatchModel( modelPattern: string, availableModels: Model[], context: ModelPreferenceContext, + options?: { modelRegistry?: CanonicalModelRegistry }, ): Model | undefined { - // Try exact reference match first (handles provider/modelId and bare id with ambiguity rejection) + // Explicit provider/model selectors always bypass canonical coalescing. const exactRefMatch = findExactModelReferenceMatch(modelPattern, availableModels); if (exactRefMatch) { return exactRefMatch; } + // Exact canonical ids coalesce provider variants before bare-id matching. + const exactCanonicalMatch = findExactCanonicalModelMatch(modelPattern, availableModels, options?.modelRegistry); + if (exactCanonicalMatch) { + return exactCanonicalMatch; + } + // Check for provider/modelId format — fuzzy match within provider const slashIndex = modelPattern.indexOf("/"); if (slashIndex !== -1) { @@ -300,10 +322,10 @@ function parseModelPatternWithContext( pattern: string, availableModels: Model[], context: ModelPreferenceContext, - options?: { allowInvalidThinkingSelectorFallback?: boolean }, + options?: { allowInvalidThinkingSelectorFallback?: boolean; modelRegistry?: CanonicalModelRegistry }, ): ParsedModelResult { // Try exact match first - const exactMatch = tryMatchModel(pattern, availableModels, context); + const exactMatch = tryMatchModel(pattern, availableModels, context, options); if (exactMatch) { return { model: exactMatch, thinkingLevel: undefined, warning: undefined, explicitThinkingLevel: false }; } @@ -357,7 +379,7 @@ export function parseModelPattern( pattern: string, availableModels: Model[], preferences?: ModelMatchPreferences, - options?: { allowInvalidThinkingSelectorFallback?: boolean }, + options?: { allowInvalidThinkingSelectorFallback?: boolean; modelRegistry?: CanonicalModelRegistry }, ): ParsedModelResult { const context = buildPreferenceContext(availableModels, preferences); return parseModelPatternWithContext(pattern, availableModels, context, options); @@ -469,7 +491,7 @@ export interface ResolvedModelRoleValue { export function resolveModelRoleValue( roleValue: string | undefined, availableModels: Model[], - options?: { settings?: Settings; matchPreferences?: ModelMatchPreferences }, + options?: { settings?: Settings; matchPreferences?: ModelMatchPreferences; modelRegistry?: CanonicalModelRegistry }, ): ResolvedModelRoleValue { if (!roleValue) { return { model: undefined, thinkingLevel: undefined, explicitThinkingLevel: false, warning: undefined }; @@ -490,7 +512,9 @@ export function resolveModelRoleValue( let warning: string | undefined; for (const effectivePattern of effectivePatterns) { - const resolved = parseModelPattern(effectivePattern, availableModels, options?.matchPreferences); + const resolved = parseModelPattern(effectivePattern, availableModels, options?.matchPreferences, { + modelRegistry: options?.modelRegistry, + }); if (resolved.model) { return { model: resolved.model, @@ -543,13 +567,14 @@ export function resolveModelFromString( value: string, available: Model[], matchPreferences?: ModelMatchPreferences, + modelRegistry?: CanonicalModelRegistry, ): Model | undefined { const parsed = parseModelString(value); if (parsed) { const exact = available.find(model => model.provider === parsed.provider && model.id === parsed.id); if (exact) return exact; } - return parseModelPattern(value, available, matchPreferences).model; + return parseModelPattern(value, available, matchPreferences, { modelRegistry }).model; } /** @@ -560,13 +585,19 @@ export function resolveModelFromSettings(options: { availableModels: Model[]; matchPreferences?: ModelMatchPreferences; roleOrder?: readonly ModelRole[]; + modelRegistry?: CanonicalModelRegistry; }): Model | undefined { - const { settings, availableModels, matchPreferences, roleOrder } = options; + const { settings, availableModels, matchPreferences, roleOrder, modelRegistry } = options; const roles = roleOrder ?? MODEL_ROLE_IDS; for (const role of roles) { const configured = settings.getModelRole(role); if (!configured) continue; - const resolved = resolveModelFromString(expandRoleAlias(configured, settings), availableModels, matchPreferences); + const resolved = resolveModelFromString( + expandRoleAlias(configured, settings), + availableModels, + matchPreferences, + modelRegistry, + ); if (resolved) return resolved; } return availableModels[0]; @@ -577,7 +608,7 @@ export function resolveModelFromSettings(options: { */ export function resolveModelOverride( modelPatterns: string[], - modelRegistry: ModelRegistry, + modelRegistry: ModelLookupRegistry, settings?: Settings, ): { model?: Model; thinkingLevel?: ThinkingLevel; explicitThinkingLevel: boolean } { if (modelPatterns.length === 0) return { explicitThinkingLevel: false }; @@ -587,6 +618,7 @@ export function resolveModelOverride( const { model, thinkingLevel, explicitThinkingLevel } = resolveModelRoleValue(pattern, availableModels, { settings, matchPreferences, + modelRegistry, }); if (model) { return { model, thinkingLevel, explicitThinkingLevel }; @@ -602,12 +634,14 @@ export function resolveRoleSelection( roles: readonly string[], settings: Settings, availableModels: Model[], + modelRegistry?: CanonicalModelRegistry, ): { model: Model; thinkingLevel?: ThinkingLevel } | undefined { const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() }; for (const role of roles) { const resolved = resolveModelRoleValue(settings.getModelRole(role), availableModels, { settings, matchPreferences, + modelRegistry, }); if (resolved.model) { return { model: resolved.model, thinkingLevel: resolved.thinkingLevel }; @@ -616,6 +650,36 @@ export function resolveRoleSelection( return undefined; } +function resolveExactCanonicalScopePattern( + pattern: string, + modelRegistry: Pick, + availableModels: Model[], +): { models: Model[]; thinkingLevel?: ThinkingLevel; explicitThinkingLevel: boolean } | undefined { + const lastColonIndex = pattern.lastIndexOf(":"); + let canonicalId = pattern; + let thinkingLevel: ThinkingLevel | undefined; + let explicitThinkingLevel = false; + + if (lastColonIndex !== -1) { + const suffix = pattern.substring(lastColonIndex + 1); + const parsedThinkingLevel = parseThinkingLevel(suffix); + if (parsedThinkingLevel) { + canonicalId = pattern.substring(0, lastColonIndex); + thinkingLevel = parsedThinkingLevel; + explicitThinkingLevel = true; + } + } + + const variants = modelRegistry + .getCanonicalVariants(canonicalId, { availableOnly: true, candidates: availableModels }) + .map(variant => variant.model); + if (variants.length === 0) { + return undefined; + } + + return { models: variants, thinkingLevel, explicitThinkingLevel }; +} + /** * Resolve model patterns to actual Model objects with optional thinking levels * Format: "pattern:level" where :level is optional @@ -629,7 +693,7 @@ export function resolveRoleSelection( */ export async function resolveModelScope( patterns: string[], - modelRegistry: ModelRegistry, + modelRegistry: Pick, preferences?: ModelMatchPreferences, ): Promise { const availableModels = modelRegistry.getAvailable(); @@ -682,10 +746,28 @@ export async function resolveModelScope( continue; } + const exactCanonical = resolveExactCanonicalScopePattern(pattern, modelRegistry, availableModels); + if (exactCanonical) { + for (const model of exactCanonical.models) { + if (!scopedModels.find(sm => modelsAreEqual(sm.model, model))) { + scopedModels.push({ + model, + thinkingLevel: exactCanonical.explicitThinkingLevel + ? (resolveThinkingLevelForModel(model, exactCanonical.thinkingLevel) ?? + exactCanonical.thinkingLevel) + : exactCanonical.thinkingLevel, + explicitThinkingLevel: exactCanonical.explicitThinkingLevel, + }); + } + } + continue; + } + const { model, thinkingLevel, warning, explicitThinkingLevel } = parseModelPatternWithContext( pattern, availableModels, context, + { modelRegistry }, ); if (warning) { @@ -714,6 +796,7 @@ export async function resolveModelScope( export interface ResolveCliModelResult { model: Model | undefined; + selector?: string; thinkingLevel?: ThinkingLevel; warning: string | undefined; error: string | undefined; @@ -725,19 +808,20 @@ export interface ResolveCliModelResult { export function resolveCliModel(options: { cliProvider?: string; cliModel?: string; - modelRegistry: ModelRegistry; + modelRegistry: CliModelRegistry; preferences?: ModelMatchPreferences; }): ResolveCliModelResult { const { cliProvider, cliModel, modelRegistry, preferences } = options; if (!cliModel) { - return { model: undefined, warning: undefined, error: undefined }; + return { model: undefined, selector: undefined, warning: undefined, error: undefined }; } const availableModels = modelRegistry.getAll(); if (availableModels.length === 0) { return { model: undefined, + selector: undefined, warning: undefined, error: "No models available. Check your installation or add models to models.json.", }; @@ -752,13 +836,15 @@ export function resolveCliModel(options: { if (cliProvider && !provider) { return { model: undefined, + selector: undefined, warning: undefined, error: `Unknown provider "${cliProvider}". Use --list-models to see available providers/models.`, }; } + const trimmedModel = cliModel.trim(); if (!provider) { - const lower = cliModel.toLowerCase(); + const lower = trimmedModel.toLowerCase(); // When input has provider/id format (e.g. "zai/glm-5"), prefer decomposed // provider+id match over flat id match. Without this, a model with id // "zai/glm-5" on provider "vercel-ai-gateway" wins over provider "zai" @@ -772,17 +858,35 @@ export function resolveCliModel(options: { model => model.provider.toLowerCase() === prefix && model.id.toLowerCase() === suffix, ); } + if (!exact && !trimmedModel.includes(":")) { + const canonicalMatch = modelRegistry.resolveCanonicalModel?.(trimmedModel, { availableOnly: false }); + if (canonicalMatch) { + return { + model: canonicalMatch, + selector: modelRegistry.getCanonicalId?.(canonicalMatch) ?? trimmedModel, + warning: undefined, + thinkingLevel: undefined, + error: undefined, + }; + } + } if (!exact) { exact = availableModels.find( model => model.id.toLowerCase() === lower || `${model.provider}/${model.id}`.toLowerCase() === lower, ); } if (exact) { - return { model: exact, warning: undefined, thinkingLevel: undefined, error: undefined }; + return { + model: exact, + selector: formatModelString(exact), + warning: undefined, + thinkingLevel: undefined, + error: undefined, + }; } } - let pattern = cliModel; + let pattern = trimmedModel; if (!provider) { const slashIndex = cliModel.indexOf("/"); @@ -804,19 +908,42 @@ export function resolveCliModel(options: { const candidates = provider ? availableModels.filter(model => model.provider === provider) : availableModels; const { model, thinkingLevel, warning } = parseModelPattern(pattern, candidates, preferences, { allowInvalidThinkingSelectorFallback: false, + modelRegistry, }); if (!model) { const display = provider ? `${provider}/${pattern}` : cliModel; return { model: undefined, + selector: undefined, thinkingLevel: undefined, warning, error: `Model "${display}" not found. Use --list-models to see available models.`, }; } - return { model, thinkingLevel, warning, error: undefined }; + let selector = provider ? formatModelString(model) : undefined; + if (!provider) { + const lastColonIndex = pattern.lastIndexOf(":"); + const canonicalCandidate = + lastColonIndex !== -1 && parseThinkingLevel(pattern.substring(lastColonIndex + 1)) + ? pattern.substring(0, lastColonIndex) + : pattern; + if (!canonicalCandidate.includes("/")) { + const canonicalResolved = modelRegistry.resolveCanonicalModel?.(canonicalCandidate, { availableOnly: false }); + if (canonicalResolved && canonicalResolved.provider === model.provider && canonicalResolved.id === model.id) { + selector = modelRegistry.getCanonicalId?.(canonicalResolved) ?? canonicalCandidate; + } + } + } + + return { + model, + selector, + thinkingLevel, + warning, + error: undefined, + }; } export interface InitialModelResult { @@ -841,7 +968,7 @@ export async function findInitialModel(options: { defaultProvider?: string; defaultModelId?: string; defaultThinkingSelector?: Effort; - modelRegistry: ModelRegistry; + modelRegistry: InitialModelRegistry; }): Promise { const { cliProvider, @@ -923,7 +1050,7 @@ export async function restoreModelFromSession( savedModelId: string, currentModel: Model | undefined, shouldPrintMessages: boolean, - modelRegistry: ModelRegistry, + modelRegistry: RestorableModelRegistry, ): Promise<{ model: Model | undefined; fallbackMessage: string | undefined }> { const restoredModel = modelRegistry.find(savedProvider, savedModelId); @@ -998,7 +1125,7 @@ export async function restoreModelFromSession( * @returns The best available smol model, or undefined if none found */ export async function findSmolModel( - modelRegistry: ModelRegistry, + modelRegistry: ModelLookupRegistry, savedModel?: string, ): Promise | undefined> { const availableModels = modelRegistry.getAvailable(); @@ -1006,11 +1133,8 @@ export async function findSmolModel( // 1. Try saved model from settings if (savedModel) { - const parsed = parseModelString(savedModel); - if (parsed) { - const match = availableModels.find(m => m.provider === parsed.provider && m.id === parsed.id); - if (match) return match; - } + const match = resolveModelFromString(savedModel, availableModels, undefined, modelRegistry); + if (match) return match; } // 2. Try priority chain @@ -1020,7 +1144,7 @@ export async function findSmolModel( if (providerMatch) return providerMatch; // Try exact match first - const exactMatch = availableModels.find(m => m.id.toLowerCase() === pattern); + const exactMatch = parseModelPattern(pattern, availableModels, undefined, { modelRegistry }).model; if (exactMatch) return exactMatch; // Try fuzzy match (substring) @@ -1041,7 +1165,7 @@ export async function findSmolModel( * @returns The best available slow model, or undefined if none found */ export async function findSlowModel( - modelRegistry: ModelRegistry, + modelRegistry: ModelLookupRegistry, savedModel?: string, ): Promise | undefined> { const availableModels = modelRegistry.getAvailable(); @@ -1049,17 +1173,14 @@ export async function findSlowModel( // 1. Try saved model from settings if (savedModel) { - const parsed = parseModelString(savedModel); - if (parsed) { - const match = availableModels.find(m => m.provider === parsed.provider && m.id === parsed.id); - if (match) return match; - } + const match = resolveModelFromString(savedModel, availableModels, undefined, modelRegistry); + if (match) return match; } // 2. Try priority chain for (const pattern of MODEL_PRIO.slow) { // Try exact match first - const exactMatch = availableModels.find(m => m.id.toLowerCase() === pattern.toLowerCase()); + const exactMatch = parseModelPattern(pattern, availableModels, undefined, { modelRegistry }).model; if (exactMatch) return exactMatch; // Try fuzzy match (substring) diff --git a/packages/coding-agent/src/config/settings-schema.ts b/packages/coding-agent/src/config/settings-schema.ts index 99f294286..a92ec344a 100644 --- a/packages/coding-agent/src/config/settings-schema.ts +++ b/packages/coding-agent/src/config/settings-schema.ts @@ -229,6 +229,8 @@ export const SETTINGS_SCHEMA = { modelTags: { type: "record", default: EMPTY_MODEL_TAGS_RECORD }, + modelProviderOrder: { type: "array", default: EMPTY_STRING_ARRAY }, + cycleOrder: { type: "array", default: DEFAULT_CYCLE_ORDER }, // ──────────────────────────────────────────────────────────────────────── diff --git a/packages/coding-agent/src/main.ts b/packages/coding-agent/src/main.ts index 593c278b1..13803dbda 100644 --- a/packages/coding-agent/src/main.ts +++ b/packages/coding-agent/src/main.ts @@ -453,7 +453,9 @@ async function buildSessionOptions( } } else if (resolved.model) { options.model = resolved.model; - settings.overrideModelRoles({ default: `${resolved.model.provider}/${resolved.model.id}` }); + settings.overrideModelRoles({ + default: resolved.selector ?? `${resolved.model.provider}/${resolved.model.id}`, + }); if (!parsed.thinking && resolved.thinkingLevel) { options.thinkingLevel = resolved.thinkingLevel; } @@ -467,6 +469,7 @@ async function buildSessionOptions( { settings, matchPreferences: modelMatchPreferences, + modelRegistry, }, ); const rememberedResolvedModel = rememberedSpec.model; diff --git a/packages/coding-agent/src/memories/index.ts b/packages/coding-agent/src/memories/index.ts index 47cd33920..075a63f33 100644 --- a/packages/coding-agent/src/memories/index.ts +++ b/packages/coding-agent/src/memories/index.ts @@ -6,7 +6,7 @@ import type { AgentMessage } from "@oh-my-pi/pi-agent-core"; import { completeSimple, Effort, type Model } from "@oh-my-pi/pi-ai"; import { getAgentDbPath, getMemoriesDir, logger, parseJsonlLenient, prompt } from "@oh-my-pi/pi-utils"; import type { ModelRegistry } from "../config/model-registry"; -import { parseModelString } from "../config/model-resolver"; +import { resolveModelRoleValue } from "../config/model-resolver"; import type { Settings } from "../config/settings"; import consolidationTemplate from "../prompts/memories/consolidation.md" with { type: "text" }; import readPathTemplate from "../prompts/memories/read-path.md" with { type: "text" }; @@ -1055,11 +1055,12 @@ async function resolveMemoryModel(options: { const { modelRegistry, session, fallbackRole } = options; const requestedModel = session.settings.getModelRole(fallbackRole) || session.settings.getModelRole("default"); if (requestedModel) { - const parsed = parseModelString(requestedModel); - if (parsed) { - const found = modelRegistry.find(parsed.provider, parsed.id); - if (found) return found; - } + const resolved = resolveModelRoleValue(requestedModel, modelRegistry.getAll(), { + settings: session.settings, + matchPreferences: { usageOrder: session.settings.getStorage()?.getModelUsageOrder() }, + modelRegistry, + }); + if (resolved.model) return resolved.model; } return session.model ?? modelRegistry.getAll()[0]; } diff --git a/packages/coding-agent/src/modes/components/model-selector.ts b/packages/coding-agent/src/modes/components/model-selector.ts index 8fca03a3b..4a8c1b2a5 100644 --- a/packages/coding-agent/src/modes/components/model-selector.ts +++ b/packages/coding-agent/src/modes/components/model-selector.ts @@ -28,10 +28,38 @@ function makeInvertedBadge(label: string, color: ThemeColor): string { return `${bgAnsi}\x1b[30m ${label} \x1b[39m\x1b[49m`; } +function normalizeSearchText(value: string): string { + return value + .toLowerCase() + .replace(/[^a-z0-9]+/g, " ") + .trim(); +} + +function compactSearchText(value: string): string { + return value.toLowerCase().replace(/[^a-z0-9]+/g, ""); +} + +function getAlphaSearchTokens(query: string): string[] { + return [...normalizeSearchText(query).matchAll(/[a-z]+/g)].map(match => match[0]).filter(token => token.length > 0); +} + interface ModelItem { + kind: "provider"; provider: string; id: string; model: Model; + selector: string; +} + +interface CanonicalModelItem { + kind: "canonical"; + id: string; + model: Model; + selector: string; + variantCount: number; + searchText: string; + normalizedSearchText: string; + compactSearchText: string; } interface ScopedModelItem { @@ -44,7 +72,7 @@ interface RoleAssignment { thinkingLevel: ThinkingLevel; } -type RoleSelectCallback = (model: Model, role: string | null, thinkingLevel?: ThinkingLevel) => void; +type RoleSelectCallback = (model: Model, role: string | null, thinkingLevel?: ThinkingLevel, selector?: string) => void; type CancelCallback = () => void; interface MenuRoleAction { label: string; @@ -52,6 +80,7 @@ interface MenuRoleAction { } const ALL_TAB = "ALL"; +const CANONICAL_TAB = "CANONICAL"; /** * Component that renders a model selector with provider tabs and context menu. @@ -68,6 +97,8 @@ export class ModelSelectorComponent extends Container { #menuContainer: Container; #allModels: ModelItem[] = []; #filteredModels: ModelItem[] = []; + #canonicalModels: CanonicalModelItem[] = []; + #filteredCanonicalModels: CanonicalModelItem[] = []; #selectedIndex: number = 0; #roles = {} as Record; #settings = null as unknown as Settings; @@ -97,7 +128,7 @@ export class ModelSelectorComponent extends Container { settings: Settings, modelRegistry: ModelRegistry, scopedModels: ReadonlyArray, - onSelect: (model: Model, role: string | null, thinkingLevel?: ThinkingLevel) => void, + onSelect: (model: Model, role: string | null, thinkingLevel?: ThinkingLevel, selector?: string) => void, onCancel: () => void, options?: { temporaryOnly?: boolean; initialSearchInput?: string }, ) { @@ -202,6 +233,7 @@ export class ModelSelectorComponent extends Container { const resolved = resolveModelRoleValue(roleValue, allModels, { settings: this.#settings, matchPreferences, + modelRegistry: this.#modelRegistry, }); if (resolved.model) { this.#roles[role] = { @@ -237,8 +269,8 @@ export class ModelSelectorComponent extends Container { const latestRe = /-latest$/; models.sort((a, b) => { - const aKey = `${a.provider}/${a.id}`; - const bKey = `${b.provider}/${b.id}`; + const aKey = a.selector; + const bKey = b.selector; const aRank = modelRank(a); const bRank = modelRank(b); @@ -289,15 +321,50 @@ export class ModelSelectorComponent extends Container { }); } + #sortCanonicalModels(models: CanonicalModelItem[]): void { + const mruOrder = this.#settings.getStorage()?.getModelUsageOrder() ?? []; + const mruIndex = new Map(mruOrder.map((key, i) => [key, i])); + + const modelRank = (model: CanonicalModelItem) => { + let i = 0; + while (i < MODEL_ROLE_IDS.length) { + const role = MODEL_ROLE_IDS[i]; + const assigned = this.#roles[role]; + if (assigned && modelsAreEqual(assigned.model, model.model)) { + break; + } + i++; + } + return i; + }; + + models.sort((a, b) => { + const aRank = modelRank(a); + const bRank = modelRank(b); + if (aRank !== bRank) return aRank - bRank; + + const aMru = mruIndex.get(`${a.model.provider}/${a.model.id}`) ?? Number.MAX_SAFE_INTEGER; + const bMru = mruIndex.get(`${b.model.provider}/${b.model.id}`) ?? Number.MAX_SAFE_INTEGER; + if (aMru !== bMru) return aMru - bMru; + + const providerCmp = a.model.provider.localeCompare(b.model.provider); + if (providerCmp !== 0) return providerCmp; + + return a.id.localeCompare(b.id); + }); + } + async #loadModels(): Promise { let models: ModelItem[]; // Use scoped models if provided via --models flag if (this.#scopedModels.length > 0) { models = this.#scopedModels.map(scoped => ({ + kind: "provider", provider: scoped.model.provider, id: scoped.model.id, model: scoped.model, + selector: `${scoped.model.provider}/${scoped.model.id}`, })); } else { // Reload config and cached discovery state without blocking on live provider refresh @@ -315,22 +382,61 @@ export class ModelSelectorComponent extends Container { try { const availableModels = this.#modelRegistry.getAvailable(); models = availableModels.map((model: Model) => ({ + kind: "provider", provider: model.provider, id: model.id, model, + selector: `${model.provider}/${model.id}`, })); } catch (error) { this.#allModels = []; this.#filteredModels = []; + this.#canonicalModels = []; + this.#filteredCanonicalModels = []; this.#errorMessage = error instanceof Error ? error.message : String(error); return; } } + const canonicalRecords = this.#modelRegistry.getCanonicalModels({ + availableOnly: this.#scopedModels.length === 0, + candidates: models.map(item => item.model), + }); + const canonicalModels = canonicalRecords + .map(record => { + const selectedModel = this.#modelRegistry.resolveCanonicalModel(record.id, { + availableOnly: this.#scopedModels.length === 0, + candidates: models.map(item => item.model), + }); + if (!selectedModel) return undefined; + const searchText = [ + record.id, + record.name, + selectedModel.provider, + selectedModel.id, + selectedModel.name, + ...record.variants.flatMap(variant => [variant.selector, variant.model.name]), + ].join(" "); + return { + kind: "canonical" as const, + id: record.id, + model: selectedModel, + selector: record.id, + variantCount: record.variants.length, + searchText, + normalizedSearchText: normalizeSearchText(searchText), + compactSearchText: compactSearchText(searchText), + }; + }) + .filter((item): item is CanonicalModelItem => item !== undefined); + this.#sortModels(models); + this.#sortCanonicalModels(canonicalModels); this.#allModels = models; this.#filteredModels = models; + this.#canonicalModels = canonicalModels; + this.#filteredCanonicalModels = canonicalModels; this.#selectedIndex = Math.min(this.#selectedIndex, Math.max(0, models.length - 1)); } @@ -343,12 +449,12 @@ export class ModelSelectorComponent extends Container { providerSet.add(provider.toUpperCase()); } const sortedProviders = Array.from(providerSet).sort(); - this.#providers = [ALL_TAB, ...sortedProviders]; + this.#providers = [ALL_TAB, CANONICAL_TAB, ...sortedProviders]; } async #refreshSelectedProvider(): Promise { const activeProvider = this.#getActiveProvider(); - if (this.#scopedModels.length > 0 || activeProvider === ALL_TAB) { + if (this.#scopedModels.length > 0 || activeProvider === ALL_TAB || activeProvider === CANONICAL_TAB) { return; } await this.#modelRegistry.refreshProvider(activeProvider.toLowerCase()); @@ -382,19 +488,25 @@ export class ModelSelectorComponent extends Container { return this.#providers[this.#activeTabIndex] ?? ALL_TAB; } + #isCanonicalTab(): boolean { + return this.#getActiveProvider() === CANONICAL_TAB; + } + #filterModels(query: string): void { const activeProvider = this.#getActiveProvider(); + const isCanonicalTab = activeProvider === CANONICAL_TAB; - // Start with all models or filter by provider + // Start with all models or filter by provider/canonical view let baseModels = this.#allModels; - if (activeProvider !== ALL_TAB) { + const baseCanonicalModels = this.#canonicalModels; + if (!isCanonicalTab && activeProvider !== ALL_TAB) { baseModels = this.#allModels.filter(m => m.provider.toUpperCase() === activeProvider); } // Apply fuzzy filter if query is present if (query.trim()) { - // If user is searching, auto-switch to ALL tab to show global results - if (activeProvider !== ALL_TAB) { + // If user is searching from a provider tab, auto-switch to ALL to show global provider results. + if (activeProvider !== ALL_TAB && !isCanonicalTab) { this.#activeTabIndex = 0; if (this.#tabBar && this.#tabBar.getActiveIndex() !== 0) { this.#tabBar.setActiveIndex(0); @@ -403,14 +515,41 @@ export class ModelSelectorComponent extends Container { this.#updateTabBar(); baseModels = this.#allModels; } - const fuzzyMatches = fuzzyFilter(baseModels, query, ({ id, provider }) => `${id} ${provider}`); - this.#sortModels(fuzzyMatches); - this.#filteredModels = fuzzyMatches; + + if (isCanonicalTab) { + const alphaTokens = getAlphaSearchTokens(query); + const alphaFiltered = + alphaTokens.length === 0 + ? baseCanonicalModels + : baseCanonicalModels.filter(item => + alphaTokens.every(token => item.normalizedSearchText.includes(token)), + ); + const compactQuery = compactSearchText(query); + const substringFiltered = + compactQuery.length === 0 + ? alphaFiltered + : alphaFiltered.filter(item => item.compactSearchText.includes(compactQuery)); + const fuzzySource = + substringFiltered.length > 0 + ? substringFiltered + : alphaFiltered.length > 0 + ? alphaFiltered + : baseCanonicalModels; + const fuzzyMatches = fuzzyFilter(fuzzySource, query, ({ searchText }) => searchText); + this.#sortCanonicalModels(fuzzyMatches); + this.#filteredCanonicalModels = fuzzyMatches; + } else { + const fuzzyMatches = fuzzyFilter(baseModels, query, ({ id, provider }) => `${id} ${provider}`); + this.#sortModels(fuzzyMatches); + this.#filteredModels = fuzzyMatches; + } } else { this.#filteredModels = baseModels; + this.#filteredCanonicalModels = baseCanonicalModels; } - this.#selectedIndex = Math.min(this.#selectedIndex, Math.max(0, this.#filteredModels.length - 1)); + const visibleCount = isCanonicalTab ? this.#filteredCanonicalModels.length : this.#filteredModels.length; + this.#selectedIndex = Math.min(this.#selectedIndex, Math.max(0, visibleCount - 1)); this.#updateList(); } @@ -433,7 +572,7 @@ export class ModelSelectorComponent extends Container { #getProviderEmptyStateMessage(): string | undefined { const activeProvider = this.#getActiveProvider(); - if (activeProvider === ALL_TAB || this.#searchInput.getValue().trim()) { + if (activeProvider === ALL_TAB || activeProvider === CANONICAL_TAB || this.#searchInput.getValue().trim()) { return undefined; } const state = this.#modelRegistry.getProviderDiscoveryState(activeProvider.toLowerCase()); @@ -459,21 +598,25 @@ export class ModelSelectorComponent extends Container { #updateList(): void { this.#listContainer.clear(); + const isCanonicalTab = this.#isCanonicalTab(); + const visibleItems = isCanonicalTab ? this.#filteredCanonicalModels : this.#filteredModels; const maxVisible = 10; const startIndex = Math.max( 0, - Math.min(this.#selectedIndex - Math.floor(maxVisible / 2), this.#filteredModels.length - maxVisible), + Math.min(this.#selectedIndex - Math.floor(maxVisible / 2), visibleItems.length - maxVisible), ); - const endIndex = Math.min(startIndex + maxVisible, this.#filteredModels.length); + const endIndex = Math.min(startIndex + maxVisible, visibleItems.length); const activeProvider = this.#getActiveProvider(); const showProvider = activeProvider === ALL_TAB; // Show visible slice of filtered models for (let i = startIndex; i < endIndex; i++) { - const item = this.#filteredModels[i]; + const item = visibleItems[i]; if (!item) continue; + const canonicalItem = isCanonicalTab ? (item as CanonicalModelItem) : undefined; + const providerItem = isCanonicalTab ? undefined : (item as ModelItem); const isSelected = i === this.#selectedIndex; @@ -502,17 +645,25 @@ export class ModelSelectorComponent extends Container { let line = ""; if (isSelected) { const prefix = theme.fg("accent", `${theme.nav.cursor} `); - if (showProvider) { - const providerPrefix = theme.fg("dim", `${item.provider}/`); - line = `${prefix}${providerPrefix}${theme.fg("accent", item.id)}${badgeText}`; + if (isCanonicalTab) { + const variants = theme.fg("dim", ` [${canonicalItem?.variantCount ?? 0}]`); + const backing = theme.fg("dim", ` -> ${item.model.provider}/${item.model.id}`); + line = `${prefix}${theme.fg("accent", item.id)}${variants}${backing}${badgeText}`; + } else if (showProvider) { + const providerPrefix = theme.fg("dim", `${providerItem?.provider ?? ""}/`); + line = `${prefix}${providerPrefix}${theme.fg("accent", providerItem?.id ?? item.id)}${badgeText}`; } else { line = `${prefix}${theme.fg("accent", item.id)}${badgeText}`; } } else { const prefix = " "; - if (showProvider) { - const providerPrefix = theme.fg("dim", `${item.provider}/`); - line = `${prefix}${providerPrefix}${item.id}${badgeText}`; + if (isCanonicalTab) { + const variants = theme.fg("dim", ` [${canonicalItem?.variantCount ?? 0}]`); + const backing = theme.fg("dim", ` -> ${item.model.provider}/${item.model.id}`); + line = `${prefix}${item.id}${variants}${backing}${badgeText}`; + } else if (showProvider) { + const providerPrefix = theme.fg("dim", `${providerItem?.provider ?? ""}/`); + line = `${prefix}${providerPrefix}${providerItem?.id ?? item.id}${badgeText}`; } else { line = `${prefix}${item.id}${badgeText}`; } @@ -522,8 +673,8 @@ export class ModelSelectorComponent extends Container { } // Add scroll indicator if needed - if (startIndex > 0 || endIndex < this.#filteredModels.length) { - const scrollInfo = theme.fg("muted", ` (${this.#selectedIndex + 1}/${this.#filteredModels.length})`); + if (startIndex > 0 || endIndex < visibleItems.length) { + const scrollInfo = theme.fg("muted", ` (${this.#selectedIndex + 1}/${visibleItems.length})`); this.#listContainer.addChild(new Text(scrollInfo, 0, 0)); } @@ -533,13 +684,21 @@ export class ModelSelectorComponent extends Container { for (const line of errorLines) { this.#listContainer.addChild(new Text(theme.fg("error", line), 0, 0)); } - } else if (this.#filteredModels.length === 0) { + } else if (visibleItems.length === 0) { const statusMessage = this.#getProviderEmptyStateMessage(); this.#listContainer.addChild(new Text(theme.fg("muted", statusMessage ?? " No matching models"), 0, 0)); } else { - const selected = this.#filteredModels[this.#selectedIndex]; + const selected = visibleItems[this.#selectedIndex]; + if (!selected) { + return; + } this.#listContainer.addChild(new Spacer(1)); - this.#listContainer.addChild(new Text(theme.fg("muted", ` Model Name: ${selected.model.name}`), 0, 0)); + const suffix = isCanonicalTab + ? ` (${selected.model.provider}/${selected.model.id}, ${(selected as CanonicalModelItem).variantCount} variants)` + : ""; + this.#listContainer.addChild( + new Text(theme.fg("muted", ` Model Name: ${selected.model.name}${suffix}`), 0, 0), + ); } } #getThinkingLevelsForModel(model: Model): ReadonlyArray { @@ -557,8 +716,14 @@ export class ModelSelectorComponent extends Container { return foundIndex >= 0 ? foundIndex : 0; } + #getSelectedItem(): ModelItem | CanonicalModelItem | undefined { + return this.#isCanonicalTab() + ? this.#filteredCanonicalModels[this.#selectedIndex] + : this.#filteredModels[this.#selectedIndex]; + } + #openMenu(): void { - if (this.#filteredModels.length === 0) return; + if (!this.#getSelectedItem()) return; this.#isMenuOpen = true; this.#menuStep = "role"; @@ -577,11 +742,11 @@ export class ModelSelectorComponent extends Container { #updateMenu(): void { this.#menuContainer.clear(); - const selectedModel = this.#filteredModels[this.#selectedIndex]; - if (!selectedModel) return; + const selectedItem = this.#getSelectedItem(); + if (!selectedItem) return; const showingThinking = this.#menuStep === "thinking" && this.#menuSelectedRole !== null; - const thinkingOptions = showingThinking ? this.#getThinkingLevelsForModel(selectedModel.model) : []; + const thinkingOptions = showingThinking ? this.#getThinkingLevelsForModel(selectedItem.model) : []; const optionLines = showingThinking ? thinkingOptions.map((thinkingLevel, index) => { const prefix = index === this.#menuSelectedIndex ? ` ${theme.nav.cursor} ` : " "; @@ -596,8 +761,8 @@ export class ModelSelectorComponent extends Container { const selectedRoleName = this.#menuSelectedRole ? getRoleInfo(this.#menuSelectedRole, this.#settings).name : ""; const headerText = showingThinking && this.#menuSelectedRole - ? ` Thinking for: ${selectedRoleName} (${selectedModel.id})` - : ` Action for: ${selectedModel.id}`; + ? ` Thinking for: ${selectedRoleName} (${selectedItem.id})` + : ` Action for: ${selectedItem.id}`; const hintText = showingThinking ? " Enter: confirm Esc: back" : " Enter: continue Esc: cancel"; const menuWidth = Math.max( visibleWidth(headerText), @@ -610,15 +775,13 @@ export class ModelSelectorComponent extends Container { if (showingThinking && this.#menuSelectedRole) { this.#menuContainer.addChild( new Text( - theme.fg("text", ` Thinking for: ${theme.bold(selectedRoleName)} (${theme.bold(selectedModel.id)})`), + theme.fg("text", ` Thinking for: ${theme.bold(selectedRoleName)} (${theme.bold(selectedItem.id)})`), 0, 0, ), ); } else { - this.#menuContainer.addChild( - new Text(theme.fg("text", ` Action for: ${theme.bold(selectedModel.id)}`), 0, 0), - ); + this.#menuContainer.addChild(new Text(theme.fg("text", ` Action for: ${theme.bold(selectedItem.id)}`), 0, 0)); } this.#menuContainer.addChild(new Spacer(1)); @@ -648,27 +811,29 @@ export class ModelSelectorComponent extends Container { // Up arrow - navigate list (wrap to bottom when at top) if (matchesKey(keyData, "up")) { - if (this.#filteredModels.length === 0) return; - this.#selectedIndex = this.#selectedIndex === 0 ? this.#filteredModels.length - 1 : this.#selectedIndex - 1; + const itemCount = this.#isCanonicalTab() ? this.#filteredCanonicalModels.length : this.#filteredModels.length; + if (itemCount === 0) return; + this.#selectedIndex = this.#selectedIndex === 0 ? itemCount - 1 : this.#selectedIndex - 1; this.#updateList(); return; } // Down arrow - navigate list (wrap to top when at bottom) if (matchesKey(keyData, "down")) { - if (this.#filteredModels.length === 0) return; - this.#selectedIndex = this.#selectedIndex === this.#filteredModels.length - 1 ? 0 : this.#selectedIndex + 1; + const itemCount = this.#isCanonicalTab() ? this.#filteredCanonicalModels.length : this.#filteredModels.length; + if (itemCount === 0) return; + this.#selectedIndex = this.#selectedIndex === itemCount - 1 ? 0 : this.#selectedIndex + 1; this.#updateList(); return; } // Enter - open context menu or select directly in temporary mode if (matchesKey(keyData, "enter") || matchesKey(keyData, "return") || keyData === "\n") { - const selectedModel = this.#filteredModels[this.#selectedIndex]; - if (selectedModel) { + const selectedItem = this.#getSelectedItem(); + if (selectedItem) { if (this.#temporaryOnly) { // In temporary mode, skip menu and select directly - this.#handleSelect(selectedModel.model, null); + this.#handleSelect(selectedItem, null); } else { this.#openMenu(); } @@ -687,12 +852,12 @@ export class ModelSelectorComponent extends Container { this.#filterModels(this.#searchInput.getValue()); } #handleMenuInput(keyData: string): void { - const selectedModel = this.#filteredModels[this.#selectedIndex]; - if (!selectedModel) return; + const selectedItem = this.#getSelectedItem(); + if (!selectedItem) return; const optionCount = this.#menuStep === "thinking" && this.#menuSelectedRole !== null - ? this.#getThinkingLevelsForModel(selectedModel.model).length + ? this.#getThinkingLevelsForModel(selectedItem.model).length : this.#menuRoleActions.length; if (optionCount === 0) return; @@ -714,16 +879,16 @@ export class ModelSelectorComponent extends Container { if (!action) return; this.#menuSelectedRole = action.role; this.#menuStep = "thinking"; - this.#menuSelectedIndex = this.#getThinkingPreselectIndex(action.role, selectedModel.model); + this.#menuSelectedIndex = this.#getThinkingPreselectIndex(action.role, selectedItem.model); this.#updateMenu(); return; } if (!this.#menuSelectedRole) return; - const thinkingOptions = this.#getThinkingLevelsForModel(selectedModel.model); + const thinkingOptions = this.#getThinkingLevelsForModel(selectedItem.model); const thinkingLevel = thinkingOptions[this.#menuSelectedIndex]; if (!thinkingLevel) return; - this.#handleSelect(selectedModel.model, this.#menuSelectedRole, thinkingLevel); + this.#handleSelect(selectedItem, this.#menuSelectedRole, thinkingLevel); this.#closeMenu(); return; } @@ -742,28 +907,20 @@ export class ModelSelectorComponent extends Container { } } - #formatRoleModelValue(model: Model, thinkingLevel: ThinkingLevel): string { - const modelKey = `${model.provider}/${model.id}`; - if (thinkingLevel === ThinkingLevel.Inherit) return modelKey; - return `${modelKey}:${thinkingLevel}`; - } - #handleSelect(model: Model, role: string | null, thinkingLevel?: ThinkingLevel): void { + #handleSelect(item: ModelItem | CanonicalModelItem, role: string | null, thinkingLevel?: ThinkingLevel): void { // For temporary role, don't save to settings - just notify caller if (role === null) { - this.#onSelectCallback(model, null); + this.#onSelectCallback(item.model, null, undefined, item.selector); return; } const selectedThinkingLevel = thinkingLevel ?? this.#getCurrentRoleThinkingLevel(role); - // Save to settings - this.#settings.setModelRole(role, this.#formatRoleModelValue(model, selectedThinkingLevel)); - // Update local state for UI - this.#roles[role] = { model, thinkingLevel: selectedThinkingLevel }; + this.#roles[role] = { model: item.model, thinkingLevel: selectedThinkingLevel }; // Notify caller (for updating agent state if needed) - this.#onSelectCallback(model, role, selectedThinkingLevel); + this.#onSelectCallback(item.model, role, selectedThinkingLevel, item.selector); // Update list to show new badges this.#updateList(); diff --git a/packages/coding-agent/src/modes/controllers/selector-controller.ts b/packages/coding-agent/src/modes/controllers/selector-controller.ts index aac5e9256..254e68b96 100644 --- a/packages/coding-agent/src/modes/controllers/selector-controller.ts +++ b/packages/coding-agent/src/modes/controllers/selector-controller.ts @@ -7,6 +7,7 @@ import { Input, Loader, Spacer, Text } from "@oh-my-pi/pi-tui"; import { getAgentDbPath, getConfigDirName, getProjectDir } from "@oh-my-pi/pi-utils"; import { invalidate as invalidateFsCache } from "../../capability/fs"; import { getRoleInfo } from "../../config/model-registry"; +import { formatModelSelectorValue } from "../../config/model-resolver"; import { settings } from "../../config/settings"; import { DebugSelectorComponent } from "../../debug"; import { disableProvider, enableProvider } from "../../discovery"; @@ -387,31 +388,38 @@ export class SelectorController { this.ctx.settings, this.ctx.session.modelRegistry, this.ctx.session.scopedModels, - async (model, role, thinkingLevel) => { + async (model, role, thinkingLevel, selector) => { try { if (role === null) { // Temporary: update agent state but don't persist to settings await this.ctx.session.setModelTemporary(model); this.ctx.statusLine.invalidate(); this.ctx.updateEditorBorderColor(); - this.ctx.showStatus(`Temporary model: ${model.id}`); + this.ctx.showStatus(`Temporary model: ${selector ?? model.id}`); done(); this.ctx.ui.requestRender(); } else if (role === "default") { // Default: update agent state and persist - await this.ctx.session.setModel(model, role); + await this.ctx.session.setModel(model, role, { + selector, + thinkingLevel, + }); if (thinkingLevel && thinkingLevel !== ThinkingLevel.Inherit) { this.ctx.session.setThinkingLevel(thinkingLevel); } this.ctx.statusLine.invalidate(); this.ctx.updateEditorBorderColor(); - this.ctx.showStatus(`Default model: ${model.id}`); + this.ctx.showStatus(`Default model: ${selector ?? model.id}`); // Don't call done() - selector stays open for role assignment } else { // Other roles (smol, slow): just update settings, not current model + this.ctx.settings.setModelRole( + role, + formatModelSelectorValue(selector ?? `${model.provider}/${model.id}`, thinkingLevel), + ); const roleInfo = getRoleInfo(role, settings); const roleLabel = roleInfo?.name ?? role; - this.ctx.showStatus(`${roleLabel} model: ${model.id}`); + this.ctx.showStatus(`${roleLabel} model: ${selector ?? model.id}`); // Don't call done() - selector stays open } } catch (error) { diff --git a/packages/coding-agent/src/sdk.ts b/packages/coding-agent/src/sdk.ts index 474861c71..07761211d 100644 --- a/packages/coding-agent/src/sdk.ts +++ b/packages/coding-agent/src/sdk.ts @@ -728,6 +728,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} resolveModelRoleValue(settings.getModelRole("default"), modelRegistry.getAvailable(), { settings, matchPreferences: modelMatchPreferences, + modelRegistry, }), ); let model = options.model; @@ -1132,7 +1133,9 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder(), }; - const { model: resolved } = parseModelPattern(options.modelPattern, availableModels, matchPreferences); + const { model: resolved } = parseModelPattern(options.modelPattern, availableModels, matchPreferences, { + modelRegistry, + }); if (resolved) { model = resolved; modelFallbackMessage = undefined; diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index 3cbef15a9..22b2ac473 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -56,6 +56,7 @@ import type { Rule } from "../capability/rule"; import { MODEL_ROLE_IDS, type ModelRegistry } from "../config/model-registry"; import { extractExplicitThinkingSelector, + formatModelSelectorValue, formatModelString, parseModelString, type ResolvedModelRoleValue, @@ -3390,7 +3391,11 @@ export class AgentSession { * Validates API key, saves to session and settings. * @throws Error if no API key available for the model */ - async setModel(model: Model, role: string = "default"): Promise { + async setModel( + model: Model, + role: string = "default", + options?: { selector?: string; thinkingLevel?: ThinkingLevel }, + ): Promise { const apiKey = await this.#modelRegistry.getApiKey(model, this.sessionId); if (!apiKey) { throw new Error(`No API key for ${model.provider}/${model.id}`); @@ -3399,7 +3404,10 @@ export class AgentSession { this.#clearActiveRetryFallback(); this.#setModelWithProviderSessionReset(model); this.sessionManager.appendModelChange(`${model.provider}/${model.id}`, role); - this.settings.setModelRole(role, this.#formatRoleModelValue(role, model)); + this.settings.setModelRole( + role, + this.#formatRoleModelValue(role, model, options?.selector, options?.thinkingLevel), + ); this.settings.getStorage()?.recordModelUsage(`${model.provider}/${model.id}`); // Re-apply the current thinking level for the newly selected model @@ -3472,6 +3480,7 @@ export class AgentSession { const resolved = resolveModelRoleValue(roleModelStr, availableModels, { settings: this.settings, matchPreferences, + modelRegistry: this.#modelRegistry, }); if (!resolved.model) continue; @@ -4588,14 +4597,21 @@ export class AgentSession { return `${model.provider}/${model.id}`; } - #formatRoleModelValue(role: string, model: Model): string { - const modelKey = `${model.provider}/${model.id}`; + #formatRoleModelValue( + role: string, + model: Model, + selectorOverride?: string, + thinkingLevelOverride?: ThinkingLevel, + ): string { + const modelKey = selectorOverride ?? `${model.provider}/${model.id}`; + if (thinkingLevelOverride !== undefined) { + return formatModelSelectorValue(modelKey, thinkingLevelOverride); + } const existingRoleValue = this.settings.getModelRole(role); if (!existingRoleValue) return modelKey; const thinkingLevel = extractExplicitThinkingSelector(existingRoleValue, this.settings); - if (thinkingLevel === undefined) return modelKey; - return `${modelKey}:${thinkingLevel}`; + return formatModelSelectorValue(modelKey, thinkingLevel); } #resolveContextPromotionConfiguredTarget(currentModel: Model, availableModels: Model[]): Model | undefined { const configuredTarget = currentModel.contextPromotionTarget?.trim(); @@ -4628,6 +4644,7 @@ export class AgentSession { return resolveModelRoleValue(roleModelStr, availableModels, { settings: this.settings, matchPreferences: { usageOrder: this.settings.getStorage()?.getModelUsageOrder() }, + modelRegistry: this.#modelRegistry, }); } diff --git a/packages/coding-agent/src/tools/inspect-image.ts b/packages/coding-agent/src/tools/inspect-image.ts index 1f3e915cd..2f246811c 100644 --- a/packages/coding-agent/src/tools/inspect-image.ts +++ b/packages/coding-agent/src/tools/inspect-image.ts @@ -79,7 +79,7 @@ export class InspectImageTool implements AgentTool | undefined => { if (!pattern) return undefined; const expanded = expandRoleAlias(pattern, this.session.settings); - return resolveModelFromString(expanded, availableModels, matchPreferences); + return resolveModelFromString(expanded, availableModels, matchPreferences, modelRegistry); }; const activeModelPattern = this.session.getActiveModelString?.() ?? this.session.getModelString?.(); diff --git a/packages/coding-agent/src/utils/commit-message-generator.ts b/packages/coding-agent/src/utils/commit-message-generator.ts index 92019c1f4..218226fde 100644 --- a/packages/coding-agent/src/utils/commit-message-generator.ts +++ b/packages/coding-agent/src/utils/commit-message-generator.ts @@ -52,6 +52,7 @@ function getSmolModelCandidates( const configuredSmol = resolveModelRoleValue(settings.getModelRole("smol"), availableModels, { settings, matchPreferences, + modelRegistry: registry, }); addCandidate(configuredSmol.model, configuredSmol.thinkingLevel); diff --git a/packages/coding-agent/src/utils/title-generator.ts b/packages/coding-agent/src/utils/title-generator.ts index f5c895fca..34a764c27 100644 --- a/packages/coding-agent/src/utils/title-generator.ts +++ b/packages/coding-agent/src/utils/title-generator.ts @@ -27,7 +27,7 @@ function getTitleModel( const availableModels = registry.getAvailable(); if (availableModels.length === 0) return undefined; - const titleModel = resolveRoleSelection(["commit", "smol"], settings, availableModels); + const titleModel = resolveRoleSelection(["commit", "smol"], settings, availableModels, registry); if (titleModel) { return { model: titleModel.model, thinkingLevel: titleModel.thinkingLevel }; } diff --git a/packages/coding-agent/test/keybindings-escape-components.test.ts b/packages/coding-agent/test/keybindings-escape-components.test.ts index 515091473..3ef7138eb 100644 --- a/packages/coding-agent/test/keybindings-escape-components.test.ts +++ b/packages/coding-agent/test/keybindings-escape-components.test.ts @@ -77,6 +77,8 @@ describe("component escape bindings", () => { const modelRegistry = { getAll: () => [model], getDiscoverableProviders: () => [], + getCanonicalModels: () => [], + resolveCanonicalModel: () => undefined, } as unknown as ModelRegistry; const ui = { requestRender: vi.fn(), diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index 85f78b54f..aa090fce0 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -103,6 +103,10 @@ describe("ModelRegistry", () => { fs.writeFileSync(modelsJsonPath, JSON.stringify({ providers })); } + function writeRawModelsConfig(config: Record) { + fs.writeFileSync(modelsJsonPath, JSON.stringify(config)); + } + function mockOpenAiCompatibleModels(url: string, modelIds: string[]) { return hookFetch(input => { const requestUrl = String(input); @@ -116,6 +120,222 @@ describe("ModelRegistry", () => { }); } + describe("canonical equivalence", () => { + test("groups dotted provider variants under the bundled canonical id", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [{ id: "anthropic/claude-sonnet-4.5" }]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const variants = registry.getCanonicalVariants("claude-sonnet-4-5"); + + expect(variants.some(variant => variant.selector === "anthropic/claude-sonnet-4-5")).toBe(true); + expect(variants.some(variant => variant.selector === "demo/anthropic/claude-sonnet-4.5")).toBe(true); + }); + + test("collapses wrapped, dated, and tuned anthropic variants under the base canonical id", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [ + { id: "anthropic/claude-opus-4.5" }, + { id: "claude-opus-4-5-20251101" }, + { id: "claude-4.5-opus-high-thinking" }, + ]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const variants = registry.getCanonicalVariants("claude-opus-4-5"); + + expect(variants.some(variant => variant.selector === "demo/anthropic/claude-opus-4.5")).toBe(true); + expect(variants.some(variant => variant.selector === "demo/claude-opus-4-5-20251101")).toBe(true); + expect(variants.some(variant => variant.selector === "demo/claude-4.5-opus-high-thinking")).toBe(true); + }); + + test("collapses gitlab duo chat wrapper ids into the upstream canonical id", () => { + writeRawModelsJson({ + "gitlab-duo": providerConfig("https://demo.example.com/v1", [{ id: "duo-chat-opus-4-6" }]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const variants = registry.getCanonicalVariants("claude-opus-4-6"); + + expect(variants.some(variant => variant.selector === "gitlab-duo/duo-chat-opus-4-6")).toBe(true); + }); + + test("collapses synthetic and vendor-prefixed glm wrappers into the upstream canonical id", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [{ id: "hf:zai-org/GLM-4.7" }, { id: "zai-glm-4.7" }]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const variants = registry.getCanonicalVariants("glm-4.7"); + + expect(variants.some(variant => variant.selector === "demo/hf:zai-org/GLM-4.7")).toBe(true); + expect(variants.some(variant => variant.selector === "demo/zai-glm-4.7")).toBe(true); + }); + + test("collapses compact and reordered claude aliases into the upstream canonical id", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [ + { id: "claude-opus-45" }, + { id: "claude-4.5-sonnet" }, + ]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const opusVariants = registry.getCanonicalVariants("claude-opus-4-5"); + const sonnetVariants = registry.getCanonicalVariants("claude-sonnet-4-5"); + + expect(opusVariants.some(variant => variant.selector === "demo/claude-opus-45")).toBe(true); + expect(sonnetVariants.some(variant => variant.selector === "demo/claude-4.5-sonnet")).toBe(true); + }); + + test("collapses anthropic latest aliases into the best upstream claude family id", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [ + { id: "anthropic/claude-opus-latest" }, + { id: "anthropic/claude-haiku-latest" }, + ]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const opusVariants = registry.getCanonicalVariants("claude-opus-4-6"); + const haikuVariants = registry.getCanonicalVariants("claude-haiku-4-5"); + + expect(opusVariants.some(variant => variant.selector === "demo/anthropic/claude-opus-latest")).toBe(true); + expect(haikuVariants.some(variant => variant.selector === "demo/anthropic/claude-haiku-latest")).toBe(true); + expect( + registry + .getCanonicalVariants("claude-haiku-4-5-20251001-thinking") + .some(variant => variant.selector === "demo/anthropic/claude-haiku-latest"), + ).toBe(false); + }); + + test("collapses wrapped gemini tool and tuning variants under the base preview id", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [ + { id: "google/gemini-3.1-pro-preview" }, + { id: "google/gemini-3.1-pro-preview-customtools" }, + { id: "google/gemini-3.1-pro-preview-high" }, + ]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const variants = registry.getCanonicalVariants("gemini-3.1-pro-preview"); + + expect(variants.some(variant => variant.selector === "demo/google/gemini-3.1-pro-preview")).toBe(true); + expect(variants.some(variant => variant.selector === "demo/google/gemini-3.1-pro-preview-customtools")).toBe( + true, + ); + expect(variants.some(variant => variant.selector === "demo/google/gemini-3.1-pro-preview-high")).toBe(true); + }); + + test("collapses compact version aliases and hardware suffixes into clean canonical ids", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [ + { id: "hf:nvidia/Kimi-K2.5-NVFP4" }, + { id: "kimi-k2-5" }, + { id: "z-ai/glm4.7" }, + { id: "z-ai/glm5" }, + ]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const kimiVariants = registry.getCanonicalVariants("kimi-k2.5"); + const glm47Variants = registry.getCanonicalVariants("glm-4.7"); + const glm5Variants = registry.getCanonicalVariants("glm-5"); + + expect(kimiVariants.some(variant => variant.selector === "demo/hf:nvidia/Kimi-K2.5-NVFP4")).toBe(true); + expect(kimiVariants.some(variant => variant.selector === "demo/kimi-k2-5")).toBe(true); + expect(glm47Variants.some(variant => variant.selector === "demo/z-ai/glm4.7")).toBe(true); + expect(glm5Variants.some(variant => variant.selector === "demo/z-ai/glm5")).toBe(true); + }); + + test("prefers clean canonical ids over bundled wrapper ids when available", () => { + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [ + { id: "zai/glm-4.6v-flash" }, + { id: "hf:deepseek-ai/DeepSeek-V3" }, + { id: "google/gemini-pro-latest" }, + ]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + + expect( + registry + .getCanonicalVariants("glm-4.6v-flash") + .some(variant => variant.selector === "demo/zai/glm-4.6v-flash"), + ).toBe(true); + expect( + registry + .getCanonicalVariants("deepseek-v3") + .some(variant => variant.selector === "demo/hf:deepseek-ai/DeepSeek-V3"), + ).toBe(true); + expect( + registry + .getCanonicalVariants("gemini-pro") + .some(variant => variant.selector === "demo/google/gemini-pro-latest"), + ).toBe(true); + }); + + test("applies explicit equivalence overrides from config", () => { + writeRawModelsConfig({ + providers: { + "p-anthropic": providerConfig("https://demo.example.com/v1", [{ id: "corp-sonnet" }]), + }, + equivalence: { + overrides: { + "p-anthropic/corp-sonnet": "claude-sonnet-4-5", + }, + }, + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const variants = registry.getCanonicalVariants("claude-sonnet-4-5"); + + expect(variants.some(variant => variant.selector === "p-anthropic/corp-sonnet")).toBe(true); + }); + + test("exclusions keep variants out of canonical grouping", () => { + writeRawModelsConfig({ + providers: { + demo: providerConfig("https://demo.example.com/v1", [{ id: "anthropic/claude-sonnet-4.5" }]), + }, + equivalence: { + exclude: ["demo/anthropic/claude-sonnet-4.5"], + }, + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const grouped = registry.getCanonicalVariants("claude-sonnet-4-5"); + const fallback = registry.getCanonicalVariants("anthropic/claude-sonnet-4.5"); + + expect(grouped.some(variant => variant.selector === "demo/anthropic/claude-sonnet-4.5")).toBe(false); + expect(fallback.some(variant => variant.selector === "demo/anthropic/claude-sonnet-4.5")).toBe(true); + }); + + test("resolves canonical models using configured provider order", async () => { + await Settings.init({ + inMemory: true, + overrides: { + modelProviderOrder: ["demo", "anthropic"], + }, + }); + writeRawModelsJson({ + demo: providerConfig("https://demo.example.com/v1", [{ id: "anthropic/claude-sonnet-4.5" }]), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const resolved = registry.resolveCanonicalModel("claude-sonnet-4-5", { + availableOnly: false, + candidates: registry.getAll(), + }); + + expect(resolved?.provider).toBe("demo"); + expect(resolved?.id).toBe("anthropic/claude-sonnet-4.5"); + }); + }); + describe("baseUrl override (no custom models)", () => { test("overriding baseUrl keeps all built-in models", () => { writeRawModelsJson({ diff --git a/packages/coding-agent/test/model-resolver.test.ts b/packages/coding-agent/test/model-resolver.test.ts index f46a5d299..27b7545c7 100644 --- a/packages/coding-agent/test/model-resolver.test.ts +++ b/packages/coding-agent/test/model-resolver.test.ts @@ -9,6 +9,7 @@ import { resolveModelFromString, resolveModelOverride, resolveModelRoleValue, + resolveModelScope, } from "@oh-my-pi/pi-coding-agent/config/model-resolver"; import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; @@ -142,6 +143,66 @@ const mockCodexOverlapModels: Model<"anthropic-messages">[] = [ }, ]; +const canonicalVariantModels: Model<"anthropic-messages">[] = [ + { + id: "claude-sonnet-4-5", + name: "Claude Sonnet 4.5", + api: "anthropic-messages", + provider: "anthropic", + baseUrl: "https://api.anthropic.com", + reasoning: true, + thinking: { + mode: "budget", + minLevel: Effort.Minimal, + maxLevel: Effort.High, + }, + input: ["text", "image"], + cost: { input: 3, output: 15, cacheRead: 0.3, cacheWrite: 3.75 }, + contextWindow: 200000, + maxTokens: 8192, + }, + { + id: "anthropic/claude-sonnet-4.5", + name: "Claude Sonnet 4.5 (Copilot)", + api: "anthropic-messages", + provider: "github-copilot", + baseUrl: "https://api.githubcopilot.com", + reasoning: true, + thinking: { + mode: "budget", + minLevel: Effort.Minimal, + maxLevel: Effort.High, + }, + input: ["text", "image"], + cost: { input: 3, output: 15, cacheRead: 0.3, cacheWrite: 3.75 }, + contextWindow: 200000, + maxTokens: 8192, + }, +]; + +const canonicalRegistry = { + resolveCanonicalModel: (canonicalId: string, options?: { candidates?: Model<"anthropic-messages">[] }) => { + if (canonicalId !== "claude-sonnet-4-5") return undefined; + const candidates = options?.candidates ?? canonicalVariantModels; + return ( + candidates.find(model => model.provider === "github-copilot") ?? + candidates.find(model => model.provider === "anthropic") + ); + }, + getCanonicalVariants: (canonicalId: string, options?: { candidates?: Model<"anthropic-messages">[] }) => { + if (canonicalId !== "claude-sonnet-4-5") return []; + const candidates = options?.candidates ?? canonicalVariantModels; + return candidates.map(model => ({ + canonicalId, + selector: `${model.provider}/${model.id}`, + model, + source: model.id === canonicalId ? "bundled" : "heuristic", + })); + }, + getCanonicalId: () => "claude-sonnet-4-5", + getAvailable: () => canonicalVariantModels, +} as unknown as Parameters[0]["modelRegistry"]; + const allModels = [...mockModels, ...mockOpenRouterModels, ...mockProviderOverlapModels, ...mockCodexOverlapModels]; describe("parseModelPattern", () => { @@ -319,6 +380,16 @@ describe("parseModelPattern", () => { expect(result.model?.id).toBe("moonshotai/kimi-k2.5"); }); }); + + describe("canonical ids", () => { + test("resolves an exact canonical id through the registry before bare-id matching", () => { + const result = parseModelPattern("claude-sonnet-4-5", canonicalVariantModels, undefined, { + modelRegistry: canonicalRegistry, + }); + expect(result.model?.provider).toBe("github-copilot"); + expect(result.model?.id).toBe("anthropic/claude-sonnet-4.5"); + }); + }); }); describe("resolveModelRoleValue", () => { @@ -484,6 +555,21 @@ describe("resolveModelOverride", () => { }); }); describe("resolveCliModel", () => { + test("resolves exact canonical ids to the preferred concrete provider", () => { + const result = resolveCliModel({ + cliModel: "claude-sonnet-4-5", + modelRegistry: { + ...canonicalRegistry, + getAll: () => canonicalVariantModels, + } as unknown as Parameters[0]["modelRegistry"], + }); + + expect(result.error).toBeUndefined(); + expect(result.selector).toBe("claude-sonnet-4-5"); + expect(result.model?.provider).toBe("github-copilot"); + expect(result.model?.id).toBe("anthropic/claude-sonnet-4.5"); + }); + test("resolves --model provider/id without --provider", () => { const registry = { getAll: () => allModels, @@ -634,6 +720,22 @@ describe("resolveCliModel", () => { }); }); +describe("resolveModelScope", () => { + test("expands exact canonical ids into all concrete variants", async () => { + const scoped = await resolveModelScope(["claude-sonnet-4-5"], { + getAvailable: () => canonicalVariantModels, + getCanonicalVariants: (canonicalId: string, options?: { candidates?: Model<"anthropic-messages">[] }) => + canonicalRegistry.getCanonicalVariants!(canonicalId, options), + } as unknown as Parameters[1]); + + expect(scoped).toHaveLength(2); + expect(scoped.map(entry => `${entry.model.provider}/${entry.model.id}`).sort()).toEqual([ + "anthropic/claude-sonnet-4-5", + "github-copilot/anthropic/claude-sonnet-4.5", + ]); + }); +}); + describe("parseModelString", () => { test("parses standard provider/id format", () => { const result = parseModelString("anthropic/claude-sonnet-4-5"); diff --git a/packages/coding-agent/test/model-selector-role-badge-thinking.test.ts b/packages/coding-agent/test/model-selector-role-badge-thinking.test.ts index 8709fc18b..89b609779 100644 --- a/packages/coding-agent/test/model-selector-role-badge-thinking.test.ts +++ b/packages/coding-agent/test/model-selector-role-badge-thinking.test.ts @@ -21,6 +21,8 @@ function createSelector(model: Model, settings: Settings): ModelSelectorComponen const modelRegistry = { getAll: () => [model], getDiscoverableProviders: () => [], + getCanonicalModels: () => [], + resolveCanonicalModel: () => undefined, } as unknown as ModelRegistry; const ui = { requestRender: vi.fn(),