mirror of
https://github.com/wu736139669/hapi.git
synced 2026-08-05 06:24:37 +00:00
unify model selection
This commit is contained in:
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
})
|
||||
|
||||
@@ -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')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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}`);
|
||||
|
||||
|
||||
@@ -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';
|
||||
|
||||
@@ -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';
|
||||
|
||||
@@ -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');
|
||||
});
|
||||
});
|
||||
@@ -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) => {
|
||||
|
||||
@@ -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: {
|
||||
|
||||
Reference in New Issue
Block a user