diff --git a/packages/app/e2e/performance/unit/mock-server.test.ts b/packages/app/e2e/performance/unit/mock-server.test.ts index 83308c0a866b..3f48df03295b 100644 --- a/packages/app/e2e/performance/unit/mock-server.test.ts +++ b/packages/app/e2e/performance/unit/mock-server.test.ts @@ -44,3 +44,49 @@ test("applies message latency after a list response gate is released", async () expect(performance.now() - released).toBeGreaterThanOrEqual(20) expect(events).toEqual(["start", "before", "page", "end", "fulfill"]) }) + +test("serves current message list and lookup response contracts", async () => { + let handler: ((route: Route) => Promise) | undefined + const responses: unknown[] = [] + const message = { + info: { id: "message", role: "user" as const, time: { created: 1 } }, + parts: [{ type: "text", text: "hello" }], + } + const page = { + route: (_url: string, callback: (route: Route) => Promise) => { + handler = callback + return Promise.resolve() + }, + } as unknown as Page + await mockOpenCodeServer(page, { + protocol: "v2", + provider: {}, + directory: "C:/OpenCode", + project: {}, + sessions: [{ id: "session" }], + pageMessages: () => ({ items: [message] }), + message: () => message, + }) + + const route = (url: string) => + ({ + request: () => ({ url: () => url, method: () => "GET" }), + fulfill: (input: { body?: string }) => { + responses.push(JSON.parse(input.body ?? "null")) + return Promise.resolve() + }, + fallback: () => Promise.resolve(), + }) as unknown as Route + + await handler!(route("http://127.0.0.1:4096/api/session/session/message")) + await handler!(route("http://127.0.0.1:4096/api/session/session/message/message")) + + expect(responses).toEqual([ + { + data: [{ id: "message", type: "user", time: { created: 1 }, text: "hello" }], + context: [], + cursor: {}, + }, + { data: { id: "message", type: "user", time: { created: 1 }, text: "hello" } }, + ]) +}) diff --git a/packages/app/e2e/utils/mock-server.ts b/packages/app/e2e/utils/mock-server.ts index 76987421b607..2e6ad3d3dd71 100644 --- a/packages/app/e2e/utils/mock-server.ts +++ b/packages/app/e2e/utils/mock-server.ts @@ -261,6 +261,15 @@ export async function mockOpenCodeServer(page: Page, config: MockServerConfig) { const projectMatch = path.match(/^\/project\/([^/]+)$/) if (projectMatch) return json(route, config.project) + const currentMessageMatch = path.match(/^\/api\/session\/([^/]+)\/message\/([^/]+)$/) + if (currentMessageMatch) { + config.onMessage?.({ sessionID: currentMessageMatch[1]!, messageID: currentMessageMatch[2]! }) + if (config.messageDelay !== undefined) await new Promise((resolve) => setTimeout(resolve, config.messageDelay)) + const message = config.message?.(currentMessageMatch[1]!, currentMessageMatch[2]!) + if (message === undefined) return json(route, { error: "Message not found" }, undefined, 404) + return json(route, { data: currentMessage(message) }) + } + const messageMatch = path.match(/^\/session\/([^/]+)\/message\/([^/]+)$/) if (messageMatch) { config.onMessage?.({ sessionID: messageMatch[1]!, messageID: messageMatch[2]! }) @@ -288,6 +297,7 @@ export async function mockOpenCodeServer(page: Page, config: MockServerConfig) { if (cursor) cursors.set(cursor, pageData.cursor!) return json(route, { data: pageData.items.map(currentMessage).reverse(), + context: [], cursor: { next: cursor }, }) } @@ -306,7 +316,12 @@ export async function mockOpenCodeServer(page: Page, config: MockServerConfig) { if (!pageData.cursor) return json(route, pageData.items) const cursor = `cursor_${++nextCursor}` cursors.set(cursor, pageData.cursor) - return json(route, pageData.items, { "x-next-cursor": cursor }) + const next = new URL(url) + next.searchParams.delete("after") + next.searchParams.delete("oldest") + next.searchParams.set("limit", limit.toString()) + next.searchParams.set("before", cursor) + return json(route, pageData.items, { "x-next-cursor": cursor, Link: `<${next.toString()}>; rel="next"` }) } if (url.port === targetPort && targetPort !== appPort) return json(route, {}) @@ -429,7 +444,7 @@ function json(route: Route, body: unknown, headers?: Record, sta contentType: "application/json", headers: { "access-control-allow-origin": "*", - "access-control-expose-headers": "x-next-cursor", + "access-control-expose-headers": "Link, X-Next-Cursor", ...headers, }, body: JSON.stringify(body ?? null), diff --git a/packages/app/src/components/session/session-context-tab.tsx b/packages/app/src/components/session/session-context-tab.tsx index a0758f3eacae..d0eff91f0a75 100644 --- a/packages/app/src/components/session/session-context-tab.tsx +++ b/packages/app/src/components/session/session-context-tab.tsx @@ -18,6 +18,7 @@ import { useLanguage } from "@/context/language" import { useProviders } from "@/hooks/use-providers" import { useSDK } from "@/context/sdk" import { useSessionLayout } from "@/pages/session/session-layout" +import { selectVisibleUserMessages } from "@/pages/session/timeline/model" import { getSessionContext } from "./session-context-metrics" import { estimateSessionContextBreakdown, type SessionContextBreakdownKey } from "./session-context-breakdown" import { createSessionContextFormatter } from "./session-context-format" @@ -113,19 +114,8 @@ export function SessionContextTab() { { equals: same }, ) - const userMessages = createMemo( - () => messages().filter((m) => m.role === "user") as UserMessage[], - emptyUserMessages, - { equals: same }, - ) - const visibleUserMessages = createMemo( - () => { - const revert = info()?.revert?.messageID - if (!revert) return userMessages() - const boundary = userMessages().findIndex((message) => message.id === revert) - return boundary < 0 ? userMessages() : userMessages().slice(0, boundary) - }, + () => selectVisibleUserMessages(messages(), info()?.revert), emptyUserMessages, { equals: same }, ) diff --git a/packages/app/src/context/directory-sync.ts b/packages/app/src/context/directory-sync.ts index befd5b61e095..3327ba8e9d52 100644 --- a/packages/app/src/context/directory-sync.ts +++ b/packages/app/src/context/directory-sync.ts @@ -17,6 +17,7 @@ const sessionFields = new Set([ "question", "message", "session_message", + "revert_preview", "part", "part_text_accum_delta", ]) diff --git a/packages/app/src/context/global-sync/bootstrap.test.ts b/packages/app/src/context/global-sync/bootstrap.test.ts index a95706a3dbe5..ccae7a260fd8 100644 --- a/packages/app/src/context/global-sync/bootstrap.test.ts +++ b/packages/app/src/context/global-sync/bootstrap.test.ts @@ -70,6 +70,7 @@ function directoryState() { limit: 5, message: {}, session_message: {}, + revert_preview: {}, part: {}, part_text_accum_delta: {}, }) diff --git a/packages/app/src/context/global-sync/child-store.ts b/packages/app/src/context/global-sync/child-store.ts index 4eaa785789fc..83f0b7270697 100644 --- a/packages/app/src/context/global-sync/child-store.ts +++ b/packages/app/src/context/global-sync/child-store.ts @@ -256,6 +256,7 @@ export function createChildStoreManager(input: { limit: 5, message: {}, session_message: {}, + revert_preview: {}, part: {}, part_text_accum_delta: {}, }) diff --git a/packages/app/src/context/global-sync/event-reducer.test.ts b/packages/app/src/context/global-sync/event-reducer.test.ts index 06d536618e05..008af05899ce 100644 --- a/packages/app/src/context/global-sync/event-reducer.test.ts +++ b/packages/app/src/context/global-sync/event-reducer.test.ts @@ -81,6 +81,7 @@ const baseState = (input: Partial = {}) => limit: 10, message: {}, session_message: {}, + revert_preview: {}, part: {}, part_text_accum_delta: {}, ...input, diff --git a/packages/app/src/context/global-sync/session-cache.test.ts b/packages/app/src/context/global-sync/session-cache.test.ts index 45fbe38abe73..29e6fb446749 100644 --- a/packages/app/src/context/global-sync/session-cache.test.ts +++ b/packages/app/src/context/global-sync/session-cache.test.ts @@ -2,6 +2,7 @@ import { describe, expect, test } from "bun:test" import type { Message, Part, PermissionRequest, QuestionRequest, SessionStatus, Todo } from "@opencode-ai/sdk/v2/client" import type { FileDiffInfo } from "@opencode-ai/client/promise" import { dropSessionCaches, pickSessionCacheEvictions } from "./session-cache" +import type { RevertPreview } from "./types" const msg = (id: string, sessionID: string) => ({ @@ -30,6 +31,7 @@ describe("app session cache", () => { todo: Record message: Record session_message: Record + revert_preview: Record part: Record permission: Record question: Record @@ -40,6 +42,7 @@ describe("app session cache", () => { todo: { ses_1: [] as Todo[] }, message: {}, session_message: {}, + revert_preview: { ses_1: { messageID: "msg_1", userCount: 0, hasMore: false, items: [] } }, part: { msg_1: [part("prt_1", "ses_1", "msg_1")] }, permission: { ses_1: [] as PermissionRequest[] }, question: { ses_1: [] as QuestionRequest[] }, @@ -56,6 +59,7 @@ describe("app session cache", () => { expect(store.session_status.ses_1).toBeUndefined() expect(store.permission.ses_1).toBeUndefined() expect(store.question.ses_1).toBeUndefined() + expect(store.revert_preview.ses_1).toBeUndefined() }) test("dropSessionCaches clears message-backed parts", () => { @@ -66,6 +70,7 @@ describe("app session cache", () => { todo: Record message: Record session_message: Record + revert_preview: Record part: Record permission: Record question: Record @@ -76,6 +81,7 @@ describe("app session cache", () => { todo: {}, message: { ses_1: [m] }, session_message: {}, + revert_preview: {}, part: { [m.id]: [part("prt_1", "ses_1", m.id)] }, permission: {}, question: {}, diff --git a/packages/app/src/context/global-sync/session-cache.ts b/packages/app/src/context/global-sync/session-cache.ts index 7d684a5a1ab0..c046c14fe0a5 100644 --- a/packages/app/src/context/global-sync/session-cache.ts +++ b/packages/app/src/context/global-sync/session-cache.ts @@ -1,6 +1,7 @@ import type { Message, Part, PermissionRequest, QuestionRequest, SessionStatus, Todo } from "@opencode-ai/sdk/v2/client" import type { FileDiffInfo } from "@opencode-ai/client/promise" import type { SessionMessageInfo } from "@opencode-ai/client/promise" +import type { RevertPreview } from "./types" export const SESSION_CACHE_LIMIT = 40 @@ -10,6 +11,7 @@ type SessionCache = { todo: Record message: Record session_message: Record + revert_preview: Record part: Record permission: Record question: Record @@ -33,6 +35,7 @@ export function dropSessionCaches(store: SessionCache, sessionIDs: Iterable + role === "assistant" + ? ({ + id, + sessionID: "ses_1", + role: "assistant", + agent: "default", + model: { providerID: "openai", modelID: "gpt-4" }, + time: { created: Number(id.slice(2)) }, + } as unknown as Message) + : ({ + id, + sessionID: "ses_1", + role: "user", + agent: "default", + model: { providerID: "openai", modelID: "gpt-4" }, + time: { created: Number(id.slice(2)) }, + } as unknown as Message) + +const textPart = (id: string, messageID: string): Extract => ({ + id, + sessionID: "ses_1", + messageID, + type: "text", + text: id, +}) + +describe("revert page helpers", () => { + test("only treats 404 boundary fetch responses as missing boundaries", () => { + const found = { info: message("m6", "user"), parts: [textPart("p6", "m6")], cursor: "boundary" } + expect(boundaryFromMessageResponse({ data: found, error: undefined, response: { status: 200 } })).toBe(found) + expect( + boundaryFromMessageResponse({ data: undefined, error: { message: "missing" }, response: { status: 404 } }), + ).toBeUndefined() + expect(() => + boundaryFromMessageResponse({ data: undefined, error: new Error("server failed"), response: { status: 500 } }), + ).toThrow("server failed") + expect(() => boundaryFromMessageResponse({ data: undefined, error: new Error("network failed") })).toThrow( + "network failed", + ) + expect(() => boundaryFromMessageResponse({ data: undefined, error: undefined, response: { status: 200 } })).toThrow( + "missing revert boundary message", + ) + }) + + test("derives an old-server preview from full V1 history", () => { + const preview = revertPreviewFromMessages( + [ + { info: message("m5", "user"), parts: [textPart("p5", "m5")] }, + { info: message("m6", "user"), parts: [textPart("p6", "m6")] }, + { info: message("m7", "assistant"), parts: [] }, + { info: message("m8", "user"), parts: [textPart("p8", "m8")] }, + ], + { messageID: "m6" }, + ) + + expect(preview).toEqual({ + userCount: 2, + nextMessageID: "m8", + partID: undefined, + items: [ + { id: "m6", text: "p6" }, + { id: "m8", text: "p8" }, + ], + }) + }) + + test("recognizes only the SDK unsupported-version response", () => { + expect( + unsupportedServerRequest( + new Error("Request is not supported by this version of OpenCode Server (Server responded with text/html)"), + ), + ).toBe(true) + expect(unsupportedServerRequest(new Error("network failed"))).toBe(false) + }) + + test("detects when the loaded page has no visible user before revert", () => { + expect(hasVisibleUserBeforeRevert([message("m6", "user"), message("m7", "assistant")], { messageID: "m6" })).toBe( + false, + ) + expect( + hasVisibleUserBeforeRevert( + [message("m5", "user"), message("m6", "user")], + { messageID: "m6" }, + message("m6", "user"), + ), + ).toBe(true) + }) + + test("treats a user boundary as visible for part-level reverts", () => { + const messages = [message("m6", "user"), message("m7", "assistant")] + expect(hasVisibleUserBeforeRevert(messages, { messageID: "m6", partID: "p6" })).toBe(true) + expect(hasVisibleUserBeforeRevert(messages, { messageID: "m6" })).toBe(false) + }) + + test("does not walk older history for a part-level revert at the first user", async () => { + const boundaryPart = textPart("p6", "m6") + let olderCalls = 0 + + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m6", "user"), message("m7", "assistant")], + part: [ + { id: "m6", part: [boundaryPart] }, + { id: "m7", part: [] }, + ], + cursor: undefined, + complete: true, + }, + revert: { messageID: "m6", partID: "p6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [boundaryPart], cursor: "boundary" }), + fetchPage: async () => { + olderCalls += 1 + return { session: [], part: [], cursor: undefined, complete: true } + }, + }) + + expect(olderCalls).toBe(0) + expect(result.session.map((item) => item.id)).toEqual(["m6", "m7"]) + }) + + test("loads and merges an older boundary window when latest page is fully reverted", async () => { + const olderPart = textPart("p5", "m5") + const boundaryPart = textPart("p6", "m6") + + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m6", "user"), message("m7", "assistant"), message("m8", "user")], + part: [ + { id: "m6", part: [boundaryPart] }, + { id: "m7", part: [] }, + { id: "m8", part: [] }, + ], + cursor: undefined, + complete: true, + }, + revert: { messageID: "m6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [boundaryPart], cursor: "boundary" }), + fetchPage: async (before) => { + expect(before).toBe("boundary") + return { + session: [message("m4", "assistant"), message("m5", "user")], + part: [ + { id: "m4", part: [] }, + { id: "m5", part: [olderPart] }, + ], + cursor: "older", + complete: false, + } + }, + }) + + expect(result.session.map((item) => item.id)).toEqual(["m4", "m5", "m6", "m7", "m8"]) + expect(result.part.find((item) => item.id === "m5")?.part).toEqual([olderPart]) + expect(result.part.find((item) => item.id === "m6")?.part).toEqual([boundaryPart]) + expect(result.cursor).toBe("older") + expect(result.complete).toBe(false) + }) + + test("merges the fetched boundary when the current page already has a visible prior user", async () => { + const boundaryPart = textPart("p6", "m6") + let olderCalls = 0 + + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m5", "user"), message("m7", "assistant")], + part: [ + { id: "m5", part: [] }, + { id: "m7", part: [] }, + ], + cursor: undefined, + complete: true, + }, + revert: { messageID: "m6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [boundaryPart], cursor: "boundary" }), + fetchPage: async () => { + olderCalls += 1 + return { session: [], part: [], cursor: undefined, complete: true } + }, + }) + + expect(olderCalls).toBe(0) + expect(result.session.map((item) => item.id)).toEqual(["m5", "m6", "m7"]) + expect(result.part.find((item) => item.id === "m6")?.part).toEqual([boundaryPart]) + }) + + test("merges the fetched boundary when no older boundary cursor exists", async () => { + const boundaryPart = textPart("p6", "m6") + + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m7", "assistant")], + part: [{ id: "m7", part: [] }], + cursor: undefined, + complete: true, + }, + revert: { messageID: "m6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [boundaryPart] }), + fetchPage: async () => ({ session: [], part: [], cursor: undefined, complete: true }), + }) + + expect(result.session.map((item) => item.id)).toEqual(["m6", "m7"]) + expect(result.part.find((item) => item.id === "m6")?.part).toEqual([boundaryPart]) + }) + + test("walks the latest cursor when an older server omits the boundary cursor", async () => { + let before: string | undefined + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m7", "assistant")], + part: [{ id: "m7", part: [] }], + cursor: "raw-hidden-older", + complete: false, + }, + revert: { messageID: "m6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [] }), + fetchPage: async (cursor) => { + before = cursor + return { session: [message("m5", "user")], part: [], cursor: undefined, complete: true } + }, + }) + + expect(before).toBe("raw-hidden-older") + expect(result.session.map((item) => item.id)).toEqual(["m5", "m6", "m7"]) + expect(result.cursor).toBeUndefined() + expect(result.complete).toBe(true) + }) + + test("keeps loading older pages until a visible user exists before revert", async () => { + const boundaryPart = textPart("p6", "m6") + const olderPart = textPart("p3", "m3") + let call = 0 + + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m6", "user"), message("m7", "assistant"), message("m8", "user")], + part: [ + { id: "m6", part: [boundaryPart] }, + { id: "m7", part: [] }, + { id: "m8", part: [] }, + ], + cursor: undefined, + complete: true, + }, + revert: { messageID: "m6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [boundaryPart], cursor: "boundary" }), + fetchPage: async () => { + call += 1 + if (call === 1) { + return { + session: [message("m4", "assistant"), message("m5", "assistant")], + part: [ + { id: "m4", part: [] }, + { id: "m5", part: [] }, + ], + cursor: "older-2", + complete: false, + } + } + return { + session: [message("m3", "user")], + part: [{ id: "m3", part: [olderPart] }], + cursor: undefined, + complete: true, + } + }, + }) + + expect(call).toBe(2) + expect(result.session.map((item) => item.id)).toEqual(["m3", "m4", "m5", "m6", "m7", "m8"]) + expect(result.part.find((item) => item.id === "m3")?.part).toEqual([olderPart]) + expect(result.cursor).toBeUndefined() + expect(result.complete).toBe(true) + }) + + test("rejects a repeated older page instead of looping", async () => { + let calls = 0 + const current = { + session: [message("m6", "user"), message("m7", "assistant")], + part: [], + cursor: "repeat", + complete: false, + } + + await expect( + loadRevertAwareLatestPage({ + current, + revert: { messageID: "m6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [], cursor: "repeat" }), + fetchPage: async () => { + calls += 1 + if (calls > 1) throw new Error("pagination did not stop") + return current + }, + }), + ).rejects.toThrow("Message pagination returned no new messages") + expect(calls).toBe(1) + }) + + test("keeps an unseen page that completes revert repair even when its cursor repeats", async () => { + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m6", "user"), message("m7", "assistant")], + part: [], + cursor: "repeat", + complete: false, + }, + revert: { messageID: "m6" }, + fetchMessage: async () => ({ info: message("m6", "user"), parts: [], cursor: "repeat" }), + fetchPage: async () => ({ + session: [message("m5", "user")], + part: [], + cursor: "repeat", + complete: false, + }), + }) + + expect(result.session.map((message) => message.id)).toEqual(["m5", "m6", "m7"]) + }) + + test("marks stale revert boundaries for clearing when the boundary message is missing", async () => { + const result = await loadRevertAwareLatestPage({ + current: { + session: [message("m7", "assistant"), message("m8", "user")], + part: [ + { id: "m7", part: [] }, + { id: "m8", part: [] }, + ], + cursor: undefined, + complete: true, + }, + revert: { messageID: "m6" }, + fetchMessage: async () => undefined, + fetchPage: async () => ({ + session: [], + part: [], + cursor: undefined, + complete: true, + }), + }) + + expect(result.clearedRevert).toBe(true) + expect(result.session.map((item) => item.id)).toEqual(["m7", "m8"]) + }) +}) diff --git a/packages/app/src/context/revert-page.ts b/packages/app/src/context/revert-page.ts new file mode 100644 index 000000000000..5182561b13b5 --- /dev/null +++ b/packages/app/src/context/revert-page.ts @@ -0,0 +1,104 @@ +import type { Message, Part } from "@opencode-ai/sdk/v2/client" + +type MessagePage = { + session: Message[] + part: { id: string; part: Part[] }[] + cursor?: string + complete: boolean + clearedRevert?: boolean +} + +type MessageWithParts = { + info: Message + parts: Part[] + cursor?: string +} + +export type RevertTarget = { + messageID: string + partID?: string +} + +const cmp = (a: string, b: string) => (a < b ? -1 : a > b ? 1 : 0) +export const compareMessages = (a: Message, b: Message) => a.time.created - b.time.created || cmp(a.id, b.id) +export const messageBefore = (message: Message, boundary: Message) => compareMessages(message, boundary) < 0 + +const sortParts = (parts: Part[]) => parts.filter((part) => !!part?.id).sort((a, b) => cmp(a.id, b.id)) + +// A part-level revert keeps the boundary message itself visible (with trimmed parts), so a +// user boundary counts as the visible user and no older history is required to render it. +export function hasVisibleUserBeforeRevert(messages: readonly Message[], revert?: RevertTarget, boundary?: Message) { + if (!revert?.messageID) return true + const resolved = boundary ?? messages.find((message) => message.id === revert.messageID) + if (!resolved) return false + return messages.some( + (message) => + message.role === "user" && (messageBefore(message, resolved) || (message.id === resolved.id && !!revert.partID)), + ) +} + +function mergeMessages(current: Message[], older: Message[], boundary: Message) { + const merged = new Map(current.filter((message) => !!message?.id).map((message) => [message.id, message] as const)) + for (const message of older) { + if (!message?.id) continue + merged.set(message.id, message) + } + merged.set(boundary.id, boundary) + return [...merged.values()].sort(compareMessages) +} + +function mergeParts(current: MessagePage["part"], older: MessagePage["part"], boundary: MessageWithParts) { + const merged = new Map(current.filter((item) => !!item?.id).map((item) => [item.id, sortParts(item.part)] as const)) + for (const item of older) { + if (!item?.id) continue + merged.set(item.id, sortParts(item.part)) + } + merged.set(boundary.info.id, sortParts(boundary.parts)) + return [...merged.entries()].sort((a, b) => cmp(a[0], b[0])).map(([id, part]) => ({ id, part })) +} + +export async function loadRevertAwareLatestPage(input: { + current: MessagePage + revert?: RevertTarget + fetchMessage: (messageID: string) => Promise + fetchPage: (before: string) => Promise +}) { + if (!input.revert?.messageID) return input.current + + const boundary = await input.fetchMessage(input.revert.messageID) + if (!boundary) return { ...input.current, clearedRevert: true } + const current = { + ...input.current, + session: mergeMessages(input.current.session, [], boundary.info), + part: mergeParts(input.current.part, [], boundary), + cursor: boundary.cursor ?? input.current.cursor, + complete: !boundary.cursor && !input.current.cursor, + } + if (hasVisibleUserBeforeRevert(current.session, input.revert, boundary.info)) return current + if (!current.cursor) return current + + const cursors = new Set() + const ids = new Set(current.session.map((message) => message.id)) + let older: MessagePage = current + let session = current.session + let part = current.part + let cursor: string | undefined = current.cursor + while (!hasVisibleUserBeforeRevert(session, input.revert, boundary.info) && cursor) { + if (cursors.has(cursor)) throw new Error("Message pagination cursor did not advance") + cursors.add(cursor) + older = await input.fetchPage(cursor) + const unseen = older.session.filter((message) => !ids.has(message.id)) + unseen.forEach((message) => ids.add(message.id)) + session = mergeMessages(session, older.session, boundary.info) + part = mergeParts(part, older.part, boundary) + cursor = older.cursor + if (!hasVisibleUserBeforeRevert(session, input.revert, boundary.info) && cursor && unseen.length === 0) + throw new Error("Message pagination returned no new messages") + } + return { + session, + part, + cursor: older.cursor, + complete: older.complete, + } +} diff --git a/packages/app/src/context/server-session-v2-reducer.test.ts b/packages/app/src/context/server-session-v2-reducer.test.ts index 00cc37cf5240..d664e4eafc3f 100644 --- a/packages/app/src/context/server-session-v2-reducer.test.ts +++ b/packages/app/src/context/server-session-v2-reducer.test.ts @@ -153,4 +153,66 @@ describe("v2 session reducer", () => { expect(result).toMatchObject({ sessionID: "ses_1", missing: "msg_user", touched: [] }) }) + + test("continues canonical hydrated shell and pending tool messages", () => { + const reducer = createV2SessionReducer() + let messages = [ + { + id: "msg_shell", + type: "shell", + callID: "shell_1", + command: "pwd", + output: "", + time: { created: 1 }, + }, + { + id: "msg_assistant", + type: "assistant", + agent: "build", + model: { id: "model", providerID: "provider" }, + content: [ + { + type: "tool", + id: "call_1", + name: "read", + state: { status: "pending", input: '{"filePath":"REA' }, + time: { created: 1 }, + }, + ], + time: { created: 1 }, + }, + ] as unknown as SessionMessageInfo[] + const apply = (input: object) => { + const result = reducer.reduce(messages, event(input)) + if (result) messages = result.messages + } + + apply({ + ...base, + id: "evt_shell_end", + type: "session.shell.ended", + data: { + sessionID: "ses_1", + shell: { id: "shell_1", status: "exited", exit: 0 }, + output: { output: "/repo", cursor: 5, size: 5, truncated: false }, + }, + }) + apply({ + ...base, + id: "evt_tool_delta", + type: "session.tool.input.delta", + data: { sessionID: "ses_1", assistantMessageID: "msg_assistant", callID: "call_1", delta: 'DME.md"}' }, + }) + + expect(messages[0]).toMatchObject({ + type: "shell", + status: "exited", + output: { output: "/repo" }, + time: { completed: 1 }, + }) + expect(messages[1]).toMatchObject({ + type: "assistant", + content: [{ type: "tool", state: { input: '{"filePath":"README.md"}' } }], + }) + }) }) diff --git a/packages/app/src/context/server-session-v2-reducer.ts b/packages/app/src/context/server-session-v2-reducer.ts index b34ab3985ec1..80fb441650fb 100644 --- a/packages/app/src/context/server-session-v2-reducer.ts +++ b/packages/app/src/context/server-session-v2-reducer.ts @@ -105,7 +105,9 @@ export function createV2SessionReducer() { case "session.shell.ended": return updateMessage( source, - (item): item is Shell => item.type === "shell" && item.shellID === event.data.shell.id, + (item): item is Shell => + item.type === "shell" && + (("callID" in item && item.callID === event.data.shell.id) || item.shellID === event.data.shell.id), (item) => ({ ...item, status: event.data.shell.status, @@ -255,15 +257,15 @@ export function createV2SessionReducer() { ], })) case "session.tool.input.delta": - return updateTool(source, event.data.assistantMessageID, event.data.callID, sessionID, (tool) => - tool.state.status === "streaming" - ? { ...tool, state: { ...tool.state, input: tool.state.input + event.data.delta } } - : tool, - ) + return updateTool(source, event.data.assistantMessageID, event.data.callID, sessionID, (tool) => { + if (!["streaming", "pending"].includes(tool.state.status) || typeof tool.state.input !== "string") return tool + return { ...tool, state: { status: "streaming", input: tool.state.input + event.data.delta } } + }) case "session.tool.input.ended": - return updateTool(source, event.data.assistantMessageID, event.data.callID, sessionID, (tool) => - tool.state.status === "streaming" ? { ...tool, state: { ...tool.state, input: event.data.text } } : tool, - ) + return updateTool(source, event.data.assistantMessageID, event.data.callID, sessionID, (tool) => { + if (!["streaming", "pending"].includes(tool.state.status) || typeof tool.state.input !== "string") return tool + return { ...tool, state: { status: "streaming", input: event.data.text } } + }) case "session.tool.called": return updateTool(source, event.data.assistantMessageID, event.data.callID, sessionID, (tool) => ({ ...tool, @@ -303,7 +305,7 @@ export function createV2SessionReducer() { }) case "session.tool.failed": return updateTool(source, event.data.assistantMessageID, event.data.callID, sessionID, (tool) => { - if (tool.state.status !== "streaming" && tool.state.status !== "running") return tool + if (typeof tool.state.input !== "string" && tool.state.status !== "running") return tool return { ...tool, executed: event.data.executed || tool.executed === true, diff --git a/packages/app/src/context/server-session.test.ts b/packages/app/src/context/server-session.test.ts index 2ebf5f88ef0b..e0b5307d2105 100644 --- a/packages/app/src/context/server-session.test.ts +++ b/packages/app/src/context/server-session.test.ts @@ -140,9 +140,12 @@ const retryImmediately: typeof retry = async (task, options = {}) => { } } +const tick = () => new Promise((resolve) => setTimeout(resolve, 0)) + function setup(sessions: Record) { const get: unknown[] = [] const messages: unknown[] = [] + const message: unknown[] = [] const client = { session: { get: async (input: unknown) => { @@ -154,11 +157,39 @@ function setup(sessions: Record) { messages.push(input) return response() }, + message: async (input: unknown) => { + message.push(input) + return { data: { info: userMessage("boundary", { sessionID: "root", time: { created: 2 } }), parts: [] } } + }, diff: async () => ({ data: [] }), todo: async () => ({ data: [] }), }, } as unknown as OpencodeClient - return { get, messages, store: createServerSession(client) } + return { get, message, messages, store: createServerSession(client) } +} + +function revertClient( + pages: MessageResponse[], + boundary = userMessage("boundary", { sessionID: "root", time: { created: 2 } }), + info: Session = session("root"), +) { + let index = 0 + const messages: unknown[] = [] + const message: unknown[] = [] + const client = { + session: { + get: async () => ({ data: info }), + messages: async (input: unknown) => { + messages.push(input) + return pages[index++] ?? response() + }, + message: async (input: unknown) => { + message.push(input) + return { data: { info: boundary, parts: [] } } + }, + }, + } as unknown as OpencodeClient + return { client, message, messages } } describe("server session", () => { @@ -242,7 +273,7 @@ describe("server session", () => { agent: "build", model: { id: "model", providerID: "provider" }, content: [{ type: "text", text: "hi" }], - time: { created: 2, completed: 3 }, + time: { created: 1, completed: 3 }, } const client = { session: { @@ -263,6 +294,16 @@ describe("server session", () => { await store.sync("root") expect(requests).toEqual([{ sessionID: "root", limit: 20, order: "desc" }]) + expect(store.data.session_message.root.map((message) => message.id)).toEqual([user.id, assistant.id]) + + store.applyV2({ + id: "evt_text_delta", + created: 4, + type: "session.text.delta", + location: { directory: "/repo" }, + data: { sessionID: "root", assistantMessageID: assistant.id, ordinal: 0, delta: " again" }, + }) + expect(store.data.session_message.root.map((message) => message.id)).toEqual([user.id, assistant.id]) expect(store.data.message.root.map((message) => message.id)).toEqual([user.id, assistant.id]) }) @@ -309,9 +350,225 @@ describe("server session", () => { expect(assistants.map((item) => store.data.part[item.id]?.[0]?.type)).toEqual(["text", "text", "text"]) }) + test("walks current pages until a whole-message revert has a visible user", async () => { + const older = { id: "msg_1_user", type: "user", text: "older", time: { created: 1 } } as const + const boundary = { id: "msg_2_user", type: "user", text: "boundary", time: { created: 2 } } as const + const pages = [ + { data: [boundary], cursor: { previous: null, next: "older" } }, + { data: [older], cursor: { previous: null, next: null } }, + ] + const requests: unknown[] = [] + const messageApi = { + list: async (input: unknown) => { + requests.push(input) + return pages.shift()! + }, + } as unknown as MessageApi + const store = createServerSession({} as OpencodeClient, {} as SessionApi, messageApi) + store.remember({ ...session("root"), revert: { messageID: boundary.id } }) + + await store.sync("root") + + expect(requests).toEqual([ + { sessionID: "root", limit: 20, order: "desc" }, + { sessionID: "root", limit: 20, cursor: "older" }, + ]) + expect(store.data.message.root.map((message) => message.id)).toEqual([older.id, boundary.id]) + expect(store.data.info.root?.revert).toEqual({ messageID: boundary.id }) + }) + + test("rejects repeated current pages during revert repair", async () => { + const boundary = { id: "msg_boundary", type: "user", text: "boundary", time: { created: 2 } } as const + let requests = 0 + const messageApi = { + list: async () => { + requests += 1 + if (requests > 2) throw new Error("pagination did not stop") + return { data: [boundary], cursor: { previous: null, next: "repeat" } } + }, + } as unknown as MessageApi + const store = createServerSession({} as OpencodeClient, {} as SessionApi, messageApi) + store.remember({ ...session("root"), revert: { messageID: boundary.id } }) + + await expect(store.sync("root")).rejects.toThrow("Message pagination returned no new messages") + expect(requests).toBe(2) + }) + + test("rejects current revert pages that contain no unseen messages", async () => { + const boundary = { id: "msg_boundary", type: "user", text: "boundary", time: { created: 2 } } as const + let requests = 0 + const messageApi = { + list: async () => { + requests += 1 + return { data: [boundary], cursor: { previous: null, next: `older-${requests}` } } + }, + } as unknown as MessageApi + const store = createServerSession({} as OpencodeClient, {} as SessionApi, messageApi) + store.remember({ ...session("root"), revert: { messageID: boundary.id } }) + + await expect(store.sync("root")).rejects.toThrow("Message pagination returned no new messages") + expect(requests).toBe(2) + }) + + test("keeps an unseen current page that completes revert repair when its cursor repeats", async () => { + const older = { id: "msg_older", type: "user", text: "older", time: { created: 1 } } as const + const boundary = { id: "msg_boundary", type: "user", text: "boundary", time: { created: 2 } } as const + const pages = [ + { data: [boundary], cursor: { previous: null, next: "repeat" } }, + { data: [older], cursor: { previous: null, next: "repeat" } }, + ] + const store = createServerSession( + {} as OpencodeClient, + {} as SessionApi, + { + list: async () => pages.shift()!, + } as unknown as MessageApi, + ) + store.remember({ ...session("root"), revert: { messageID: boundary.id } }) + + await store.sync("root") + + expect(store.data.message.root.map((message) => message.id)).toEqual([older.id, boundary.id]) + }) + + test("uses current page context and revert preview without walking older pages", async () => { + const older = { id: "msg_z_older", type: "user", text: "older", time: { created: 1 } } as const + const boundary = { id: "msg_a_boundary", type: "user", text: "boundary", time: { created: 1 } } as const + const requests: unknown[] = [] + const messageApi = { + list: async (input: unknown) => { + requests.push(input) + return { + data: [boundary], + context: [older], + contextCursor: "visible-older", + revert: { + messageID: boundary.id, + userCount: 1, + hasMore: false, + items: [{ id: boundary.id, text: boundary.text }], + }, + cursor: { previous: null, next: "raw-older" }, + } + }, + } as unknown as MessageApi + const store = createServerSession({} as OpencodeClient, {} as SessionApi, messageApi) + store.remember({ ...session("root"), revert: { messageID: boundary.id } }) + + await store.sync("root") + + expect(requests).toEqual([{ sessionID: "root", limit: 20, order: "desc" }]) + expect(store.data.message.root.map((message) => message.id)).toEqual([older.id, boundary.id]) + expect(store.data.revert_preview.root).toEqual({ + messageID: boundary.id, + userCount: 1, + hasMore: false, + items: [{ id: boundary.id, text: boundary.text }], + }) + expect(store.history.more("root")).toBe(true) + + await store.history.loadMore("root") + expect(requests[1]).toEqual({ sessionID: "root", limit: 200, cursor: "visible-older" }) + }) + + test("clears a cached current preview when the revert boundary changes", () => { + const store = createServerSession({} as OpencodeClient) + store.remember({ ...session("root"), revert: { messageID: "msg_a" } }) + store.set("revert_preview", "root", { + messageID: "msg_a", + userCount: 1, + hasMore: false, + items: [{ id: "msg_a", text: "a" }], + }) + + store.remember({ ...session("root"), revert: { messageID: "msg_b" } }) + + expect(store.data.revert_preview.root).toBeUndefined() + }) + + test("keeps the committed current prefix when refresh fails", async () => { + const boundary = { id: "msg_boundary", type: "user", text: "keep", time: { created: 1 } } as const + const removed = { id: "msg_removed", type: "user", text: "remove", time: { created: 2 } } as const + const gets: unknown[] = [] + const requests: unknown[] = [] + const sessionApi = { + get: async (input: unknown) => { + gets.push(input) + return session("root") + }, + } as unknown as SessionApi + const messageApi = { + list: async (input: unknown) => { + requests.push(input) + throw new Error("refresh failed") + }, + } as unknown as MessageApi + const store = createServerSession({} as OpencodeClient, sessionApi, messageApi) + store.remember({ ...session("root"), revert: { messageID: boundary.id } }) + store.set("message", "root", [ + userMessage(boundary.id, { sessionID: "root", time: boundary.time }), + userMessage(removed.id, { sessionID: "root", time: removed.time }), + ]) + store.set("session_message", "root", [boundary, removed]) + + store.applyV2({ + id: "evt_commit", + created: 3, + type: "session.next.revert.committed", + durable: { aggregateID: "root", seq: 3, version: 1 }, + data: { sessionID: "root", messageID: boundary.id, timestamp: 3 }, + } as unknown as OpenCodeEvent) + await tick() + await tick() + + expect(gets).toEqual([{ sessionID: "root" }]) + expect(requests).toEqual([{ sessionID: "root", limit: 20, order: "desc" }]) + expect(store.data.session_message.root.map((message) => message.id)).toEqual([boundary.id]) + expect(store.data.message.root.map((message) => message.id)).toEqual([boundary.id]) + }) + + test("clears committed revert state when the session refresh fails", async () => { + const boundary = { id: "msg_boundary", type: "user", text: "keep", time: { created: 1 } } as const + const removed = { id: "msg_removed", type: "user", text: "remove", time: { created: 2 } } as const + const requests: unknown[] = [] + const sessionApi = { + get: async () => { + throw new Error("session refresh failed") + }, + } as unknown as SessionApi + const messageApi = { + list: async (input: unknown) => { + requests.push(input) + return { data: [boundary], context: [], cursor: { previous: null, next: null } } + }, + } as unknown as MessageApi + const store = createServerSession({} as OpencodeClient, sessionApi, messageApi) + store.remember({ ...session("root"), revert: { messageID: boundary.id } }) + store.set("message", "root", [ + userMessage(boundary.id, { sessionID: "root", time: boundary.time }), + userMessage(removed.id, { sessionID: "root", time: removed.time }), + ]) + store.set("session_message", "root", [boundary, removed]) + + store.applyV2({ + id: "evt_commit", + created: 3, + type: "session.next.revert.committed", + durable: { aggregateID: "root", seq: 3, version: 1 }, + data: { sessionID: "root", messageID: boundary.id, timestamp: 3 }, + } as unknown as OpenCodeEvent) + await tick() + + expect(store.data.info.root?.revert).toBeUndefined() + expect(store.data.message.root.map((message) => message.id)).toEqual([boundary.id]) + + await store.sync("root") + expect(requests).toEqual([{ sessionID: "root", limit: 20, order: "desc" }]) + }) + test("indexes V1 messages for the current timeline projection", async () => { - const user = userMessage("message-1", { sessionID: "root" }) - const assistant = assistantMessage("message-2", user.id, { sessionID: "root" }) + const user = userMessage("message-z", { sessionID: "root", time: { created: 1 } }) + const assistant = assistantMessage("message-a", user.id, { sessionID: "root", time: { created: 2, completed: 2 } }) const client = messageClient( response([ { info: user, parts: [textPart(user.id, { sessionID: "root" })] }, @@ -336,9 +593,9 @@ describe("server session", () => { { id: assistant.id, type: "assistant" }, ]) - const next = userMessage("message-3", { sessionID: "root" }) + const next = userMessage("message-0", { sessionID: "root", time: { created: 0 } }) store.apply({ type: "message.updated", properties: { info: next } }) - expect(store.data.session_message.root.map((message) => message.id)).toEqual([user.id, assistant.id, next.id]) + expect(store.data.session_message.root.map((message) => message.id)).toEqual([next.id, user.id, assistant.id]) store.apply({ type: "message.removed", properties: { sessionID: "root", messageID: next.id } }) expect(store.data.session_message.root.map((message) => message.id)).toEqual([user.id, assistant.id]) @@ -619,6 +876,149 @@ describe("server session", () => { expect(store.data.part[assistant.id]).toEqual([live]) }) + test("keeps fetched messages in chronological order", async () => { + const first = userMessage("msg_9", { time: { created: 1 } }) + const second = userMessage("msg_2", { time: { created: 2 } }) + const third = userMessage("msg_1", { time: { created: 3 } }) + const store = createServerSession( + messageClient( + response([ + { info: third, parts: [] }, + { info: first, parts: [] }, + { info: second, parts: [] }, + ]), + ), + ) + + await store.sync("child") + + expect(store.data.message.child.map((message) => message.id)).toEqual(["msg_9", "msg_2", "msg_1"]) + }) + + test("reloads cached messages when remembered revert boundary changes", async () => { + const ctx = setup({}) + ctx.store.remember(session("root")) + ctx.store.optimistic.add({ sessionID: "root", message: userMessage("cached", { sessionID: "root" }), parts: [] }) + + ctx.store.remember({ ...session("root"), revert: { messageID: "boundary", partID: "part" } }) + await tick() + + expect(ctx.messages).toEqual([{ sessionID: "root", limit: 20, before: undefined }]) + expect(ctx.message).toEqual([{ sessionID: "root", messageID: "boundary" }]) + }) + + test("repairs cached revert sessions that do not contain a visible user before the boundary", async () => { + const cached = userMessage("cached", { sessionID: "root", time: { created: 3 } }) + const older = userMessage("older", { sessionID: "root", time: { created: 1 } }) + const { client, message, messages } = revertClient([ + response([{ info: cached, parts: [] }]), + response([{ info: older, parts: [] }]), + ]) + const store = createServerSession(client) + await store.sync("root") + + store.set("info", "root", { ...session("root"), revert: { messageID: "boundary", partID: "part" } }) + await store.sync("root") + + expect(messages).toEqual([ + { sessionID: "root", limit: 20, before: undefined }, + { sessionID: "root", limit: 1, before: undefined }, + ]) + expect(message).toEqual([{ sessionID: "root", messageID: "boundary" }]) + }) + + test("does not repair a complete part-level revert cache when the boundary user is visible", async () => { + const boundary = userMessage("boundary", { sessionID: "root", time: { created: 2 } }) + const { client, message, messages } = revertClient([response([{ info: boundary, parts: [] }])], boundary) + const store = createServerSession(client) + await store.sync("root") + + store.set("info", "root", { ...session("root"), revert: { messageID: "boundary", partID: "part" } }) + await store.sync("root") + + expect(messages).toEqual([{ sessionID: "root", limit: 20, before: undefined }]) + expect(message).toEqual([]) + }) + + test("does not repeatedly repair a complete whole-message revert at the first user", async () => { + const boundary = userMessage("boundary", { sessionID: "root", time: { created: 2 } }) + const { client, message, messages } = revertClient([response([{ info: boundary, parts: [] }])], boundary) + const store = createServerSession(client) + await store.sync("root") + + store.set("info", "root", { ...session("root"), revert: { messageID: "boundary" } }) + await store.sync("root") + + expect(messages).toEqual([{ sessionID: "root", limit: 20, before: undefined }]) + expect(message).toEqual([]) + }) + + test("repairs uncached sessions when resolve discovers a revert boundary after the first load", async () => { + const cached = userMessage("cached", { sessionID: "root", time: { created: 3 } }) + const older = userMessage("older", { sessionID: "root", time: { created: 1 } }) + const { client, message, messages } = revertClient( + [response([{ info: cached, parts: [] }]), response([{ info: older, parts: [] }])], + undefined, + { ...session("root"), revert: { messageID: "boundary" } }, + ) + const store = createServerSession(client) + + await store.sync("root") + + expect(messages).toEqual([ + { sessionID: "root", limit: 20, before: undefined }, + { sessionID: "root", limit: 1, before: undefined }, + ]) + expect(message).toEqual([{ sessionID: "root", messageID: "boundary" }]) + }) + + test("does not expose raw hidden cursors after terminal revert repair", async () => { + const newer = userMessage("newer", { sessionID: "root", time: { created: 3 } }) + const { client, message, messages } = revertClient( + [response([{ info: newer, parts: [] }], "raw-hidden-older")], + undefined, + { ...session("root"), revert: { messageID: "boundary" } }, + ) + const store = createServerSession(client) + + await store.sync("root") + + expect(messages).toEqual([ + { sessionID: "root", limit: 20, before: undefined }, + { sessionID: "root", limit: 1, before: undefined }, + ]) + expect(message).toEqual([{ sessionID: "root", messageID: "boundary" }]) + expect(store.history.more("root")).toBe(false) + }) + + test("uses a nonzero page size when repairing an empty cached revert session", async () => { + const ctx = setup({ root: session("root") }) + await ctx.store.sync("root") + + ctx.store.set("info", "root", { ...session("root"), revert: { messageID: "boundary" } }) + await ctx.store.sync("root") + + expect(ctx.messages).toEqual([ + { sessionID: "root", limit: 20, before: undefined }, + { sessionID: "root", limit: 20, before: undefined }, + ]) + expect(ctx.message).toEqual([{ sessionID: "root", messageID: "boundary" }]) + }) + + test("uses a nonzero page size when remember reloads an empty cached revert session", async () => { + const ctx = setup({ root: session("root") }) + await ctx.store.sync("root") + + ctx.store.remember({ ...session("root"), revert: { messageID: "boundary" } }) + await tick() + + expect(ctx.messages).toEqual([ + { sessionID: "root", limit: 20, before: undefined }, + { sessionID: "root", limit: 20, before: undefined }, + ]) + expect(ctx.message).toEqual([{ sessionID: "root", messageID: "boundary" }]) + }) + test("merges live events into the initial page", async () => { const pending = deferredResponse() const user = userMessage("message-1") @@ -636,6 +1036,26 @@ describe("server session", () => { expect(store.data.part[live.id]).toEqual([livePart]) }) + test("merges live events by chronology, not ID", async () => { + const pending = deferredResponse() + const first = userMessage("msg_9", { time: { created: 1 } }) + const live = userMessage("msg_2", { time: { created: 2 } }) + const third = userMessage("msg_1", { time: { created: 3 } }) + const store = createServerSession(messageClient(pending.promise)) + const loading = store.sync("child") + + store.apply({ type: "message.updated", properties: { info: live } }) + pending.resolve( + response([ + { info: first, parts: [] }, + { info: third, parts: [] }, + ]), + ) + await loading + + expect(store.data.message.child.map((message) => message.id)).toEqual(["msg_9", "msg_2", "msg_1"]) + }) + test("preserves same-ID live updates over the initial page", async () => { const pending = deferredResponse() const fetched = userMessage("message") @@ -1420,6 +1840,22 @@ describe("server session", () => { expect(store.data.message.child).toEqual([older, latest]) }) + test("stops history pagination when a page makes no progress", async () => { + const latest = userMessage("message-2", { time: { created: 2 } }) + const store = createServerSession( + messageClient( + response([{ info: latest, parts: [] }], "repeat"), + response([{ info: latest, parts: [] }], "repeat"), + ), + ) + await store.sync("child") + + await store.history.loadMore("child") + + expect(store.data.message.child).toEqual([latest]) + expect(store.history.more("child")).toBe(false) + }) + test("preserves loaded history during an incomplete refresh", async () => { const older = userMessage("message-1") const latest = userMessage("message-2", { time: { created: 2 } }) diff --git a/packages/app/src/context/server-session.ts b/packages/app/src/context/server-session.ts index c5d98682f03c..823ff050c5b4 100644 --- a/packages/app/src/context/server-session.ts +++ b/packages/app/src/context/server-session.ts @@ -1,4 +1,5 @@ import { Binary } from "@opencode-ai/core/util/binary" +import { linkParam, parseLinkHeader } from "@opencode-ai/core/util/link-header" import { retry } from "@opencode-ai/core/util/retry" import type { OpenCodeEvent, SessionApi, SessionMessageInfo } from "@opencode-ai/client/promise" import type { @@ -19,11 +20,26 @@ import { sessionNotFoundError } from "@/utils/server-errors" import { rootSession } from "@/utils/session-route" import { normalizeSessionInfo } from "@/utils/session" import { compareMessages, messageKey, normalizeSessionMessages } from "@/utils/session-message" +import { boundaryFromMessageResponse } from "@opencode-ai/core/util/revert-boundary" import { dropSessionCaches, pickSessionCacheEvictions, SESSION_CACHE_LIMIT } from "./global-sync/session-cache" import { createV2SessionReducer, type V2SessionReduction } from "./server-session-v2-reducer" import type { ServerApi } from "@/utils/server" +import { hasVisibleUserBeforeRevert, loadRevertAwareLatestPage } from "./revert-page" +import type { RevertPreview } from "./global-sync/types" -type MessageApi = ServerApi["message"] +// The vendored client snapshot predates the pagination fields on SessionMessagesResponse +// (context, contextCursor, revert). This widens the pinned types to the current wire +// contract; delete it when the vendor snapshot is refreshed. +type MessageList = ServerApi["message"]["list"] +type MessageApi = { + list: (...input: Parameters) => Promise< + Awaited> & { + context?: SessionMessageInfo[] + contextCursor?: string + revert?: RevertPreview + } + > +} const cmp = (a: string, b: string) => (a < b ? -1 : a > b ? 1 : 0) const SKIP_PARTS = new Set(["patch", "step-start", "step-finish"]) @@ -38,11 +54,25 @@ function needsOlderTurnRoot(source: readonly SessionMessageInfo[]) { message.type === "user" || message.type === "shell" || message.type === "assistant" || - (message.type === "synthetic" && message.description?.trim()), + (message.type === "synthetic" && (message.description?.trim() || message.text.trim())), ) return boundary?.type === "assistant" } +function nextBefore(link: string | null, fallback?: string | null) { + const links = parseLinkHeader(link ?? "") + return linkParam(links.next, "before") ?? linkParam(links.prev, "before") ?? fallback ?? undefined +} + +function messageIndex(messages: readonly Message[] | undefined, id: string) { + return messages?.findIndex((message) => message.id === id) ?? -1 +} + +function insertMessage(messages: Message[], message: Message) { + const result = Binary.search(messages, messageKey(message), messageKey) + if (!result.found) messages.splice(result.index, 0, message) +} + type OptimisticItem = { message: Message parts: Part[] @@ -50,6 +80,11 @@ type OptimisticItem = { confirmedMessage?: boolean } +type OptimisticStore = { + message: Record + part: Record +} + type MessagePage = { session: Message[] part: { id: string; part: Part[] }[] @@ -58,6 +93,8 @@ type MessagePage = { projectSource?: boolean cursor?: string complete: boolean + clearedRevert?: boolean + revertPreview?: RevertPreview } function legacyMessageSource(items: { info: Message; parts: Part[] }[]): SessionMessageInfo[] { @@ -104,23 +141,35 @@ type MessageLoadBaseline = Pick< "touchedMessages" | "retainedMessages" | "touchedParts" | "clearedMessageParts" > -function mergeOptimisticPage(page: MessagePage, items: OptimisticItem[]) { - if (items.length === 0) return { ...page, observed: [] as { messageID: string; parts: Part[] }[] } +const hasParts = (parts: Part[] | undefined, want: Part[]) => { + if (!parts) return want.length === 0 + return want.every((part) => Binary.search(parts, part.id, (item) => item.id).found) +} + +export function mergeOptimisticPage(page: MessagePage, items: OptimisticItem[]) { + if (items.length === 0) + return { ...page, observed: [] as { messageID: string; parts: Part[] }[], confirmed: [] as string[] } const session = [...page.session] const part = new Map(page.part.map((item) => [item.id, item.part])) const observed: { messageID: string; parts: Part[] }[] = [] + const confirmed: string[] = [] for (const item of items) { - const result = Binary.search(session, messageKey(item.message), messageKey) - const found = result.found - if (!found) session.splice(result.index, 0, item.message) + const found = messageIndex(session, item.message.id) !== -1 + if (!found) { + if (page.projectSource) session.push(item.message) + if (!page.projectSource) insertMessage(session, item.message) + } const current = part.get(item.message.id) - const confirmed = found ? item.parts.filter((part) => current?.some((value) => value.id === part.id)) : [] - if (found) observed.push({ messageID: item.message.id, parts: confirmed }) + if (found && hasParts(current, item.parts)) confirmed.push(item.message.id) + const confirmedParts = found + ? item.parts.filter((part) => Binary.search(current ?? [], part.id, (value) => value.id).found) + : [] + if (found) observed.push({ messageID: item.message.id, parts: confirmedParts }) part.set( item.message.id, merge( found ? (current ?? []) : merge(item.confirmedParts ?? [], current ?? []), - item.parts.filter((part) => !confirmed.includes(part)), + item.parts.filter((part) => !confirmedParts.includes(part)), ), ) } @@ -129,7 +178,27 @@ function mergeOptimisticPage(page: MessagePage, items: OptimisticItem[]) { session, part: [...part.entries()].sort((a, b) => cmp(a[0], b[0])).map(([id, parts]) => ({ id, part: parts })), observed, + confirmed, + } +} + +export function applyOptimisticAdd(draft: OptimisticStore, input: OptimisticItem & { sessionID: string }) { + const messages = draft.message[input.sessionID] + if (messages) { + if (messageIndex(messages, input.message.id) === -1) insertMessage(messages, input.message) + } else { + draft.message[input.sessionID] = [input.message] } + draft.part[input.message.id] = input.parts.filter((part) => !!part?.id).sort((a, b) => cmp(a.id, b.id)) +} + +export function applyOptimisticRemove(draft: OptimisticStore, input: { sessionID: string; messageID: string }) { + const messages = draft.message[input.sessionID] + if (messages) { + const index = messageIndex(messages, input.messageID) + if (index !== -1) messages.splice(index, 1) + } + delete draft.part[input.messageID] } function runInflight(map: Map>, key: string, task: () => Promise) { @@ -148,6 +217,32 @@ function merge(a: readonly T[], b: readonly T[]) { return [...items.values()].sort((x, y) => cmp(x.id, y.id)) } +function mergeMessages(a: readonly Message[], b: readonly Message[]) { + const items = new Map(a.map((item) => [item.id, item] as const)) + for (const item of b) items.set(item.id, item) + return [...items.values()].sort(compareMessages) +} + +function mergeProjectedMessages( + current: readonly Message[], + updated: readonly Message[], + source?: SessionMessageInfo[], +) { + if (!source) return mergeMessages(current, updated) + const items = new Map(current.map((message) => [message.id, message] as const)) + updated.forEach((message) => items.set(message.id, message)) + const order = normalizeSessionMessages(updated[0]?.sessionID ?? current[0]?.sessionID ?? "", source).messages + return [ + ...order.flatMap((message) => { + const item = items.get(message.id) + if (!item) return [] + items.delete(message.id) + return [item] + }), + ...items.values(), + ] +} + function reconcileFetched( fetched: T[], current: readonly T[], @@ -157,6 +252,7 @@ function reconcileFetched( removed?: ReadonlySet preserveUnfetched?: boolean | ((item: T) => boolean) compare?: (a: T, b: T) => number + preserveOrder?: boolean } = {}, ) { const result = new Map(fetched.map((item) => [item.id, item])) @@ -179,8 +275,8 @@ function reconcileFetched( if (!item) result.delete(id) } for (const id of options.removed ?? emptyIDs) result.delete(id) - const items = [...result.values()] - return options.compare ? items.sort(options.compare) : items + const values = [...result.values()] + return options.preserveOrder ? values : values.sort(options.compare ?? ((a, b) => cmp(a.id, b.id))) } type ServerSessionOptions = { retry?: typeof retry; protocol?: Promise<"v1" | "v2"> } @@ -202,6 +298,7 @@ export function createServerSession( question: {} as Record, message: {} as Record, session_message: {} as Record, + revert_preview: {} as Record, part: {} as Record, part_text_accum_delta: {} as Record, session_working(id: string) { @@ -253,12 +350,20 @@ export function createServerSession( setData( "session_message", message.sessionID, - reconcile([...current, ...legacyMessageSource([{ info: message, parts: [] }])]), + reconcile([...current, ...legacyMessageSource([{ info: message, parts: [] }])].sort(compareMessages)), ) } const remember = (session: Session) => { + const previous = data.info[session.id] + const reloadMessages = + !!previous && + data.message[session.id] !== undefined && + !session.time.archived && + (previous?.revert?.messageID !== session.revert?.messageID || previous?.revert?.partID !== session.revert?.partID) setData("info", session.id, reconcile(session)) + if (data.revert_preview[session.id]?.messageID !== session.revert?.messageID) + setData("revert_preview", session.id, undefined) infoSeen.delete(session.id) infoSeen.add(session.id) if (infoSeen.size > sessionInfoLimit) { @@ -298,6 +403,21 @@ export function createServerSession( produce((draft) => stale.forEach((sessionID) => delete draft[sessionID])), ) } + if (reloadMessages) { + generations.set(session.id, {}) + inflight.delete(session.id) + setMeta("loading", session.id, false) + const limit = meta.limit[session.id] ?? initialMessagePageSize + // Fire-and-forget reload: on failure the cached window stays visible and the next + // sync() detects the unrepaired revert boundary and retries. + void loadMessages( + session.id, + limit === 0 ? initialMessagePageSize : limit, + undefined, + undefined, + session.revert, + ).catch(() => {}) + } return session } @@ -534,7 +654,13 @@ export function createServerSession( pickSessionCacheEvictions({ seen, keep: sessionID, limit: SESSION_CACHE_LIMIT, preserve: protectedSessions() }), ) - const fetchMessages = async (sessionID: string, limit: number, before?: string, onAttempt?: () => void) => { + const fetchMessages = async ( + sessionID: string, + limit: number, + before?: string, + onAttempt?: () => void, + revert?: Session["revert"], + ) => { if (messageApi && (await options?.protocol) !== "v1") { const request = (cursor?: string) => (options?.retry ?? retry)(() => { @@ -543,41 +669,115 @@ export function createServerSession( }) const first = await request(before) const pages = [first] - while (pages.at(-1)?.cursor.next && needsOlderTurnRoot(pages.flatMap((page) => page.data).toReversed())) { - const response = await request(pages.at(-1)!.cursor.next ?? undefined) + const project = () => { + const seen = new Set() + const source = pages.toReversed().flatMap((page) => + [...(page.context ?? []), ...page.data.toReversed()].filter((message) => { + if (seen.has(message.id)) return false + seen.add(message.id) + return true + }), + ) + return { source, normalized: normalizeSessionMessages(sessionID, source) } + } + const cursors = new Set() + const seen = new Set(first.data.map((message) => message.id)) + let progressed = true + while (pages.at(-1)?.cursor.next) { + const current = project() + const hasContext = "context" in first + const hasPreview = "revert" in first + if ( + (hasContext || !needsOlderTurnRoot(current.source)) && + (hasPreview || hasVisibleUserBeforeRevert(current.normalized.messages, revert)) + ) + break + if (!progressed) throw new Error("Message pagination returned no new messages") + const cursor = pages.at(-1)!.cursor.next! + if (cursors.has(cursor)) throw new Error("Message pagination cursor did not advance") + cursors.add(cursor) + const response = await request(cursor) + const unseen = response.data.filter((message) => !seen.has(message.id)) + unseen.forEach((message) => seen.add(message.id)) + progressed = unseen.length > 0 pages.push(response) if (!response.data.length) break } const response = pages.at(-1)! - const source = pages.flatMap((page) => page.data).toReversed() - const normalized = normalizeSessionMessages(sessionID, source) + const current = project() + const complete = !response.cursor.next return { - session: normalized.messages.sort(compareMessages), - part: [...normalized.parts.entries()] + session: current.normalized.messages, + part: [...current.normalized.parts.entries()] .map(([id, part]) => ({ id, part: part.sort((a, b) => cmp(a.id, b.id)) })) .sort((a, b) => cmp(a.id, b.id)), - source, + source: current.source, sourceMode: before ? ("older" as const) : ("latest" as const), projectSource: true, - cursor: response.cursor.next ?? undefined, - complete: response.data.length === 0, + cursor: first.contextCursor ?? response.cursor.next ?? undefined, + complete, + revertPreview: "revert" in first ? first.revert : undefined, + clearedRevert: + !!revert?.messageID && + !("revert" in first) && + complete && + !current.normalized.messages.some((message) => message.id === revert.messageID) + ? true + : undefined, + } + } + const toPage = (response: Awaited>) => { + const items = (response.data ?? []).filter((item) => !!item?.info?.id) + const cursor = nextBefore( + response.response.headers.get("Link"), + response.response.headers.get("X-Next-Cursor") ?? response.response.headers.get("x-next-cursor"), + ) + return { + session: items.map((item) => cleanMessage(item.info)).sort(compareMessages), + part: items.map((item) => ({ + id: item.info.id, + part: item.parts.filter((part) => !!part?.id).sort((a, b) => cmp(a.id, b.id)), + })), + cursor, + complete: !cursor, } } + const response = await (options?.retry ?? retry)(() => { onAttempt?.() return client.session.messages({ sessionID, limit, before }) }) - const items = (response.data ?? []).filter((item) => !!item?.info?.id) + const page = toPage(response) + const result = before + ? page + : await loadRevertAwareLatestPage({ + current: page, + revert, + fetchMessage: (messageID) => + (options?.retry ?? retry)(() => + client.session.message({ sessionID, messageID }, { throwOnError: false }), + ).then((result) => { + const boundary = boundaryFromMessageResponse(result) + if (!boundary) return undefined + const cursor = "cursor" in boundary ? boundary.cursor : undefined + return { + info: cleanMessage(boundary.info), + parts: boundary.parts.filter((part) => !!part?.id).sort((a, b) => cmp(a.id, b.id)), + cursor: typeof cursor === "string" ? cursor : undefined, + } + }), + // Boundary repair walks assistant-heavy stretches; the sync limit (often the tiny + // initial page size) would turn that walk into dozens of sequential requests. + fetchPage: (cursor) => + (options?.retry ?? retry)(() => + client.session.messages({ sessionID, limit: Math.max(limit, historyMessagePageSize), before: cursor }), + ).then(toPage), + }) + const parts = new Map(result.part.map((item) => [item.id, item.part])) return { - session: items.map((item) => cleanMessage(item.info)).sort(compareMessages), - part: items.map((item) => ({ - id: item.info.id, - part: item.parts.filter((part) => !!part?.id).sort((a, b) => cmp(a.id, b.id)), - })), - source: legacyMessageSource(items), + ...result, + source: legacyMessageSource(result.session.map((info) => ({ info, parts: parts.get(info.id) ?? [] }))), sourceMode: before ? ("older" as const) : ("latest" as const), - cursor: response.response.headers.get("x-next-cursor") ?? undefined, - complete: !response.response.headers.get("x-next-cursor"), } } @@ -682,8 +882,14 @@ export function createServerSession( const existing = data.session_message[sessionID] ?? [] const current = existing.filter((message) => !incoming.has(message.id)) const live = new Map(existing.map((message) => [message.id, message])) - return (page.sourceMode === "older" ? [...page.source, ...current] : [...current, ...page.source]).map( - (message) => (load?.touchedSource.has(message.id) ? (live.get(message.id) ?? message) : message), + const messages = (() => { + if (page.sourceMode === "older") return [...page.source, ...current] + const before = current.filter((message) => !load?.touchedSource.has(message.id)) + const after = current.filter((message) => load?.touchedSource.has(message.id)) + return [...before, ...page.source, ...after] + })() + return messages.map((message) => + load?.touchedSource.has(message.id) ? (live.get(message.id) ?? message) : message, ) })() : undefined @@ -693,7 +899,7 @@ export function createServerSession( const normalized = normalizeSessionMessages(sessionID, source) return { ...page, - session: normalized.messages.sort(compareMessages), + session: normalized.messages, part: [...normalized.parts.entries()] .map(([id, part]) => ({ id, part: part.sort((a, b) => cmp(a.id, b.id)) })) .sort((a, b) => cmp(a.id, b.id)), @@ -710,12 +916,15 @@ export function createServerSession( retained: load?.retainedMessages, removed: load?.removedMessages, preserveUnfetched, - compare: compareMessages, + compare: page.projectSource ? undefined : compareMessages, + preserveOrder: page.projectSource, }) batch(() => { if (source) setData("session_message", sessionID, reconcile(source)) + if (merged.revertPreview !== undefined) setData("revert_preview", sessionID, merged.revertPreview) const messageIDs = replaceMessages(sessionID, messages) replaceParts(sessionID, merged.part, messageIDs, load) + if (merged.clearedRevert) setData("info", sessionID, (info) => (info ? { ...info, revert: undefined } : info)) const orphans = orphanParts.get(sessionID) if (cleanupOrphans && page.complete && orphans) { for (const messageID of orphans) { @@ -730,7 +939,13 @@ export function createServerSession( }) } - const loadMessages = async (sessionID: string, limit: number, before?: string, mode?: "replace" | "prepend") => { + const loadMessages = async ( + sessionID: string, + limit: number, + before?: string, + mode?: "replace" | "prepend", + revert?: Session["revert"], + ) => { if (meta.loading[sessionID]) return const active = generation(sessionID) const load: MessageLoadState = { @@ -750,11 +965,25 @@ export function createServerSession( setMeta("loading", sessionID, true) let applied = false try { - const page = await fetchMessages(sessionID, limit, before, () => resetMessageLoad(sessionID, load)) - const first = page.session.reduce( - (oldest, message) => (!oldest || compareMessages(message, oldest) < 0 ? message : oldest), - undefined, + const fetched = await fetchMessages(sessionID, limit, before, () => resetMessageLoad(sessionID, load), revert) + const incoming = fetched.projectSource ? (fetched.source ?? []) : fetched.session + const existing = new Set( + (fetched.projectSource ? (data.session_message[sessionID] ?? []) : (data.message[sessionID] ?? [])).map( + (message) => message.id, + ), ) + const page = + mode === "prepend" && + before && + (fetched.cursor === before || (!!fetched.cursor && !incoming.some((message) => !existing.has(message.id)))) + ? { ...fetched, cursor: undefined, complete: true } + : fetched + const first = page.projectSource + ? page.session[0] + : page.session.reduce( + (oldest, message) => (!oldest || compareMessages(message, oldest) < 0 ? message : oldest), + undefined, + ) if (generations.get(sessionID) !== active) return const parents = [] as Awaited>[] @@ -799,10 +1028,10 @@ export function createServerSession( ? page : { ...page, - session: merge( + session: mergeMessages( page.session, parents.map((parent) => parent.message), - ).sort(compareMessages), + ), part: merge( page.part, parents.map((parent) => ({ id: parent.message.id, part: parent.parts })), @@ -836,14 +1065,38 @@ export function createServerSession( const sync = (sessionID: string, options?: { force?: boolean; messageLimit?: number }) => { touch(sessionID) return runInflight(inflight, sessionID, async () => { + const sessionInfo = data.info[sessionID] const cached = data.message[sessionID] !== undefined && meta.limit[sessionID] !== undefined - if (cached && data.info[sessionID] && !options?.force) return - await Promise.all([ - resolve(sessionID, options), - cached && !options?.force + const needsRevertRepair = (revert?: Session["revert"]) => { + if (!revert?.messageID || data.message[sessionID] === undefined || meta.limit[sessionID] === undefined) + return false + const messages = data.message[sessionID] ?? [] + const boundaryLoaded = messages.some((message) => message.id === revert.messageID) + return !hasVisibleUserBeforeRevert(messages, revert) && !(meta.complete[sessionID] && boundaryLoaded) + } + const needsExistingRevertRepair = needsRevertRepair(sessionInfo?.revert) + const canReuseCached = !!(cached && sessionInfo && !needsExistingRevertRepair) + const messageLimit = () => options?.messageLimit ?? meta.limit[sessionID] ?? initialMessagePageSize + const repairMessageLimit = () => { + const limit = messageLimit() + return limit === 0 ? initialMessagePageSize : limit + } + if (canReuseCached && !options?.force) return + const resolving = sessionInfo && !options?.force ? Promise.resolve(sessionInfo) : resolve(sessionID, options) + const loading = + cached && !options?.force && !needsExistingRevertRepair ? Promise.resolve() - : loadMessages(sessionID, options?.messageLimit ?? meta.limit[sessionID] ?? initialMessagePageSize), - ]) + : loadMessages( + sessionID, + needsExistingRevertRepair ? repairMessageLimit() : messageLimit(), + undefined, + undefined, + sessionInfo?.revert, + ) + const nextSession = await resolving + await loading + if (!sessionInfo?.revert?.messageID && nextSession.revert?.messageID && needsRevertRepair(nextSession.revert)) + await loadMessages(sessionID, repairMessageLimit(), undefined, undefined, nextSession.revert) }) } @@ -855,7 +1108,9 @@ export function createServerSession( (meta.complete[sessionID] || (data.message[sessionID]?.length ?? 0) >= limit) ) return - await runInflight(inflight, sessionID, () => loadMessages(sessionID, limit)) + await runInflight(inflight, sessionID, () => + loadMessages(sessionID, limit, undefined, undefined, data.info[sessionID]?.revert), + ) } const eventSessionID = (event: { type: string; properties?: unknown }) => { @@ -888,7 +1143,10 @@ export function createServerSession( const touched = new Set(reduction.touched) let parentID: string | undefined for (const message of reduction.messages) { - if (message.type === "user" || (message.type === "synthetic" && message.description?.trim())) + if ( + message.type === "user" || + (message.type === "synthetic" && (message.description?.trim() || message.text.trim())) + ) parentID = message.id if (message.type === "shell") { if (touched.has(message.id)) touched.add(`${message.id}:assistant`) @@ -921,16 +1179,15 @@ export function createServerSession( }) } - const hydrateV2Message = (sessionID: string, messageID: string) => { - if (!sessionApi) return - void sessionApi - .message({ sessionID, messageID }) - .then((message) => { - const current = data.session_message[sessionID] ?? [] - const messages = [...current.filter((item) => item.id !== message.id), message].sort(compareMessages) - projectV2({ sessionID, messages, touched: [message.id] }) - }) - .catch(() => {}) + const hydrateV2Message = (sessionID: string, _messageID: string) => { + if (!sessionApi || !messageApi) return + const refresh = () => sync(sessionID, { force: true }) + const pending = inflight.get(sessionID) + if (pending) { + void pending.finally(refresh).catch(() => {}) + return + } + void refresh().catch(() => {}) } const applyV2 = (event: OpenCodeEvent) => { @@ -975,11 +1232,45 @@ export function createServerSession( next: event.data.at, }) if (event.type === "session.forked") void resolve(sessionID, { force: true }).catch(() => {}) - if ( - event.type === "session.revert.staged" || - event.type === "session.revert.cleared" || - event.type === "session.revert.committed" - ) + const eventType: string = event.type + if (eventType === "session.next.revert.committed") { + const messageID = + "messageID" in event.data && typeof event.data.messageID === "string" ? event.data.messageID : undefined + const source = data.session_message[sessionID] ?? [] + const boundary = messageID ? source.findIndex((message) => message.id === messageID) : -1 + if (boundary !== -1) { + const committed = source.slice(0, boundary + 1) + const normalized = normalizeSessionMessages(sessionID, committed) + const messageIDs = new Set(normalized.messages.map((message) => message.id)) + const dropped = (data.message[sessionID] ?? []).filter((message) => !messageIDs.has(message.id)) + batch(() => { + setData( + "session_message", + sessionID, + produce((draft) => draft.splice(0, draft.length, ...committed)), + ) + setData( + "message", + sessionID, + produce((draft) => draft.splice(0, draft.length, ...normalized.messages)), + ) + setData(produce((draft) => dropped.forEach((message) => deleteMessageParts(draft, message.id)))) + replaceParts( + sessionID, + [...normalized.parts].map(([id, part]) => ({ id, part })), + messageIDs, + ) + }) + } + setData("revert_preview", sessionID, undefined) + setData("info", sessionID, (info) => (info ? { ...info, revert: undefined } : info)) + setMeta("limit", sessionID, undefined) + setMeta("cursor", sessionID, undefined) + setMeta("complete", sessionID, undefined) + setMeta("at", sessionID, undefined) + } + if (eventType === "session.next.revert.committed") void sync(sessionID, { force: true }).catch(() => {}) + if (eventType === "session.next.revert.staged" || eventType === "session.next.revert.cleared") void resolve(sessionID, { force: true }).catch(() => {}) } @@ -1050,14 +1341,13 @@ export function createServerSession( setData("message", info.sessionID, [info]) return } - const result = Binary.search(messages, messageKey(info), messageKey) - if (result.found) setData("message", info.sessionID, result.index, reconcile(info)) - if (!result.found) - setData("message", info.sessionID, (value = []) => { - const next = value.slice() - next.splice(result.index, 0, info) - return next - }) + setData("message", info.sessionID, (value = []) => + mergeProjectedMessages( + value.filter((message) => message.id !== info.id), + [info], + data.session_message[info.sessionID], + ), + ) return } case "message.removed": { @@ -1340,7 +1630,9 @@ export function createServerSession( if (items) items.set(input.message.id, { ...input, parts, confirmedParts: [] }) if (!items) optimistic.set(input.sessionID, new Map([[input.message.id, { ...input, parts, confirmedParts: [] }]])) - setData("message", input.sessionID, (messages = []) => merge(messages, [input.message]).sort(compareMessages)) + setData("message", input.sessionID, (messages = []) => + mergeProjectedMessages(messages, [input.message], data.session_message[input.sessionID]), + ) setData( "part_text_accum_delta", produce((draft) => { diff --git a/packages/app/src/context/sync-optimistic.test.ts b/packages/app/src/context/sync-optimistic.test.ts index d7ac9fd9641f..ca2c31c67034 100644 --- a/packages/app/src/context/sync-optimistic.test.ts +++ b/packages/app/src/context/sync-optimistic.test.ts @@ -1,6 +1,6 @@ import { describe, expect, test } from "bun:test" import type { Message, Part } from "@opencode-ai/sdk/v2/client" -import { applyOptimisticAdd, applyOptimisticRemove, mergeOptimisticPage } from "./sync" +import { applyOptimisticAdd, applyOptimisticRemove, mergeOptimisticPage } from "./server-session" type Text = Extract @@ -87,6 +87,24 @@ describe("sync optimistic reducers", () => { expect(page.session.map((message) => message.id)).toEqual(["msg_a", "msg_z"]) }) + test("mergeOptimisticPage appends pending messages after current protocol source order", () => { + const sessionID = "ses_1" + const durable = userMessage("msg_z_durable", sessionID) + const optimistic = userMessage("msg_a_optimistic", sessionID) + const page = mergeOptimisticPage( + { + session: [durable], + part: [{ id: durable.id, part: [] }], + source: [{ id: durable.id, type: "user", text: "durable", time: { created: 1 } }], + projectSource: true, + complete: true, + }, + [{ message: optimistic, parts: [] }], + ) + + expect(page.session.map((message) => message.id)).toEqual([durable.id, optimistic.id]) + }) + test("mergeOptimisticPage keeps missing optimistic parts until the server has them", () => { const sessionID = "ses_1" const page = mergeOptimisticPage( diff --git a/packages/app/src/pages/session.tsx b/packages/app/src/pages/session.tsx index 2e647e0f4789..0721210b4dab 100644 --- a/packages/app/src/pages/session.tsx +++ b/packages/app/src/pages/session.tsx @@ -3,7 +3,6 @@ import { getFilename } from "@opencode-ai/core/util/path" import { useDialog } from "@opencode-ai/ui/context/dialog" import { createQuery, skipToken, useMutation, useQueryClient } from "@tanstack/solid-query" import { - batch, ErrorBoundary, onCleanup, Suspense, @@ -76,6 +75,7 @@ import { createTimelineModel } from "@/pages/session/timeline/model" import { type DiffStyle, SessionReviewTab, type SessionReviewTabProps } from "@/pages/session/review-tab" import { useSessionLayout } from "@/pages/session/session-layout" import { restorePromptModel, syncPromptModel, syncSessionModel } from "@/pages/session/session-model-helpers" +import { runPromptRollbackMutation } from "@/pages/session/session-prompt-rollback" import { clampSessionPanelWidth, SESSION_PANEL_WIDTH_MIN, @@ -100,12 +100,22 @@ import { Persist, persisted } from "@/utils/persist" import { extractPromptFromParts } from "@/utils/prompt" import { formatServerError, isLocalSessionNotFoundError, isSessionNotFoundError } from "@/utils/server-errors" import { legacySessionHref, requireServerKey, sessionHref } from "@/utils/session-route" +import { normalizeSessionMessages } from "@/utils/session-message" +import { restoreTarget, type RevertPreviewState } from "@/pages/session/timeline/revert" +import { revertPreviewFromMessages, unsupportedServerRequest } from "@opencode-ai/core/util/revert-boundary" import { useUsageExceededDialogs } from "./session/usage-exceeded-dialogs" import { createSessionOwnership } from "./session/session-ownership" import { createSessionLineage } from "./session/session-lineage" type FollowupItem = FollowupDraft & { id: string } type FollowupEdit = Pick +type RevertPreview = { + userCount: number + nextMessageID?: string + hasMore?: boolean + continuationMessageID?: string + items: { id: string; text: string }[] +} const emptyFollowups: FollowupItem[] = [] type ChangeMode = "git" | "branch" | "turn" @@ -121,29 +131,6 @@ function isCurrentSessionNotFoundError(error: unknown, sessionID: string | undef return isSessionNotFoundError(error, sessionID) || isLocalSessionNotFoundError(error, sessionID) } -async function runPromptRollbackMutation(input: { - capturePrompt: () => { current: () => T[]; set: (value: T[]) => void; reset: () => void } - optimistic: (prompt: { set: (value: T[]) => void; reset: () => void }) => void - request: () => Promise - complete: (result: R) => void - rollback: () => void - fail: (error: unknown) => void -}) { - const prompt = input.capturePrompt() - const previous = prompt.current().slice() - batch(() => input.optimistic(prompt)) - await input - .request() - .then(input.complete) - .catch((error) => { - batch(() => { - input.rollback() - prompt.set(previous) - }) - input.fail(error) - }) -} - export function SessionPage() { return ( @@ -544,8 +531,9 @@ export default function Page() { }) const activeTab = tabState.activeTab const activeFileTab = tabState.activeFileTab - const revertMessageID = createMemo(() => info()?.revert?.messageID) - const timeline = createTimelineModel({ sessionID: () => params.id, revertMessageID }) + const revertInfo = createMemo(() => info()?.revert) + const revertMessageID = createMemo(() => revertInfo()?.messageID) + const timeline = createTimelineModel({ sessionID: () => params.id, revert: revertInfo }) const historyLoading = timeline.history.loading const historyMore = timeline.history.more const lastUserMessage = timeline.lastUserMessage @@ -554,6 +542,11 @@ export default function Page() { const sessionSync = timeline.resource const userMessages = timeline.userMessages const visibleUserMessages = timeline.visibleUserMessages + const currentRevertPreview = createMemo(() => { + const id = params.id + if (!id) return + return sync().data.revert_preview[id] + }) createEffect(() => { const tab = activeFileTab() @@ -891,6 +884,42 @@ export default function Page() { const hasScrollGesture = () => Date.now() - ui.scrollGesture < scrollGestureWindowMs + const [revertPreview, setRevertPreview] = createSignal() + let revertPreviewEpoch = 0 + const loadV1RevertPreview = async (sessionID: string) => { + const revert = info()?.revert + if (!revert) return + try { + const preview = await sdk().client.session.revertPreview({ sessionID }, { throwOnError: false }) + if (preview.response.status !== 404 && preview.response.status !== 405) { + if (preview.error) throw preview.error + return preview.data ?? undefined + } + } catch (error) { + if (!unsupportedServerRequest(error)) throw error + } + const result = await sdk().client.session.messages({ sessionID }, { throwOnError: true }) + return revertPreviewFromMessages(result.data ?? [], revert) + } + createEffect( + on( + () => [params.id, revertMessageID(), serverSDK().protocolKind()] as const, + ([sessionID, revert, protocol]) => { + const epoch = ++revertPreviewEpoch + setRevertPreview(undefined) + if (!sessionID || !revert || protocol !== "v1") return + void loadV1RevertPreview(sessionID) + .then((result) => { + if (epoch !== revertPreviewEpoch) return + setRevertPreview(result) + }) + .catch(() => { + if (epoch !== revertPreviewEpoch) return + setRevertPreview(undefined) + }) + }, + ), + ) createEffect( on( () => { @@ -1137,12 +1166,26 @@ export default function Page() { } useComposerCommands() + const redoPreview = createMemo(() => { + const protocol = serverSDK().protocolKind() + if (!protocol) return { ready: false } + if (protocol === "v1") { + const preview = revertPreview() + if (!preview) return { ready: false } + return { ready: true, nextMessageID: preview.nextMessageID } + } + const preview = currentRevertPreview() + if (!preview) return { ready: false } + return { ready: true, nextMessageID: preview.nextMessageID } + }) + useSessionCommands({ navigateMessageByOffset, setActiveMessage, focusInput, review: reviewTab, fileBrowser: () => newSessionDesign() && isDesktop() && !!params.id, + revertPreview: redoPreview, }) command.register("session-palette", () => [ { @@ -1668,22 +1711,6 @@ export default function Page() { ), ) - const draft = (id: string) => - extractPromptFromParts(sync().data.part[id] ?? [], { - directory: sdk().directory, - attachmentName: language.t("common.attachment"), - }) - - const line = (id: string) => { - const text = draft(id) - .map((part) => (part.type === "image" ? `[image:${part.filename}]` : part.content)) - .join("") - .replace(/\s+/g, " ") - .trim() - if (text) return text - return `[${language.t("common.attachment")}]` - } - const fail = (err: unknown) => { showToast({ variant: "error", @@ -1692,15 +1719,15 @@ export default function Page() { }) } - const merge = (next: NonNullable>, target = sync()) => target.session.remember(next) - - const roll = (sessionID: string, next: NonNullable>["revert"], target = sync()) => { - const session = target.session.get(sessionID) - if (!session) return - target.session.remember({ ...session, revert: next }) - } - const busy = (sessionID: string) => sync().data.session_working(sessionID) + const rememberRevert = ( + state: ReturnType, + sessionID: string, + revert: NonNullable>["revert"], + ) => { + const current = state.session.get(sessionID) + if (current) state.session.remember({ ...current, revert }) + } const queuedFollowups = createMemo(() => { const id = params.id @@ -1817,87 +1844,144 @@ export default function Page() { setFollowup("edit", id, undefined) } - const halt = (sessionID: string) => - busy(sessionID) - ? sdk() - .api.session.interrupt({ sessionID }) - .catch(() => {}) - : Promise.resolve() + const captureRollbackPrompt = () => { + const target = prompt.capture() + return { + current: target.current, + cursor: target.cursor, + set: target.set, + reset: target.reset, + setCursor: (cursor: number | undefined) => target.store[1]("cursor", cursor), + } + } + + const captureRevert = (input: { sessionID: string; messageID: string }) => { + const client = sdk() + const state = sync() + return { + ...input, + state, + session: client.api.session, + prompt: captureRollbackPrompt(), + draft: extractPromptFromParts(state.data.part[input.messageID] ?? [], { + directory: client.directory, + attachmentName: language.t("common.attachment"), + }), + } + } + + const captureRestore = (sessionID: string, id: string) => { + const client = sdk() + return { + id, + sessionID, + client, + currentApi: serverSDK().currentApi, + state: sync(), + session: client.api.session, + legacyPreview: revertPreview(), + currentPreview: currentRevertPreview(), + directory: client.directory, + attachmentName: language.t("common.attachment"), + prompt: captureRollbackPrompt(), + } + } const revertMutation = useMutation(() => ({ - mutationFn: async (input: { sessionID: string; messageID: string }) => { - const session = sdk().api.session - const target = sync() - const last = target.session.get(input.sessionID)?.revert - const value = draft(input.messageID) + mutationFn: async (input: ReturnType) => { await runPromptRollbackMutation({ - capturePrompt: prompt.capture, - optimistic: (prompt) => { - roll(input.sessionID, { messageID: input.messageID }, target) - prompt.set(value) + prompt: input.prompt, + prepare: () => input.draft, + optimistic: (prompt, value) => prompt.set(value), + request: async () => { + if (input.state.data.session_working(input.sessionID)) { + await input.session.interrupt({ sessionID: input.sessionID }).catch(() => {}) + } + return input.session.revert.stage({ sessionID: input.sessionID, messageID: input.messageID }) }, - request: () => halt(input.sessionID).then(() => session.revert.stage(input)), - complete: () => undefined, - rollback: () => roll(input.sessionID, last, target), + complete: (revert) => rememberRevert(input.state, input.sessionID, revert), + rollback: () => {}, fail, }) }, })) const restoreMutation = useMutation(() => ({ - mutationFn: async (id: string) => { - const sessionID = params.id - if (!sessionID) return - - const session = sdk().api.session - const target = sync() - const index = userMessages().findIndex((item) => item.id === id) - if (index < 0) return - const next = userMessages()[index + 1] - const last = target.session.get(sessionID)?.revert - + mutationFn: async (input: ReturnType) => { await runPromptRollbackMutation({ - capturePrompt: prompt.capture, - optimistic: (promptSession) => { - roll(sessionID, next ? { messageID: next.id } : undefined, target) - if (next) { - promptSession.set(draft(next.id)) - return + prompt: input.prompt, + prepare: async () => { + const protocol = await input.client.protocol + const preview = protocol === "v1" ? input.legacyPreview : input.currentPreview + if (!preview?.items.length) throw new Error("Failed to load revert preview") + const next = restoreTarget(preview, input.id) + if (next === undefined) throw new Error("Restore target missing from revert preview") + if (!next) return { next, draft: undefined } + const draft = await (async () => { + const parts = input.state.data.part[next] + if (parts) + return extractPromptFromParts(parts, { + directory: input.directory, + attachmentName: input.attachmentName, + }) + if (protocol !== "v1") { + const message = await input.currentApi.session.message({ + sessionID: input.sessionID, + messageID: next, + }) + const normalized = normalizeSessionMessages(input.sessionID, [message]) + return extractPromptFromParts(normalized.parts.get(next) ?? [], { + directory: input.directory, + attachmentName: input.attachmentName, + }) + } + const result = await input.client.client.session.message({ sessionID: input.sessionID, messageID: next }) + return extractPromptFromParts(result.data?.parts ?? [], { + directory: input.directory, + attachmentName: input.attachmentName, + }) + })().catch(() => undefined) + if (!draft) throw new Error("Failed to load next restore draft") + return { next, draft } + }, + optimistic: (promptSession, prepared) => { + if (prepared.next) promptSession.set(prepared.draft ?? []) + else promptSession.reset() + }, + request: async (prepared) => { + if (input.state.data.session_working(input.sessionID)) { + await input.session.interrupt({ sessionID: input.sessionID }).catch(() => {}) + } + if (!prepared.next) { + await input.session.revert.clear({ sessionID: input.sessionID }) + return undefined } - promptSession.reset() + return input.session.revert.stage({ sessionID: input.sessionID, messageID: prepared.next }) }, - request: () => - !next - ? halt(sessionID).then(() => session.revert.clear({ sessionID })) - : halt(sessionID).then(() => session.revert.stage({ sessionID, messageID: next.id }).then(() => undefined)), - complete: () => undefined, - rollback: () => roll(sessionID, last, target), + complete: (revert) => rememberRevert(input.state, input.sessionID, revert), + rollback: () => {}, fail, }) }, })) const reverting = createMemo(() => revertMutation.isPending || restoreMutation.isPending) - const restoring = createMemo(() => (restoreMutation.isPending ? restoreMutation.variables : undefined)) + const restoring = createMemo(() => (restoreMutation.isPending ? restoreMutation.variables?.id : undefined)) const revert = (input: { sessionID: string; messageID: string }) => { if (reverting()) return - return revertMutation.mutateAsync(input) + return revertMutation.mutateAsync(captureRevert(input)) } const restore = (id: string) => { - if (!params.id || reverting()) return - return restoreMutation.mutateAsync(id) + const sessionID = params.id + if (!sessionID || reverting()) return + return restoreMutation.mutateAsync(captureRestore(sessionID, id)) } const rolled = createMemo(() => { - const id = revertMessageID() - if (!id) return [] - const index = userMessages().findIndex((item) => item.id === id) - if (index < 0) return [] - return userMessages() - .slice(index) - .map((item) => ({ id: item.id, text: line(item.id) })) + if (serverSDK().protocolKind() === "v1") return revertPreview()?.items ?? [] + return currentRevertPreview()?.items ?? [] }) // attachment bytes are embedded as a data URL, so downloading always works; diff --git a/packages/app/src/pages/session/session-prompt-rollback.test.ts b/packages/app/src/pages/session/session-prompt-rollback.test.ts new file mode 100644 index 000000000000..cd7e80e285b3 --- /dev/null +++ b/packages/app/src/pages/session/session-prompt-rollback.test.ts @@ -0,0 +1,118 @@ +import { expect, test } from "bun:test" +import { runPromptRollbackMutation } from "./session-prompt-rollback" + +const prompt = (initial: string[], initialCursor?: number) => { + let value = initial + let cursor = initialCursor + return { + current: () => value, + cursor: () => cursor, + set: (next: string[], nextCursor?: number) => { + value = next + if (nextCursor !== undefined) cursor = nextCursor + }, + setCursor: (next: number | undefined) => { + cursor = next + }, + reset: () => { + value = [] + cursor = 0 + }, + } +} + +test("captures the initiating prompt before asynchronous restore preparation", async () => { + const first = prompt(["first draft"]) + const second = prompt(["second draft"]) + const prepared = Promise.withResolvers() + let active = first + const rollback = active + + const mutation = runPromptRollbackMutation({ + prompt: rollback, + prepare: () => prepared.promise, + optimistic: (target, draft) => target.set(draft), + request: async () => {}, + complete: () => {}, + rollback: () => {}, + fail: () => {}, + }) + + active = second + second.set(["edited second draft"]) + prepared.resolve(["restored first draft"]) + await mutation + + expect(first.current()).toEqual(["restored first draft"]) + expect(second.current()).toEqual(["edited second draft"]) +}) + +test("restores the initiating prompt cursor when the request fails", async () => { + const target = prompt(["original draft"], 7) + + await runPromptRollbackMutation({ + prompt: target, + prepare: () => undefined, + optimistic: (prompt) => prompt.reset(), + request: async () => { + throw new Error("request failed") + }, + complete: () => {}, + rollback: () => {}, + fail: () => {}, + }) + + expect(target.current()).toEqual(["original draft"]) + expect(target.cursor()).toBe(7) +}) + +test("rolls back prompt edits made while restore preparation is pending", async () => { + const target = prompt(["original draft"], 7) + const prepared = Promise.withResolvers() + const mutation = runPromptRollbackMutation({ + prompt: target, + prepare: () => prepared.promise, + optimistic: (prompt) => prompt.reset(), + request: async () => { + throw new Error("request failed") + }, + complete: () => {}, + rollback: () => {}, + fail: () => {}, + }) + + target.set(["edited during preparation"], 9) + prepared.resolve() + await mutation + + expect(target.current()).toEqual(["edited during preparation"]) + expect(target.cursor()).toBe(9) +}) + +test("reports preparation failures without changing the prompt", async () => { + const target = prompt(["original draft"], 7) + const error = new Error("prepare failed") + let failed: unknown + let optimistic = false + + await runPromptRollbackMutation({ + prompt: target, + prepare: async () => { + throw error + }, + optimistic: () => { + optimistic = true + }, + request: async () => {}, + complete: () => {}, + rollback: () => {}, + fail: (cause) => { + failed = cause + }, + }) + + expect(failed).toBe(error) + expect(optimistic).toBe(false) + expect(target.current()).toEqual(["original draft"]) + expect(target.cursor()).toBe(7) +}) diff --git a/packages/app/src/pages/session/session-prompt-rollback.ts b/packages/app/src/pages/session/session-prompt-rollback.ts new file mode 100644 index 000000000000..e6023e944c09 --- /dev/null +++ b/packages/app/src/pages/session/session-prompt-rollback.ts @@ -0,0 +1,41 @@ +import { batch } from "solid-js" + +type PromptTarget = { + current: () => T[] + cursor: () => number | undefined + set: (value: T[]) => void + setCursor: (value: number | undefined) => void + reset: () => void +} + +export async function runPromptRollbackMutation(input: { + prompt: PromptTarget + prepare: () => P | Promise

