diff --git a/src/config/index.ts b/src/config/index.ts index ed250d6bc..a108ecf29 100644 --- a/src/config/index.ts +++ b/src/config/index.ts @@ -650,6 +650,13 @@ export interface Config { * id; `"pick"` opens the interactive picker. Omitted for a fresh session. */ resumeMode?: "id" | "pick"; + /** + * True when this invocation passed --provider and/or --model. Resume keeps + * the stored session's model unless this is set; the override is a + * parse-time fact, never inferred by comparing launch values against the + * stored record. + */ + modelOverride?: boolean; // Deprecated no-op retained for CLI compatibility. noWorkflow: boolean; @@ -1181,6 +1188,9 @@ export async function loadConfig( } : {}), ...(resumePicker ? { resumePicker: true } : {}), + ...(provider !== undefined || model !== undefined + ? { modelOverride: true as const } + : {}), ...(settings?.defaultProvider !== undefined ? { globalDefaultProvider: settings.defaultProvider } : {}), diff --git a/src/tui/session-start.test.ts b/src/tui/session-start.test.ts index c223c4910..fd3ed9601 100644 --- a/src/tui/session-start.test.ts +++ b/src/tui/session-start.test.ts @@ -2,7 +2,10 @@ import { describe, expect, test } from "bun:test"; import { withMockedModuleDuring } from "../../tests/helpers/mock-module.js"; import { setActiveRun, clearActiveRun } from "../session/active-run.js"; +import type { RunStateHandle } from "../session/active-run.js"; +import type { Config } from "../config/index.js"; import type { RunState } from "../session/state.js"; +import type { Telemetry } from "../telemetry/index.js"; describe("createTUICrashGuard", () => { test("finalizeOnCrash writes live session id and provider:model after bindLiveSession", async () => { @@ -99,3 +102,218 @@ describe("createTUICrashGuard", () => { expect(ran).toBe(true); }); }); + +interface CapturedSave { + cwd: string; + sessionId: string; + state: RunState; +} + +function storedRunState(overrides: Partial = {}): RunState { + return { + status: "running", + turnsUsed: 7, + task: "stored task", + startedAt: 111, + model: "stored-p:stored-m", + ...overrides, + }; +} + +function launchConfig(overrides: Partial = {}): Config { + return { + configured: true, + apiKey: "key", + baseURL: "https://example.test", + model: "launch-m", + providerName: "launch-p", + cwd: "/cwd", + task: "", + dangerouslySkipPermissions: false, + anthropicCachePrompt: false, + skipPermissionsFromSettings: false, + auto: true, + command: "tui", + globalSettingsPath: "/settings", + providers: [], + mcpServerEntries: [], + sessionId: "launch-session", + noWorkflow: false, + ...overrides, + } as Config; +} + +async function runPrepareTUISession( + config: Config, + opts: { pickSessionId?: string; stored?: RunState | null }, +): Promise<{ + prepared: Awaited< + ReturnType + >; + saves: CapturedSave[]; + activeRuns: RunStateHandle[]; +}> { + const saves: CapturedSave[] = []; + const activeRuns: RunStateHandle[] = []; + const stored = opts.stored === undefined ? storedRunState() : opts.stored; + const prepared = await withMockedModuleDuring( + import.meta.resolve("../session/assemble-runtime.js"), + (real: typeof import("../session/assemble-runtime.js")) => ({ + ...real, + assembleInferenceBase: async () => ({}), + assembleSessionTrust: async () => ({}), + }), + async () => + withMockedModuleDuring( + import.meta.resolve("./pick-session.js"), + (real: typeof import("./pick-session.js")) => ({ + ...real, + pickSession: async () => + opts.pickSessionId === undefined + ? null + : { + sessionId: opts.pickSessionId, + task: "picked task", + startedAt: 1, + updatedAt: 2, + status: "running" as const, + }, + }), + async () => + withMockedModuleDuring( + import.meta.resolve("../session/state.js"), + (real: typeof import("../session/state.js")) => ({ + ...real, + loadState: async () => + stored === null + ? { kind: "missing" as const } + : { kind: "ok" as const, state: stored }, + saveState: async ( + cwd: string, + sessionId: string, + state: RunState, + ) => { + saves.push({ cwd, sessionId, state }); + }, + }), + async () => + withMockedModuleDuring( + import.meta.resolve("../session/index.js"), + (real: typeof import("../session/index.js")) => ({ + ...real, + initSessionDir: async () => "/dir", + sessionContextDir: () => "/workdir", + }), + async () => + withMockedModuleDuring( + import.meta.resolve("../session/active-run.js"), + (real: typeof import("../session/active-run.js")) => ({ + ...real, + setActiveRun: (handle: RunStateHandle) => { + activeRuns.push(handle); + }, + }), + async () => { + const { prepareTUISession } = + await import("./session-start.js"); + return prepareTUISession(config, {} as Telemetry); + }, + ), + ), + ), + ), + ); + return { prepared, saves, activeRuns }; +} + +describe("prepareTUISession resume model", () => { + test("picker resume without flags restores the stored provider:model", async () => { + const { prepared, saves } = await runPrepareTUISession( + launchConfig({ resumePicker: true }), + { pickSessionId: "picked-session" }, + ); + + expect(prepared?.config.providerName).toBe("stored-p"); + expect(prepared?.config.model).toBe("stored-m"); + expect(saves).toHaveLength(1); + expect(saves[0]?.state.model).toBe("stored-p:stored-m"); + }); + + test("id resume without flags restores the stored provider:model", async () => { + const { prepared, saves, activeRuns } = await runPrepareTUISession( + launchConfig({ resumeMode: "id", sessionId: "resume-id" }), + {}, + ); + + expect(prepared?.config.providerName).toBe("stored-p"); + expect(prepared?.config.model).toBe("stored-m"); + expect(prepared?.resumeSeed.storedModel).toEqual({ + providerName: "stored-p", + model: "stored-m", + }); + expect(saves).toHaveLength(1); + expect(saves[0]?.sessionId).toBe("resume-id"); + expect(saves[0]?.state.model).toBe("stored-p:stored-m"); + expect(activeRuns).toHaveLength(1); + expect(activeRuns[0]?.model).toBe("stored-p:stored-m"); + }); + + test("explicit flags win over the stored model on both resume branches", async () => { + for (const config of [ + launchConfig({ resumePicker: true, modelOverride: true }), + launchConfig({ + resumeMode: "id", + sessionId: "resume-id", + modelOverride: true, + }), + ]) { + const { prepared, saves } = await runPrepareTUISession(config, { + pickSessionId: "picked-session", + }); + + expect(prepared?.config.providerName).toBe("launch-p"); + expect(prepared?.config.model).toBe("launch-m"); + expect(saves).toHaveLength(1); + expect(saves[0]?.state.model).toBe("launch-p:launch-m"); + } + }); + + test("legacy model-less records keep the launch default", async () => { + const { model: _dropped, ...legacy } = storedRunState(); + const { prepared, saves } = await runPrepareTUISession( + launchConfig({ resumePicker: true }), + { pickSessionId: "picked-session", stored: legacy }, + ); + + expect(prepared?.config.providerName).toBe("launch-p"); + expect(prepared?.config.model).toBe("launch-m"); + expect(prepared?.resumeSeed.storedModel).toBeUndefined(); + expect(saves).toHaveLength(1); + expect(saves[0]?.state.model).toBe("launch-p:launch-m"); + }); + + test("resolveResumeSeed carries the stored model and tolerates malformed values", async () => { + const { resolveResumeSeed } = await import("./session-start.js"); + + expect(resolveResumeSeed(null)).toEqual({ + turnsUsed: 0, + mcpServers: [], + activatedTools: [], + }); + expect(resolveResumeSeed(storedRunState()).storedModel).toEqual({ + providerName: "stored-p", + model: "stored-m", + }); + expect( + resolveResumeSeed(storedRunState({ model: "stored-p:org:model-v2" })) + .storedModel, + ).toEqual({ providerName: "stored-p", model: "org:model-v2" }); + for (const model of ["nocolon", ":empty-provider", "provider:", ""]) { + expect( + resolveResumeSeed(storedRunState({ model })).storedModel, + ).toBeUndefined(); + } + const { model: _dropped, ...legacy } = storedRunState(); + expect(resolveResumeSeed(legacy).storedModel).toBeUndefined(); + }); +}); diff --git a/src/tui/session-start.ts b/src/tui/session-start.ts index e1b429e89..b5c47278a 100644 --- a/src/tui/session-start.ts +++ b/src/tui/session-start.ts @@ -50,6 +50,9 @@ export interface ResumeSeed { // Present only when the prior run stamped an Anthropic-protocol cache write. lastCacheWriteAt?: number; cacheWriteModel?: string; + // The prior run's resolved provider:model, split for restore. Absent for a + // fresh run and for legacy records that predate the model field. + storedModel?: { providerName: string; model: string }; } const FRESH_RESUME_SEED: ResumeSeed = { @@ -58,6 +61,23 @@ const FRESH_RESUME_SEED: ResumeSeed = { activatedTools: [], }; +/** + * Split a run.json `provider:model` identity on its first colon so model ids + * containing colons survive. Absent or malformed values yield undefined and + * the caller keeps the launch default; this never throws. + */ +function splitStoredModel( + value: string | undefined, +): { providerName: string; model: string } | undefined { + if (value === undefined) return undefined; + const colon = value.indexOf(":"); + if (colon <= 0 || colon === value.length - 1) return undefined; + return { + providerName: value.slice(0, colon), + model: value.slice(colon + 1), + }; +} + /** * Fold a resumed session's run.json into a concrete seed once, at the * resume boundary, so every downstream reader (the run sink, the @@ -68,10 +88,12 @@ const FRESH_RESUME_SEED: ResumeSeed = { */ export function resolveResumeSeed(pickedState: RunState | null): ResumeSeed { if (pickedState === null) return FRESH_RESUME_SEED; + const storedModel = splitStoredModel(pickedState.model); return { turnsUsed: pickedState.turnsUsed, mcpServers: pickedState.mcpServers ?? [], activatedTools: pickedState.activatedTools ?? [], + ...(storedModel !== undefined ? { storedModel } : {}), ...(pickedState.lastCacheWriteAt !== undefined ? { lastCacheWriteAt: pickedState.lastCacheWriteAt, @@ -290,6 +312,19 @@ export async function prepareTUISession( } } + // A resume without explicit --provider/--model keeps the stored session's + // model: the launch default would otherwise clobber it on both resume + // branches above. An explicit flag wins, so the restore is gated on the + // parse-time signal rather than any value comparison. Fresh runs and + // legacy model-less records carry no storedModel and keep the default. + if (config.modelOverride !== true && resumeSeed.storedModel !== undefined) { + config = { + ...config, + providerName: resumeSeed.storedModel.providerName, + model: resumeSeed.storedModel.model, + }; + } + const workdir = sessionContextDir(config.cwd, sessionId); await initSessionDir(config.cwd, sessionId);