mirror of
https://github.com/wu736139669/hapi.git
synced 2026-08-05 06:24:37 +00:00
feat(gemini): support mid-session model change (#379)
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
import { spawn, type ChildProcessWithoutNullStreams } from 'node:child_process';
|
||||
import { logger } from '@/ui/logger';
|
||||
import { killProcessByChildProcess } from '@/utils/process';
|
||||
import { GEMINI_MODEL_PRESETS } from '@hapi/protocol';
|
||||
|
||||
interface JsonRpcRequest {
|
||||
jsonrpc: '2.0';
|
||||
@@ -303,7 +304,7 @@ export class AcpStdioTransport {
|
||||
if (lowerText.includes('status 404') || lowerText.includes('model not found') || lowerText.includes('not_found')) {
|
||||
this.stderrErrorHandler({
|
||||
type: 'model_not_found',
|
||||
message: 'Model not found. Available models: gemini-3.1-pro-preview, gemini-3-flash-preview, gemini-2.5-pro, gemini-2.5-flash, gemini-2.5-flash-lite',
|
||||
message: `Model not found. Available models: ${GEMINI_MODEL_PRESETS.join(', ')}`,
|
||||
raw: text
|
||||
});
|
||||
return;
|
||||
|
||||
@@ -46,17 +46,22 @@ export async function geminiLoop(opts: GeminiLoopOptions): Promise<void> {
|
||||
session.onSessionFound(opts.resumeSessionId);
|
||||
}
|
||||
|
||||
const getCurrentModel = (): string | undefined => {
|
||||
const sessionModel = session.getModel();
|
||||
return sessionModel != null ? sessionModel : opts.model;
|
||||
};
|
||||
|
||||
await runLocalRemoteSession({
|
||||
session,
|
||||
startingMode: opts.startingMode,
|
||||
logTag: 'gemini-loop',
|
||||
runLocal: (instance) => geminiLocalLauncher(instance, {
|
||||
model: opts.model,
|
||||
model: getCurrentModel(),
|
||||
allowedTools: opts.allowedTools,
|
||||
hookSettingsPath: opts.hookSettingsPath
|
||||
}),
|
||||
runRemote: (instance) => geminiRemoteLauncher(instance, {
|
||||
model: opts.model,
|
||||
model: getCurrentModel(),
|
||||
hookSettingsPath: opts.hookSettingsPath
|
||||
}),
|
||||
onSessionReady: opts.onSessionReady
|
||||
|
||||
@@ -1,5 +1,11 @@
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
|
||||
const mockGeminiSession = vi.hoisted(() => ({
|
||||
setModel: vi.fn(),
|
||||
setPermissionMode: vi.fn(),
|
||||
stopKeepAlive: vi.fn()
|
||||
}));
|
||||
|
||||
const harness = vi.hoisted(() => ({
|
||||
bootstrapArgs: [] as Array<Record<string, unknown>>,
|
||||
geminiLoopArgs: [] as Array<Record<string, unknown>>,
|
||||
@@ -24,6 +30,10 @@ vi.mock('@/agent/sessionFactory', () => ({
|
||||
vi.mock('./loop', () => ({
|
||||
geminiLoop: vi.fn(async (options: Record<string, unknown>) => {
|
||||
harness.geminiLoopArgs.push(options);
|
||||
const onSessionReady = options.onSessionReady as ((session: unknown) => void) | undefined;
|
||||
if (onSessionReady) {
|
||||
onSessionReady(mockGeminiSession);
|
||||
}
|
||||
})
|
||||
}));
|
||||
|
||||
@@ -78,6 +88,8 @@ describe('runGemini', () => {
|
||||
beforeEach(() => {
|
||||
harness.bootstrapArgs.length = 0;
|
||||
harness.geminiLoopArgs.length = 0;
|
||||
mockGeminiSession.setModel.mockReset();
|
||||
mockGeminiSession.setPermissionMode.mockReset();
|
||||
harness.session.onUserMessage.mockReset();
|
||||
harness.session.rpcHandlerManager.registerHandler.mockReset();
|
||||
resolveGeminiRuntimeConfigMock.mockReset();
|
||||
@@ -97,13 +109,125 @@ describe('runGemini', () => {
|
||||
|
||||
it('does not persist the hardcoded default fallback model', async () => {
|
||||
resolveGeminiRuntimeConfigMock.mockReturnValue({
|
||||
model: 'gemini-2.5-pro',
|
||||
model: 'gemini-3-flash-preview',
|
||||
modelSource: 'default'
|
||||
});
|
||||
|
||||
await runGemini({});
|
||||
|
||||
expect(harness.bootstrapArgs[0]?.model).toBeUndefined();
|
||||
expect(harness.geminiLoopArgs[0]?.model).toBe('gemini-3-flash-preview');
|
||||
});
|
||||
|
||||
it('applies model change via set-session-config RPC', async () => {
|
||||
resolveGeminiRuntimeConfigMock.mockReturnValue({
|
||||
model: 'gemini-3-flash-preview',
|
||||
modelSource: 'default'
|
||||
});
|
||||
|
||||
await runGemini({});
|
||||
|
||||
const registerCalls = harness.session.rpcHandlerManager.registerHandler.mock.calls;
|
||||
const configHandler = registerCalls.find(
|
||||
(call: unknown[]) => call[0] === 'set-session-config'
|
||||
);
|
||||
expect(configHandler).toBeDefined();
|
||||
|
||||
const handler = configHandler![1] as (payload: unknown) => Promise<unknown>;
|
||||
const result = await handler({ model: 'gemini-2.5-flash' }) as Record<string, unknown>;
|
||||
const applied = result.applied as Record<string, unknown>;
|
||||
expect(applied.model).toBe('gemini-2.5-flash');
|
||||
});
|
||||
|
||||
it('rejects invalid model in set-session-config RPC', async () => {
|
||||
resolveGeminiRuntimeConfigMock.mockReturnValue({
|
||||
model: 'gemini-3-flash-preview',
|
||||
modelSource: 'default'
|
||||
});
|
||||
|
||||
await runGemini({});
|
||||
|
||||
const registerCalls = harness.session.rpcHandlerManager.registerHandler.mock.calls;
|
||||
const configHandler = registerCalls.find(
|
||||
(call: unknown[]) => call[0] === 'set-session-config'
|
||||
);
|
||||
const handler = configHandler![1] as (payload: unknown) => Promise<unknown>;
|
||||
await expect(handler({ model: 123 })).rejects.toThrow();
|
||||
});
|
||||
|
||||
it('accepts null model (Auto) in set-session-config RPC', async () => {
|
||||
resolveGeminiRuntimeConfigMock.mockReturnValue({
|
||||
model: 'gemini-3-flash-preview',
|
||||
modelSource: 'default'
|
||||
});
|
||||
|
||||
await runGemini({});
|
||||
|
||||
const registerCalls = harness.session.rpcHandlerManager.registerHandler.mock.calls;
|
||||
const configHandler = registerCalls.find(
|
||||
(call: unknown[]) => call[0] === 'set-session-config'
|
||||
);
|
||||
const handler = configHandler![1] as (payload: unknown) => Promise<unknown>;
|
||||
const result = await handler({ model: null }) as Record<string, unknown>;
|
||||
const applied = result.applied as Record<string, unknown>;
|
||||
// null (Default) should be passed through to hub for DB clearing
|
||||
expect(applied.model).toBeNull();
|
||||
});
|
||||
|
||||
it('only includes changed fields in applied response', async () => {
|
||||
resolveGeminiRuntimeConfigMock.mockReturnValue({
|
||||
model: 'gemini-3-flash-preview',
|
||||
modelSource: 'default'
|
||||
});
|
||||
|
||||
await runGemini({});
|
||||
|
||||
const registerCalls = harness.session.rpcHandlerManager.registerHandler.mock.calls;
|
||||
const configHandler = registerCalls.find(
|
||||
(call: unknown[]) => call[0] === 'set-session-config'
|
||||
);
|
||||
const handler = configHandler![1] as (payload: unknown) => Promise<unknown>;
|
||||
const result = await handler({ permissionMode: 'default' }) as Record<string, unknown>;
|
||||
const applied = result.applied as Record<string, unknown>;
|
||||
expect(applied.permissionMode).toBe('default');
|
||||
expect(applied).not.toHaveProperty('model');
|
||||
});
|
||||
|
||||
it('stores null model in session on Default selection for keepalive', async () => {
|
||||
resolveGeminiRuntimeConfigMock.mockReturnValue({
|
||||
model: 'gemini-2.5-pro',
|
||||
modelSource: 'default'
|
||||
});
|
||||
|
||||
await runGemini({});
|
||||
|
||||
const registerCalls = harness.session.rpcHandlerManager.registerHandler.mock.calls;
|
||||
const configHandler = registerCalls.find(
|
||||
(call: unknown[]) => call[0] === 'set-session-config'
|
||||
);
|
||||
const handler = configHandler![1] as (payload: unknown) => Promise<unknown>;
|
||||
|
||||
// First set an explicit model
|
||||
await handler({ model: 'gemini-2.5-flash' });
|
||||
expect(mockGeminiSession.setModel).toHaveBeenLastCalledWith('gemini-2.5-flash');
|
||||
|
||||
// Then select Default (null) — session should store null, not concrete model
|
||||
await handler({ model: null });
|
||||
expect(mockGeminiSession.setModel).toHaveBeenLastCalledWith(null);
|
||||
});
|
||||
|
||||
it('passes machine default (not startup model) to geminiLoop for fallback', async () => {
|
||||
// Session started with explicit model, but machine default differs
|
||||
resolveGeminiRuntimeConfigMock.mockImplementation((opts?: { model?: string }) => {
|
||||
if (opts?.model) {
|
||||
return { model: opts.model, modelSource: 'explicit' };
|
||||
}
|
||||
return { model: 'gemini-2.5-pro', modelSource: 'default' };
|
||||
});
|
||||
|
||||
await runGemini({ model: 'gemini-2.5-flash' });
|
||||
|
||||
// geminiLoop should receive machine default as fallback, not the explicit startup model
|
||||
expect(harness.geminiLoopArgs[0]?.model).toBe('gemini-2.5-pro');
|
||||
});
|
||||
|
||||
|
||||
@@ -37,6 +37,7 @@ export async function runGemini(opts: {
|
||||
controlledByUser: false
|
||||
};
|
||||
|
||||
const machineDefault = resolveGeminiRuntimeConfig().model;
|
||||
const runtimeConfig = resolveGeminiRuntimeConfig({ model: opts.model });
|
||||
const persistedModel = runtimeConfig.modelSource === 'default'
|
||||
? undefined
|
||||
@@ -62,7 +63,8 @@ export async function runGemini(opts: {
|
||||
|
||||
const sessionWrapperRef: { current: GeminiSession | null } = { current: null };
|
||||
let currentPermissionMode: PermissionMode = opts.permissionMode ?? 'default';
|
||||
const resolvedModel = runtimeConfig.model;
|
||||
let sessionModel: string | null = persistedModel ?? null;
|
||||
let resolvedModel = sessionModel ?? machineDefault;
|
||||
|
||||
const hookServer = await startHookServer({
|
||||
onSessionHook: (sessionId, data) => {
|
||||
@@ -105,7 +107,8 @@ export async function runGemini(opts: {
|
||||
return;
|
||||
}
|
||||
sessionInstance.setPermissionMode(currentPermissionMode);
|
||||
logger.debug(`[gemini] Synced session permission mode for keepalive: ${currentPermissionMode}`);
|
||||
sessionInstance.setModel(sessionModel);
|
||||
logger.debug(`[gemini] Synced session config for keepalive: permissionMode=${currentPermissionMode}, model=${resolvedModel}`);
|
||||
};
|
||||
|
||||
session.onUserMessage((message) => {
|
||||
@@ -125,18 +128,36 @@ export async function runGemini(opts: {
|
||||
return parsed.data as PermissionMode;
|
||||
};
|
||||
|
||||
const resolveModel = (value: unknown): string | null => {
|
||||
if (value === null) {
|
||||
return null;
|
||||
}
|
||||
if (typeof value !== 'string' || value.trim().length === 0) {
|
||||
throw new Error('Invalid model');
|
||||
}
|
||||
return value.trim();
|
||||
};
|
||||
|
||||
session.rpcHandlerManager.registerHandler('set-session-config', async (payload: unknown) => {
|
||||
if (!payload || typeof payload !== 'object') {
|
||||
throw new Error('Invalid session config payload');
|
||||
}
|
||||
const config = payload as { permissionMode?: unknown };
|
||||
const config = payload as { permissionMode?: unknown; model?: unknown };
|
||||
const applied: Record<string, unknown> = {};
|
||||
|
||||
if (config.permissionMode !== undefined) {
|
||||
currentPermissionMode = resolvePermissionMode(config.permissionMode);
|
||||
applied.permissionMode = currentPermissionMode;
|
||||
}
|
||||
|
||||
if (config.model !== undefined) {
|
||||
sessionModel = resolveModel(config.model);
|
||||
resolvedModel = sessionModel ?? machineDefault;
|
||||
applied.model = sessionModel;
|
||||
}
|
||||
|
||||
syncSessionMode();
|
||||
return { applied: { permissionMode: currentPermissionMode } };
|
||||
return { applied };
|
||||
});
|
||||
|
||||
try {
|
||||
@@ -148,7 +169,7 @@ export async function runGemini(opts: {
|
||||
session,
|
||||
api,
|
||||
permissionMode: currentPermissionMode,
|
||||
model: resolvedModel,
|
||||
model: machineDefault,
|
||||
hookSettingsPath,
|
||||
resumeSessionId: opts.resumeSessionId,
|
||||
onModeChange: createModeChangeHandler(session),
|
||||
|
||||
@@ -78,6 +78,10 @@ export class GeminiSession extends AgentSessionBase<GeminiMode> {
|
||||
this.permissionMode = mode;
|
||||
};
|
||||
|
||||
setModel = (model: string | null): void => {
|
||||
this.model = model;
|
||||
};
|
||||
|
||||
recordLocalLaunchFailure = (message: string, exitReason: LocalLaunchExitReason): void => {
|
||||
this.localLaunchFailure = { message, exitReason };
|
||||
};
|
||||
|
||||
@@ -2,11 +2,12 @@ import { existsSync, readFileSync } from 'node:fs';
|
||||
import { homedir } from 'node:os';
|
||||
import { join } from 'node:path';
|
||||
import { logger } from '@/ui/logger';
|
||||
import { DEFAULT_GEMINI_MODEL } from '@hapi/protocol';
|
||||
|
||||
export const GEMINI_API_KEY_ENV = 'GEMINI_API_KEY';
|
||||
export const GOOGLE_API_KEY_ENV = 'GOOGLE_API_KEY';
|
||||
export const GEMINI_MODEL_ENV = 'GEMINI_MODEL';
|
||||
export const DEFAULT_GEMINI_MODEL = 'gemini-2.5-pro';
|
||||
export { DEFAULT_GEMINI_MODEL };
|
||||
|
||||
export type GeminiLocalConfig = {
|
||||
token?: string;
|
||||
@@ -91,7 +92,7 @@ export function resolveGeminiRuntimeConfig(opts: {
|
||||
const local = readGeminiLocalConfig();
|
||||
|
||||
let modelSource: GeminiModelSource = 'default';
|
||||
let model = DEFAULT_GEMINI_MODEL;
|
||||
let model: string = DEFAULT_GEMINI_MODEL;
|
||||
|
||||
if (opts.model) {
|
||||
model = opts.model;
|
||||
|
||||
Reference in New Issue
Block a user