+ optimistic: (prompt: PromptTarget, prepared: P) => void + request: (prepared: P) => Promise + complete: (result: R) => void + rollback: () => void + fail: (error: unknown) => void +}) { + const prepared = await Promise.resolve() + .then(input.prepare) + .then( + (value) => ({ ok: true as const, value }), + (error) => ({ ok: false as const, error }), + ) + if (!prepared.ok) return input.fail(prepared.error) + const previous = input.prompt.current().slice() + const cursor = input.prompt.cursor() + batch(() => input.optimistic(input.prompt, prepared.value)) + await input + .request(prepared.value) + .then(input.complete) + .catch((error) => { + batch(() => { + input.rollback() + input.prompt.set(previous) + input.prompt.setCursor(cursor) + }) + input.fail(error) + }) +} diff --git a/packages/app/src/pages/session/timeline/message-timeline.tsx b/packages/app/src/pages/session/timeline/message-timeline.tsx index 25abdc9e7341..047b36d9cdaf 100644 --- a/packages/app/src/pages/session/timeline/message-timeline.tsx +++ b/packages/app/src/pages/session/timeline/message-timeline.tsx @@ -75,7 +75,8 @@ import { sessionTitle } from "@/utils/session-title" import { scheduleConnectedMeasure } from "./measure" import { observeElementOffsetReconnectAware } from "./observe-element-offset" import { createTimelineProjection } from "./projection" -import { MessageComment, SummaryDiff, TimelineRow, TimelineRowMap } from "./rows" +import { visiblePartsForMessage } from "./revert" +import { MessageComment, SummaryDiff, TimelineRow, TimelineRowMap, visibleUserMessageContent } from "./rows" import { filterVirtualIndexes } from "./virtual-items" const emptyMessages: MessageType[] = [] @@ -287,9 +288,8 @@ export function MessageTimeline(props: { const visible = new Set(props.userMessages.map((message) => message.id)) const boundary = sessionMessages().find((message) => message.role === "user" && !visible.has(message.id))?.id const messages = sync().data.session_message[id] ?? [] - if (!boundary) return messages - const index = messages.findIndex((message) => message.id === boundary) - return index < 0 ? messages : messages.slice(0, index) + const index = boundary ? messages.findIndex((message) => message.id === boundary) : -1 + return index === -1 ? messages : messages.slice(0, index) }) const info = createMemo(() => { const id = sessionID() @@ -313,6 +313,8 @@ export function MessageTimeline(props: { }) const parentTitle = createMemo(() => sessionTitle(parent()?.title) ?? language.t("command.session.new")) const getMsgParts = (msgId: string) => sync().data.part[msgId] ?? emptyParts + const getVisibleMsgParts = (messageID: string) => + visiblePartsForMessage(messageID, getMsgParts(messageID), info()?.revert) const getMsgPart = (messageID: string, partID: string) => getMsgParts(messageID).find((part) => part.id === partID) const childTaskDescription = createMemo(() => { const id = sessionID() @@ -338,6 +340,7 @@ export function MessageTimeline(props: { status: sessionStatus, showReasoningSummaries: settings.general.showReasoningSummaries, inlineComments: settings.general.newLayoutDesigns, + revert: () => info()?.revert, }) const activeMessageID = projection.activeMessageID const assistantMessagesByParent = projection.assistantMessagesByParent @@ -1011,7 +1014,7 @@ export function MessageTimeline(props: { const message = messages[i] if (!message) continue - const parts = getMsgParts(message.id) + const parts = getVisibleMsgParts(message.id) for (let j = parts.length - 1; j >= 0; j--) { const part = parts[j] if (!part || part.type !== "text" || !part.text?.trim()) continue @@ -1123,8 +1126,13 @@ export function MessageTimeline(props: { return