ソースを参照

fix(tui): render submitted prompts optimistically

Kit Langton 4 週間 前
コミット
7af31118da

+ 36 - 11
packages/tui/src/component/prompt/index.tsx

@@ -53,6 +53,7 @@ import { readLocalAttachment } from "./local-attachment"
 import { useData } from "../../context/data"
 import { useLocation } from "../../context/location"
 import { contextUsage } from "../../util/session"
+import { SessionMessage } from "@opencode-ai/core/session/message"
 
 registerOpencodeSpinner()
 
@@ -1003,8 +1004,10 @@ export function Prompt(props: PromptProps) {
 
     // Capture mode before it gets reset
     const currentMode = store.mode
+    const submittedPrompt = structuredClone(unwrap(store.prompt))
     const editorSelection = editorContext()
     const pendingEditorSelection = editorSelection && editor.labelState() === "pending" ? editorSelection : undefined
+    let cleared = false
 
     if (store.mode === "shell") {
       move.startSubmit()
@@ -1097,31 +1100,53 @@ export function Prompt(props: PromptProps) {
           return false
         }
       }
+      const messageID = SessionMessage.ID.create()
+      data.session.message.optimistic.add(sessionID, {
+        id: messageID,
+        type: "user",
+        text: inputText,
+        time: { created: Date.now() },
+      })
+      history.append({
+        ...submittedPrompt,
+        mode: currentMode,
+      })
+      input.extmarks.clear()
+      setStore("prompt", emptyPrompt())
+      setStore("extmarkToPart", new Map())
+      props.onSubmit?.()
+      input.clear()
+      cleared = true
+
       const error = await sdk.api.session
         .prompt({
           sessionID,
+          id: messageID,
           text: inputText,
-          files: store.prompt.files,
-          agents: store.prompt.agents,
+          files: submittedPrompt.files,
+          agents: submittedPrompt.agents,
         })
         .then(
           () => undefined,
           (error) => error,
         )
       if (error) {
+        data.session.message.optimistic.remove(sessionID, messageID)
         toast.show({ title: "Failed to send prompt", message: errorMessage(error), variant: "error" })
         return false
       }
       if (pendingEditorSelection) editor.markSelectionSent()
     }
-    history.append({
-      ...store.prompt,
-      mode: currentMode,
-    })
-    input.extmarks.clear()
-    setStore("prompt", emptyPrompt())
-    setStore("extmarkToPart", new Map())
-    props.onSubmit?.()
+    if (!cleared) {
+      history.append({
+        ...submittedPrompt,
+        mode: currentMode,
+      })
+      input.extmarks.clear()
+      setStore("prompt", emptyPrompt())
+      setStore("extmarkToPart", new Map())
+      props.onSubmit?.()
+    }
 
     // temporary hack to make sure the message is sent
     if (!props.sessionID) {
@@ -1133,7 +1158,7 @@ export function Prompt(props: PromptProps) {
         })
       }, 50)
     }
-    input.clear()
+    if (!cleared) input.clear()
     if (finishMoveProgress) move.finishSubmit()
     return true
   }

+ 27 - 5
packages/tui/src/context/data.tsx

@@ -384,9 +384,7 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
               event.data.inputID,
             ])
           message.update(event.data.sessionID, (draft, index) => {
-            message.append(
-              draft,
-              index,
+            const item: SessionMessageInfo =
               event.data.input.type === "user"
                 ? {
                     id: event.data.inputID,
@@ -399,8 +397,10 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
                     type: "synthetic",
                     ...event.data.input.data,
                     time: { created: event.created },
-                  },
-            )
+                  }
+            const position = index.get(event.data.inputID)
+            if (position === undefined) return message.append(draft, index, item)
+            draft[position] = item
           })
           break
         case "session.instructions.updated":
@@ -928,6 +928,28 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
             const position = messageIndex.get(sessionID)?.get(messageID)
             return position === undefined ? undefined : messages?.[position]
           },
