Compare commits

...

1 Commits

Author SHA1 Message Date
Kit Langton a841e828ce fix(tui): wait for session model hydration 2026-08-04 16:17:33 -04:00
4 changed files with 150 additions and 95 deletions
+1 -4
View File
@@ -327,10 +327,7 @@ export function Prompt(props: PromptProps) {
if (!session) return if (!session) return
const agent = session.agent && local.agent.list().find((agent) => agent.id === session.agent) const agent = session.agent && local.agent.list().find((agent) => agent.id === session.agent)
if (agent && !args.agent) local.agent.set(agent.id) if (agent && !args.agent) local.agent.set(agent.id)
if (session.model) { if (!local.model.hydrate(session.model)) return
local.model.set({ providerID: session.model.providerID, modelID: session.model.id })
local.model.variant.set(session.model.variant)
}
syncedSessionID = sessionID syncedSessionID = sessionID
}) })
+45 -18
View File
@@ -6,7 +6,6 @@ import { useEvent } from "./event"
import path from "path" import path from "path"
import { useTuiPaths } from "./runtime" import { useTuiPaths } from "./runtime"
import { useArgs } from "./args" import { useArgs } from "./args"
import { useClient } from "./client"
import { RGBA } from "@opentui/core" import { RGBA } from "@opentui/core"
import { readJson, writeJsonAtomic } from "../util/persistence" import { readJson, writeJsonAtomic } from "../util/persistence"
import { import {
@@ -48,7 +47,6 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
name: "Local", name: "Local",
init: () => { init: () => {
const data = useData() const data = useData()
const client = useClient()
const toast = useToast() const toast = useToast()
const theme = useTheme() const theme = useTheme()
const { mode } = useThemes() const { mode } = useThemes()
@@ -210,6 +208,47 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
) )
}) })
function select(model: ModelPreferenceModel, options?: { recent?: boolean }) {
batch(() => {
if (!isModelValid(model)) return
const a = agent.current()
if (!a) return
setModelStore("model", a.id, model)
if (!options?.recent) return
setModelStore("recent", recentModels(model, modelStore.recent))
save()
})
}
function selectVariant(value: string | undefined) {
const model = currentModel()
if (!model) return
const key = modelPreferenceKey(model)
const variant = normalizeModelVariant(value)
if (modelStore.variant[key] === variant) return
setModelStore("variant", key, variant)
save()
}
function matches(model?: { providerID: string; id: string }) {
if (!modelStore.ready) return false
const current = currentModel()
if (!current) return false
if (!model) return true
return current.providerID === model.providerID && current.modelID === model.id
}
function hydrate(model?: { providerID: string; id: string; variant?: string }) {
if (!modelStore.ready) return false
if (!model) return true
if (data.location.model.list() === undefined) return false
const selected = { providerID: model.providerID, modelID: model.id }
if (!isModelValid(selected)) return false
select(selected)
selectVariant(model.variant)
return matches(model)
}
return { return {
current: currentModel, current: currentModel,
get ready() { get ready() {
@@ -221,6 +260,8 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
favorite() { favorite() {
return modelStore.favorite return modelStore.favorite
}, },
hydrate,
matches,
parsed: createMemo(() => { parsed: createMemo(() => {
const value = currentModel() const value = currentModel()
if (!value) { if (!value) {
@@ -285,18 +326,7 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
setModelStore("recent", recentModels(next, modelStore.recent)) setModelStore("recent", recentModels(next, modelStore.recent))
save() save()
}, },
set(model: { providerID: string; modelID: string }, options?: { recent?: boolean }) { set: select,
batch(() => {
if (!isModelValid(model)) return
const a = agent.current()
if (!a) return
setModelStore("model", a.id, model)
if (options?.recent) {
setModelStore("recent", recentModels(model, modelStore.recent))
save()
}
})
},
toggleFavorite(model: { providerID: string; modelID: string }) { toggleFavorite(model: { providerID: string; modelID: string }) {
batch(() => { batch(() => {
if (!isModelValid(model)) return if (!isModelValid(model)) return
@@ -333,10 +363,7 @@ export const { use: useLocal, provider: LocalProvider } = createSimpleContext({
return info?.variants?.map((variant) => variant.id) ?? [] return info?.variants?.map((variant) => variant.id) ?? []
}, },
set(value: string | undefined) { set(value: string | undefined) {
const m = currentModel() selectVariant(value)
if (!m) return
setModelStore("variant", modelPreferenceKey(m), normalizeModelVariant(value))
save()
}, },
cycle() { cycle() {
const variants = this.list() const variants = this.list()
+1 -1
View File
@@ -360,7 +360,7 @@ export function Session() {
createEffect(() => { createEffect(() => {
const current = prompt() const current = prompt()
if (sent || !current || !synced() || !local.model.ready) return if (sent || !current || !synced() || !local.model.ready) return
if (!local.agent.current() || !local.model.current()) return if (!local.agent.current() || !local.model.matches(session()?.model)) return
if (!args.prompt || route.prompt?.text !== args.prompt || current.current.text !== args.prompt) return if (!args.prompt || route.prompt?.text !== args.prompt || current.current.text !== args.prompt) return
sent = true sent = true
current.submit() current.submit()
+37 -6
View File
@@ -226,30 +226,42 @@ test("session title generated while an untitled session is loading remains visib
} }
}) })
test("session startup prompt is submitted exactly once", async () => { for (const scenario of [
{ name: "session startup prompt is submitted exactly once", delayed: false },
{ name: "session model hydration retries after the model catalog loads", delayed: true },
]) {
test(scenario.name, async () => {
const setup = await createTestRenderer({ width: 80, height: 24, useThread: false }) const setup = await createTestRenderer({ width: 80, height: 24, useThread: false })
const core = await import("@opentui/core") const core = await import("@opentui/core")
mock.module("@opentui/core", () => ({ ...core, createCliRenderer: async () => setup.renderer })) mock.module("@opentui/core", () => ({ ...core, createCliRenderer: async () => setup.renderer }))
const events = createEventStream() const events = createEventStream()
const cwd = process.cwd() const cwd = process.cwd()
const location = { directory: cwd, project: { id: "project", directory: cwd } } const location = { directory: cwd, project: { id: "project", directory: cwd } }
const selected = scenario.delayed ? "selected" : "model"
const session = { const session = {
id: "dummy", id: "dummy",
title: "Demo session", title: "Demo session",
projectID: "project", projectID: "project",
location: { directory: cwd }, location: { directory: cwd },
agent: "build", agent: "build",
model: { providerID: "provider", id: "model" }, model: { providerID: "provider", id: selected },
cost: 0, cost: 0,
tokens: { input: 0, output: 0, reasoning: 0, cache: { read: 0, write: 0 } }, tokens: { input: 0, output: 0, reasoning: 0, cache: { read: 0, write: 0 } },
time: { created: 0, updated: 0 }, time: { created: 0, updated: 0 },
} }
const bodies: unknown[] = [] const sessionRequested = Promise.withResolvers<void>()
const modelRequested = Promise.withResolvers<void>()
const releaseModels = Promise.withResolvers<void>()
const promptSubmitted = Promise.withResolvers<void>() const promptSubmitted = Promise.withResolvers<void>()
const bodies: unknown[] = []
const selections: unknown[] = []
const calls = createFetch(async (url, request) => { const calls = createFetch(async (url, request) => {
if (url.pathname === "/api/location") return json(location) if (url.pathname === "/api/location") return json(location)
if (url.pathname === "/api/session") return json({ data: [session], cursor: {} }) if (url.pathname === "/api/session") return json({ data: [session], cursor: {} })
if (url.pathname === "/api/session/dummy") return json({ data: session }) if (url.pathname === "/api/session/dummy") {
sessionRequested.resolve()
return json({ data: session })
}
if (url.pathname === "/api/session/dummy/message") return json({ data: [], cursor: {} }) if (url.pathname === "/api/session/dummy/message") return json({ data: [], cursor: {} })
if (url.pathname === "/api/session/dummy/pending") return json({ data: [] }) if (url.pathname === "/api/session/dummy/pending") return json({ data: [] })
if (url.pathname === "/api/session/dummy/permission") return json({ data: [] }) if (url.pathname === "/api/session/dummy/permission") return json({ data: [] })
@@ -258,11 +270,21 @@ test("session startup prompt is submitted exactly once", async () => {
location, location,
data: [{ id: "build", mode: "primary", hidden: false, permissions: [] }], data: [{ id: "build", mode: "primary", hidden: false, permissions: [] }],
}) })
if (url.pathname === "/api/model") if (url.pathname === "/api/model") {
modelRequested.resolve()
if (scenario.delayed) await releaseModels.promise
return json({ return json({
location, location,
data: [{ id: "model", providerID: "provider", name: "Model", variants: [] }], data: [
...(scenario.delayed ? [{ id: "fallback", providerID: "provider", name: "Fallback", variants: [] }] : []),
{ id: selected, providerID: "provider", name: "Selected", variants: [] },
],
}) })
}
if (url.pathname === "/api/session/dummy/model") {
selections.push(await request.json())
return new Response(null, { status: 204 })
}
if (url.pathname === "/api/session/dummy/prompt") { if (url.pathname === "/api/session/dummy/prompt") {
bodies.push(await request.json()) bodies.push(await request.json())
promptSubmitted.resolve() promptSubmitted.resolve()
@@ -284,6 +306,12 @@ test("session startup prompt is submitted exactly once", async () => {
}).pipe(Effect.provide(AppNodeBuilder.build(Global.node)), Effect.provide(FileSystem.layerNoop({}))), }).pipe(Effect.provide(AppNodeBuilder.build(Global.node)), Effect.provide(FileSystem.layerNoop({}))),
) )
if (scenario.delayed) {
await Promise.all([sessionRequested.promise, modelRequested.promise])
await Bun.sleep(20)
expect(bodies).toEqual([])
releaseModels.resolve()
}
await Promise.race([ await Promise.race([
promptSubmitted.promise, promptSubmitted.promise,
Bun.sleep(2000).then(() => { Bun.sleep(2000).then(() => {
@@ -296,9 +324,12 @@ test("session startup prompt is submitted exactly once", async () => {
expect(bodies).toHaveLength(1) expect(bodies).toHaveLength(1)
expect(bodies[0]).toMatchObject({ text: "RESUME_READY" }) expect(bodies[0]).toMatchObject({ text: "RESUME_READY" })
expect(selections).toEqual([])
} finally { } finally {
releaseModels.resolve()
if (!setup.renderer.isDestroyed) setup.renderer.destroy() if (!setup.renderer.isDestroyed) setup.renderer.destroy()
await server.stop() await server.stop()
mock.restore() mock.restore()
} }
}) })
}