Explorar o código

fix(tui): track pending session inputs

Dax Raad hai 1 mes
pai
achega
9c2842b9b0

+ 28 - 4
packages/tui/src/context/data.tsx

@@ -55,6 +55,7 @@ type Data = {
     family: Record<string, string[]>
     status: Record<string, DataSessionStatus>
     message: Record<string, SessionMessage[]>
+    input: Record<string, string[]>
     permission: Record<string, PermissionV2Request[]>
     // Pending forms keyed by session ID.
     form: Record<string, FormInfo[]>
@@ -90,6 +91,7 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
         family: {},
         status: {},
         message: {},
+        input: {},
         permission: {},
         form: {},
       },
@@ -276,10 +278,8 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
             const position = index.get(event.data.inputID)
             if (position === undefined) return
             const existing = draft[position]
-            if (existing?.type === "user" && existing.metadata?.queued === true) {
+            if (existing?.type === "user" && store.session.input[event.data.sessionID]?.includes(event.data.inputID)) {
               existing.time.created = event.created
-              delete existing.metadata.queued
-              if (Object.keys(existing.metadata).length === 0) existing.metadata = undefined
               draft.splice(position, 1)
               draft.push(existing)
               index.clear()
@@ -287,9 +287,20 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
               return
             }
           })
+          setStore(
+            "session",
+            "input",
+            event.data.sessionID,
+            (store.session.input[event.data.sessionID] ?? []).filter((id) => id !== event.data.inputID),
+          )
           break
         }
         case "session.prompt.admitted":
+          if (!store.session.input[event.data.sessionID]?.includes(event.data.inputID))
+            setStore("session", "input", event.data.sessionID, [
+              ...(store.session.input[event.data.sessionID] ?? []),
+              event.data.inputID,
+            ])
           message.update(event.data.sessionID, (draft, index) => {
             message.append(draft, index, {
               id: event.data.inputID,
@@ -297,7 +308,6 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
               text: event.data.prompt.text,
               files: event.data.prompt.files,
               agents: event.data.prompt.agents,
-              metadata: { queued: true },
               time: { created: event.created },
             })
           })
@@ -559,6 +569,12 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
         case "session.revert.committed":
           if (store.session.info[event.data.sessionID])
             setStore("session", "info", event.data.sessionID, "revert", undefined)
