fix: honor per-role thinking in modelRoles helpers

Fixes #186
This commit is contained in:
can1357
2026-03-09 14:16:37 +01:00
parent 36d2fed5d9
commit 4ed4cfb27c
19 changed files with 283 additions and 88 deletions
+3 -1
View File
@@ -245,7 +245,9 @@ Supported model roles:
- `default`, `smol`, `slow`, `plan`, `commit`
Role aliases like `pi/smol` expand through `settings.modelRoles`.
Role aliases like `pi/smol` expand through `settings.modelRoles`. Each role value can also append a thinking selector such as `:minimal`, `:low`, `:medium`, or `:high`.
If a role points at another role, the target model still inherits normally and any explicit suffix on the referring role wins for that role-specific use.
Related settings:
+4
View File
@@ -11,6 +11,10 @@
- Updated tool documentation to clarify that `path` parameter accepts files, directories, glob patterns, or comma/space-separated path lists
- Refactored path resolution logic in `find`, `grep`, `ast_grep`, and `ast_edit` tools to use unified multi-path handling
### Fixed
- Per-role `modelRoles` thinking selectors now propagate through commit/title helper model selection, legacy commit analysis, and agentic commit sessions while preserving default thinking inheritance when no role override is configured
## [13.10.1] - 2026-03-10
### Added
@@ -1,4 +1,4 @@
import { INTENT_FIELD } from "@oh-my-pi/pi-agent-core";
import { INTENT_FIELD, type ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, Model } from "@oh-my-pi/pi-ai";
import { Markdown } from "@oh-my-pi/pi-tui";
import chalk from "chalk";
@@ -20,6 +20,7 @@ export interface CommitAgentInput {
cwd: string;
git: ControlledGit;
model: Model<Api>;
thinkingLevel?: ThinkingLevel;
settings: Settings;
modelRegistry: ModelRegistry;
authStorage: AuthStorage;
@@ -61,6 +62,7 @@ export async function runCommitAgentSession(input: CommitAgentInput): Promise<Co
modelRegistry: input.modelRegistry,
settings: input.settings,
model: input.model,
thinkingLevel: input.thinkingLevel,
systemPrompt,
customTools: tools,
enableLsp: false,
@@ -48,7 +48,12 @@ export async function runAgenticCommit(args: CommitCommandArgs): Promise<void> {
const { model: primaryModel, apiKey: primaryApiKey } = primaryModelResult;
process.stdout.write(` └─ ${primaryModel.name}\n`);
const { model: agentModel } = await resolveSmolModel(settings, modelRegistry, primaryModel, primaryApiKey);
const { model: agentModel, thinkingLevel: agentThinkingLevel } = await resolveSmolModel(
settings,
modelRegistry,
primaryModel,
primaryApiKey,
);
if (stagedFiles.length === 0) {
process.stderr.write("No changes to commit.\n");
@@ -126,6 +131,7 @@ export async function runAgenticCommit(args: CommitCommandArgs): Promise<void> {
cwd,
git,
model: agentModel,
thinkingLevel: agentThinkingLevel,
settings,
modelRegistry,
authStorage,
@@ -1,3 +1,4 @@
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, AssistantMessage, Model } from "@oh-my-pi/pi-ai";
import { completeSimple, validateToolCall } from "@oh-my-pi/pi-ai";
import { Type } from "@sinclair/typebox";
@@ -5,6 +6,7 @@ import analysisSystemPrompt from "../../commit/prompts/analysis-system.md" with
import analysisUserPrompt from "../../commit/prompts/analysis-user.md" with { type: "text" };
import type { ChangelogCategory, ConventionalAnalysis } from "../../commit/types";
import { renderPromptTemplate } from "../../config/prompt-templates";
import { toReasoningEffort } from "../../thinking";
import { extractTextContent, extractToolCall, normalizeAnalysis, parseJsonPayload } from "../utils";
const ConventionalAnalysisTool = {
@@ -49,6 +51,7 @@ const ConventionalAnalysisTool = {
export interface ConventionalAnalysisInput {
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
contextFiles?: Array<{ path: string; content: string }>;
userContext?: string;
typesDescription?: string;
@@ -64,6 +67,7 @@ export interface ConventionalAnalysisInput {
export async function generateConventionalAnalysis({
model,
apiKey,
thinkingLevel,
contextFiles,
userContext,
typesDescription,
@@ -89,7 +93,7 @@ export async function generateConventionalAnalysis({
messages: [{ role: "user", content: prompt, timestamp: Date.now() }],
tools: [ConventionalAnalysisTool],
},
{ apiKey, maxTokens: 2400 },
{ apiKey, maxTokens: 2400, reasoning: toReasoningEffort(thinkingLevel) },
);
return parseAnalysisFromResponse(response);
@@ -1,3 +1,4 @@
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, AssistantMessage, Model } from "@oh-my-pi/pi-ai";
import { completeSimple, validateToolCall } from "@oh-my-pi/pi-ai";
import { Type } from "@sinclair/typebox";
@@ -5,6 +6,7 @@ import summarySystemPrompt from "../../commit/prompts/summary-system.md" with {
import summaryUserPrompt from "../../commit/prompts/summary-user.md" with { type: "text" };
import type { CommitSummary } from "../../commit/types";
import { renderPromptTemplate } from "../../config/prompt-templates";
import { toReasoningEffort } from "../../thinking";
import { extractTextContent, extractToolCall } from "../utils";
const SummaryTool = {
@@ -18,6 +20,7 @@ const SummaryTool = {
export interface SummaryInput {
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
commitType: string;
scope: string | null;
details: string[];
@@ -32,6 +35,7 @@ export interface SummaryInput {
export async function generateSummary({
model,
apiKey,
thinkingLevel,
commitType,
scope,
details,
@@ -53,7 +57,7 @@ export async function generateSummary({
messages: [{ role: "user", content: userPrompt, timestamp: Date.now() }],
tools: [SummaryTool],
},
{ apiKey, maxTokens: 200 },
{ apiKey, maxTokens: 200, reasoning: toReasoningEffort(thinkingLevel) },
);
return parseSummaryFromResponse(response, commitType, scope);
@@ -1,3 +1,4 @@
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, AssistantMessage, Model } from "@oh-my-pi/pi-ai";
import { completeSimple, validateToolCall } from "@oh-my-pi/pi-ai";
import { type TSchema, Type } from "@sinclair/typebox";
@@ -5,6 +6,7 @@ import changelogSystemPrompt from "../../commit/prompts/changelog-system.md" wit
import changelogUserPrompt from "../../commit/prompts/changelog-user.md" with { type: "text" };
import { CHANGELOG_CATEGORIES, type ChangelogCategory, type ChangelogGenerationResult } from "../../commit/types";
import { renderPromptTemplate } from "../../config/prompt-templates";
import { toReasoningEffort } from "../../thinking";
import { extractTextContent, extractToolCall, parseJsonPayload } from "../utils";
const changelogEntryProperties = CHANGELOG_CATEGORIES.reduce<Record<ChangelogCategory, TSchema>>(
@@ -28,6 +30,7 @@ export const changelogTool = {
export interface ChangelogPromptInput {
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
changelogPath: string;
isPackageChangelog: boolean;
existingEntries?: string;
@@ -38,6 +41,7 @@ export interface ChangelogPromptInput {
export async function generateChangelogEntries({
model,
apiKey,
thinkingLevel,
changelogPath,
isPackageChangelog,
existingEntries,
@@ -58,7 +62,7 @@ export async function generateChangelogEntries({
messages: [{ role: "user", content: prompt, timestamp: Date.now() }],
tools: [changelogTool],
},
{ apiKey, maxTokens: 1200 },
{ apiKey, maxTokens: 1200, reasoning: toReasoningEffort(thinkingLevel) },
);
const parsed = parseChangelogResponse(response);
@@ -1,4 +1,5 @@
import * as path from "node:path";
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, Model } from "@oh-my-pi/pi-ai";
import { logger } from "@oh-my-pi/pi-utils";
import type { ControlledGit } from "../../commit/git";
@@ -16,6 +17,7 @@ export interface ChangelogFlowInput {
cwd: string;
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
stagedFiles: string[];
dryRun: boolean;
maxDiffChars?: number;
@@ -42,6 +44,7 @@ export async function runChangelogFlow({
cwd,
model,
apiKey,
thinkingLevel,
stagedFiles,
dryRun,
maxDiffChars,
@@ -72,6 +75,7 @@ export async function runChangelogFlow({
const generated = await generateChangelogEntries({
model,
apiKey,
thinkingLevel,
changelogPath: boundary.changelogPath,
isPackageChangelog,
existingEntries: existingEntries || undefined,
@@ -1,3 +1,4 @@
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, Model } from "@oh-my-pi/pi-ai";
import { $env } from "@oh-my-pi/pi-utils";
import { parseFileDiffs } from "../../commit/git/diff";
@@ -21,8 +22,10 @@ export interface MapReduceSettings {
export interface MapReduceInput {
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
smolModel: Model<Api>;
smolApiKey: string;
smolThinkingLevel?: ThinkingLevel;
diff: string;
stat: string;
scopeCandidates: string;
@@ -50,12 +53,14 @@ export async function runMapReduceAnalysis(input: MapReduceInput): Promise<Conve
const observations = await runMapPhase({
model: input.smolModel,
apiKey: input.smolApiKey,
thinkingLevel: input.smolThinkingLevel,
files: fileDiffs,
config: input.settings,
});
return runReducePhase({
model: input.model,
apiKey: input.apiKey,
thinkingLevel: input.thinkingLevel,
observations,
stat: input.stat,
scopeCandidates: input.scopeCandidates,
@@ -1,3 +1,4 @@
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, AssistantMessage, Message, Model } from "@oh-my-pi/pi-ai";
import { completeSimple } from "@oh-my-pi/pi-ai";
import fileObserverSystemPrompt from "../../commit/prompts/file-observer-system.md" with { type: "text" };
@@ -5,6 +6,7 @@ import fileObserverUserPrompt from "../../commit/prompts/file-observer-user.md"
import type { FileDiff, FileObservation } from "../../commit/types";
import { isExcludedFile } from "../../commit/utils/exclusions";
import { renderPromptTemplate } from "../../config/prompt-templates";
import { toReasoningEffort } from "../../thinking";
import { truncateToTokenLimit } from "./utils";
const MAX_FILE_TOKENS = 50_000;
@@ -17,6 +19,7 @@ const RETRY_BACKOFF_MS = 1000;
export interface MapPhaseInput {
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
files: FileDiff[];
config?: {
maxFileTokens?: number;
@@ -27,7 +30,7 @@ export interface MapPhaseInput {
};
}
export async function runMapPhase({ model, apiKey, files, config }: MapPhaseInput): Promise<FileObservation[]> {
export async function runMapPhase({ model, apiKey, thinkingLevel, files, config }: MapPhaseInput): Promise<FileObservation[]> {
const filtered = files.filter(file => !isExcludedFile(file.filename));
const systemPrompt = renderPromptTemplate(fileObserverSystemPrompt);
const maxFileTokens = config?.maxFileTokens ?? MAX_FILE_TOKENS;
@@ -58,7 +61,7 @@ export async function runMapPhase({ model, apiKey, files, config }: MapPhaseInpu
};
const response = await withRetry(
() => completeSimple(model, request, { apiKey, maxTokens: 400, signal: AbortSignal.timeout(timeoutMs) }),
() => completeSimple(model, request, { apiKey, maxTokens: 400, reasoning: toReasoningEffort(thinkingLevel), signal: AbortSignal.timeout(timeoutMs) }),
maxRetries,
retryBackoffMs,
);
@@ -1,3 +1,4 @@
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, AssistantMessage, Model } from "@oh-my-pi/pi-ai";
import { completeSimple, validateToolCall } from "@oh-my-pi/pi-ai";
import { Type } from "@sinclair/typebox";
@@ -5,6 +6,7 @@ import reduceSystemPrompt from "../../commit/prompts/reduce-system.md" with { ty
import reduceUserPrompt from "../../commit/prompts/reduce-user.md" with { type: "text" };
import type { ChangelogCategory, ConventionalAnalysis, FileObservation } from "../../commit/types";
import { renderPromptTemplate } from "../../config/prompt-templates";
import { toReasoningEffort } from "../../thinking";
import { extractTextContent, extractToolCall, normalizeAnalysis, parseJsonPayload } from "../utils";
const ReduceTool = {
@@ -49,6 +51,7 @@ const ReduceTool = {
export interface ReducePhaseInput {
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
observations: FileObservation[];
stat: string;
scopeCandidates: string;
@@ -58,6 +61,7 @@ export interface ReducePhaseInput {
export async function runReducePhase({
model,
apiKey,
thinkingLevel,
observations,
stat,
scopeCandidates,
@@ -76,7 +80,7 @@ export async function runReducePhase({
messages: [{ role: "user", content: prompt, timestamp: Date.now() }],
tools: [ReduceTool],
},
{ apiKey, maxTokens: 2400 },
{ apiKey, maxTokens: 2400, reasoning: toReasoningEffort(thinkingLevel) },
);
return parseAnalysisResponse(response);
@@ -1,14 +1,34 @@
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 {
expandRoleAlias,
parseModelPattern,
resolveModelFromSettings,
resolveModelFromString,
} from "../config/model-resolver";
import { expandRoleAlias, parseModelPattern, resolveModelRoleValue } from "../config/model-resolver";
import type { Settings } from "../config/settings";
import MODEL_PRIO from "../priority.json" with { type: "json" };
export interface ResolvedCommitModel {
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
}
function resolveRoleSelection(
roles: readonly string[],
settings: Settings,
availableModels: Model<Api>[],
): { 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,
});
if (resolved.model) {
return { model: resolved.model, thinkingLevel: resolved.thinkingLevel };
}
}
return undefined;
}
export async function resolvePrimaryModel(
override: string | undefined,
settings: Settings,
@@ -16,18 +36,13 @@ export async function resolvePrimaryModel(
getAvailable: () => Model<Api>[];
getApiKey: (model: Model<Api>) => Promise<string | undefined>;
},
): Promise<{ model: Model<Api>; apiKey: string }> {
): Promise<ResolvedCommitModel> {
const available = modelRegistry.getAvailable();
const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() };
const roleOrder = ["commit", "smol", ...MODEL_ROLE_IDS] as const;
const model = override
? resolveModelFromString(expandRoleAlias(override, settings), available, matchPreferences)
: resolveModelFromSettings({
settings,
availableModels: available,
matchPreferences,
roleOrder,
});
const resolved = override
? resolveModelRoleValue(override, available, { settings, matchPreferences })
: resolveRoleSelection(["commit", "smol", ...MODEL_ROLE_IDS], settings, available);
const model = resolved?.model;
if (!model) {
throw new Error("No model available for commit generation");
}
@@ -35,7 +50,7 @@ export async function resolvePrimaryModel(
if (!apiKey) {
throw new Error(`No API key available for model ${model.provider}/${model.id}`);
}
return { model, apiKey };
return { model, apiKey, thinkingLevel: resolved?.thinkingLevel };
}
export async function resolveSmolModel(
@@ -46,18 +61,15 @@ export async function resolveSmolModel(
},
fallbackModel: Model<Api>,
fallbackApiKey: string,
): Promise<{ model: Model<Api>; apiKey: string }> {
): Promise<ResolvedCommitModel> {
const available = modelRegistry.getAvailable();
const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() };
const role = settings.getModelRole("smol");
const roleModel = role
? resolveModelFromString(expandRoleAlias(role, settings), available, matchPreferences)
: undefined;
if (roleModel) {
const apiKey = await modelRegistry.getApiKey(roleModel);
if (apiKey) return { model: roleModel, apiKey };
const resolvedSmol = resolveRoleSelection(["smol"], settings, available);
if (resolvedSmol?.model) {
const apiKey = await modelRegistry.getApiKey(resolvedSmol.model);
if (apiKey) return { model: resolvedSmol.model, apiKey, thinkingLevel: resolvedSmol.thinkingLevel };
}
const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() };
for (const pattern of MODEL_PRIO.smol) {
const candidate = parseModelPattern(pattern, available, matchPreferences).model;
if (!candidate) continue;
+14 -2
View File
@@ -1,4 +1,5 @@
import * as path from "node:path";
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, Model } from "@oh-my-pi/pi-ai";
import { getProjectDir, logger } from "@oh-my-pi/pi-utils";
import { ModelRegistry } from "../config/model-registry";
@@ -45,12 +46,12 @@ async function runLegacyCommitCommand(args: CommitCommandArgs): Promise<void> {
const modelRegistry = new ModelRegistry(authStorage);
await modelRegistry.refresh();
const { model: primaryModel, apiKey: primaryApiKey } = await resolvePrimaryModel(
const { model: primaryModel, apiKey: primaryApiKey, thinkingLevel: primaryThinkingLevel } = await resolvePrimaryModel(
args.model,
settings,
modelRegistry,
);
const { model: smolModel, apiKey: smolApiKey } = await resolveSmolModel(
const { model: smolModel, apiKey: smolApiKey, thinkingLevel: smolThinkingLevel } = await resolveSmolModel(
settings,
modelRegistry,
primaryModel,
@@ -75,6 +76,7 @@ async function runLegacyCommitCommand(args: CommitCommandArgs): Promise<void> {
cwd,
model: primaryModel,
apiKey: primaryApiKey,
thinkingLevel: primaryThinkingLevel,
stagedFiles,
dryRun: args.dryRun,
maxDiffChars: commitSettings.changelogMaxDiffChars,
@@ -101,8 +103,10 @@ async function runLegacyCommitCommand(args: CommitCommandArgs): Promise<void> {
userContext: args.context,
primaryModel,
primaryApiKey,
primaryThinkingLevel,
smolModel,
smolApiKey,
smolThinkingLevel,
commitSettings,
});
@@ -116,6 +120,7 @@ async function runLegacyCommitCommand(args: CommitCommandArgs): Promise<void> {
stat,
model: primaryModel,
apiKey: primaryApiKey,
thinkingLevel: primaryThinkingLevel,
userContext: args.context,
});
@@ -144,8 +149,10 @@ async function generateAnalysis(input: {
userContext?: string;
primaryModel: Model<Api>;
primaryApiKey: string;
primaryThinkingLevel?: ThinkingLevel;
smolModel: Model<Api>;
smolApiKey: string;
smolThinkingLevel?: ThinkingLevel;
commitSettings: {
mapReduceEnabled: boolean;
mapReduceMinFiles: number;
@@ -166,8 +173,10 @@ async function generateAnalysis(input: {
return runMapReduceAnalysis({
model: input.primaryModel,
apiKey: input.primaryApiKey,
thinkingLevel: input.primaryThinkingLevel,
smolModel: input.smolModel,
smolApiKey: input.smolApiKey,
smolThinkingLevel: input.smolThinkingLevel,
diff: input.diff,
stat: input.stat,
scopeCandidates: input.scopeCandidates,
@@ -185,6 +194,7 @@ async function generateAnalysis(input: {
return generateConventionalAnalysis({
model: input.primaryModel,
apiKey: input.primaryApiKey,
thinkingLevel: input.primaryThinkingLevel,
contextFiles: input.contextFiles,
userContext: input.userContext,
typesDescription: TYPES_DESCRIPTION,
@@ -200,6 +210,7 @@ async function generateSummaryWithRetry(input: {
stat: string;
model: Model<Api>;
apiKey: string;
thinkingLevel?: ThinkingLevel;
userContext?: string;
}): Promise<{ summary: string }> {
let context = input.userContext;
@@ -207,6 +218,7 @@ async function generateSummaryWithRetry(input: {
const result = await generateSummary({
model: input.model,
apiKey: input.apiKey,
thinkingLevel: input.thinkingLevel,
commitType: input.analysis.type,
scope: input.analysis.scope,
details: input.analysis.details.map(detail => detail.text),
@@ -318,8 +318,7 @@ export class InputController {
const hasUserMessages = this.ctx.session.messages.some((m: AgentMessage) => m.role === "user");
if (!hasUserMessages && !this.ctx.sessionManager.getSessionName() && !$env.PI_NO_TITLE) {
const registry = this.ctx.session.modelRegistry;
const smolModel = this.ctx.settings.getModelRole("smol");
generateSessionTitle(text, registry, smolModel, this.ctx.session.sessionId)
generateSessionTitle(text, registry, this.ctx.settings, this.ctx.session.sessionId)
.then(async title => {
if (title) {
await this.ctx.sessionManager.setSessionName(title);
+2 -4
View File
@@ -834,11 +834,10 @@ export class TaskTool implements AgentTool<TaskSchema, TaskToolDetails, Theme> {
const commitMsg =
commitStyle === "ai" && this.session.modelRegistry
? async (diff: string) => {
const smolModel = this.session.settings.getModelRole("smol");
return generateCommitMessage(
diff,
this.session.modelRegistry!,
smolModel,
this.session.settings,
this.session.getSessionId?.() ?? undefined,
);
}
@@ -1081,11 +1080,10 @@ export class TaskTool implements AgentTool<TaskSchema, TaskToolDetails, Theme> {
const commitMsg =
commitStyle === "ai" && this.session.modelRegistry
? async (diff: string) => {
const smolModel = this.session.settings.getModelRole("smol");
return generateCommitMessage(
diff,
this.session.modelRegistry!,
smolModel,
this.session.settings,
this.session.getSessionId?.() ?? undefined,
);
}
@@ -2,12 +2,15 @@
* Generate commit messages from diffs using a smol, fast model.
* Follows the same pattern as title-generator.ts.
*/
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, Model } from "@oh-my-pi/pi-ai";
import { completeSimple } from "@oh-my-pi/pi-ai";
import { logger } 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 { renderPromptTemplate } from "../config/prompt-templates";
import { toReasoningEffort } from "../thinking";
import type { Settings } from "../config/settings";
import MODEL_PRIO from "../priority.json" with { type: "json" };
import commitSystemPrompt from "../prompts/system/commit-message-system.md" with { type: "text" };
@@ -32,24 +35,26 @@ function filterDiffNoise(diff: string): string {
return filtered.join("\n");
}
function getSmolModelCandidates(registry: ModelRegistry, savedSmolModel?: string): Model<Api>[] {
function getSmolModelCandidates(
registry: ModelRegistry,
settings: Settings,
): Array<{ model: Model<Api>; thinkingLevel?: ThinkingLevel }> {
const availableModels = registry.getAvailable();
if (availableModels.length === 0) return [];
const candidates: Model<Api>[] = [];
const addCandidate = (model?: Model<Api>): void => {
const candidates: Array<{ model: Model<Api>; thinkingLevel?: ThinkingLevel }> = [];
const addCandidate = (model?: Model<Api>, thinkingLevel?: ThinkingLevel): void => {
if (!model) return;
if (candidates.some(c => c.provider === model.provider && c.id === model.id)) return;
candidates.push(model);
if (candidates.some(c => c.model.provider === model.provider && c.model.id === model.id)) return;
candidates.push({ model, thinkingLevel });
};
if (savedSmolModel) {
const parsed = parseModelString(savedSmolModel);
if (parsed) {
const match = availableModels.find(m => m.provider === parsed.provider && m.id === parsed.id);
addCandidate(match);
}
}
const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() };
const configuredSmol = resolveModelRoleValue(settings.getModelRole("smol"), availableModels, {
settings,
matchPreferences,
});
addCandidate(configuredSmol.model, configuredSmol.thinkingLevel);
for (const pattern of MODEL_PRIO.smol) {
const needle = pattern.toLowerCase();
@@ -71,10 +76,10 @@ function getSmolModelCandidates(registry: ModelRegistry, savedSmolModel?: string
export async function generateCommitMessage(
diff: string,
registry: ModelRegistry,
savedSmolModel?: string,
settings: Settings,
sessionId?: string,
): Promise<string | null> {
const candidates = getSmolModelCandidates(registry, savedSmolModel);
const candidates = getSmolModelCandidates(registry, settings);
if (candidates.length === 0) {
logger.debug("commit-msg-generator: no smol model found");
return null;
@@ -89,22 +94,22 @@ export async function generateCommitMessage(
}
const userMessage = `<diff>\n${truncatedDiff}\n</diff>`;
for (const model of candidates) {
const apiKey = await registry.getApiKey(model, sessionId);
for (const candidate of candidates) {
const apiKey = await registry.getApiKey(candidate.model, sessionId);
if (!apiKey) continue;
try {
const response = await completeSimple(
model,
candidate.model,
{
systemPrompt: COMMIT_SYSTEM_PROMPT,
messages: [{ role: "user", content: userMessage, timestamp: Date.now() }],
},
{ apiKey, maxTokens: 60 },
{ apiKey, maxTokens: 60, reasoning: toReasoningEffort(candidate.thinkingLevel) },
);
if (response.stopReason === "error") {
logger.debug("commit-msg-generator: error", { model: model.id, error: response.errorMessage });
logger.debug("commit-msg-generator: error", { model: candidate.model.id, error: response.errorMessage });
continue;
}
@@ -118,11 +123,11 @@ export async function generateCommitMessage(
// Clean up: remove wrapping quotes, backticks, trailing period
msg = msg.replace(/^[`"']|[`"']$/g, "").replace(/\.$/, "");
logger.debug("commit-msg-generator: generated", { model: model.id, msg });
logger.debug("commit-msg-generator: generated", { model: candidate.model.id, msg });
return msg;
} catch (err) {
logger.debug("commit-msg-generator: error", {
model: model.id,
model: candidate.model.id,
error: err instanceof Error ? err.message : String(err),
});
}
@@ -1,12 +1,15 @@
/**
* Generate session titles using a smol, fast model.
*/
import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core";
import type { Api, Model } from "@oh-my-pi/pi-ai";
import { completeSimple } from "@oh-my-pi/pi-ai";
import { logger } 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 { renderPromptTemplate } from "../config/prompt-templates";
import { toReasoningEffort } from "../thinking";
import type { Settings } from "../config/settings";
import MODEL_PRIO from "../priority.json" with { type: "json" };
import titleSystemPrompt from "../prompts/system/title-system.md" with { type: "text" };
@@ -14,26 +17,28 @@ const TITLE_SYSTEM_PROMPT = renderPromptTemplate(titleSystemPrompt);
const MAX_INPUT_CHARS = 2000;
function getTitleModelCandidates(registry: ModelRegistry, savedSmolModel?: string): Model<Api>[] {
function getTitleModelCandidates(
registry: ModelRegistry,
settings: Settings,
): Array<{ model: Model<Api>; thinkingLevel?: ThinkingLevel }> {
const availableModels = registry.getAvailable();
if (availableModels.length === 0) return [];
const candidates: Model<Api>[] = [];
const addCandidate = (model?: Model<Api>): void => {
const candidates: Array<{ model: Model<Api>; thinkingLevel?: ThinkingLevel }> = [];
const addCandidate = (model?: Model<Api>, thinkingLevel?: ThinkingLevel): void => {
if (!model) return;
const exists = candidates.some(candidate => candidate.provider === model.provider && candidate.id === model.id);
const exists = candidates.some(candidate => candidate.model.provider === model.provider && candidate.model.id === model.id);
if (!exists) {
candidates.push(model);
candidates.push({ model, thinkingLevel });
}
};
if (savedSmolModel) {
const parsed = parseModelString(savedSmolModel);
if (parsed) {
const match = availableModels.find(model => model.provider === parsed.provider && model.id === parsed.id);
addCandidate(match);
}
}
const matchPreferences = { usageOrder: settings.getStorage()?.getModelUsageOrder() };
const configuredSmol = resolveModelRoleValue(settings.getModelRole("smol"), availableModels, {
settings,
matchPreferences,
});
addCandidate(configuredSmol.model, configuredSmol.thinkingLevel);
for (const pattern of MODEL_PRIO.smol) {
const needle = pattern.toLowerCase();
@@ -56,16 +61,16 @@ function getTitleModelCandidates(registry: ModelRegistry, savedSmolModel?: strin
*
* @param firstMessage The first user message
* @param registry Model registry
* @param savedSmolModel Optional saved smol model from settings (provider/modelId format)
* @param settings Settings used to resolve the smol role, including per-role thinking
* @param sessionId Optional session id for sticky API key selection
*/
export async function generateSessionTitle(
firstMessage: string,
registry: ModelRegistry,
savedSmolModel?: string,
settings: Settings,
sessionId?: string,
): Promise<string | null> {
const candidates = getTitleModelCandidates(registry, savedSmolModel);
const candidates = getTitleModelCandidates(registry, settings);
if (candidates.length === 0) {
logger.debug("title-generator: no smol model found");
return null;
@@ -74,17 +79,19 @@ export async function generateSessionTitle(
// Truncate message if too long
const truncatedMessage =
firstMessage.length > MAX_INPUT_CHARS ? `${firstMessage.slice(0, MAX_INPUT_CHARS)}…` : firstMessage;
const userMessage = `<user-message>\n${truncatedMessage}\n</user-message>`;
const userMessage = `<user-message>
${truncatedMessage}
</user-message>`;
for (const model of candidates) {
const apiKey = await registry.getApiKey(model, sessionId);
for (const candidate of candidates) {
const apiKey = await registry.getApiKey(candidate.model, sessionId);
if (!apiKey) {
logger.debug("title-generator: no API key for model", { provider: model.provider, id: model.id });
logger.debug("title-generator: no API key for model", { provider: candidate.model.provider, id: candidate.model.id });
continue;
}
const request = {
model: `${model.provider}/${model.id}`,
model: `${candidate.model.provider}/${candidate.model.id}`,
systemPrompt: TITLE_SYSTEM_PROMPT,
userMessage,
maxTokens: 30,
@@ -93,7 +100,7 @@ export async function generateSessionTitle(
try {
const response = await completeSimple(
model,
candidate.model,
{
systemPrompt: request.systemPrompt,
messages: [{ role: "user", content: request.userMessage, timestamp: Date.now() }],
@@ -101,6 +108,7 @@ export async function generateSessionTitle(
{
apiKey,
maxTokens: 30,
reasoning: toReasoningEffort(candidate.thinkingLevel),
},
);
@@ -151,5 +159,5 @@ export async function generateSessionTitle(
*/
export function setTerminalTitle(title: string): void {
// OSC 2 sets the window title
process.stdout.write(`\x1b]2;${title}\x07`);
process.stdout.write(`]2;${title}`);
}
@@ -0,0 +1,51 @@
import { describe, expect, it } from "bun:test";
import { Effort, getBundledModel } from "@oh-my-pi/pi-ai";
import { resolvePrimaryModel, resolveSmolModel } from "../src/commit/model-selection";
function getModelOrThrow(id: string) {
const model = getBundledModel("anthropic", id);
if (!model) throw new Error(`Expected model ${id}`);
return model;
}
function createSettings(modelRoles: Record<string, string>) {
return {
getModelRole(role: string) {
return modelRoles[role];
},
getStorage() {
return undefined;
},
setModelRole(role: string, value: string) {
modelRoles[role] = value;
},
get(path: string) {
if (path === "modelRoles") return modelRoles;
return undefined;
},
} as never;
}
describe("commit role thinking selection", () => {
it("returns explicit thinking for commit and smol roles, including alias overrides", async () => {
const defaultModel = getModelOrThrow("claude-sonnet-4-5");
const commitModel = getModelOrThrow("claude-opus-4-5");
const settings = createSettings({
default: `${defaultModel.provider}/${defaultModel.id}:high`,
commit: `${commitModel.provider}/${commitModel.id}:low`,
smol: "pi/default:minimal",
});
const registry = {
getAvailable: () => [defaultModel, commitModel],
getApiKey: async () => "test-key",
};
const primary = await resolvePrimaryModel(undefined, settings, registry);
expect(primary.model.id).toBe(commitModel.id);
expect(primary.thinkingLevel).toBe(Effort.Low);
const smol = await resolveSmolModel(settings, registry, commitModel, "fallback-key");
expect(smol.model.id).toBe(defaultModel.id);
expect(smol.thinkingLevel).toBe(Effort.Minimal);
});
});
@@ -0,0 +1,68 @@
import { afterEach, describe, expect, it, vi } from "bun:test";
import * as ai from "@oh-my-pi/pi-ai";
import { Effort, getBundledModel } from "@oh-my-pi/pi-ai";
import { generateCommitMessage } from "../src/utils/commit-message-generator";
import { generateSessionTitle } from "../src/utils/title-generator";
function getModelOrThrow(id: string) {
const model = getBundledModel("anthropic", id);
if (!model) throw new Error(`Expected model ${id}`);
return model;
}
function createSettings(modelRoles: Record<string, string>) {
return {
getModelRole(role: string) {
return modelRoles[role];
},
getStorage() {
return undefined;
},
} as never;
}
afterEach(() => {
vi.restoreAllMocks();
});
describe("role thinking helper propagation", () => {
it("passes smol-role thinking to commit message generation", async () => {
const model = getModelOrThrow("claude-sonnet-4-5");
const settings = createSettings({
default: `${model.provider}/${model.id}:high`,
smol: "pi/default:minimal",
});
const registry = {
getAvailable: () => [model],
getApiKey: async () => "test-key",
};
const completeSimpleMock = vi.spyOn(ai, "completeSimple").mockResolvedValue({
stopReason: "end_turn",
content: [{ type: "text", text: "fix scope handling" }],
} as never);
const message = await generateCommitMessage(`diff --git a/x b/x\n+change\n`, registry as never, settings);
expect(message).toBe("fix scope handling");
expect(completeSimpleMock.mock.calls[0]?.[2]).toMatchObject({ reasoning: Effort.Minimal });
});
it("passes smol-role thinking to title generation", async () => {
const model = getModelOrThrow("claude-sonnet-4-5");
const settings = createSettings({
default: `${model.provider}/${model.id}:high`,
smol: "pi/default:low",
});
const registry = {
getAvailable: () => [model],
getApiKey: async () => "test-key",
};
const completeSimpleMock = vi.spyOn(ai, "completeSimple").mockResolvedValue({
stopReason: "end_turn",
content: [{ type: "text", text: "Investigate resolver" }],
} as never);
const title = await generateSessionTitle("Investigate resolver", registry as never, settings);
expect(title).toBe("Investigate resolver");
expect(completeSimpleMock.mock.calls[0]?.[2]).toMatchObject({ reasoning: Effort.Low });
});
});