unify model selection

This commit is contained in:
weishu
2026-03-16 18:29:09 +08:00
parent 1f02bb7de1
commit 02c8e12e80
27 changed files with 454 additions and 41 deletions
+3 -1
View File
@@ -21,6 +21,7 @@ export type SessionBootstrapOptions = {
workingDirectory?: string
tag?: string
agentState?: AgentState | null
model?: string
}
export type SessionBootstrapResult = {
@@ -124,7 +125,8 @@ export async function bootstrapSession(options: SessionBootstrapOptions): Promis
const sessionInfo = await api.getOrCreateSession({
tag: sessionTag,
metadata,
state: agentState
state: agentState,
model: options.model
})
const session = api.sessionSyncClient(sessionInfo)
+4 -1
View File
@@ -18,13 +18,15 @@ export class ApiClient {
tag: string
metadata: Metadata
state: AgentState | null
model?: string
}): Promise<Session> {
const response = await axios.post<CreateSessionResponse>(
`${configuration.apiUrl}/cli/sessions`,
{
tag: opts.tag,
metadata: opts.metadata,
agentState: opts.state
agentState: opts.state,
model: opts.model
},
{
headers: {
@@ -69,6 +71,7 @@ export class ApiClient {
thinking: raw.thinking,
thinkingAt: raw.thinkingAt,
todos: raw.todos,
model: raw.model,
permissionMode: raw.permissionMode,
modelMode: raw.modelMode
}
+1
View File
@@ -96,6 +96,7 @@ export const CreateSessionResponseSchema = z.object({
thinking: z.boolean(),
thinkingAt: z.number(),
todos: TodosSchema.optional(),
model: z.string().optional(),
permissionMode: PermissionModeSchema.optional(),
modelMode: ModelModeSchema.optional()
})
+17 -1
View File
@@ -1,5 +1,5 @@
import { describe, expect, it } from 'vitest'
import { resolveClaudeSessionModelMode } from './modelMode'
import { resolveClaudePersistedModel, resolveClaudeSessionModelMode } from './modelMode'
describe('resolveClaudeSessionModelMode', () => {
it('returns default when model is missing', () => {
@@ -21,3 +21,19 @@ describe('resolveClaudeSessionModelMode', () => {
expect(resolveClaudeSessionModelMode('opus[1m]')).toBe('opus[1m]')
})
})
describe('resolveClaudePersistedModel', () => {
it('skips missing, auto, default, and representable mode names', () => {
expect(resolveClaudePersistedModel()).toBeUndefined()
expect(resolveClaudePersistedModel('')).toBeUndefined()
expect(resolveClaudePersistedModel('auto')).toBeUndefined()
expect(resolveClaudePersistedModel('default')).toBeUndefined()
expect(resolveClaudePersistedModel('sonnet')).toBeUndefined()
expect(resolveClaudePersistedModel('opus[1m]')).toBeUndefined()
})
it('persists unsupported custom Claude model strings', () => {
expect(resolveClaudePersistedModel('claude-3-7-sonnet-latest')).toBe('claude-3-7-sonnet-latest')
expect(resolveClaudePersistedModel(' claude-opus-4-1-20250805 ')).toBe('claude-opus-4-1-20250805')
})
})
+15 -3
View File
@@ -8,11 +8,23 @@ const CLAUDE_SESSION_MODEL_MODES = new Set<SessionModelMode>([
])
export function resolveClaudeSessionModelMode(model?: string): SessionModelMode {
if (!model) {
const trimmedModel = model?.trim()
if (!trimmedModel) {
return 'default'
}
return CLAUDE_SESSION_MODEL_MODES.has(model as SessionModelMode)
? model as SessionModelMode
return CLAUDE_SESSION_MODEL_MODES.has(trimmedModel as SessionModelMode)
? trimmedModel as SessionModelMode
: 'default'
}
export function resolveClaudePersistedModel(model?: string): string | undefined {
const trimmedModel = model?.trim()
if (!trimmedModel || trimmedModel === 'auto' || trimmedModel === 'default') {
return undefined
}
return resolveClaudeSessionModelMode(trimmedModel) === 'default'
? trimmedModel
: undefined
}
+3 -2
View File
@@ -17,7 +17,7 @@ import { createModeChangeHandler, createRunnerLifecycle, setControlledByUser } f
import { isModelModeAllowedForFlavor, isPermissionModeAllowedForFlavor } from '@hapi/protocol';
import { ModelModeSchema, PermissionModeSchema } from '@hapi/protocol/schemas';
import { formatMessageWithAttachments } from '@/utils/attachmentFormatter';
import { resolveClaudeSessionModelMode } from './modelMode';
import { resolveClaudePersistedModel, resolveClaudeSessionModelMode } from './modelMode';
export interface StartOptions {
model?: string
@@ -50,7 +50,8 @@ export async function runClaude(options: StartOptions = {}): Promise<void> {
flavor: 'claude',
startedBy,
workingDirectory,
agentState: initialState
agentState: initialState,
model: resolveClaudePersistedModel(options.model)
});
logger.debug(`Session created: ${sessionInfo.id}`);
+2 -1
View File
@@ -33,7 +33,8 @@ export async function runCodex(opts: {
flavor: 'codex',
startedBy,
workingDirectory,
agentState: state
agentState: state,
model: opts.model
});
const startingMode: 'local' | 'remote' = startedBy === 'runner' ? 'remote' : 'local';
+2 -1
View File
@@ -38,7 +38,8 @@ export async function runCursor(opts: {
flavor: 'cursor',
startedBy,
workingDirectory,
agentState: state
agentState: state,
model: opts.model
});
const startingMode: 'local' | 'remote' = startedBy === 'runner' ? 'remote' : 'local';
+109
View File
@@ -0,0 +1,109 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
const harness = vi.hoisted(() => ({
bootstrapArgs: [] as Array<Record<string, unknown>>,
geminiLoopArgs: [] as Array<Record<string, unknown>>,
session: {
onUserMessage: vi.fn(),
rpcHandlerManager: {
registerHandler: vi.fn()
}
}
}));
vi.mock('@/agent/sessionFactory', () => ({
bootstrapSession: vi.fn(async (options: Record<string, unknown>) => {
harness.bootstrapArgs.push(options);
return {
api: {},
session: harness.session
};
})
}));
vi.mock('./loop', () => ({
geminiLoop: vi.fn(async (options: Record<string, unknown>) => {
harness.geminiLoopArgs.push(options);
})
}));
vi.mock('@/claude/registerKillSessionHandler', () => ({
registerKillSessionHandler: vi.fn()
}));
vi.mock('@/agent/runnerLifecycle', () => ({
createModeChangeHandler: vi.fn(() => vi.fn()),
createRunnerLifecycle: vi.fn(() => ({
registerProcessHandlers: vi.fn(),
cleanupAndExit: vi.fn(async () => {}),
markCrash: vi.fn(),
setExitCode: vi.fn(),
setArchiveReason: vi.fn()
})),
setControlledByUser: vi.fn()
}));
vi.mock('@/claude/utils/startHookServer', () => ({
startHookServer: vi.fn(async () => ({
port: 1234,
token: 'token',
stop: vi.fn()
}))
}));
vi.mock('@/modules/common/hooks/generateHookSettings', () => ({
cleanupHookSettingsFile: vi.fn(),
generateHookSettingsFile: vi.fn(() => '/tmp/gemini-hooks.json')
}));
const resolveGeminiRuntimeConfigMock = vi.hoisted(() => vi.fn());
vi.mock('./utils/config', () => ({
resolveGeminiRuntimeConfig: resolveGeminiRuntimeConfigMock
}));
vi.mock('@/ui/logger', () => ({
logger: {
debug: vi.fn()
}
}));
vi.mock('@/utils/attachmentFormatter', () => ({
formatMessageWithAttachments: vi.fn((text: string) => text)
}));
import { runGemini } from './runGemini';
describe('runGemini', () => {
beforeEach(() => {
harness.bootstrapArgs.length = 0;
harness.geminiLoopArgs.length = 0;
harness.session.onUserMessage.mockReset();
harness.session.rpcHandlerManager.registerHandler.mockReset();
resolveGeminiRuntimeConfigMock.mockReset();
});
it('persists a resolved config model before bootstrapping the session', async () => {
resolveGeminiRuntimeConfigMock.mockReturnValue({
model: 'gemini-3-pro-preview',
modelSource: 'local'
});
await runGemini({});
expect(harness.bootstrapArgs[0]?.model).toBe('gemini-3-pro-preview');
expect(harness.geminiLoopArgs[0]?.model).toBe('gemini-3-pro-preview');
});
it('does not persist the hardcoded default fallback model', async () => {
resolveGeminiRuntimeConfigMock.mockReturnValue({
model: 'gemini-2.5-pro',
modelSource: 'default'
});
await runGemini({});
expect(harness.bootstrapArgs[0]?.model).toBeUndefined();
expect(harness.geminiLoopArgs[0]?.model).toBe('gemini-2.5-pro');
});
});
+8 -2
View File
@@ -35,11 +35,17 @@ export async function runGemini(opts: {
controlledByUser: false
};
const runtimeConfig = resolveGeminiRuntimeConfig({ model: opts.model });
const persistedModel = runtimeConfig.modelSource === 'default'
? undefined
: runtimeConfig.model;
const { api, session } = await bootstrapSession({
flavor: 'gemini',
startedBy,
workingDirectory,
agentState: initialState
agentState: initialState,
model: persistedModel
});
const startingMode: 'local' | 'remote' = opts.startingMode
@@ -54,7 +60,7 @@ export async function runGemini(opts: {
const sessionWrapperRef: { current: GeminiSession | null } = { current: null };
let currentPermissionMode: PermissionMode = opts.permissionMode ?? 'default';
const resolvedModel = resolveGeminiRuntimeConfig({ model: opts.model }).model;
const resolvedModel = runtimeConfig.model;
const hookServer = await startHookServer({
onSessionHook: (sessionId, data) => {
+17 -6
View File
@@ -13,6 +13,8 @@ export type GeminiLocalConfig = {
model?: string;
};
export type GeminiModelSource = 'explicit' | 'env' | 'local' | 'default';
const GEMINI_DIR = join(homedir(), '.gemini');
const SETTINGS_PATH = join(GEMINI_DIR, 'settings.json');
const CONFIG_PATH = join(GEMINI_DIR, 'config.json');
@@ -85,20 +87,29 @@ export function readGeminiLocalConfig(): GeminiLocalConfig {
export function resolveGeminiRuntimeConfig(opts: {
model?: string;
token?: string;
} = {}): { model: string; token?: string } {
} = {}): { model: string; token?: string; modelSource: GeminiModelSource } {
const local = readGeminiLocalConfig();
const model = opts.model
?? process.env[GEMINI_MODEL_ENV]
?? local.model
?? DEFAULT_GEMINI_MODEL;
let modelSource: GeminiModelSource = 'default';
let model = DEFAULT_GEMINI_MODEL;
if (opts.model) {
model = opts.model;
modelSource = 'explicit';
} else if (process.env[GEMINI_MODEL_ENV]) {
model = process.env[GEMINI_MODEL_ENV]!;
modelSource = 'env';
} else if (local.model) {
model = local.model;
modelSource = 'local';
}
const token = opts.token
?? process.env[GEMINI_API_KEY_ENV]
?? process.env[GOOGLE_API_KEY_ENV]
?? local.token;
return { model, token };
return { model, token, modelSource };
}
export function buildGeminiEnv(opts: {