+          setStore(
+            "session",
+            "input",
+            event.data.sessionID,
+            (store.session.input[event.data.sessionID] ?? []).filter((id) => id < event.data.messageID),
+          )
           message.update(event.data.sessionID, (draft, index) => {
             const position = draft.findIndex((item) => item.id >= event.data.messageID)
             if (position === -1) return
@@ -686,6 +702,14 @@ export const { use: useData, provider: DataProvider } = createSimpleContext({
         status(sessionID: string) {
           return store.session.status[sessionID] ?? "idle"
         },
+        input: {
+          list(sessionID: string) {
+            return store.session.input[sessionID] ?? []
+          },
+          has(sessionID: string, inputID: string) {
+            return store.session.input[sessionID]?.includes(inputID) ?? false
+          },
+        },
         async refresh(sessionID: string) {
           setStore("session", "info", sessionID, mutable(await sdk.api.session.get({ sessionID })))
           registerSession(sessionID)

+ 15 - 20
packages/tui/src/routes/session/index.tsx

@@ -22,7 +22,7 @@ import { useData } from "../../context/data"
 import { SplitBorder } from "../../ui/border"
 import { useTuiPaths, useTuiTerminalEnvironment } from "../../context/runtime"
 import { Spinner } from "../../component/spinner"
-import { createSyntaxStyleMemo, generateSubtleSyntax, selectedForeground, useTheme } from "../../context/theme"
+import { createSyntaxStyleMemo, generateSubtleSyntax, useTheme } from "../../context/theme"
 import { BoxRenderable, ScrollBoxRenderable, addDefaultParsers, TextAttributes, RGBA } from "@opentui/core"
 import { Prompt, type PromptRef } from "../../component/prompt"
 import type {
@@ -1384,9 +1384,9 @@ function UserMessage(props: { message: SessionMessageUser }) {
   const { theme } = useTheme()
   const [hover, setHover] = createSignal(false)
   const color = createMemo(() => local.agent.color(data.session.get(ctx.sessionID)?.agent ?? "build"))
-  const queued = createMemo(() => props.message.metadata?.queued === true)
-  const queuedFg = createMemo(() => selectedForeground(theme, color()))
-  const metadataVisible = createMemo(() => queued() || ctx.showTimestamps())
+  const queued = createMemo(
+    () => data.session.status(ctx.sessionID) === "running" && data.session.input.has(ctx.sessionID, props.message.id),
+  )
   const dialog = useDialog()
   const renderer = useRenderer()
 
@@ -1395,7 +1395,7 @@ function UserMessage(props: { message: SessionMessageUser }) {
       <box
         id={props.message.id}
         border={["left"]}
-        borderColor={color()}
+        borderColor={queued() ? theme.textMuted : color()}
         customBorderChars={SplitBorder.customBorderChars}
       >
         <box
@@ -1417,15 +1417,19 @@ function UserMessage(props: { message: SessionMessageUser }) {
         >
           <text fg={theme.text}>{props.message.text}</text>
           <Show when={files().length}>
-            <box flexDirection="row" paddingBottom={metadataVisible() ? 1 : 0} paddingTop={1} gap={1} flexWrap="wrap">
+            <box
+              flexDirection="row"
+              paddingBottom={ctx.showTimestamps() ? 1 : 0}
+              paddingTop={1}
+              gap={1}
+              flexWrap="wrap"
+            >
               <For each={files()}>
                 {(file) => {
                   const label = file.mime === "application/x-directory" ? "Directory" : file.mime
                   return (
                     <text fg={theme.text}>
-                      <span style={{ bg: theme.secondary, fg: theme.background }}>
-                        {` ${label} `}
-                      </span>
+                      <span style={{ bg: theme.secondary, fg: theme.background }}>{` ${label} `}</span>
                       <span style={{ bg: theme.backgroundElement, fg: theme.textMuted }}>
                         {" "}
                         {file.name ?? (file.source.type === "uri" ? file.source.uri : "attachment")}{" "}
@@ -1436,18 +1440,9 @@ function UserMessage(props: { message: SessionMessageUser }) {
               </For>
             </box>
           </Show>
-          <Show
-            when={queued()}
-            fallback={
-              <Show when={ctx.showTimestamps()}>
-                <text fg={theme.textMuted}>
-                  <span style={{ fg: theme.textMuted }}>{Locale.todayTimeOrDateTime(props.message.time.created)}</span>
-                </text>
-              </Show>
-            }
-          >
+          <Show when={ctx.showTimestamps()}>
             <text fg={theme.textMuted}>
-              <span style={{ bg: color(), fg: queuedFg(), bold: true }}> QUEUED </span>
+              <span style={{ fg: theme.textMuted }}>{Locale.todayTimeOrDateTime(props.message.time.created)}</span>
             </text>
           </Show>
         </box>

+ 31 - 27
packages/tui/src/routes/session/rows.ts

@@ -27,8 +27,9 @@ export function createSessionRows(sessionID: Accessor<string>) {
 
   function reduce() {
     const messages = data.session.message.list(sessionID())
+    const inputs = new Set(data.session.input.list(sessionID()))
     const boundary = revertBoundary()
-    const rows = reduceSessionRows(boundary ? messages.filter((message) => message.id < boundary) : messages)
+    const rows = reduceSessionRows(boundary ? messages.filter((message) => message.id < boundary) : messages, inputs)
     partitionPending(rows, pendingPermissions())
     return rows
   }
@@ -77,7 +78,13 @@ export function createSessionRows(sessionID: Accessor<string>) {
           .list(sessionID())
           .flatMap((message) =>
             message.type === "user"
-              ? [{ id: message.id, created: message.time.created, queued: message.metadata?.queued === true }]
+              ? [
+                  {
+                    id: message.id,
+                    created: message.time.created,
+                    input: data.session.input.has(sessionID(), message.id),
+                  },
+                ]
               : [],
           ),
       () => setRows(reconcile(reduce())),
@@ -132,8 +139,7 @@ export function createSessionRows(sessionID: Accessor<string>) {
     )
 
   const isQueued = (messageID: string) => {
-    const message = data.session.message.get(sessionID(), messageID)
-    return message?.type === "user" && message.metadata?.queued === true
+    return data.session.input.has(sessionID(), messageID)
   }
 
   const queuedStart = (rows: SessionRow[]) => {
@@ -191,30 +197,28 @@ export function createSessionRows(sessionID: Accessor<string>) {
   return rows
 }
 
-export function reduceSessionRows(messages: SessionMessage[]) {
-  return [...messages.filter((message) => !isQueuedMessage(message)), ...messages.filter(isQueuedMessage)].reduce<
-    SessionRow[]
-  >((rows, message) => {
-    if (message.type !== "assistant") {
-      if (message.type === "synthetic" && !message.description?.trim()) return rows
-      if (!isQueuedMessage(message)) completePrevious(rows)
-      rows.push({ type: "message", messageID: message.id })
+export function reduceSessionRows(messages: SessionMessage[], inputs = new Set<string>()) {
+  const isInput = (message: SessionMessage) => inputs.has(message.id)
+  return [...messages.filter((message) => !isInput(message)), ...messages.filter(isInput)].reduce<SessionRow[]>(
+    (rows, message) => {
+      if (message.type !== "assistant") {
+        if (message.type === "synthetic" && !message.description?.trim()) return rows
+        if (!isInput(message)) completePrevious(rows)
+        rows.push({ type: "message", messageID: message.id })
+        return rows
+      }
+      message.content.forEach((part) => {
+        if ((part.type === "text" || part.type === "reasoning") && !part.text.trim()) return
+        append(rows, { messageID: message.id, partID: part.id }, part)
+      })
+      if ((message.finish && !["tool-calls", "unknown"].includes(message.finish)) || message.error) {
+        completePrevious(rows)
+        rows.push({ type: "assistant-footer", messageID: message.id })
+      }
       return rows
-    }
-    message.content.forEach((part) => {
-      if ((part.type === "text" || part.type === "reasoning") && !part.text.trim()) return
-      append(rows, { messageID: message.id, partID: part.id }, part)
-    })
-    if ((message.finish && !["tool-calls", "unknown"].includes(message.finish)) || message.error) {
-      completePrevious(rows)
-      rows.push({ type: "assistant-footer", messageID: message.id })
-    }
-    return rows
-  }, [])
-}
-
-function isQueuedMessage(message: SessionMessage) {
-  return message.type === "user" && message.metadata?.queued === true
+    },
+    [],
+  )
 }
 
 function append(rows: SessionRow[], ref: PartRef, part: SessionMessageAssistant["content"][number]) {

+ 8 - 6
packages/tui/test/cli/tui/data.test.tsx

@@ -403,8 +403,7 @@ test("connectedOnce is false until first connect and persists across disconnect"
 test("tracks session status from active sessions and execution events", async () => {
   const events = createEventStream()
   const calls = createFetch((url) => {
-    if (url.pathname === "/api/session/active")
-      return json({ data: { "session-active": { type: "running" } } })
+    if (url.pathname === "/api/session/active") return json({ data: { "session-active": { type: "running" } } })
   }, events)
   let data!: ReturnType<typeof useData>
 
@@ -1082,7 +1081,7 @@ test("settles pending tools when a live failure arrives", async () => {
   }
 })
 
-test("renders admitted prompts immediately with queued marker and clears when promoted", async () => {
+test("renders admitted prompts immediately and tracks them until promoted", async () => {
   const events = createEventStream()
   const sessionID = "session-1"
   const messageID = "msg_user_1"
@@ -1135,10 +1134,12 @@ test("renders admitted prompts immediately with queued marker and clears when pr
     })
     await wait(() => sync.session.message.list(sessionID)?.length === 1)
     const admitted = sync.session.message.list(sessionID)?.[0]
-    expect(admitted).toMatchObject({ id: messageID, type: "user", text: "hello", metadata: { queued: true } })
+    expect(admitted).toMatchObject({ id: messageID, type: "user", text: "hello" })
+    expect(admitted?.metadata).toBeUndefined()
+    expect(sync.session.input.list(sessionID)).toEqual([messageID])
 
     await sync.session.message.refresh(sessionID)
-    expect(sync.session.message.list(sessionID)?.[0]?.metadata?.queued).toBeUndefined()
+    expect(sync.session.message.list(sessionID)?.[0]?.metadata).toBeUndefined()
 
     emitEvent(events, {
       id: "evt_prompted_1",
@@ -1158,7 +1159,8 @@ test("renders admitted prompts immediately with queued marker and clears when pr
     expect(message?.type).toBe("user")
     if (message?.type !== "user") return
     expect(message).toMatchObject({ id: messageID, text: "hello" })
-    expect(message.metadata?.queued).toBeUndefined()
+    expect(message.metadata).toBeUndefined()
+    expect(sync.session.input.list(sessionID)).toEqual([])
     expect(sync.session.message.ids(sessionID)).toEqual([messageID])
     expect(sync.session.message.ids("missing")).toEqual([])
     expect(sync.session.message.get(sessionID, messageID)).toBe(message)