Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .husky/pre-commit
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
set -e
bunx lint-staged
bun run typecheck
bun run test
44 changes: 42 additions & 2 deletions bun.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion package.json
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@
"husky": "^9.1.7",
"lint-staged": "^17.3.0",
"prettier": "^3.9.6",
"typescript": "^5"
"typescript": "^7"
},
"keywords": [
"oh-my-pi",
Expand Down
10 changes: 6 additions & 4 deletions src/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,14 @@ import type { ModelAxis, PromptWeight } from "./model-axis.js";
import { renderTuningBlock, TUNING_HEADER } from "./render.js";
import { readTuningSettings } from "./settings.js";

/** `configHome` is the config root passed to `readTuningSettings`; omp's loader
* calls this with `pi` alone, so the default keeps the production path intact.
* Tests pass a temp directory to isolate the lockfile read. */
export default async function dynamicSystemPromptPlugin(
pi: ExtensionAPI,
configHome?: string,
): Promise<void> {
const settings = await readTuningSettings();
const settings = await readTuningSettings(configHome);

if (!settings.enabled || settings.weight === "off") return;
// After this guard, weight is "auto" | PromptWeight
Expand All @@ -25,9 +29,7 @@ export default async function dynamicSystemPromptPlugin(

const base = resolveModelAxis(model.id);
const axis: ModelAxis =
requestedWeight === "auto"
? base
: { ...base, weight: requestedWeight as PromptWeight };
requestedWeight === "auto" ? base : { ...base, weight: requestedWeight };

if (settings.showHeader && ctx.hasUI && axis.profile !== lastProfile) {
lastProfile = axis.profile;
Expand Down
184 changes: 100 additions & 84 deletions src/model-axis.ts
Original file line number Diff line number Diff line change
@@ -1,17 +1,15 @@
import {
bareModelId,
parseAnthropicModel,
parseGeminiModel,
parseGlmModel,
parseOpenAIModel,
} from "@oh-my-pi/pi-catalog/identity";
import {
isClaudeModelId,
isDeepseekModelIdOrName,
isKimiModelId,
isMinimaxM2FamilyModelId,
isMinimaxM3FamilyModelId,
isQwenModelId,
parseAnthropicModel,
parseGeminiModel,
parseGlmModel,
parseOpenAIModel,
} from "@oh-my-pi/pi-catalog/identity";

// ---------------------------------------------------------------------------
Expand Down Expand Up @@ -77,9 +75,27 @@ const VERSION_PATTERNS: Record<string, RegExp> = {
deepseek: /deepseek.*?[vr](\d+)(?:[.\-_](\d+))?/,
grok: /grok.*?(\d+)(?:[.\-_](\d+))?/,
qwen: /qwen.*?(\d+)(?:[.\-_](\d+))?/,
// Loose on purpose: "m<digit>" would match many ids, but this regex only runs
// inside the minimax branch, which is gated by isMinimaxM2/M3FamilyModelId, so
// non-minimax ids never reach it.
minimax: /m(\d+)(?:[.\-_](\d+))?/,
};

/** Combine a major digit string and optional minor digit string into a single
* version number, e.g. ("2", "7") -> 2.7. Shared by regex-based version
* extraction wherever a family lacks a structured catalog parser. */
function versionFromDigits(
majorStr: string,
minorStr: string | undefined,
): number | null {
const major = Number(majorStr);
if (Number.isNaN(major)) return null;
if (minorStr === undefined) return major;
const minor = Number(minorStr);
if (Number.isNaN(minor)) return major;
return major + minor / 10 ** minorStr.length;
}

function parseVersionFromPattern(
family: string,
normalized: string,
Expand All @@ -88,13 +104,7 @@ function parseVersionFromPattern(
if (!pattern) return null;
const m = pattern.exec(normalized);
if (!m) return null;
const major = Number(m[1]);
if (Number.isNaN(major)) return null;
if (m[2] === undefined) return major;
const minorStr = m[2];
const minor = Number(minorStr);
if (Number.isNaN(minor)) return major;
return major + minor / 10 ** minorStr.length;
return versionFromDigits(m[1], m[2]);
}

// ---------------------------------------------------------------------------
Expand All @@ -114,13 +124,21 @@ const SMALL_TOKENS = [
];
const MID_TOKENS = ["chat", "instruct", "plus", "next"];

// Invariant: tier tokens are separator-delimited in real catalog ids (e.g.
// "deepseek-v4-flash", "kimi-k2-turbo"), never embedded in a larger alphanumeric
// run. Match at a word boundary — start or one of [-._/] on each side — so "mini"
// does not fire inside "minimax", "air" does not fire inside a hypothetical "fair"
// sku, and "fast" does not fire inside "fastest". Tokens are lowercase letters
// only, so no regex escaping is needed.
function tokenBoundaryRegexes(tokens: readonly string[]): RegExp[] {
return tokens.map((t) => new RegExp(`(^|[-._/])${t}($|[-._/])`));
}
const SMALL_TOKEN_RE = tokenBoundaryRegexes(SMALL_TOKENS);
const MID_TOKEN_RE = tokenBoundaryRegexes(MID_TOKENS);

function tierFromTokens(normalized: string): Tier {
for (const tok of SMALL_TOKENS) {
if (normalized.includes(tok)) return "small";
}
for (const tok of MID_TOKENS) {
if (normalized.includes(tok)) return "mid";
}
if (SMALL_TOKEN_RE.some((re) => re.test(normalized))) return "small";
if (MID_TOKEN_RE.some((re) => re.test(normalized))) return "mid";
return "flagship";
}

Expand Down Expand Up @@ -175,6 +193,51 @@ function deriveWeight(
return (["light", "medium", "heavy"] as const)[idx];
}

// ---------------------------------------------------------------------------
// Family-matching helpers
// ---------------------------------------------------------------------------

/** Catalog family predicates are case/namespace sensitive in places, so a model
* may pass on the original id but not the normalized form or vice versa. Test
* both to avoid misclassifying a recognized family as "unknown". */
function matchesFamily(
pred: (id: string) => boolean,
id: string,
normalized: string,
): boolean {
return pred(id) || pred(normalized);
}

/** Assemble the derived fields (weight, profile) that every resolver branch
* computes from (family, version, tier). "unknown" has no version, so its
* profile drops that segment; every other family always has one. */
function buildAxis(
family: Family,
version: number | null,
tier: Tier,
): ModelAxis {
const weight = deriveWeight(family, version, tier);
const profile =
version !== null ? `${family}-${version}-${tier}` : `${family}-${tier}`;
return { family, version, tier, weight, profile };
}

/** pi-catalog's parseAnthropicModel omits the haiku kind, so haiku ids fall
* through the structured parser. Recover them with an isClaudeModelId gate and
* a haiku-version regex; returns null when the id is not a haiku Claude model. */
function resolveAnthropicHaikuFallback(
modelId: string,
normalized: string,
): ModelAxis | null {
if (!normalized.includes("haiku")) return null;
if (!matchesFamily(isClaudeModelId, modelId, normalized)) return null;
const m = /haiku[-._]?(\d+)(?:[.\-_](\d+))?/.exec(normalized);
if (!m) return null;
const version = versionFromDigits(m[1], m[2]);
if (version === null) return null;
return buildAxis("anthropic", version, "small");
}

// ---------------------------------------------------------------------------
// Resolver
// ---------------------------------------------------------------------------
Expand All @@ -193,9 +256,7 @@ export function resolveModelAxis(modelId: string): ModelAxis {
openai.variant === "codex-spark"
? "small"
: "flagship";
const weight = deriveWeight("openai", version, tier);
const profile = `${"openai"}-${version}-${tier}`;
return { family: "openai", version, tier, weight, profile };
return buildAxis("openai", version, tier);
}

const anthropic = parseAnthropicModel(normalized);
Expand All @@ -215,32 +276,13 @@ export function resolveModelAxis(modelId: string): ModelAxis {
// keep the branch for forward compatibility
tier = "small";
}
const weight = deriveWeight("anthropic", version, tier);
const profile = `${"anthropic"}-${version}-${tier}`;
return { family: "anthropic", version, tier, weight, profile };
return buildAxis("anthropic", version, tier);
}

// Fallback for Anthropic Haiku, which pi-catalog's parseAnthropicModel
// does not yet recognize (kind enum omits haiku). Use modelFamilyToken's
// fallback logic: isClaudeModelId.
if (
normalized.includes("haiku") &&
(isClaudeModelId(modelId) || isClaudeModelId(normalized))
) {
const m = /haiku[-._]?(\d+)(?:[.\-_](\d+))?/.exec(normalized);
if (m) {
const major = Number(m[1]);
const minorStr = m[2];
const version =
minorStr !== undefined
? major + Number(minorStr) / 10 ** minorStr.length
: major;
const tier: Tier = "small";
const weight = deriveWeight("anthropic", version, tier);
const profile = `anthropic-${version}-${tier}`;
return { family: "anthropic", version, tier, weight, profile };
}
}
// does not yet recognize (kind enum omits haiku).
const haiku = resolveAnthropicHaikuFallback(modelId, normalized);
if (haiku) return haiku;

const gemini = parseGeminiModel(normalized);
if (gemini) {
Expand All @@ -254,9 +296,7 @@ export function resolveModelAxis(modelId: string): ModelAxis {
} else {
tier = "mid";
}
const weight = deriveWeight("gemini", version, tier);
const profile = `${"gemini"}-${version}-${tier}`;
return { family: "gemini", version, tier, weight, profile };
return buildAxis("gemini", version, tier);
}

const glm = parseGlmModel(normalized);
Expand All @@ -271,75 +311,51 @@ export function resolveModelAxis(modelId: string): ModelAxis {
// flash, flashx, preview
tier = "small";
}
const weight = deriveWeight("glm", version, tier);
const profile = `${"glm"}-${version}-${tier}`;
return { family: "glm", version, tier, weight, profile };
return buildAxis("glm", version, tier);
}

// Boolean-only families — gate on the catalog predicate, then apply our own regex.
// Grok has no family predicate in the catalog (isGrokReasoningEffortCapable is a
// capability check, not a family check), so match a grok- prefix on the normalized id.

if (isKimiModelId(modelId) || isKimiModelId(normalized)) {
if (matchesFamily(isKimiModelId, modelId, normalized)) {
const version = parseVersionFromPattern("kimi", normalized);
if (version !== null) {
const tier = tierFromTokens(normalized);
const weight = deriveWeight("kimi", version, tier);
const profile = `kimi-${version}-${tier}`;
return { family: "kimi", version, tier, weight, profile };
return buildAxis("kimi", version, tierFromTokens(normalized));
}
}

if (isDeepseekModelIdOrName(modelId) || isDeepseekModelIdOrName(normalized)) {
if (matchesFamily(isDeepseekModelIdOrName, modelId, normalized)) {
const version = parseVersionFromPattern("deepseek", normalized);
if (version !== null) {
const tier = tierFromTokens(normalized);
const weight = deriveWeight("deepseek", version, tier);
const profile = `deepseek-${version}-${tier}`;
return { family: "deepseek", version, tier, weight, profile };
return buildAxis("deepseek", version, tierFromTokens(normalized));
}
}

if (normalized.startsWith("grok-") || normalized.includes("/grok-")) {
const version = parseVersionFromPattern("grok", normalized);
if (version !== null) {
const tier = tierFromTokens(normalized);
const weight = deriveWeight("grok", version, tier);
const profile = `grok-${version}-${tier}`;
return { family: "grok", version, tier, weight, profile };
return buildAxis("grok", version, tierFromTokens(normalized));
}
}

if (isQwenModelId(modelId) || isQwenModelId(normalized)) {
if (matchesFamily(isQwenModelId, modelId, normalized)) {
const version = parseVersionFromPattern("qwen", normalized);
if (version !== null) {
const tier = tierFromTokens(normalized);
const weight = deriveWeight("qwen", version, tier);
const profile = `qwen-${version}-${tier}`;
return { family: "qwen", version, tier, weight, profile };
return buildAxis("qwen", version, tierFromTokens(normalized));
}
}

if (
isMinimaxM2FamilyModelId(modelId) ||
isMinimaxM3FamilyModelId(modelId) ||
isMinimaxM2FamilyModelId(normalized) ||
isMinimaxM3FamilyModelId(normalized)
matchesFamily(isMinimaxM2FamilyModelId, modelId, normalized) ||
matchesFamily(isMinimaxM3FamilyModelId, modelId, normalized)
) {
const version = parseVersionFromPattern("minimax", normalized);
if (version !== null) {
const tier = tierFromTokens(normalized);
const weight = deriveWeight("minimax", version, tier);
const profile = `minimax-${version}-${tier}`;
return { family: "minimax", version, tier, weight, profile };
return buildAxis("minimax", version, tierFromTokens(normalized));
}
}

// Fallback: unknown family.
{
const tier: Tier = "flagship";
const weight = deriveWeight("unknown", null, tier);
const profile = `unknown-${tier}`;
return { family: "unknown", version: null, tier, weight, profile };
}
return buildAxis("unknown", null, "flagship");
}
Loading