feat(coding-agent): added canonical aliases for model role resolution

- Added canonical model equivalence types, cache helpers, and registry APIs for provider variant lookup.
- Changed model resolution to apply canonical ID overrides/excludes with provider order before fallback matching.
- Added canonical and provider model views in list-models and selector UI with canonical sorting/persistence.
- Updated role/model persistence to store selectors while runtime now resolves concrete canonical-backed provider models.
This commit is contained in:
can1357
2026-04-11 08:28:50 +02:00
parent 0dab29c07c
commit 5277e44139
22 changed files with 1876 additions and 221 deletions
+19 -4
View File
@@ -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`
+106
View File
@@ -29,10 +29,20 @@ Legacy behavior still present:
providers:
<provider-id>:
# provider-level config
equivalence:
overrides:
<provider-id>/<model-id>: <canonical-model-id>
exclude:
- <provider-id>/<model-id>
```
`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.
+12 -1
View File
@@ -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
+96 -57
View File
@@ -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<T extends Record<string, string>>(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<keyof T, number>;
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<Api>[] = 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",
});
}
}
@@ -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<Api>) => Promise<string | undefined>;
};
export async function resolvePrimaryModel(
override: string | undefined,
settings: Settings,
modelRegistry: {
getAvailable: () => Model<Api>[];
getApiKey: (model: Model<Api>) => Promise<string | undefined>;
},
modelRegistry: CommitModelRegistry,
): Promise<ResolvedCommitModel> {
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<Api>[];
getApiKey: (model: Model<Api>) => Promise<string | undefined>;
},
modelRegistry: CommitModelRegistry,
fallbackModel: Model<Api>,
fallbackApiKey: string,
): Promise<ResolvedCommitModel> {
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 };
@@ -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<string, string>;
exclude?: string[];
}
export interface CanonicalModelVariant {
canonicalId: string;
selector: string;
model: Model<Api>;
source: CanonicalModelSource;
}
export interface CanonicalModelRecord {
id: string;
name: string;
variants: CanonicalModelVariant[];
}
export interface CanonicalModelIndex {
records: CanonicalModelRecord[];
byId: Map<string, CanonicalModelRecord>;
bySelector: Map<string, string>;
}
interface CanonicalReferenceData {
references: Map<string, Model<Api>>;
officialIds: Set<string>;
}
interface CompiledEquivalenceConfig {
overrides: Map<string, string>;
exclude: Set<string>;
}
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<Api> | undefined, candidate: Model<Api>): 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<string, Model<Api>>();
for (const provider of getBundledProviders()) {
for (const model of getBundledModels(provider as Parameters<typeof getBundledModels>[0])) {
const candidate = model as Model<Api>;
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<Api>): string {
return `${model.provider}/${model.id}`;
}
function buildOverrideMap(overrides: Record<string, string> | undefined): Map<string, string> {
const result = new Map<string, string>();
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<string> {
const result = new Set<string>();
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<string>, 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<string>();
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<string>();
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<string>();
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<string>();
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>): 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>): 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<string>();
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<string>();
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<string>();
const queue = [modelId];
const visited = new Set<string>();
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<Api>,
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<Api>[],
equivalence?: ModelEquivalenceConfig,
): CanonicalModelIndex {
const referenceData = createCanonicalReferenceData();
const compiledEquivalence = compileEquivalenceConfig(equivalence);
const byId = new Map<string, CanonicalModelRecord>();
const bySelector = new Map<string, string>();
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 };
}
@@ -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<typeof ModelsConfigSchema>;
@@ -356,7 +372,7 @@ function validateProviderConfiguration(
export const ModelsConfigFile = new ConfigFile<ModelsConfig>("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<Api>[];
}
/** Result of loading custom models from models.json */
interface CustomModelsResult {
models?: CustomModelOverlay[];
@@ -413,6 +434,7 @@ interface CustomModelsResult {
keylessProviders?: Set<string>;
discoverableProviders?: DiscoveryProviderConfig[];
configuredProviders?: Set<string>;
equivalence?: ModelEquivalenceConfig;
error?: ConfigError;
found: boolean;
}
@@ -739,17 +761,27 @@ function getDisabledProviderIdsFromSettings(): Set<string> {
}
}
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<Api>[] = [];
#canonicalIndex: CanonicalModelIndex = { records: [], byId: new Map(), bySelector: new Map() };
#customProviderApiKeys: Map<string, string> = new Map();
#keylessProviders: Set<string> = new Set();
#discoverableProviders: DiscoveryProviderConfig[] = [];
#customModelOverlays: CustomModelOverlay[] = [];
#providerOverrides: Map<string, ProviderOverride> = new Map();
#modelOverrides: Map<string, Map<string, ModelOverride>> = new Map();
#equivalenceConfig: ModelEquivalenceConfig | undefined;
#configError: ConfigError | undefined = undefined;
#modelsConfigFile: ConfigFile<ModelsConfig>;
#registeredProviderSources: Set<string> = 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<string, Map<string, ModelOverride>>();
const keylessProviders = new Set<string>();
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<Api>): 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<Api>[]): Map<string, number> {
const configuredProviders = getConfiguredProviderOrderFromSettings();
const result = new Map<string, number>();
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<Api>[],
): CanonicalModelVariant | undefined {
if (variants.length === 0) {
return undefined;
}
const providerRank = this.#providerRank(allCandidates);
const modelOrder = new Map<string, number>();
for (let index = 0; index < allCandidates.length; index += 1) {
modelOrder.set(formatCanonicalVariantSelector(allCandidates[index]!), index);
}
const sourceRank: Record<CanonicalModelVariant["source"], number> = {
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<Api> | 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<Api>): 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<Api>[] {
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();
}
}
@@ -58,6 +58,10 @@ export function formatModelString(model: Model<Api>): 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<ModelRegistry, "resolveCanonicalModel" | "getCanonicalVariants" | "getCanonicalId">
>;
export type ModelLookupRegistry = Pick<ModelRegistry, "getAvailable"> & Partial<CanonicalModelRegistry>;
type CliModelRegistry = Pick<ModelRegistry, "getAll"> & Partial<CanonicalModelRegistry>;
type InitialModelRegistry = Pick<ModelRegistry, "getAvailable" | "find">;
type RestorableModelRegistry = Pick<ModelRegistry, "getAvailable" | "find" | "getApiKey">;
interface ModelPreferenceContext {
modelUsageRank: Map<string, number>;
providerUsageRank: Map<string, number>;
@@ -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<Api>[],
modelRegistry: CanonicalModelRegistry | undefined,
): Model<Api> | 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<Api>[],
context: ModelPreferenceContext,
options?: { modelRegistry?: CanonicalModelRegistry },
): Model<Api> | 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<Api>[],
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<Api>[],
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<Api>[],
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<Api>[],
matchPreferences?: ModelMatchPreferences,
modelRegistry?: CanonicalModelRegistry,
): Model<Api> | 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<Api>[];
matchPreferences?: ModelMatchPreferences;
roleOrder?: readonly ModelRole[];
modelRegistry?: CanonicalModelRegistry;
}): Model<Api> | 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<Api>; 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<Api>[],
modelRegistry?: CanonicalModelRegistry,
): { model: Model<Api>; 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<ModelRegistry, "getCanonicalVariants">,
availableModels: Model<Api>[],
): { models: Model<Api>[]; 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<ModelRegistry, "getAvailable" | "getCanonicalVariants">,
preferences?: ModelMatchPreferences,
): Promise<ScopedModel[]> {
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<Api> | 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<InitialModelResult> {
const {
cliProvider,
@@ -923,7 +1050,7 @@ export async function restoreModelFromSession(
savedModelId: string,
currentModel: Model<Api> | undefined,
shouldPrintMessages: boolean,
modelRegistry: ModelRegistry,
modelRegistry: RestorableModelRegistry,
): Promise<{ model: Model<Api> | 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<Model<Api> | 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<Model<Api> | 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)
@@ -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 },
// ────────────────────────────────────────────────────────────────────────
+4 -1
View File
@@ -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;
+7 -6
View File
@@ -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];
}
@@ -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<string, RoleAssignment | undefined>;
#settings = null as unknown as Settings;
@@ -97,7 +128,7 @@ export class ModelSelectorComponent extends Container {
settings: Settings,
modelRegistry: ModelRegistry,
scopedModels: ReadonlyArray<ScopedModelItem>,
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<void> {
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<void> {
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<ThinkingLevel> {
@@ -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();
@@ -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) {
+4 -1
View File
@@ -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;
@@ -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<void> {
async setModel(
model: Model,
role: string = "default",
options?: { selector?: string; thinkingLevel?: ThinkingLevel },
): Promise<void> {
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,
});
}
@@ -79,7 +79,7 @@ export class InspectImageTool implements AgentTool<typeof inspectImageSchema, In
const resolvePattern = (pattern: string | undefined): Model<Api> | 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?.();
@@ -52,6 +52,7 @@ function getSmolModelCandidates(
const configuredSmol = resolveModelRoleValue(settings.getModelRole("smol"), availableModels, {
settings,
matchPreferences,
modelRegistry: registry,
});
addCandidate(configuredSmol.model, configuredSmol.thinkingLevel);
@@ -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 };
}
@@ -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(),
@@ -103,6 +103,10 @@ describe("ModelRegistry", () => {
fs.writeFileSync(modelsJsonPath, JSON.stringify({ providers }));
}
function writeRawModelsConfig(config: Record<string, unknown>) {
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({
@@ -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<typeof resolveCliModel>[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<typeof resolveCliModel>[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<typeof resolveModelScope>[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");
@@ -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(),