+          optimistic: {
+            add(sessionID: string, item: Extract<SessionMessageInfo, { type: "user" }>) {
+              if (!store.session.input[sessionID]?.includes(item.id))
+                setStore("session", "input", sessionID, [...(store.session.input[sessionID] ?? []), item.id])
+              message.update(sessionID, (draft, index) => message.append(draft, index, item))
+            },
+            remove(sessionID: string, messageID: string) {
+              setStore(
+                "session",
+                "input",
+                sessionID,
+                (store.session.input[sessionID] ?? []).filter((id) => id !== messageID),
+              )
+              message.update(sessionID, (draft, index) => {
+                const position = index.get(messageID)
+                if (position === undefined) return
+                draft.splice(position, 1)
+                index.clear()
+                draft.forEach((item, itemIndex) => index.set(item.id, itemIndex))
+              })
+            },
+          },
           async refresh(sessionID: string) {
             const messages = (await sdk.api.message.list({ sessionID, limit: 200, order: "desc" })).data.toReversed()
             messageIndex.set(sessionID, new Map(messages.map((message, index) => [message.id, index])))

+ 69 - 0
packages/tui/test/cli/tui/data.test.tsx

@@ -653,6 +653,75 @@ test("completes exploration when a queued prompt is promoted", async () => {
   }
 })
 
+test("shows optimistic prompts immediately and reconciles admission", async () => {
+  const events = createEventStream()
+  const sessionID = "session-optimistic"
+  const calls = createFetch((url) => {
+    if (url.pathname === `/api/session/${sessionID}/message`) return json({ data: [], cursor: {} })
+  }, events)
+  let data!: ReturnType<typeof useData>
+
+  function Probe() {
+    data = useData()
+    return <box />
+  }
+
+  const app = await testRender(() => (
+    <TestTuiContexts>
+      <SDKProvider client={createClient(calls.fetch)} api={createApi(calls.fetch)}>
+        <ProjectProvider>
+          <DataProvider>
+            <Probe />
+          </DataProvider>
+        </ProjectProvider>
+      </SDKProvider>
+    </TestTuiContexts>
+  ))
+
+  try {
+    const messageID = SessionMessage.ID.create()
+    data.session.message.optimistic.add(sessionID, {
+      id: messageID,
+      type: "user",
+      text: "Steer now",
+      time: { created: 1 },
+    })
+
+    const optimistic = data.session.message.get(sessionID, messageID)
+    expect(optimistic?.type).toBe("user")
+    expect(optimistic?.type === "user" ? optimistic.text : undefined).toBe("Steer now")
+    expect(data.session.input.has(sessionID, messageID)).toBe(true)
+
+    emitEvent(events, {
+      id: EventV2.ID.create(),
+      created: 2,
+      type: "session.input.admitted",
+      durable: durable(sessionID),
+      data: {
+        sessionID,
+        inputID: messageID,
+        input: { type: "user", data: { text: "Steer now", metadata: { admitted: true } }, delivery: "steer" },
+      },
+    })
+
+    await wait(() => data.session.message.get(sessionID, messageID)?.metadata?.admitted === true)
+    expect(data.session.message.list(sessionID)).toHaveLength(1)
+
+    const failedID = SessionMessage.ID.create()
+    data.session.message.optimistic.add(sessionID, {
+      id: failedID,
+      type: "user",
+      text: "Fails",
+      time: { created: 3 },
+    })
+    data.session.message.optimistic.remove(sessionID, failedID)
+    expect(data.session.message.get(sessionID, failedID)).toBeUndefined()
+    expect(data.session.input.has(sessionID, failedID)).toBe(false)
+  } finally {
+    app.renderer.destroy()
+  }
+})
+
 test("removes committed revert messages from local state", async () => {
   const events = createEventStream()
   const sessionID = "session-revert"