diff --git a/packages/tui/src/context/local.tsx b/packages/tui/src/context/local.tsx index 7e6da553f47e..79a21accdc23 100644 --- a/packages/tui/src/context/local.tsx +++ b/packages/tui/src/context/local.tsx @@ -12,8 +12,10 @@ import { readJson, writeJsonAtomic } from "../util/persistence" import { createModelPreferenceRepository, cycleModelVariant, + favoriteModels, modelPreferenceKey, normalizeModelVariant, + recentModels, type ModelPreference, type ModelPreferenceModel, } from "../model-preference" @@ -32,19 +34,6 @@ export function parseModel(model: string) { } } -export function recentModels(model: ModelPreferenceModel, recent: ModelPreferenceModel[]) { - const seen = new Set() - return [model, ...recent] - .filter((item) => { - const key = modelPreferenceKey(item) - if (seen.has(key)) return false - seen.add(key) - return true - }) - .slice(0, 10) - .map((item) => ({ providerID: item.providerID, modelID: item.modelID })) -} - export const { use: useLocal, provider: LocalProvider } = createSimpleContext({ name: "Local", init: () => { @@ -151,37 +140,16 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({ const pendingSelectionCommits = new Map() const selectionKey = (value: ModelSelection) => `${modelPreferenceKey(value)}:${normalizeModelVariant(value.variant) ?? "default"}` - const saveState = { - pending: false, - } - - function savePreferences() { - if (!preferences.ready) { - saveState.pending = true - return - } - saveState.pending = false - void repository - .patch({ - recent: preferences.recent, - favorite: preferences.favorite, - variant: preferences.variant, - }) - .catch(() => undefined) - } - - repository - .load() - .then((value) => { + function applyPreferences(value: ModelPreference) { + batch(() => { setPreferences("recent", value.recent) setPreferences("favorite", value.favorite) setPreferences("variant", value.variant) - }) - .catch(() => {}) - .finally(() => { setPreferences("ready", true) - if (saveState.pending) savePreferences() }) + } + + onCleanup(repository.subscribe(applyPreferences)) const fallbackModel = createMemo(() => { if (args.model) { @@ -389,7 +357,7 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({ if (!next) return if (!selectModel({ ...next })) return setPreferences("recent", recentModels(next, preferences.recent)) - savePreferences() + void repository.addRecent(next).catch(() => undefined) }, set(model: { providerID: string; modelID: string }, options?: { recent?: boolean }) { batch(() => { @@ -397,7 +365,7 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({ if (!selectModel(model)) return if (options?.recent) { setPreferences("recent", recentModels(model, preferences.recent)) - savePreferences() + void repository.addRecent(model).catch(() => undefined) } }) }, @@ -407,14 +375,8 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({ const exists = preferences.favorite.some( (x) => x.providerID === model.providerID && x.modelID === model.modelID, ) - const next = exists - ? preferences.favorite.filter((x) => x.providerID !== model.providerID || x.modelID !== model.modelID) - : [model, ...preferences.favorite] - setPreferences( - "favorite", - next.map((x) => ({ providerID: x.providerID, modelID: x.modelID })), - ) - savePreferences() + setPreferences("favorite", favoriteModels(model, preferences.favorite, !exists)) + void repository.setFavorite(model, !exists).catch(() => undefined) }) }, variant: { @@ -439,7 +401,7 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({ setSessionDraft(route.data.sessionID, { ...m, variant: normalizeModelVariant(value) }) } setPreferences("variant", modelPreferenceKey(m), normalizeModelVariant(value)) - savePreferences() + void repository.saveVariant(m, value).catch(() => undefined) }, cycle() { const variants = this.list() diff --git a/packages/tui/src/model-preference.ts b/packages/tui/src/model-preference.ts index 5f9ce566c052..39486d53b561 100644 --- a/packages/tui/src/model-preference.ts +++ b/packages/tui/src/model-preference.ts @@ -1,5 +1,7 @@ import { readJson, writeJsonAtomic } from "./util/persistence" import { isRecord } from "./util/record" +import { watch } from "node:fs" +import path from "node:path" export type ModelPreferenceModel = { providerID: string @@ -43,6 +45,24 @@ export function modelPreferenceKey(model: ModelPreferenceModel) { return `${model.providerID}/${model.modelID}` } +export function recentModels(model: ModelPreferenceModel, recent: ModelPreferenceModel[]) { + const seen = new Set() + return [model, ...recent] + .filter((item) => { + const key = modelPreferenceKey(item) + if (seen.has(key)) return false + seen.add(key) + return true + }) + .slice(0, 10) + .map((item) => ({ providerID: item.providerID, modelID: item.modelID })) +} + +export function favoriteModels(model: ModelPreferenceModel, favorite: ModelPreferenceModel[], enabled: boolean) { + const current = favorite.filter((item) => modelPreferenceKey(item) !== modelPreferenceKey(model)) + return enabled ? [model, ...current] : current +} + export function cycleModelVariant(current: string | undefined, variants: string[]) { const named = variants.filter((variant) => variant !== "default") if (named.length === 0) return undefined @@ -80,32 +100,78 @@ function patch(value: Partial) { } export function createModelPreferenceRepository(filePath: string) { - const state = { - pending: Promise.resolve(), - } + let pending = Promise.resolve() + let revision = 0 + let watcher: ReturnType | undefined + let reload: ReturnType | undefined + const listeners = new Set<(value: ModelPreference) => void>() const read = () => readJson(filePath) .then(decodeModelPreference) .catch(() => decodeModelPreference(undefined)) function update(change: (current: ModelPreference) => Partial) { - const result = state.pending.then(async () => { - const current = await read() - const next = { ...current, ...patch(change(preference(current))) } - await writeJsonAtomic(filePath, next) + const result = pending.then(async () => { + const { Flock } = await import("@opencode-ai/util/flock") + return Flock.withLock( + filePath, + async () => { + const current = await read() + const next = { ...current, ...patch(change(preference(current))) } + await writeJsonAtomic(filePath, next) + }, + { dir: path.join(path.dirname(filePath), "locks") }, + ) }) - state.pending = result.catch(() => undefined) + pending = result.then( + () => undefined, + () => undefined, + ) return result } function load() { - return state.pending.then(read).then(preference) + return pending.then(read).then(preference) + } + + async function refresh() { + const current = ++revision + const value = await load() + if (current !== revision) return + listeners.forEach((listener) => listener(value)) } return { load, - patch(value: Partial) { - return update(() => value) + addRecent(model: ModelPreferenceModel) { + return update((current) => ({ recent: recentModels(model, current.recent) })) + }, + setFavorite(model: ModelPreferenceModel, enabled: boolean) { + return update((current) => ({ favorite: favoriteModels(model, current.favorite, enabled) })) + }, + subscribe(listener: (value: ModelPreference) => void) { + listeners.add(listener) + void refresh() + if (!watcher) { + watcher = watch(path.dirname(filePath), (_event, filename) => { + const changed = filename?.toString() + const name = path.basename(filePath) + if (changed !== undefined && changed !== name && !changed.startsWith(name + ".")) return + clearTimeout(reload) + reload = setTimeout(() => void refresh(), 50) + }) + watcher.on("error", () => { + watcher?.close() + watcher = undefined + }) + } + return () => { + listeners.delete(listener) + if (listeners.size > 0) return + clearTimeout(reload) + watcher?.close() + watcher = undefined + } }, async resolveVariant(model: ModelPreferenceModel) { return (await load()).variant[modelPreferenceKey(model)] diff --git a/packages/tui/test/context/local.test.ts b/packages/tui/test/context/local.test.ts index e2f1e45f75a9..056afc7daf21 100644 --- a/packages/tui/test/context/local.test.ts +++ b/packages/tui/test/context/local.test.ts @@ -1,5 +1,6 @@ import { expect, test } from "bun:test" -import { parseModel, recentModels } from "../../src/context/local" +import { parseModel } from "../../src/context/local" +import { recentModels } from "../../src/model-preference" test("parses model IDs containing slashes", () => { expect(parseModel("provider/family/model")).toEqual({ diff --git a/packages/tui/test/model-preference.test.ts b/packages/tui/test/model-preference.test.ts index 5fd7d14bd382..b4183ca63fcb 100644 --- a/packages/tui/test/model-preference.test.ts +++ b/packages/tui/test/model-preference.test.ts @@ -19,7 +19,7 @@ test("repairs known model preferences and preserves unrelated fields", () => { }) }) -test("atomically serializes patches and variant updates", async () => { +test("atomically serializes model preference updates", async () => { await using tmp = await tmpdir() const file = path.join(tmp.path, "model.json") await Bun.write(file, JSON.stringify({ unrelated: "keep", favorite: [], variant: {} })) @@ -28,7 +28,7 @@ test("atomically serializes patches and variant updates", async () => { const anthropic = { providerID: "anthropic", modelID: "claude/sonnet" } await Promise.all([ - repository.patch({ recent: [openai] }), + repository.addRecent(openai), repository.saveVariant(openai, "high"), repository.saveVariant(anthropic, "low"), ]) @@ -43,3 +43,47 @@ test("atomically serializes patches and variant updates", async () => { expect(await repository.resolveVariant(openai)).toBeUndefined() expect((await Bun.file(file).json()).variant).toEqual({ "anthropic/claude/sonnet": "low" }) }) + +test("serializes updates across repositories", async () => { + await using tmp = await tmpdir() + const file = path.join(tmp.path, "model.json") + await Bun.write(file, JSON.stringify({ recent: [], favorite: [], variant: {} })) + const first = createModelPreferenceRepository(file) + const second = createModelPreferenceRepository(file) + const openai = { providerID: "openai", modelID: "gpt-5" } + const anthropic = { providerID: "anthropic", modelID: "claude-sonnet" } + + await Promise.all([first.setFavorite(openai, true), second.addRecent(anthropic), second.saveVariant(openai, "high")]) + + expect(await first.load()).toEqual({ + recent: [anthropic], + favorite: [openai], + variant: { "openai/gpt-5": "high" }, + }) +}) + +test("subscribes to updates from another repository", async () => { + await using tmp = await tmpdir() + const file = path.join(tmp.path, "model.json") + await Bun.write(file, JSON.stringify({ recent: [], favorite: [], variant: {} })) + const first = createModelPreferenceRepository(file) + const second = createModelPreferenceRepository(file) + const openai = { providerID: "openai", modelID: "gpt-5" } + const changed = Promise.withResolvers() + const unsubscribe = first.subscribe((value) => { + if (value.favorite.some((item) => item.providerID === openai.providerID && item.modelID === openai.modelID)) + changed.resolve() + }) + + try { + await second.setFavorite(openai, true) + await Promise.race([ + changed.promise, + Bun.sleep(2_000).then(() => { + throw new Error("timed out waiting for model preference update") + }), + ]) + } finally { + unsubscribe() + } +})