diff --git a/.github/workflows/pr-build.yml b/.github/workflows/pr-build.yml index 02ccb32cf..1dcb32fb8 100644 --- a/.github/workflows/pr-build.yml +++ b/.github/workflows/pr-build.yml @@ -104,11 +104,18 @@ jobs: - name: Test changed runnable UI behavior run: >- node --import tsx --test + packages/ui/src/components/prompt-input/submitPrompt.test.ts packages/ui/src/components/session-list-visibility.test.ts packages/ui/src/components/unified-picker-path.test.ts packages/ui/src/lib/hooks/use-app-session-capture.test.ts + packages/ui/src/lib/global-cache.test.ts packages/ui/src/lib/launch-errors.test.ts packages/ui/src/lib/message-selection-position.test.ts + packages/ui/src/lib/native/desktop-events-reconnect.test.ts + packages/ui/src/lib/native/desktop-events.test.ts + packages/ui/src/lib/opencode-api.test.ts + packages/ui/src/lib/session-memory-budget.test.ts + packages/ui/src/lib/session-search.test.ts packages/ui/src/lib/trailing-resync.test.ts packages/ui/src/stores/abort-created-workspace-cleanup.test.ts packages/ui/src/stores/app-session-reconciliation.test.ts @@ -119,8 +126,14 @@ jobs: packages/ui/src/stores/restore-workspace-commit-gates.test.ts packages/ui/src/stores/client-state-codec.test.ts packages/ui/src/stores/client-state.test.ts + packages/ui/src/stores/commands.test.ts + packages/ui/src/stores/delta-buffer.test.ts packages/ui/src/stores/instances-restore-cancellation.test.ts + packages/ui/src/stores/permission-replies.test.ts + packages/ui/src/stores/request-locations.test.ts packages/ui/src/stores/message-v2/message-hydration-authority.test.ts + packages/ui/src/stores/message-v2/instance-store.test.ts + packages/ui/src/stores/message-prompt-display.test.ts packages/ui/src/stores/session-generation-recovery.test.ts packages/ui/src/stores/session-metadata.test.ts packages/ui/src/stores/session-pagination.test.ts @@ -131,6 +144,20 @@ jobs: node --conditions=browser --import tsx --test --test-force-exit packages/ui/src/stores/instances-restore-ownership.test.ts packages/ui/src/stores/permission-lifecycle.test.ts + packages/ui/src/stores/session-request-authority.test.ts + packages/ui/src/stores/session-memory.test.ts + packages/ui/src/components/tool-call/search-text.test.ts + packages/ui/src/components/tool-call/diagnostics.test.ts + packages/ui/src/components/tool-call/permission-constants.test.ts + packages/ui/src/components/tool-call/renderers/task-transcript-load.test.ts + packages/ui/src/components/tool-call/utils.test.ts + packages/ui/src/lib/sse-manager-generation.test.ts + packages/ui/src/stores/session-state-purge.test.ts + packages/ui/src/stores/instances-reconnect-resync.test.ts + packages/ui/src/stores/session-revision-recovery.test.ts + packages/ui/src/stores/session-actions-delivery.test.ts + packages/ui/src/stores/opencode-workspaces.test.ts + packages/ui/src/components/virtual-follow-behavior.test.ts - name: Test server run: node --import tsx --test "packages/server/src/**/*.test.ts" @@ -176,4 +203,4 @@ jobs: - name: Test Tauri crate on Windows working-directory: packages/tauri-app/src-tauri - run: cargo test --locked + run: cargo test --locked -- --test-threads=1 diff --git a/packages/server/src/api-types.ts b/packages/server/src/api-types.ts index d61b12f0e..bfd87b6f2 100644 --- a/packages/server/src/api-types.ts +++ b/packages/server/src/api-types.ts @@ -482,8 +482,8 @@ export type WorkspaceEventPayload = | { type: "storage.configChanged"; owner: SettingsOwner; value: SettingsBucket } | { type: "storage.stateChanged"; owner: SettingsOwner; value: SettingsBucket } | { type: "instance.dataChanged"; instanceId: string; data: InstanceData } - | { type: "instance.event"; instanceId: string; event: InstanceStreamEvent } - | { type: "instance.eventStatus"; instanceId: string; status: InstanceStreamStatus; reason?: string } + | { type: "instance.event"; instanceId: string; streamId?: string; event: InstanceStreamEvent } + | { type: "instance.eventStatus"; instanceId: string; streamId?: string; status: InstanceStreamStatus; reason?: string } | { type: "yolo.stateChanged"; instanceId: string; sessionId: string; enabled: boolean } | { type: "yolo.autoAccepted"; instanceId: string; sessionId: string; permissionId: string } diff --git a/packages/server/src/permissions/auto-accept-manager.test.ts b/packages/server/src/permissions/auto-accept-manager.test.ts index d803587be..50d729e26 100644 --- a/packages/server/src/permissions/auto-accept-manager.test.ts +++ b/packages/server/src/permissions/auto-accept-manager.test.ts @@ -20,8 +20,8 @@ const noopLogger: Logger = { }, } as unknown as Logger -function publishInstanceEvent(bus: EventBus, instanceId: string, event: Record) { - bus.publish({ type: "instance.event", instanceId, event: { ...event } as InstanceStreamEvent }) +function publishInstanceEvent(bus: EventBus, instanceId: string, event: Record, streamId?: string) { + bus.publish({ type: "instance.event", instanceId, streamId, event: { ...event } as InstanceStreamEvent }) } /** Publish a `session.*` event using the real OpenCode shape (`properties.info`). */ @@ -129,7 +129,7 @@ describe("AutoAcceptManager persistence", () => { const writes: unknown[][] = [] const persistence: AutoAcceptPersistence = { async loadSessions() { return [{ id: "root", parentId: null, yoloEnabled: false }] }, - async persist(...args) { writes.push(args); await gate }, + async persist(...args) { writes.push(args.slice(0, 4)); await gate }, } const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) const toggle = manager.toggle("inst", "root") @@ -160,8 +160,10 @@ describe("AutoAcceptManager persistence", () => { const bus = new EventBus(noopLogger) const replier = makeRecordingReplier() let attempts = 0 + const signals: AbortSignal[] = [] const persistence: AutoAcceptPersistence = { - async loadSessions() { + async loadSessions(_instanceId, signal) { + signals.push(signal) if (++attempts === 1) throw new Error("temporary failure") return [{ id: "root", parentId: null, yoloEnabled: true }] }, @@ -175,12 +177,217 @@ describe("AutoAcceptManager persistence", () => { }) await flushMicrotasks() assert.equal(replier.calls.length, 0) + assert.equal(signals[0]?.aborted, true) await manager.hydrateInstance("inst") await flushMicrotasks() + assert.notEqual(signals[1], signals[0]) assert.equal(replier.calls.length, 1) manager.stop() }) + it("releases generation controllers after arbitrary hydration failures", async () => { + const bus = new EventBus(noopLogger) + const signals: AbortSignal[] = [] + const persistence: AutoAcceptPersistence = { + async loadSessions(instanceId, signal) { + signals.push(signal) + throw new Error(`unknown workspace: ${instanceId}`) + }, + async persist() {}, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + + await assert.rejects(manager.hydrateInstance("invalid-a"), /unknown workspace/) + await assert.rejects(manager.hydrateInstance("invalid-b"), /unknown workspace/) + + assert.equal(signals.every((signal) => signal.aborted), true) + assert.equal( + (manager as unknown as { generationControllers: Map }).generationControllers.size, + 0, + ) + }) + + it("deduplicates one bounded hydration retry run and fails closed at the bounded event queue", { timeout: 2_000 }, async () => { + const bus = new EventBus(noopLogger) + const replier = makeRecordingReplier() + let loads = 0 + let failLoads = true + const persistence: AutoAcceptPersistence = { + async loadSessions() { + loads += 1 + if (failLoads) throw new Error("temporary hydration outage") + return [{ id: "root", parentId: null, yoloEnabled: true }] + }, + async persist() {}, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier, persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "current", status: "connecting" }) + + for (let index = 0; index < 700; index += 1) { + publishSession(bus, "inst", "session.updated", { id: `session-${index}`, parentID: null }) + } + for (let index = 0; index < 100; index += 1) { + publishSession(bus, "inst", "session.updated", { id: "same-session", parentID: null, title: String(index) }) + } + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "open", sessionID: "root" }, + }) + publishInstanceEvent(bus, "inst", { + type: "permission.updated", + properties: { id: "open", sessionID: "root", detail: "latest" }, + }) + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "closed", sessionID: "root" }, + }) + publishInstanceEvent(bus, "inst", { + type: "permission.v2.replied", + properties: { id: "closed" }, + }) + + await new Promise((resolve) => setTimeout(resolve, 100)) + assert.equal(loads, 4, "all queued events must share the initial load and three retries") + const queued = (manager as unknown as { + queuedEvents: Map> + }).queuedEvents.get("inst") + assert.equal(queued?.size, 512) + assert.equal(queued?.has("session:same-session"), true) + assert.equal(queued?.has("permission:open"), false) + assert.equal(queued?.has("permission:closed"), false) + + publishSession(bus, "inst", "session.updated", { id: "after-exhaustion", parentID: null }) + await new Promise((resolve) => setTimeout(resolve, 30)) + assert.equal(loads, 4, "new events must not start one load each after retry exhaustion") + + failLoads = false + await manager.hydrateInstance("inst") + await flushMicrotasks() + + assert.equal(loads, 6, "overflow must force a second authoritative load before replay") + assert.deepEqual(replier.calls, []) + assert.equal((manager as unknown as { queuedEvents: Map }).queuedEvents.has("inst"), false) + manager.stop() + }) + + it("replays deduplicated session topology before queued permissions", async () => { + const bus = new EventBus(noopLogger) + const replier = makeRecordingReplier() + let release!: () => void + const gate = new Promise((resolve) => { release = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions() { + await gate + return [ + { id: "enabled", parentId: null, yoloEnabled: true }, + { id: "disabled", parentId: null, yoloEnabled: false }, + ] + }, + async persist() {}, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier, persistence }) + manager.start() + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "enabled" }) + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "permission", sessionID: "child" }, + }) + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "disabled" }) + + const hydration = manager.hydrateInstance("inst") + release() + await hydration + await flushMicrotasks() + + assert.deepEqual(replier.calls, []) + assert.equal(manager.isEnabled("inst", "child"), false) + manager.stop() + }) + + it("keeps permissions queued until microtask-delayed topology replay is authoritative", async () => { + const bus = new EventBus(noopLogger) + const replier = makeRecordingReplier() + const persistence: AutoAcceptPersistence = { + async loadSessions() { + return [ + { id: "enabled", parentId: null, yoloEnabled: true }, + { id: "disabled", parentId: null, yoloEnabled: false }, + { id: "child", parentId: "enabled", yoloEnabled: false }, + ] + }, + async persist() {}, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier, persistence }) + let permissionScheduled = false + bus.on("yolo.stateChanged", () => { + if (permissionScheduled) return + permissionScheduled = true + queueMicrotask(() => publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "permission", sessionID: "child" }, + })) + }) + manager.start() + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "disabled" }) + await manager.hydrateInstance("inst") + await flushMicrotasks() + + assert.deepEqual(replier.calls, []) + assert.equal(manager.isEnabled("inst", "child"), false) + manager.stop() + }) + + it("drops permission auto-accept and rehydrates when the 513th queued event reaches the cap", async () => { + const bus = new EventBus(noopLogger) + const replier = makeRecordingReplier() + let release!: () => void + let loads = 0 + const gate = new Promise((resolve) => { release = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions() { + const load = ++loads + if (load === 1) await gate + return [ + { id: "enabled", parentId: null, yoloEnabled: true }, + { id: "disabled", parentId: null, yoloEnabled: false }, + { id: "child", parentId: load === 1 ? "enabled" : "disabled", yoloEnabled: false }, + ] + }, + async persist() {}, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier, persistence }) + manager.start() + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "disabled" }) + for (let index = 0; index < 511; index += 1) { + publishSession(bus, "inst", "session.updated", { id: `filler-${index}`, parentID: null }) + } + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "permission", sessionID: "child" }, + }) + + const queued = (manager as unknown as { + queuedEvents: Map> + }).queuedEvents.get("inst") + assert.equal(queued?.size, 512) + assert.equal(queued?.has("session:child"), true) + assert.equal(queued?.has("permission:permission"), false) + + const hydration = manager.hydrateInstance("inst") + release() + await hydration + await flushMicrotasks() + + assert.equal(loads, 2) + assert.deepEqual(replier.calls, []) + assert.equal(manager.isEnabled("inst", "child"), false) + manager.stop() + }) + it("does not re-enable memory when a persisted toggle finishes after cleanup", async () => { const bus = new EventBus(noopLogger) const changes: Record[] = [] @@ -188,23 +395,177 @@ describe("AutoAcceptManager persistence", () => { let release!: () => void const gate = new Promise((resolve) => { release = resolve }) let writes = 0 + let staleSignal: AbortSignal | undefined const persistence: AutoAcceptPersistence = { async loadSessions() { return [{ id: "root", parentId: null, yoloEnabled: false }] }, - async persist() { writes += 1; await gate }, + async persist(_instanceId, _rootSessionId, _enabled, _workspaceId, signal) { + staleSignal = signal + await gate + if (signal.aborted) throw signal.reason + writes += 1 + }, } const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) const toggle = manager.toggle("inst", "root") const queued = manager.toggle("inst", "root") await flushMicrotasks() manager.clearInstance("inst") + assert.equal(staleSignal?.aborted, true) release() assert.equal(await toggle, false) assert.equal(await queued, false) - assert.equal(writes, 1) + assert.equal(writes, 0) assert.equal(manager.isEnabled("inst", "root"), false) assert.equal(changes.length, 0) }) + it("reports a committed toggle after runtime rotation and hydrates the durable value", async () => { + const bus = new EventBus(noopLogger) + const changes: Array<{ enabled?: boolean }> = [] + let durableEnabled = false + const persistence: AutoAcceptPersistence = { + async loadSessions() { + return [{ id: "root", parentId: null, yoloEnabled: durableEnabled }] + }, + async persist(_instanceId, _rootSessionId, enabled) { + durableEnabled = enabled + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "new", status: "connecting" }) + }, + } + bus.on("yolo.stateChanged", (event) => changes.push(event)) + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + + assert.equal(await manager.toggle("inst", "root"), true) + await manager.hydrateInstance("inst") + + assert.equal(durableEnabled, true) + assert.equal(manager.isEnabled("inst", "root"), true) + assert.deepEqual(changes.map((event) => event.enabled), [true]) + manager.stop() + }) + + it("rebinds the requested toggle when rotation drops the first persistence attempt", async () => { + const bus = new EventBus(noopLogger) + const changes: Array<{ enabled?: boolean }> = [] + const writes: Array<{ enabled: boolean; signal: AbortSignal }> = [] + let durableEnabled = false + let loads = 0 + const persistence: AutoAcceptPersistence = { + async loadSessions() { + loads += 1 + return [{ id: "root", parentId: null, yoloEnabled: durableEnabled }] + }, + async persist(_instanceId, _rootSessionId, enabled, _workspaceId, signal) { + writes.push({ enabled, signal }) + if (writes.length === 1) { + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + return + } + durableEnabled = enabled + }, + } + bus.on("yolo.stateChanged", (event) => changes.push(event)) + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + + assert.equal(await manager.toggle("inst", "root"), true) + + assert.equal(loads, 2) + assert.deepEqual(writes.map(({ enabled }) => enabled), [true, true]) + assert.equal(writes[0]?.signal.aborted, true) + assert.equal(writes[1]?.signal.aborted, false) + assert.equal(durableEnabled, true) + assert.equal(manager.isEnabled("inst", "root"), true) + assert.deepEqual(changes.map(({ enabled }) => enabled), [true]) + manager.stop() + }) + + it("publishes the authoritative disabled state when rotation aborts persistence after commit", async () => { + const bus = new EventBus(noopLogger) + const changes: Array<{ sessionId?: string; enabled?: boolean }> = [] + let durableEnabled = true + let loads = 0 + let persistSignal: AbortSignal | undefined + const persistence: AutoAcceptPersistence = { + async loadSessions() { + loads += 1 + return [ + { id: "root", parentId: null, yoloEnabled: durableEnabled }, + ...(loads === 1 ? [] : [{ id: "unrelated", parentId: null, yoloEnabled: true }]), + ] + }, + async persist(_instanceId, _rootSessionId, enabled, _workspaceId, signal) { + durableEnabled = enabled + persistSignal = signal + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + throw signal.reason + }, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + bus.on("yolo.stateChanged", (event) => changes.push(event)) + + assert.equal(await manager.toggle("inst", "root"), false) + + assert.equal(persistSignal?.aborted, true) + assert.equal(loads, 2) + assert.equal(manager.isEnabled("inst", "root"), false) + assert.equal(manager.isEnabled("inst", "unrelated"), true) + assert.deepEqual(changes.map(({ sessionId, enabled }) => [sessionId, enabled]), [ + ["unrelated", true], + ["root", false], + ]) + manager.stop() + }) + + it("rebinds a toggle behind replacement hydration when runtime rotation interrupts initial hydration", async () => { + const bus = new EventBus(noopLogger) + const changes: Array<{ enabled?: boolean }> = [] + const writes: boolean[] = [] + let releaseOld!: () => void + let releaseReplacement!: () => void + let loads = 0 + const oldGate = new Promise((resolve) => { releaseOld = resolve }) + const replacementGate = new Promise((resolve) => { releaseReplacement = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions() { + await (++loads === 1 ? oldGate : replacementGate) + return [{ id: "root", parentId: null, yoloEnabled: false }] + }, + async persist(_instanceId, _rootSessionId, enabled) { writes.push(enabled) }, + } + bus.on("yolo.stateChanged", (event) => changes.push(event)) + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + const oldHydration = manager.hydrateInstance("inst") + const toggle = manager.toggle("inst", "root") + + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + releaseOld() + await oldHydration + await flushMicrotasks() + + assert.equal(loads, 2) + assert.deepEqual(writes, [], "stale authority must not persist before replacement hydration") + assert.equal(changes.length, 0, "stale authority must not publish replacement state") + assert.equal(manager.isEnabled("inst", "root"), false) + + releaseReplacement() + assert.equal(await toggle, true) + assert.deepEqual(writes, [true]) + assert.equal(manager.isEnabled("inst", "root"), true) + assert.deepEqual(changes.map((event) => event.enabled), [true]) + manager.stop() + }) + it("moves persisted Yolo state when late ancestry changes the family root", async () => { const bus = new EventBus(noopLogger) const writes: unknown[][] = [] @@ -215,7 +576,7 @@ describe("AutoAcceptManager persistence", () => { { id: "child", parentId: null, workspaceId: "workspace", yoloEnabled: true }, ] }, - async persist(...args) { writes.push(args) }, + async persist(...args) { writes.push(args.slice(0, 4)) }, } const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) manager.start() @@ -230,6 +591,368 @@ describe("AutoAcceptManager persistence", () => { manager.stop() }) + it("re-resolves the authoritative root when topology advances during migration rebind", async () => { + const bus = new EventBus(noopLogger) + const changes: Array<{ sessionId?: string; enabled?: boolean }> = [] + const durableEnabled = new Map([["grandparent", false], ["parent", false], ["child", true]]) + const writes: Array<{ rootSessionId: string; enabled: boolean; signal: AbortSignal }> = [] + let rotated = false + let releaseStale!: () => void + let markStaleStarted!: () => void + let markMigrated!: () => void + const staleGate = new Promise((resolve) => { releaseStale = resolve }) + const staleStarted = new Promise((resolve) => { markStaleStarted = resolve }) + const migrated = new Promise((resolve) => { markMigrated = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions() { + return [ + { id: "grandparent", parentId: null, workspaceId: "workspace", yoloEnabled: durableEnabled.get("grandparent") ?? false }, + { id: "parent", parentId: null, workspaceId: "workspace", yoloEnabled: durableEnabled.get("parent") ?? false }, + { + id: "child", + parentId: rotated ? "grandparent" : null, + workspaceId: "workspace", + yoloEnabled: durableEnabled.get("child") ?? false, + }, + ] + }, + async persist(_instanceId, rootSessionId, enabled, _workspaceId, signal) { + writes.push({ rootSessionId, enabled, signal }) + if (!rotated) { + rotated = true + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + markStaleStarted() + await staleGate + return + } + durableEnabled.set(rootSessionId, enabled) + if (rootSessionId === "child" && !enabled) markMigrated() + }, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + bus.on("yolo.stateChanged", (event) => changes.push(event)) + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "parent", workspaceID: "workspace" }) + await staleStarted + + assert.equal(writes[0]?.signal.aborted, true) + assert.equal(manager.isEnabled("inst", "parent"), false) + assert.equal(changes.length, 0, "stale migration must not publish replacement state") + + releaseStale() + await migrated + await flushMicrotasks() + + assert.deepEqual(writes.map(({ rootSessionId, enabled }) => [rootSessionId, enabled]), [ + ["parent", true], + ["grandparent", true], + ["child", false], + ]) + assert.notEqual(writes[1]?.signal, writes[0]?.signal) + assert.equal(writes[1]?.signal.aborted, false) + assert.equal(durableEnabled.get("grandparent"), true) + assert.equal(durableEnabled.get("parent"), false) + assert.equal(durableEnabled.get("child"), false) + assert.equal(manager.isEnabled("inst", "grandparent"), true) + assert.equal(manager.isEnabled("inst", "parent"), false) + assert.deepEqual(changes.map(({ sessionId, enabled }) => [sessionId, enabled]), [ + ["grandparent", true], + ]) + manager.stop() + }) + + it("automatically retries replacement hydration for the same current generation", { timeout: 1_000 }, async () => { + const bus = new EventBus(noopLogger) + const changes: Array<{ sessionId?: string; enabled?: boolean }> = [] + const durableEnabled = new Map([["parent", false], ["child", true]]) + const writes: Array<{ rootSessionId: string; enabled: boolean; signal: AbortSignal }> = [] + const loadSignals: AbortSignal[] = [] + let loads = 0 + let releaseFailedHydration!: () => void + let markFailedHydrationStarted!: () => void + let markMigrated!: () => void + const failedHydrationGate = new Promise((resolve) => { releaseFailedHydration = resolve }) + const failedHydrationStarted = new Promise((resolve) => { markFailedHydrationStarted = resolve }) + const migrated = new Promise((resolve) => { markMigrated = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions(_instanceId, signal) { + loadSignals.push(signal) + loads += 1 + if (loads === 2) { + markFailedHydrationStarted() + await failedHydrationGate + throw new Error("temporary hydration failure") + } + return [ + { id: "parent", parentId: null, workspaceId: "workspace", yoloEnabled: durableEnabled.get("parent") ?? false }, + { + id: "child", + parentId: loads === 1 ? null : "parent", + workspaceId: "workspace", + yoloEnabled: durableEnabled.get("child") ?? false, + }, + ] + }, + async persist(_instanceId, rootSessionId, enabled, _workspaceId, signal) { + writes.push({ rootSessionId, enabled, signal }) + if (writes.length === 1) { + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + return + } + durableEnabled.set(rootSessionId, enabled) + if (rootSessionId === "child" && !enabled) markMigrated() + }, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + bus.on("yolo.stateChanged", (event) => changes.push(event)) + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "parent", workspaceID: "workspace" }) + await markFailedHydrationStarted + releaseFailedHydration() + await flushMicrotasks() + + assert.equal(loads, 2) + assert.equal(manager.isEnabled("inst", "parent"), false) + assert.equal(changes.length, 0) + + await migrated + await flushMicrotasks() + + assert.equal(loads, 3) + assert.equal(loadSignals[1]?.aborted, true) + assert.notEqual(loadSignals[2], loadSignals[1]) + assert.equal(loadSignals[2]?.aborted, false) + assert.deepEqual(writes.map(({ rootSessionId, enabled }) => [rootSessionId, enabled]), [ + ["parent", true], + ["parent", true], + ["child", false], + ]) + assert.equal(writes[0]?.signal.aborted, true) + assert.notEqual(writes[1]?.signal, writes[0]?.signal) + assert.equal(durableEnabled.get("parent"), true) + assert.equal(durableEnabled.get("child"), false) + assert.equal(manager.isEnabled("inst", "parent"), true) + assert.deepEqual(changes.map(({ sessionId, enabled }) => [sessionId, enabled]), [["parent", true]]) + manager.stop() + }) + + it("retries a current-generation rebound migration after a transient persistence outage", { timeout: 1_000 }, async () => { + const bus = new EventBus(noopLogger) + const changes: Array<{ sessionId?: string; enabled?: boolean }> = [] + const durableEnabled = new Map([["parent", false], ["child", true]]) + const writes: Array<{ rootSessionId: string; enabled: boolean; signal: AbortSignal }> = [] + let loads = 0 + let outageInjected = false + let markMigrated!: () => void + const migrated = new Promise((resolve) => { markMigrated = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions() { + loads += 1 + return [ + { id: "parent", parentId: null, workspaceId: "workspace", yoloEnabled: durableEnabled.get("parent") ?? false }, + { + id: "child", + parentId: loads === 1 ? null : "parent", + workspaceId: "workspace", + yoloEnabled: durableEnabled.get("child") ?? false, + }, + ] + }, + async persist(_instanceId, rootSessionId, enabled, _workspaceId, signal) { + writes.push({ rootSessionId, enabled, signal }) + if (writes.length === 1) { + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + return + } + if (!outageInjected) { + outageInjected = true + throw new Error("temporary persistence outage") + } + durableEnabled.set(rootSessionId, enabled) + if (rootSessionId === "child" && !enabled) markMigrated() + }, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + bus.on("yolo.stateChanged", (event) => changes.push(event)) + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "parent", workspaceID: "workspace" }) + await migrated + await flushMicrotasks() + + assert.equal(loads, 3, "recovery must not require manual hydration") + assert.deepEqual(writes.map(({ rootSessionId, enabled }) => [rootSessionId, enabled]), [ + ["parent", true], + ["parent", true], + ["parent", true], + ["child", false], + ]) + assert.equal(writes[0]?.signal.aborted, true) + assert.equal(writes[1]?.signal.aborted, true) + assert.notEqual(writes[2]?.signal, writes[1]?.signal) + assert.equal(writes[2]?.signal.aborted, false) + assert.equal(writes[3]?.signal, writes[2]?.signal) + assert.equal(durableEnabled.get("parent"), true) + assert.equal(durableEnabled.get("child"), false) + assert.equal(manager.isEnabled("inst", "parent"), true) + assert.equal(manager.isEnabled("inst", "child"), true) + assert.deepEqual(changes.map(({ sessionId, enabled }) => [sessionId, enabled]), [["parent", true]]) + manager.stop() + }) + + it("does not let a repaired migration override a newer successful disable", { timeout: 1_000 }, async () => { + const bus = new EventBus(noopLogger) + const durableEnabled = new Map([["parent", false], ["child", true]]) + const writes: Array<{ rootSessionId: string; enabled: boolean }> = [] + let loads = 0 + let childClearAttempts = 0 + let markChildClearFailed!: () => void + let markChildCleared!: () => void + const childClearFailed = new Promise((resolve) => { markChildClearFailed = resolve }) + const childCleared = new Promise((resolve) => { markChildCleared = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions() { + loads += 1 + return [ + { id: "parent", parentId: null, workspaceId: "workspace", yoloEnabled: durableEnabled.get("parent") ?? false }, + { + id: "child", + parentId: loads === 1 ? null : "parent", + workspaceId: "workspace", + yoloEnabled: durableEnabled.get("child") ?? false, + }, + ] + }, + async persist(_instanceId, rootSessionId, enabled) { + writes.push({ rootSessionId, enabled }) + durableEnabled.set(rootSessionId, enabled) + if (writes.length === 1) { + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + return + } + if (rootSessionId === "child" && !enabled && ++childClearAttempts === 1) { + durableEnabled.set(rootSessionId, true) + markChildClearFailed() + throw new Error("temporary child clear failure") + } + if (rootSessionId === "child" && !enabled) markChildCleared() + }, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "parent", workspaceID: "workspace" }) + await childClearFailed + await childCleared + + assert.equal(await manager.toggle("inst", "parent"), false) + await new Promise((resolve) => setTimeout(resolve, 50)) + + const disableIndex = writes.findIndex(({ rootSessionId, enabled }) => rootSessionId === "parent" && !enabled) + assert.notEqual(disableIndex, -1) + assert.equal( + writes.slice(disableIndex + 1).some(({ rootSessionId, enabled }) => rootSessionId === "parent" && enabled), + false, + ) + assert.equal(durableEnabled.get("parent"), false) + assert.equal(durableEnabled.get("child"), false) + assert.equal(manager.isEnabled("inst", "parent"), false) + manager.stop() + }) + + it("self-heals a partial migration after retry exhaustion and process restart", { timeout: 1_000 }, async () => { + const bus = new EventBus(noopLogger) + const durableEnabled = new Map([ + ["parent", false], + ["child", true], + ["unrelated", false], + ["isolated-fork", true], + ]) + let loads = 0 + let childClearAttempts = 0 + let markRetryBudgetExhausted!: () => void + const retryBudgetExhausted = new Promise((resolve) => { markRetryBudgetExhausted = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions() { + loads += 1 + return [ + { id: "parent", parentId: null, workspaceId: "workspace", yoloEnabled: durableEnabled.get("parent") ?? false }, + { + id: "child", + parentId: loads === 1 ? null : "parent", + workspaceId: "workspace", + yoloEnabled: durableEnabled.get("child") ?? false, + }, + { id: "unrelated", parentId: null, workspaceId: "workspace", yoloEnabled: false }, + { + id: "isolated-fork", + parentId: "unrelated", + revert: { snapshot: "fork" }, + workspaceId: "workspace", + yoloEnabled: durableEnabled.get("isolated-fork") ?? false, + }, + ] + }, + async persist(_instanceId, rootSessionId, enabled) { + durableEnabled.set(rootSessionId, enabled) + if (rootSessionId === "parent" && enabled && loads === 1) { + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "replacement", status: "connecting" }) + return + } + if (rootSessionId === "child" && !enabled) { + childClearAttempts += 1 + if (childClearAttempts <= 4) { + durableEnabled.set(rootSessionId, true) + if (childClearAttempts === 4) markRetryBudgetExhausted() + throw new Error("child clear unavailable") + } + } + }, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier: makeRecordingReplier(), persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + + publishSession(bus, "inst", "session.updated", { id: "child", parentID: "parent", workspaceID: "workspace" }) + await retryBudgetExhausted + await flushMicrotasks() + assert.equal(durableEnabled.get("parent"), true) + assert.equal(durableEnabled.get("child"), true) + + manager.stop() + + const restartedBus = new EventBus(noopLogger) + const restartedManager = new AutoAcceptManager({ + eventBus: restartedBus, + logger: noopLogger, + replier: makeRecordingReplier(), + persistence, + }) + restartedManager.start() + restartedBus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "after-restart", status: "connecting" }) + await restartedManager.hydrateInstance("inst") + + assert.equal(childClearAttempts, 5) + assert.equal(durableEnabled.get("parent"), true) + assert.equal(durableEnabled.get("child"), false) + assert.equal(restartedManager.isEnabled("inst", "parent"), true) + assert.equal(restartedManager.isEnabled("inst", "unrelated"), false) + assert.equal(restartedManager.isEnabled("inst", "isolated-fork"), true) + assert.equal(durableEnabled.get("isolated-fork"), true) + restartedManager.stop() + }) + it("does not re-enable a family when ancestry repeatedly changes during a disable", async () => { const bus = new EventBus(noopLogger) let releaseFirst!: () => void @@ -504,6 +1227,166 @@ describe("AutoAcceptManager lifecycle", () => { assert.equal(replier.calls.length, 0) }) + + it("stop() fences late hydration and refuses new work until restarted", async () => { + const bus = new EventBus(noopLogger) + const changes: Record[] = [] + const signals: AbortSignal[] = [] + let release!: () => void + let loads = 0 + const gate = new Promise((resolve) => { release = resolve }) + const persistence: AutoAcceptPersistence = { + async loadSessions(_instanceId, signal) { + loads += 1 + signals.push(signal) + if (loads === 1) await gate + return [{ id: "root", parentId: null, yoloEnabled: true }] + }, + async persist() {}, + } + bus.on("yolo.stateChanged", (event) => changes.push(event)) + const manager = new AutoAcceptManager({ + eventBus: bus, + logger: noopLogger, + replier: makeRecordingReplier(), + persistence, + }) + manager.start() + + const hydration = manager.hydrateInstance("inst") + manager.stop() + assert.equal(signals[0]?.aborted, true) + release() + await hydration + await manager.hydrateInstance("inst") + + assert.equal(loads, 1, "stopped managers must not create new generation work") + assert.equal(manager.isEnabled("inst", "root"), false) + assert.equal(changes.length, 0) + + manager.start() + await manager.hydrateInstance("inst") + assert.equal(loads, 2) + assert.equal(signals[1]?.aborted, false) + assert.equal(manager.isEnabled("inst", "root"), true) + manager.stop() + }) + + it("stop() ignores a reply that settles after cancellation", async () => { + const bus = new EventBus(noopLogger) + const accepted: Record[] = [] + let signal: AbortSignal | undefined + let release!: () => void + const replier: PermissionReplier = (_reply, replySignal) => { + signal = replySignal + return new Promise((resolve) => { release = resolve }) + } + bus.on("yolo.autoAccepted", (event) => accepted.push(event)) + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier }) + manager.start() + publishSession(bus, "inst", "session.updated", { id: "root", parentID: null }) + manager.toggle("inst", "root") + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "permission", sessionID: "root" }, + }) + + manager.stop() + assert.equal(signal?.aborted, true) + release() + await flushMicrotasks() + + assert.equal(accepted.length, 0) + }) +}) + +describe("AutoAcceptManager runtime rotation", () => { + it("clears runtime state and rehydrates persisted roots for the new stream", async () => { + const bus = new EventBus(noopLogger) + const replier = makeRecordingReplier() + let loads = 0 + const persistence: AutoAcceptPersistence = { + async loadSessions() { + loads += 1 + return loads === 1 + ? [ + { id: "old-root", parentId: null, yoloEnabled: true }, + { id: "old-child", parentId: "old-root", yoloEnabled: false }, + { id: "stale", parentId: null, yoloEnabled: false }, + ] + : [{ id: "new-root", parentId: null, yoloEnabled: true }] + }, + async persist() {}, + } + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier, persistence }) + manager.start() + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + await manager.hydrateInstance("inst") + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "stale-permission", sessionID: "stale" }, + }, "old") + + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "new", status: "connecting" }) + await manager.hydrateInstance("inst") + await flushMicrotasks() + + assert.equal(loads, 2) + assert.equal(manager.isEnabled("inst", "old-child"), false) + assert.equal(manager.isEnabled("inst", "new-root"), true) + assert.equal(replier.calls.length, 0, "stale pending permission must not drain into the new runtime") + manager.stop() + }) + + it("ignores an old reply completion without disturbing the new in-flight reply", async () => { + const bus = new EventBus(noopLogger) + const completions: Array<() => void> = [] + const calls: AutoAcceptReply[] = [] + const accepted: Record[] = [] + const signals: AbortSignal[] = [] + const replier: PermissionReplier = (reply, signal) => { + calls.push(reply) + signals.push(signal) + return new Promise((resolve) => completions.push(resolve)) + } + bus.on("yolo.autoAccepted", (event) => accepted.push(event)) + const manager = new AutoAcceptManager({ eventBus: bus, logger: noopLogger, replier }) + manager.start() + + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "old", status: "connecting" }) + publishSession(bus, "inst", "session.updated", { id: "root", parentID: null }) + manager.toggle("inst", "root") + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "same-permission", sessionID: "root" }, + }, "old") + + bus.publish({ type: "instance.eventStatus", instanceId: "inst", streamId: "new", status: "connecting" }) + assert.equal(signals[0]?.aborted, true) + publishSession(bus, "inst", "session.updated", { id: "root", parentID: null }) + manager.toggle("inst", "root") + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "same-permission", sessionID: "root" }, + }, "new") + assert.equal(calls.length, 2) + assert.equal(signals[1]?.aborted, false) + + completions[0]() + await flushMicrotasks() + assert.equal(accepted.length, 0) + publishInstanceEvent(bus, "inst", { + type: "permission.v2.asked", + properties: { id: "same-permission", sessionID: "root" }, + }, "new") + assert.equal(calls.length, 2, "old completion must not clear the new in-flight marker") + + completions[1]() + await flushMicrotasks() + assert.equal(accepted.length, 1) + assert.equal(accepted[0].permissionId, "same-permission") + manager.stop() + }) }) describe("AutoAcceptManager pending permissions drain", () => { diff --git a/packages/server/src/permissions/auto-accept-manager.ts b/packages/server/src/permissions/auto-accept-manager.ts index 3fc7e7e60..1bff20069 100644 --- a/packages/server/src/permissions/auto-accept-manager.ts +++ b/packages/server/src/permissions/auto-accept-manager.ts @@ -27,7 +27,7 @@ export interface AutoAcceptReply { reply: PermissionReplyValue } -export type PermissionReplier = (reply: AutoAcceptReply) => Promise +export type PermissionReplier = (reply: AutoAcceptReply, signal: AbortSignal) => Promise interface PendingPermission { permissionId: string @@ -48,17 +48,30 @@ export interface PersistedAutoAcceptSession extends AutoAcceptSessionInfo { } export interface AutoAcceptPersistence { - loadSessions(instanceId: string): Promise - persist(instanceId: string, rootSessionId: string, enabled: boolean, workspaceId?: string): Promise + loadSessions(instanceId: string, signal: AbortSignal): Promise + persist( + instanceId: string, + rootSessionId: string, + enabled: boolean, + workspaceId: string | undefined, + signal: AbortSignal, + ): Promise } const PERMISSION_ASK_TYPES = new Set(["permission.v2.asked", "permission.asked", "permission.updated"]) const PERMISSION_REPLIED_TYPES = new Set(["permission.v2.replied", "permission.replied"]) const SESSION_UPSERT_TYPES = new Set(["session.updated", "session.created"]) const SESSION_REMOVE_TYPES = new Set(["session.deleted"]) +const REBIND_TOGGLE_AFTER_ROTATION = new Error("Rebind Yolo toggle after runtime rotation") export class AutoAcceptManager { private static readonly MAX_REPLY_ATTEMPTS = 3 + private static readonly MAX_HYDRATION_RETRY_ATTEMPTS = 3 + private static readonly HYDRATION_RETRY_MS = 10 + private static readonly MAX_QUEUED_EVENTS_PER_INSTANCE = 512 + private static readonly MAX_PENDING_PER_INSTANCE = 512 + private static readonly MAX_REBOUND_MIGRATION_ATTEMPTS = 3 + private static readonly REBOUND_MIGRATION_RETRY_MS = 10 private readonly store = new AutoAcceptStore() /** instanceId:permissionId entries currently being replied, to dedupe re-emissions */ private readonly inFlight = new Set() @@ -68,23 +81,38 @@ export class AutoAcceptManager { private readonly replyAttempts = new Map() private readonly hydratedInstances = new Set() private readonly hydration = new Map>() - private readonly queuedEvents = new Map() + private readonly failedHydrations = new Set() + private readonly queuedEvents = new Map>() + private readonly overflowedEventQueues = new Set() private readonly instanceGeneration = new Map() + private readonly generationControllers = new Map() + private readonly streamIds = new Map() private readonly sessionWorkspaces = new Map>() + private readonly persistedEnabledSessions = new Map>() + private readonly reboundRootMigrations = new Map>() + private readonly reboundRootMigrationAttempts = new Map() + private readonly reboundRootMigrationTimers = new Map>() + private readonly reboundRootMigrationRuns = new Set() private readonly mutations = new Map>() private unsubscribe?: () => void + private stopped = false constructor(private readonly deps: AutoAcceptManagerDeps) {} start(): void { if (this.unsubscribe) return - const handler = (payload: { instanceId?: string; event?: InstanceStreamPayload }) => { + this.stopped = false + const handler = (payload: { instanceId?: string; streamId?: string; event?: InstanceStreamPayload }) => { if (!payload || !payload.instanceId || !payload.event) return + if (payload.streamId) { + const current = this.streamIds.get(payload.instanceId) + if (current && current !== payload.streamId) return + this.streamIds.set(payload.instanceId, payload.streamId) + } if (this.deps.persistence && !this.hydratedInstances.has(payload.instanceId)) { - const queued = this.queuedEvents.get(payload.instanceId) ?? [] - queued.push(payload.event) - this.queuedEvents.set(payload.instanceId, queued) - void this.hydrateInstance(payload.instanceId).catch((error) => { + this.queueEvent(payload.instanceId, payload.event) + if (this.failedHydrations.has(payload.instanceId)) return + void this.beginHydration(payload.instanceId).catch((error) => { this.deps.logger.warn({ instanceId: payload.instanceId, err: error }, "Failed to hydrate persisted Yolo state") }) return @@ -104,12 +132,27 @@ export class AutoAcceptManager { const onError = (event: { workspace?: { id?: string } }) => { if (event?.workspace?.id) this.clearInstance(event.workspace.id) } + const onStreamStatus = (event: { instanceId?: string; streamId?: string; status?: string }) => { + if (!event.instanceId || !event.streamId || event.status !== "connecting") return + const current = this.streamIds.get(event.instanceId) + if (current && current !== event.streamId) { + this.invalidateRuntimeState(event.instanceId) + this.streamIds.set(event.instanceId, event.streamId) + void this.hydrateInstance(event.instanceId).catch((error) => { + this.deps.logger.warn({ instanceId: event.instanceId, err: error }, "Failed to hydrate persisted Yolo state") + }) + return + } + this.streamIds.set(event.instanceId, event.streamId) + } this.deps.eventBus.on("instance.event", handler) + this.deps.eventBus.on("instance.eventStatus", onStreamStatus) this.deps.eventBus.on("workspace.started", onStarted) this.deps.eventBus.on("workspace.stopped", onStopped) this.deps.eventBus.on("workspace.error", onError) this.unsubscribe = () => { this.deps.eventBus.off("instance.event", handler) + this.deps.eventBus.off("instance.eventStatus", onStreamStatus) this.deps.eventBus.off("workspace.started", onStarted) this.deps.eventBus.off("workspace.stopped", onStopped) this.deps.eventBus.off("workspace.error", onError) @@ -119,6 +162,26 @@ export class AutoAcceptManager { stop(): void { this.unsubscribe?.() this.unsubscribe = undefined + this.stopped = true + const instanceIds = new Set([ + ...this.generationControllers.keys(), + ...this.instanceGeneration.keys(), + ...this.streamIds.keys(), + ...this.hydratedInstances, + ...this.hydration.keys(), + ...this.failedHydrations, + ...this.queuedEvents.keys(), + ...this.overflowedEventQueues, + ...this.sessionWorkspaces.keys(), + ...this.persistedEnabledSessions.keys(), + ...this.mutations.keys(), + ...this.pending.keys(), + ]) + for (const instanceId of instanceIds) this.invalidateRuntimeState(instanceId) + for (const instanceId of this.reboundRootMigrationTimers.keys()) this.clearReboundRootMigrationRetry(instanceId) + this.reboundRootMigrations.clear() + this.reboundRootMigrationRuns.clear() + this.streamIds.clear() } isEnabled(instanceId: string, sessionId: string): boolean { @@ -126,41 +189,118 @@ export class AutoAcceptManager { } hydrateInstance(instanceId: string): Promise { - if (!this.deps.persistence || this.hydratedInstances.has(instanceId)) return Promise.resolve() + this.failedHydrations.delete(instanceId) + return this.beginHydration(instanceId) + } + + private beginHydration(instanceId: string): Promise { + if (this.stopped || !this.deps.persistence || this.hydratedInstances.has(instanceId)) return Promise.resolve() const existing = this.hydration.get(instanceId) if (existing) return existing const generation = this.instanceGeneration.get(instanceId) ?? 0 - let hydrated = false - const pending = this.deps.persistence.loadSessions(instanceId).then((sessions) => { - if ((this.instanceGeneration.get(instanceId) ?? 0) !== generation) return - this.store.clearInstance(instanceId) - const workspaces = new Map() - for (const session of sessions) { - this.store.upsertSession(instanceId, session) - if (session.workspaceId) workspaces.set(session.id, session.workspaceId) + const pending = this.runHydrationSequence(instanceId, generation).then(async (hydrated) => { + while (hydrated && this.overflowedEventQueues.delete(instanceId)) { + hydrated = await this.runHydrationSequence(instanceId, generation) } - this.sessionWorkspaces.set(instanceId, workspaces) - for (const session of sessions) { - if (!session.yoloEnabled || this.store.familyRoot(instanceId, session.id) !== session.id) continue - this.store.setEnabled(instanceId, session.id, true) - this.deps.eventBus.publish({ type: "yolo.stateChanged", instanceId, sessionId: session.id, enabled: true }) - this.drainPending(instanceId, session.id) + if (!hydrated) return + const queued = this.queuedEvents.get(instanceId) + for (const event of queued?.values() ?? []) { + if (isSessionEvent(event)) this.handleInstanceEvent(instanceId, event) } + if (!this.hasRuntimeAuthority(instanceId, generation, this.generationSignal(instanceId))) return this.hydratedInstances.add(instanceId) - hydrated = true + this.queuedEvents.delete(instanceId) + for (const event of queued?.values() ?? []) { + if (!isSessionEvent(event)) this.handleInstanceEvent(instanceId, event) + } + void this.retryReboundRootMigrations(instanceId) }).finally(() => { - if ((this.instanceGeneration.get(instanceId) ?? 0) !== generation) return + if (this.hydration.get(instanceId) !== pending) return this.hydration.delete(instanceId) - if (!hydrated) return - const queued = this.queuedEvents.get(instanceId) ?? [] - this.queuedEvents.delete(instanceId) - for (const event of queued) this.handleInstanceEvent(instanceId, event) }) this.hydration.set(instanceId, pending) return pending } + private async runHydrationSequence(instanceId: string, generation: number): Promise { + for (let attempt = 0; attempt <= AutoAcceptManager.MAX_HYDRATION_RETRY_ATTEMPTS; attempt += 1) { + try { + const hydrated = await this.hydrateInstanceOnce(instanceId, generation) + if (hydrated) this.failedHydrations.delete(instanceId) + return hydrated + } catch (error) { + if (this.stopped || (this.instanceGeneration.get(instanceId) ?? 0) !== generation) return false + if (attempt === AutoAcceptManager.MAX_HYDRATION_RETRY_ATTEMPTS) { + this.failedHydrations.add(instanceId) + throw error + } + this.deps.logger.warn({ instanceId, err: error, attempt: attempt + 1 }, "Failed to hydrate persisted Yolo state; retrying") + await new Promise((resolve) => setTimeout(resolve, AutoAcceptManager.HYDRATION_RETRY_MS * (attempt + 1))) + } + } + return false + } + + private async hydrateInstanceOnce(instanceId: string, generation: number): Promise { + const signal = this.generationSignal(instanceId) + const priorMutation = this.mutations.get(instanceId) + const load = priorMutation + ? priorMutation.catch(() => false).then(() => { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return undefined + return this.deps.persistence!.loadSessions(instanceId, signal) + }) + : this.deps.persistence!.loadSessions(instanceId, signal) + try { + const sessions = await load + if (!sessions) return false + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + this.store.clearInstance(instanceId) + const workspaces = new Map() + for (const session of sessions) { + this.store.upsertSession(instanceId, session) + if (session.workspaceId) workspaces.set(session.id, session.workspaceId) + } + this.sessionWorkspaces.set(instanceId, workspaces) + const persistedEnabled = new Set(sessions.filter((session) => session.yoloEnabled).map((session) => session.id)) + for (const formerRootSessionId of Array.from(persistedEnabled)) { + const currentRootSessionId = this.store.familyRoot(instanceId, formerRootSessionId) + if (currentRootSessionId === formerRootSessionId) continue + await this.deps.persistence!.persist( + instanceId, + currentRootSessionId, + true, + workspaces.get(currentRootSessionId), + signal, + ) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + persistedEnabled.add(currentRootSessionId) + await this.deps.persistence!.persist( + instanceId, + formerRootSessionId, + false, + workspaces.get(formerRootSessionId), + signal, + ) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + persistedEnabled.delete(formerRootSessionId) + } + this.persistedEnabledSessions.set(instanceId, persistedEnabled) + for (const rootSessionId of persistedEnabled) { + if (this.store.familyRoot(instanceId, rootSessionId) !== rootSessionId) continue + this.store.setEnabled(instanceId, rootSessionId, true) + this.deps.eventBus.publish({ type: "yolo.stateChanged", instanceId, sessionId: rootSessionId, enabled: true }) + this.drainPending(instanceId, rootSessionId) + } + return true + } catch (error) { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + this.releaseFailedHydration(instanceId, signal) + throw error + } + } + toggle(instanceId: string, sessionId: string): boolean | Promise { + if (this.stopped) return this.store.isEnabled(instanceId, sessionId) if (this.deps.persistence) return this.togglePersisted(instanceId, sessionId) const enabled = this.store.toggle(instanceId, sessionId) this.deps.eventBus.publish({ type: "yolo.stateChanged", instanceId, sessionId, enabled }) @@ -170,26 +310,33 @@ export class AutoAcceptManager { return enabled } - private async togglePersisted(instanceId: string, sessionId: string): Promise { - await this.hydrateInstance(instanceId) + private async togglePersisted(instanceId: string, sessionId: string, requestedEnabled?: boolean): Promise { const generation = this.instanceGeneration.get(instanceId) ?? 0 + const signal = this.generationSignal(instanceId) + let targetEnabled = requestedEnabled + await this.hydrateInstance(instanceId) const mutation = (this.mutations.get(instanceId) ?? Promise.resolve(false)).catch(() => false).then(async () => { - if ((this.instanceGeneration.get(instanceId) ?? 0) !== generation) { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) { + if (!this.stopped && this.streamIds.has(instanceId)) throw REBIND_TOGGLE_AFTER_ROTATION return this.store.isEnabled(instanceId, sessionId) } const rootSessionId = this.store.familyRoot(instanceId, sessionId) const traversedRootSessionIds = new Set([rootSessionId]) - const enabled = !this.store.isEnabled(instanceId, rootSessionId) + const enabled = targetEnabled ?? !this.store.isEnabled(instanceId, rootSessionId) + targetEnabled = enabled await this.deps.persistence!.persist( instanceId, rootSessionId, enabled, this.sessionWorkspaces.get(instanceId)?.get(rootSessionId), + signal, ) - if ((this.instanceGeneration.get(instanceId) ?? 0) !== generation) { - return this.store.isEnabled(instanceId, rootSessionId) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) { + return enabled } + this.recordPersistedState(instanceId, rootSessionId, enabled) let persistedRootSessionId = rootSessionId + const persistedDisabledRootSessionIds = enabled ? undefined : new Set([rootSessionId]) let currentRootSessionId = this.store.familyRoot(instanceId, sessionId) while (currentRootSessionId !== persistedRootSessionId) { traversedRootSessionIds.add(currentRootSessionId) @@ -198,21 +345,48 @@ export class AutoAcceptManager { currentRootSessionId, enabled, this.sessionWorkspaces.get(instanceId)?.get(currentRootSessionId), + signal, ) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return enabled + this.recordPersistedState(instanceId, currentRootSessionId, enabled) + persistedDisabledRootSessionIds?.add(currentRootSessionId) if (enabled) { await this.deps.persistence!.persist( instanceId, persistedRootSessionId, false, this.sessionWorkspaces.get(instanceId)?.get(persistedRootSessionId), + signal, ) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return enabled + this.recordPersistedState(instanceId, persistedRootSessionId, false) } persistedRootSessionId = currentRootSessionId currentRootSessionId = this.store.familyRoot(instanceId, sessionId) } - if ((this.instanceGeneration.get(instanceId) ?? 0) !== generation) { - return this.store.isEnabled(instanceId, currentRootSessionId) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) { + return enabled } + if (!enabled) { + const persistedEnabled = this.persistedEnabledSessions.get(instanceId) ?? new Set() + const migrations = this.reboundRootMigrations.get(instanceId) ?? new Set() + for (const durableRootSessionId of new Set([...persistedEnabled, ...migrations])) { + if ( + persistedDisabledRootSessionIds!.has(durableRootSessionId) || + this.store.familyRoot(instanceId, durableRootSessionId) !== currentRootSessionId + ) continue + await this.deps.persistence!.persist( + instanceId, + durableRootSessionId, + false, + this.sessionWorkspaces.get(instanceId)?.get(durableRootSessionId), + signal, + ) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return enabled + this.recordPersistedState(instanceId, durableRootSessionId, false) + } + } + this.reconcileReboundRootMigrationsAfterToggle(instanceId, currentRootSessionId, enabled) if (enabled) { this.store.setEnabled(instanceId, currentRootSessionId, true) } else { @@ -223,21 +397,70 @@ export class AutoAcceptManager { this.deps.eventBus.publish({ type: "yolo.stateChanged", instanceId, sessionId: currentRootSessionId, enabled }) if (enabled) this.drainPending(instanceId, currentRootSessionId) return enabled + }).catch((error) => { + if (error === REBIND_TOGGLE_AFTER_ROTATION) throw error + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) { + if (!this.stopped && this.streamIds.has(instanceId)) throw error + return this.store.isEnabled(instanceId, sessionId) + } + throw error }) const settled = mutation.finally(() => { if (this.mutations.get(instanceId) === settled) this.mutations.delete(instanceId) }) this.mutations.set(instanceId, settled) - return settled + try { + const enabled = await settled + if (!this.hasRuntimeAuthority(instanceId, generation, signal) && !this.stopped && this.streamIds.has(instanceId)) { + return this.rebindToggleAfterRotation(instanceId, sessionId, targetEnabled) + } + return enabled + } catch (error) { + if (error === REBIND_TOGGLE_AFTER_ROTATION) return this.togglePersisted(instanceId, sessionId, targetEnabled) + if (!this.hasRuntimeAuthority(instanceId, generation, signal) && !this.stopped && this.streamIds.has(instanceId)) { + return this.rebindToggleAfterRotation(instanceId, sessionId, targetEnabled) + } + throw error + } + } + + private async rebindToggleAfterRotation( + instanceId: string, + sessionId: string, + requestedEnabled: boolean | undefined, + ): Promise { + await this.hydrateInstance(instanceId) + if (requestedEnabled === undefined) return this.togglePersisted(instanceId, sessionId) + const rootSessionId = this.store.familyRoot(instanceId, sessionId) + if (this.store.isEnabled(instanceId, rootSessionId) !== requestedEnabled) { + return this.togglePersisted(instanceId, sessionId, requestedEnabled) + } + // Hydration publishes enabled roots; disabled roots need an explicit UI update. + if (!requestedEnabled) { + this.deps.eventBus.publish({ type: "yolo.stateChanged", instanceId, sessionId: rootSessionId, enabled: false }) + } + return requestedEnabled } clearInstance(instanceId: string): void { + this.invalidateRuntimeState(instanceId) + this.clearReboundRootMigrationRetry(instanceId) + this.reboundRootMigrations.delete(instanceId) + this.reboundRootMigrationRuns.delete(instanceId) + this.streamIds.delete(instanceId) + } + + private invalidateRuntimeState(instanceId: string): void { + this.generationControllers.get(instanceId)?.abort() + this.generationControllers.delete(instanceId) this.instanceGeneration.set(instanceId, (this.instanceGeneration.get(instanceId) ?? 0) + 1) + this.failedHydrations.delete(instanceId) this.hydratedInstances.delete(instanceId) this.hydration.delete(instanceId) this.queuedEvents.delete(instanceId) + this.overflowedEventQueues.delete(instanceId) this.sessionWorkspaces.delete(instanceId) - this.mutations.delete(instanceId) + this.persistedEnabledSessions.delete(instanceId) this.store.clearInstance(instanceId) this.pending.delete(instanceId) const prefix = `${instanceId}:` @@ -250,7 +473,7 @@ export class AutoAcceptManager { } handleInstanceEvent(instanceId: string, event: InstanceStreamPayload): void { - if (!event || typeof event.type !== "string") return + if (this.stopped || !event || typeof event.type !== "string") return if (SESSION_UPSERT_TYPES.has(event.type)) { this.ingestSession(instanceId, event.properties) @@ -306,31 +529,213 @@ export class AutoAcceptManager { const added = after.filter((id) => !before.includes(id)) if (removed.length === 0 && added.length === 0) return const generation = this.instanceGeneration.get(instanceId) ?? 0 + const signal = this.generationSignal(instanceId) const mutation = (this.mutations.get(instanceId) ?? Promise.resolve(false)).catch(() => false).then(async () => { - if ((this.instanceGeneration.get(instanceId) ?? 0) !== generation) return false + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false const enabledRoots = new Set(this.store.enabledRoots(instanceId)) + const migratedRoots = new Set() for (const rootSessionId of added) { - if (!enabledRoots.has(rootSessionId)) continue + const currentRootSessionId = this.store.familyRoot(instanceId, rootSessionId) + if (!enabledRoots.has(currentRootSessionId)) continue await this.deps.persistence!.persist( - instanceId, rootSessionId, true, this.sessionWorkspaces.get(instanceId)?.get(rootSessionId), + instanceId, + currentRootSessionId, + true, + this.sessionWorkspaces.get(instanceId)?.get(currentRootSessionId), + signal, ) + migratedRoots.add(currentRootSessionId) } for (const rootSessionId of removed) { - if (enabledRoots.has(rootSessionId)) continue - if ((this.instanceGeneration.get(instanceId) ?? 0) !== generation) return false + if (enabledRoots.has(rootSessionId) || migratedRoots.has(rootSessionId)) continue + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false await this.deps.persistence!.persist( - instanceId, rootSessionId, false, this.sessionWorkspaces.get(instanceId)?.get(rootSessionId), + instanceId, rootSessionId, false, this.sessionWorkspaces.get(instanceId)?.get(rootSessionId), signal, ) } + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + const persistedEnabled = this.persistedEnabledSessions.get(instanceId) + for (const rootSessionId of removed) persistedEnabled?.delete(rootSessionId) + for (const rootSessionId of migratedRoots) { + persistedEnabled?.add(rootSessionId) + if (this.store.isEnabled(instanceId, rootSessionId)) continue + this.store.setEnabled(instanceId, rootSessionId, true) + this.deps.eventBus.publish({ type: "yolo.stateChanged", instanceId, sessionId: rootSessionId, enabled: true }) + this.drainPending(instanceId, rootSessionId) + } return false }) const settled = mutation.finally(() => { if (this.mutations.get(instanceId) === settled) this.mutations.delete(instanceId) }) this.mutations.set(instanceId, settled) - void settled.catch((error) => { - this.deps.logger.warn({ instanceId, err: error }, "Failed to migrate persisted Yolo family root") - }) + void settled.then( + () => { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) { + this.rebindRootMigration(instanceId, before, after) + } + }, + (error) => { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) { + this.rebindRootMigration(instanceId, before, after) + return + } + this.deps.logger.warn({ instanceId, err: error }, "Failed to migrate persisted Yolo family root") + }, + ) + } + + private rebindRootMigration(instanceId: string, before: readonly string[], after: readonly string[]): void { + if (this.stopped || !this.streamIds.has(instanceId)) return + const removed = before.filter((id) => !after.includes(id)) + if (removed.length === 0) return + const migrations = this.reboundRootMigrations.get(instanceId) ?? new Set() + const priorSize = migrations.size + for (const rootSessionId of removed) migrations.add(rootSessionId) + this.reboundRootMigrations.set(instanceId, migrations) + if (migrations.size !== priorSize) this.reboundRootMigrationAttempts.delete(instanceId) + void this.retryReboundRootMigrations(instanceId) + } + + private async retryReboundRootMigrations(instanceId: string): Promise { + if ( + this.stopped || + !this.streamIds.has(instanceId) || + !this.reboundRootMigrations.has(instanceId) || + this.reboundRootMigrationRuns.has(instanceId) || + this.reboundRootMigrationTimers.has(instanceId) + ) return + this.reboundRootMigrationRuns.add(instanceId) + let retryAfterRotation = false + try { + try { + await this.hydrateInstance(instanceId) + } catch { + return + } + if (this.stopped || !this.hydratedInstances.has(instanceId) || !this.streamIds.has(instanceId)) return + const generation = this.instanceGeneration.get(instanceId) ?? 0 + const signal = this.generationSignal(instanceId) + const mutation = (this.mutations.get(instanceId) ?? Promise.resolve(false)).catch(() => false).then(async () => { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + const migrations = this.reboundRootMigrations.get(instanceId) + const persistedEnabled = this.persistedEnabledSessions.get(instanceId) + if (!migrations || !persistedEnabled) return false + for (const formerRootSessionId of Array.from(migrations)) { + if (!persistedEnabled.has(formerRootSessionId)) { + migrations.delete(formerRootSessionId) + continue + } + let persistedRootSessionId = formerRootSessionId + let currentRootSessionId = this.store.familyRoot(instanceId, formerRootSessionId) + while (currentRootSessionId !== persistedRootSessionId) { + await this.deps.persistence!.persist( + instanceId, + currentRootSessionId, + true, + this.sessionWorkspaces.get(instanceId)?.get(currentRootSessionId), + signal, + ) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + this.recordPersistedState(instanceId, currentRootSessionId, true) + await this.deps.persistence!.persist( + instanceId, + persistedRootSessionId, + false, + this.sessionWorkspaces.get(instanceId)?.get(persistedRootSessionId), + signal, + ) + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return false + this.recordPersistedState(instanceId, persistedRootSessionId, false) + persistedRootSessionId = currentRootSessionId + currentRootSessionId = this.store.familyRoot(instanceId, formerRootSessionId) + } + migrations.delete(formerRootSessionId) + if (this.store.isEnabled(instanceId, currentRootSessionId)) continue + this.store.setEnabled(instanceId, currentRootSessionId, true) + this.deps.eventBus.publish({ + type: "yolo.stateChanged", + instanceId, + sessionId: currentRootSessionId, + enabled: true, + }) + this.drainPending(instanceId, currentRootSessionId) + } + if (migrations.size === 0) this.reboundRootMigrations.delete(instanceId) + return false + }) + const settled = mutation.finally(() => { + if (this.mutations.get(instanceId) === settled) this.mutations.delete(instanceId) + }) + this.mutations.set(instanceId, settled) + try { + await settled + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) retryAfterRotation = true + else this.reboundRootMigrationAttempts.delete(instanceId) + } catch (error) { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) { + retryAfterRotation = true + return + } + const attempt = (this.reboundRootMigrationAttempts.get(instanceId) ?? 0) + 1 + if (attempt >= AutoAcceptManager.MAX_REBOUND_MIGRATION_ATTEMPTS) { + this.reboundRootMigrationAttempts.delete(instanceId) + this.deps.logger.warn( + { instanceId, err: error, attempt }, + "Failed to migrate persisted Yolo family root after retries; retaining recovery intent", + ) + return + } + this.reboundRootMigrationAttempts.set(instanceId, attempt) + this.deps.logger.warn({ instanceId, err: error, attempt }, "Failed to migrate persisted Yolo family root; retrying") + this.scheduleReboundRootMigrationRetry(instanceId, attempt) + } + } finally { + this.reboundRootMigrationRuns.delete(instanceId) + if (retryAfterRotation) void this.retryReboundRootMigrations(instanceId) + } + } + + private scheduleReboundRootMigrationRetry(instanceId: string, attempt: number): void { + if (this.reboundRootMigrationTimers.has(instanceId)) return + const timer = setTimeout(() => { + if (this.reboundRootMigrationTimers.get(instanceId) !== timer) return + this.reboundRootMigrationTimers.delete(instanceId) + void this.retryReboundRootMigrations(instanceId) + }, AutoAcceptManager.REBOUND_MIGRATION_RETRY_MS * attempt) + this.reboundRootMigrationTimers.set(instanceId, timer) + } + + private clearReboundRootMigrationRetry(instanceId: string): void { + const timer = this.reboundRootMigrationTimers.get(instanceId) + if (timer) clearTimeout(timer) + this.reboundRootMigrationTimers.delete(instanceId) + this.reboundRootMigrationAttempts.delete(instanceId) + } + + private recordPersistedState(instanceId: string, rootSessionId: string, enabled: boolean): void { + const persistedEnabled = this.persistedEnabledSessions.get(instanceId) + if (enabled) persistedEnabled?.add(rootSessionId) + else persistedEnabled?.delete(rootSessionId) + } + + private reconcileReboundRootMigrationsAfterToggle( + instanceId: string, + rootSessionId: string, + enabled: boolean, + ): void { + const migrations = this.reboundRootMigrations.get(instanceId) + if (!migrations) return + const persistedEnabled = this.persistedEnabledSessions.get(instanceId) + for (const formerRootSessionId of Array.from(migrations)) { + if (this.store.familyRoot(instanceId, formerRootSessionId) !== rootSessionId) continue + if (!enabled || !persistedEnabled?.has(formerRootSessionId) || formerRootSessionId === rootSessionId) { + migrations.delete(formerRootSessionId) + } + } + if (migrations.size > 0) return + this.reboundRootMigrations.delete(instanceId) + this.clearReboundRootMigrationRetry(instanceId) } private handlePermissionRequest(instanceId: string, eventType: string, permission: unknown): void { @@ -379,25 +784,29 @@ export class AutoAcceptManager { if (this.inFlight.has(key)) return const attempts = this.replyAttempts.get(key) ?? 0 if (attempts >= AutoAcceptManager.MAX_REPLY_ATTEMPTS) return + const generation = this.instanceGeneration.get(instanceId) ?? 0 + const signal = this.generationSignal(instanceId) this.inFlight.add(key) this.replyAttempts.set(key, attempts + 1) const reply: AutoAcceptReply = { instanceId, permissionId, sessionId, source, reply: "once" } - void this.deps.replier(reply) + void this.deps.replier(reply, signal) .then(() => { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return this.replyAttempts.delete(key) this.removePending(instanceId, permissionId) this.deps.eventBus.publish({ type: "yolo.autoAccepted", instanceId, sessionId, permissionId }) }) .catch((error) => { + if (!this.hasRuntimeAuthority(instanceId, generation, signal)) return this.deps.logger.error({ instanceId, permissionId, err: error, attempt: attempts + 1 }, "Yolo auto-accept reply failed") if (attempts + 1 >= AutoAcceptManager.MAX_REPLY_ATTEMPTS) { this.removePending(instanceId, permissionId) } }) .finally(() => { - this.inFlight.delete(key) + if (this.hasRuntimeAuthority(instanceId, generation, signal)) this.inFlight.delete(key) }) } @@ -419,7 +828,12 @@ export class AutoAcceptManager { instancePending = new Map() this.pending.set(instanceId, instancePending) } + if (instancePending.has(entry.permissionId)) instancePending.delete(entry.permissionId) instancePending.set(entry.permissionId, entry) + if (instancePending.size > AutoAcceptManager.MAX_PENDING_PER_INSTANCE) { + const oldestPermissionId = instancePending.keys().next().value + if (oldestPermissionId) this.removePending(instanceId, oldestPermissionId) + } } private removePending(instanceId: string, permissionId: string): void { @@ -441,6 +855,55 @@ export class AutoAcceptManager { } if (instancePending.size === 0) this.pending.delete(instanceId) } + + private generationSignal(instanceId: string): AbortSignal { + if (this.stopped) { + const controller = new AbortController() + controller.abort() + return controller.signal + } + let controller = this.generationControllers.get(instanceId) + if (!controller || controller.signal.aborted) { + controller = new AbortController() + this.generationControllers.set(instanceId, controller) + } + return controller.signal + } + + private releaseFailedHydration(instanceId: string, signal: AbortSignal): void { + const controller = this.generationControllers.get(instanceId) + if (controller?.signal !== signal) return + controller.abort() + this.generationControllers.delete(instanceId) + } + + private queueEvent(instanceId: string, event: InstanceStreamPayload): void { + const key = queuedEventKey(event) + if (!key) return + const queued = this.queuedEvents.get(instanceId) ?? new Map() + const existing = queued.get(key) + let next = event + if (existing && PERMISSION_ASK_TYPES.has(existing.type ?? "") && event.type === "permission.updated") { + next = { ...event, type: existing.type, properties: { ...existing.properties, ...event.properties } } + } + queued.delete(key) + queued.set(key, next) + if (queued.size > AutoAcceptManager.MAX_QUEUED_EVENTS_PER_INSTANCE) { + this.overflowedEventQueues.add(instanceId) + for (const [queuedKey, queuedEvent] of queued) { + if (!isSessionEvent(queuedEvent)) queued.delete(queuedKey) + } + while (queued.size > AutoAcceptManager.MAX_QUEUED_EVENTS_PER_INSTANCE) { + const oldestKey = queued.keys().next().value + if (oldestKey) queued.delete(oldestKey) + } + } + this.queuedEvents.set(instanceId, queued) + } + + private hasRuntimeAuthority(instanceId: string, generation: number, signal: AbortSignal): boolean { + return !this.stopped && !signal.aborted && (this.instanceGeneration.get(instanceId) ?? 0) === generation + } } interface InstanceStreamPayload { @@ -473,3 +936,26 @@ interface PermissionRepliedProperties { function readString(value: unknown): string | undefined { return typeof value === "string" && value.length > 0 ? value : undefined } + +function queuedEventKey(event: InstanceStreamPayload): string | undefined { + if (SESSION_UPSERT_TYPES.has(event.type ?? "") || SESSION_REMOVE_TYPES.has(event.type ?? "")) { + const info = (event.properties as { info?: SessionProperties } | undefined)?.info + const sessionId = readString(info?.id) ?? readString(event.properties?.id) + return sessionId ? `session:${sessionId}` : undefined + } + if (PERMISSION_ASK_TYPES.has(event.type ?? "") || PERMISSION_REPLIED_TYPES.has(event.type ?? "")) { + const properties = event.properties as PermissionRepliedProperties | undefined + const permissionId = + readString(properties?.id) ?? + readString(properties?.requestID) ?? + readString(properties?.permissionID) ?? + readString(properties?.requestId) ?? + readString(properties?.permissionId) + return permissionId ? `permission:${permissionId}` : undefined + } + return undefined +} + +function isSessionEvent(event: InstanceStreamPayload): boolean { + return SESSION_UPSERT_TYPES.has(event.type ?? "") || SESSION_REMOVE_TYPES.has(event.type ?? "") +} diff --git a/packages/server/src/permissions/opencode-replier.test.ts b/packages/server/src/permissions/opencode-replier.test.ts new file mode 100644 index 000000000..6708240cf --- /dev/null +++ b/packages/server/src/permissions/opencode-replier.test.ts @@ -0,0 +1,49 @@ +import assert from "node:assert/strict" +import test from "node:test" + +import type { Logger } from "../logger" +import type { WorkspaceManager } from "../workspaces/manager" +import { createOpencodePermissionReplier } from "./opencode-replier" + +test("aborts a permission reply when its runtime generation is revoked", async () => { + const original = globalThis.fetch + let deadline: ReturnType | undefined + let requestSignal: AbortSignal | undefined + globalThis.fetch = (async (_input: unknown, init?: RequestInit) => { + requestSignal = init?.signal ?? undefined + return new Promise((_resolve, reject) => { + requestSignal?.addEventListener("abort", () => reject(requestSignal?.reason), { once: true }) + }) + }) as typeof fetch + + try { + const workspaceManager = { + getInstancePort: () => 4321, + getInstanceAuthorizationHeader: () => undefined, + get: () => ({ path: "/repo" }), + } as unknown as WorkspaceManager + const replier = createOpencodePermissionReplier({ workspaceManager, logger: {} as Logger }) + const controller = new AbortController() + const reply = replier({ + instanceId: "instance", + permissionId: "permission", + sessionId: "session", + source: "v2", + reply: "once", + }, controller.signal) + + await new Promise((resolve) => setImmediate(resolve)) + controller.abort() + + await assert.rejects(Promise.race([ + reply, + new Promise((_, reject) => { + deadline = setTimeout(() => reject(new Error("Permission reply cancellation test timed out")), 1_000) + }), + ]), (error: Error) => error.name === "AbortError") + assert.equal(requestSignal?.aborted, true) + } finally { + if (deadline) clearTimeout(deadline) + globalThis.fetch = original + } +}) diff --git a/packages/server/src/permissions/opencode-replier.ts b/packages/server/src/permissions/opencode-replier.ts index 2ed9bfc9e..c893a603e 100644 --- a/packages/server/src/permissions/opencode-replier.ts +++ b/packages/server/src/permissions/opencode-replier.ts @@ -17,13 +17,14 @@ interface OpencodeReplierDeps { * for the installed SDK version — no hand-assembled URLs. */ export function createOpencodePermissionReplier(deps: OpencodeReplierDeps): PermissionReplier { - return async (reply: AutoAcceptReply) => { + return async (reply: AutoAcceptReply, signal: AbortSignal) => { + if (signal.aborted) throw signal.reason ?? new DOMException("Operation aborted", "AbortError") const client = createInstanceClient(deps.workspaceManager, reply.instanceId) if (!client) { throw new Error(`Yolo: instance ${reply.instanceId} has no open port`) } - const opts = { throwOnError: true } as const + const opts = { throwOnError: true, signal } as const if (reply.source === "v2") { await client.v2.session.permission.reply( diff --git a/packages/server/src/permissions/opencode-yolo-metadata.test.ts b/packages/server/src/permissions/opencode-yolo-metadata.test.ts index d31d56523..04af16b61 100644 --- a/packages/server/src/permissions/opencode-yolo-metadata.test.ts +++ b/packages/server/src/permissions/opencode-yolo-metadata.test.ts @@ -30,8 +30,9 @@ describe("OpenCode Yolo metadata", () => { }, } const persistence = createOpencodeYoloPersistence({} as never, () => client as never) - const [session] = await persistence.loadSessions("instance") - await persistence.persist("instance", "root", true, session?.workspaceId) + const signal = new AbortController().signal + const [session] = await persistence.loadSessions("instance", signal) + await persistence.persist("instance", "root", true, session?.workspaceId, signal) assert.equal(session?.workspaceId, "workspace") assert.equal(calls[0]?.workspace, "workspace") assert.equal(calls[1]?.workspace, "workspace") @@ -49,8 +50,9 @@ describe("OpenCode Yolo metadata", () => { }, } const persistence = createOpencodeYoloPersistence({} as never, () => client as never) + const signal = new AbortController().signal await Promise.all([ - persistence.persist("instance-a", "root", true), + persistence.persist("instance-a", "root", true, undefined, signal), persistence.setWorktreeSlug("instance-b", "root", "feature"), ]) assert.deepEqual(metadata, { @@ -58,4 +60,36 @@ describe("OpenCode Yolo metadata", () => { codenomad: { version: 1, yolo: { enabled: true, rootSessionId: "root" }, worktreeSlug: "feature" }, }) }) + + it("does not start a queued metadata write after its generation is aborted", async () => { + let release!: () => void + const gate = new Promise((resolve) => { release = resolve }) + let gets = 0 + let updates = 0 + const client = { + session: { + async get() { + gets += 1 + if (gets === 1) await gate + return { data: { metadata: {} } } + }, + async update() { + updates += 1 + return { data: {} } + }, + }, + } + const persistence = createOpencodeYoloPersistence({} as never, () => client as never) + const blocker = persistence.setWorktreeSlug("instance", "root", "feature") + await new Promise((resolve) => setImmediate(resolve)) + const controller = new AbortController() + const stale = persistence.persist("instance", "root", true, undefined, controller.signal) + controller.abort() + release() + + await blocker + await assert.rejects(stale, (error: Error) => error.name === "AbortError") + assert.equal(gets, 1) + assert.equal(updates, 1) + }) }) diff --git a/packages/server/src/permissions/opencode-yolo-metadata.ts b/packages/server/src/permissions/opencode-yolo-metadata.ts index bca929015..308116ada 100644 --- a/packages/server/src/permissions/opencode-yolo-metadata.ts +++ b/packages/server/src/permissions/opencode-yolo-metadata.ts @@ -62,14 +62,18 @@ export function createOpencodeYoloPersistence( sessionId: string, workspaceId: string | undefined, update: (metadata: unknown) => Metadata, + signal?: AbortSignal, ): Promise => { const writeKey = sessionId const write = (writes.get(writeKey) ?? Promise.resolve()).catch(() => undefined).then(async () => { + throwIfAborted(signal) const client = clientFor(instanceId) const scope = { sessionID: sessionId, ...(workspaceId ? { workspace: workspaceId } : {}) } - const { data: session } = await client.session.get(scope, { throwOnError: true }) + const options = { throwOnError: true, ...(signal ? { signal } : {}) } as const + const { data: session } = await client.session.get(scope, options) + throwIfAborted(signal) const metadata = update(session.metadata) - const { data } = await client.session.update({ ...scope, metadata }, { throwOnError: true }) + const { data } = await client.session.update({ ...scope, metadata }, options) return record(data?.metadata ?? metadata) }) const settled = write.finally(() => { @@ -79,10 +83,11 @@ export function createOpencodeYoloPersistence( return settled } return { - async loadSessions(instanceId): Promise { + async loadSessions(instanceId, signal): Promise { + throwIfAborted(signal) const { data } = await clientFor(instanceId).session.list( { scope: "project", limit: SESSION_LIST_LIMIT }, - { throwOnError: true }, + { throwOnError: true, signal }, ) return (data ?? []).map((session) => ({ id: session.id, @@ -92,9 +97,9 @@ export function createOpencodeYoloPersistence( yoloEnabled: hasPersistedYolo(session.id, session.metadata), })) }, - persist(instanceId, rootSessionId, enabled, workspaceId): Promise { + persist(instanceId, rootSessionId, enabled, workspaceId, signal): Promise { return updateMetadata(instanceId, rootSessionId, workspaceId, - (metadata) => mergePersistedYolo(metadata, rootSessionId, enabled)).then(() => undefined) + (metadata) => mergePersistedYolo(metadata, rootSessionId, enabled), signal).then(() => undefined) }, async hasProjectSession(instanceId, sessionId): Promise { const { data } = await clientFor(instanceId).session.list( @@ -109,3 +114,7 @@ export function createOpencodeYoloPersistence( }, } } + +function throwIfAborted(signal?: AbortSignal): void { + if (signal?.aborted) throw signal.reason ?? new DOMException("Operation aborted", "AbortError") +} diff --git a/packages/server/src/workspaces/instance-client.test.ts b/packages/server/src/workspaces/instance-client.test.ts index 636bc9f1a..4799459ef 100644 --- a/packages/server/src/workspaces/instance-client.test.ts +++ b/packages/server/src/workspaces/instance-client.test.ts @@ -172,4 +172,142 @@ describe("createInstanceClient", () => { globalThis.fetch = original } }) + + it("preserves cancellation from the SDK Request signal", async () => { + const original = globalThis.fetch + let effectiveSignal: AbortSignal | null | undefined + globalThis.fetch = (async (_input: any, init: any) => { + effectiveSignal = (init as RequestInit | undefined)?.signal + return new Promise((_resolve, reject) => { + effectiveSignal?.addEventListener("abort", () => reject(effectiveSignal?.reason), { once: true }) + }) + }) as typeof fetch + try { + const manager = makeManager({ getInstancePort: () => 4321, get: () => ({ path: "/repo" }) }) + const client = createInstanceClient(manager as unknown as WorkspaceManager, "ws-1", { timeoutMs: 1_000 }) + const controller = new AbortController() + const cancellation = new Error("caller cancelled") + const request = client!.global.health({ signal: controller.signal }) + await new Promise((resolve) => setImmediate(resolve)) + controller.abort(cancellation) + + const result = await Promise.race([ + request, + new Promise((_resolve, reject) => setTimeout(() => reject(new Error("cancellation was not preserved")), 200)), + ]) + assert.ok(result.error) + assert.equal(effectiveSignal?.reason, cancellation) + } finally { + globalThis.fetch = original + } + }) + + it("keeps the timeout active while the response body is stalled", async () => { + const original = globalThis.fetch + let effectiveSignal: AbortSignal | null | undefined + globalThis.fetch = (async (_input: any, init: any) => { + effectiveSignal = (init as RequestInit | undefined)?.signal + return new Response(new ReadableStream({ + start(controller) { + effectiveSignal?.addEventListener("abort", () => controller.error(effectiveSignal?.reason), { once: true }) + }, + }), { + status: 200, + headers: { "content-type": "application/json" }, + }) + }) as typeof fetch + try { + const manager = makeManager({ getInstancePort: () => 4321, get: () => ({ path: "/repo" }) }) + const client = createInstanceClient(manager as unknown as WorkspaceManager, "ws-1", { timeoutMs: 10 }) + + await assert.rejects(Promise.race([ + client!.global.health(), + new Promise((_resolve, reject) => setTimeout(() => reject(new Error("response body outlived timeout")), 200)), + ]), (error: Error) => error.name === "TimeoutError") + assert.equal(effectiveSignal?.aborted, true) + } finally { + globalThis.fetch = original + } + }) + + it("cleans timeout and caller cancellation after the response body completes", async () => { + const original = globalThis.fetch + let effectiveSignal: AbortSignal | null | undefined + globalThis.fetch = (async (_input: any, init: any) => { + effectiveSignal = (init as RequestInit | undefined)?.signal + return new Response(JSON.stringify({ healthy: true }), { + status: 200, + headers: { "content-type": "application/json" }, + }) + }) as typeof fetch + try { + const manager = makeManager({ getInstancePort: () => 4321, get: () => ({ path: "/repo" }) }) + const client = createInstanceClient(manager as unknown as WorkspaceManager, "ws-1", { timeoutMs: 10 }) + const controller = new AbortController() + + await client!.global.health({ signal: controller.signal }) + await new Promise((resolve) => setTimeout(resolve, 20)) + controller.abort(new Error("late cancellation")) + + assert.equal(effectiveSignal?.aborted, false) + } finally { + globalThis.fetch = original + } + }) + + it("cleans timeout when a zero-length response body is not consumed", async () => { + const original = globalThis.fetch + let effectiveSignal: AbortSignal | null | undefined + globalThis.fetch = (async (_input: any, init: any) => { + effectiveSignal = (init as RequestInit | undefined)?.signal + return new Response(new ReadableStream({ start: (controller) => controller.close() }), { + status: 200, + headers: { "content-length": "0" }, + }) + }) as typeof fetch + try { + const manager = makeManager({ getInstancePort: () => 4321, get: () => ({ path: "/repo" }) }) + const client = createInstanceClient(manager as unknown as WorkspaceManager, "ws-1", { timeoutMs: 10 }) + const controller = new AbortController() + + await client!.global.health({ signal: controller.signal }) + await new Promise((resolve) => setTimeout(resolve, 20)) + controller.abort(new Error("late cancellation")) + + assert.equal(effectiveSignal?.aborted, false) + } finally { + globalThis.fetch = original + } + }) + + it("cleans an unread response when the SDK rejects it after headers", async () => { + const original = globalThis.fetch + let effectiveSignal: AbortSignal | null | undefined + let bodyCancelled = false + globalThis.fetch = (async (_input: any, init: any) => { + effectiveSignal = (init as RequestInit | undefined)?.signal + return new Response(new ReadableStream({ + cancel() { + bodyCancelled = true + }, + }), { + status: 200, + headers: { "content-type": "text/html" }, + }) + }) as typeof fetch + try { + const manager = makeManager({ getInstancePort: () => 4321, get: () => ({ path: "/repo" }) }) + const client = createInstanceClient(manager as unknown as WorkspaceManager, "ws-1", { timeoutMs: 10 }) + const controller = new AbortController() + + await assert.rejects(client!.global.health({ signal: controller.signal, throwOnError: true })) + await new Promise((resolve) => setTimeout(resolve, 20)) + controller.abort(new Error("late cancellation")) + + assert.equal(bodyCancelled, true) + assert.equal(effectiveSignal?.aborted, false) + } finally { + globalThis.fetch = original + } + }) }) diff --git a/packages/server/src/workspaces/instance-client.ts b/packages/server/src/workspaces/instance-client.ts index 5ee323f04..90bdabd56 100644 --- a/packages/server/src/workspaces/instance-client.ts +++ b/packages/server/src/workspaces/instance-client.ts @@ -53,11 +53,107 @@ export function createInstanceClient( return createOpencodeClient({ baseUrl: `http://${LOOPBACK_HOST}:${port}/`, headers, - fetch: (url, init) => - fetch(url, { - ...(init as RequestInit), - signal: (init as RequestInit)?.signal ?? AbortSignal.timeout(timeoutMs), - }), + fetch: (url, init) => { + const requestInit = init as RequestInit | undefined + const sources = [ + ...(url instanceof Request ? [url.signal] : []), + ...(requestInit?.signal ? [requestInit.signal] : []), + ] + const { signal, cleanup } = composeRequestSignal(sources, timeoutMs) + return fetch(url, { ...requestInit, signal }).then( + (response) => responseWithSignalCleanup(response, cleanup), + (error) => { + cleanup() + throw error + }, + ) + }, ...(directory ? { directory } : {}), }) } + +function composeRequestSignal( + sources: readonly AbortSignal[], + timeoutMs: number, +): { signal: AbortSignal; cleanup: () => void } { + const controller = new AbortController() + const listeners = new Map void>() + let timeout: ReturnType | undefined + let cleaned = false + + const cleanup = () => { + if (cleaned) return + cleaned = true + if (timeout) clearTimeout(timeout) + for (const [source, listener] of listeners) source.removeEventListener("abort", listener) + listeners.clear() + } + const abort = (reason: unknown) => { + if (!controller.signal.aborted) controller.abort(reason) + cleanup() + } + + for (const source of new Set(sources)) { + if (source.aborted) { + abort(source.reason) + break + } + const listener = () => abort(source.reason) + listeners.set(source, listener) + source.addEventListener("abort", listener, { once: true }) + } + + if (!controller.signal.aborted) { + timeout = setTimeout( + () => abort(new DOMException("Request timed out", "TimeoutError")), + timeoutMs, + ) + } + + return { signal: controller.signal, cleanup } +} + +function responseWithSignalCleanup(response: Response, cleanup: () => void): Response | Promise { + if (!response.body || response.headers.get("content-length") === "0") { + cleanup() + return response + } + + // The SDK rejects this response in an interceptor before reading its body. + if (response.headers.get("content-type") === "text/html") { + return response.body.cancel().catch(() => undefined).then(() => response).finally(cleanup) + } + + const reader = response.body.getReader() + const body = new ReadableStream({ + async pull(controller) { + try { + const { done, value } = await reader.read() + if (done) { + cleanup() + controller.close() + } else { + controller.enqueue(value) + } + } catch (error) { + cleanup() + controller.error(error) + } + }, + async cancel(reason) { + cleanup() + await reader.cancel(reason) + }, + }) + const wrapped = new Response(body, { + status: response.status, + statusText: response.statusText, + headers: response.headers, + }) + Object.defineProperties(wrapped, { + redirected: { value: response.redirected }, + type: { value: response.type }, + url: { value: response.url }, + }) + return wrapped +} diff --git a/packages/server/src/workspaces/instance-events.test.ts b/packages/server/src/workspaces/instance-events.test.ts new file mode 100644 index 000000000..099872a8a --- /dev/null +++ b/packages/server/src/workspaces/instance-events.test.ts @@ -0,0 +1,66 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { EventBus } from "../events/bus" +import { InstanceEventBridge } from "./instance-events" + +test("instance event bridge parses CRLF-delimited SSE frames with a stream id", () => { + const eventBus = new EventBus() + const published: any[] = [] + eventBus.on("instance.event", (event) => published.push(event)) + const bridge = new InstanceEventBridge({ + eventBus, + workspaceManager: {} as any, + logger: { debug() {}, trace() {}, warn() {}, isLevelEnabled: () => false } as any, + }) + + const remaining = (bridge as any).flushEvents( + 'data: {"type":"session.idle","properties":{"sessionID":"s"}}\r\n\r\n', + "instance", + "stream", + ) + + assert.equal(remaining, "") + assert.equal(published[0]?.streamId, "stream") + bridge.shutdown() +}) + +test("instance event bridge parses CR-delimited SSE frames", () => { + const eventBus = new EventBus() + const published: any[] = [] + eventBus.on("instance.event", (event) => published.push(event)) + const bridge = new InstanceEventBridge({ + eventBus, + workspaceManager: {} as any, + logger: { debug() {}, trace() {}, warn() {}, isLevelEnabled: () => false } as any, + }) + + const remaining = (bridge as any).flushEvents( + 'data: {"type":"session.idle","properties":{"sessionID":"s"}}\r\r', + "instance", + "stream", + ) + + assert.equal(remaining, "") + assert.equal(published.length, 1) + bridge.shutdown() +}) + +test("instance event bridge rotates stream authority when the workspace pid changes", () => { + const eventBus = new EventBus() + const bridge = new InstanceEventBridge({ + eventBus, + workspaceManager: { getInstancePort: () => undefined } as any, + logger: { debug() {}, trace() {}, warn() {}, isLevelEnabled: () => false } as any, + }) + const workspace = { id: "instance", pid: 1 } + eventBus.publish({ type: "workspace.started", workspace } as any) + const first = (bridge as any).streams.get(workspace.id) + eventBus.publish({ type: "workspace.started", workspace } as any) + assert.equal((bridge as any).streams.get(workspace.id).streamId, first.streamId) + + eventBus.publish({ type: "workspace.started", workspace: { ...workspace, pid: 2 } } as any) + const second = (bridge as any).streams.get(workspace.id) + assert.notEqual(second.streamId, first.streamId) + assert.equal(first.controller.signal.aborted, true) + bridge.shutdown() +}) diff --git a/packages/server/src/workspaces/instance-events.ts b/packages/server/src/workspaces/instance-events.ts index 5be037007..605e17458 100644 --- a/packages/server/src/workspaces/instance-events.ts +++ b/packages/server/src/workspaces/instance-events.ts @@ -1,5 +1,5 @@ -import { Agent, fetch } from "undici" -import { Agent as UndiciAgent } from "undici" +import { Agent as UndiciAgent, fetch } from "undici" +import { randomUUID } from "node:crypto" import { EventBus } from "../events/bus" import { Logger } from "../logger" import { WorkspaceManager } from "./manager" @@ -8,6 +8,7 @@ import { InstanceStreamEvent, InstanceStreamStatus } from "../api-types" const STREAM_AGENT = new UndiciAgent({ bodyTimeout: 0, headersTimeout: 0 }) const RECONNECT_DELAY_MS = 1000 +const MAX_EVENT_BUFFER_CHARACTERS = 16 * 1024 * 1024 interface InstanceEventBridgeOptions { workspaceManager: WorkspaceManager @@ -17,6 +18,8 @@ interface InstanceEventBridgeOptions { interface ActiveStream { controller: AbortController + streamId: string + runtimePid?: number task: Promise } @@ -25,7 +28,7 @@ export class InstanceEventBridge { constructor(private readonly options: InstanceEventBridgeOptions) { const bus = this.options.eventBus - bus.on("workspace.started", (event) => this.startStream(event.workspace.id)) + bus.on("workspace.started", (event) => this.startStream(event.workspace.id, event.workspace.pid)) bus.on("workspace.stopped", (event) => this.stopStream(event.workspaceId, "workspace stopped")) bus.on("workspace.error", (event) => this.stopStream(event.workspace.id, "workspace error")) } @@ -33,22 +36,25 @@ export class InstanceEventBridge { shutdown() { for (const [id, active] of this.streams) { active.controller.abort() - this.publishStatus(id, "disconnected") + this.publishStatus(id, active.streamId, "disconnected") } this.streams.clear() } - private startStream(workspaceId: string) { - if (this.streams.has(workspaceId)) { - return + private startStream(workspaceId: string, runtimePid?: number) { + const existing = this.streams.get(workspaceId) + if (existing) { + if (existing.runtimePid === runtimePid) return + this.stopStream(workspaceId, "workspace restarted") } const controller = new AbortController() - const task = this.runStream(workspaceId, controller.signal) + const streamId = randomUUID() + const task = this.runStream(workspaceId, streamId, controller.signal) .catch((error) => { if (!controller.signal.aborted) { this.options.logger.warn({ workspaceId, err: error }, "Instance event stream failed") - this.publishStatus(workspaceId, "error", error instanceof Error ? error.message : String(error)) + this.publishStatus(workspaceId, streamId, "error", error instanceof Error ? error.message : String(error)) } }) .finally(() => { @@ -58,7 +64,7 @@ export class InstanceEventBridge { } }) - this.streams.set(workspaceId, { controller, task }) + this.streams.set(workspaceId, { controller, streamId, runtimePid, task }) } private stopStream(workspaceId: string, reason?: string) { @@ -68,10 +74,10 @@ export class InstanceEventBridge { } active.controller.abort() this.streams.delete(workspaceId) - this.publishStatus(workspaceId, "disconnected", reason) + this.publishStatus(workspaceId, active.streamId, "disconnected", reason) } - private async runStream(workspaceId: string, signal: AbortSignal) { + private async runStream(workspaceId: string, streamId: string, signal: AbortSignal) { while (!signal.aborted) { const port = this.options.workspaceManager.getInstancePort(workspaceId) if (!port) { @@ -79,22 +85,23 @@ export class InstanceEventBridge { continue } - this.publishStatus(workspaceId, "connecting") + this.publishStatus(workspaceId, streamId, "connecting") try { - await this.consumeStream(workspaceId, port, signal) + await this.consumeStream(workspaceId, streamId, port, signal) + if (!signal.aborted) await this.delay(RECONNECT_DELAY_MS, signal) } catch (error) { if (signal.aborted) { break } this.options.logger.warn({ workspaceId, err: error }, "Instance event stream disconnected") - this.publishStatus(workspaceId, "error", error instanceof Error ? error.message : String(error)) + this.publishStatus(workspaceId, streamId, "error", error instanceof Error ? error.message : String(error)) await this.delay(RECONNECT_DELAY_MS, signal) } } } - private async consumeStream(workspaceId: string, port: number, signal: AbortSignal) { + private async consumeStream(workspaceId: string, streamId: string, port: number, signal: AbortSignal) { const url = `http://${LOOPBACK_HOST}:${port}/global/event` const headers: Record = { Accept: "text/event-stream" } @@ -110,40 +117,51 @@ export class InstanceEventBridge { }) if (!response.ok || !response.body) { + await response.body?.cancel().catch(() => undefined) throw new Error(`Instance event stream unavailable (${response.status})`) } - this.publishStatus(workspaceId, "connected") + this.publishStatus(workspaceId, streamId, "connected") const reader = response.body.getReader() const decoder = new TextDecoder() let buffer = "" - while (!signal.aborted) { - const { done, value } = await reader.read() - if (done || !value) { - break + try { + while (!signal.aborted) { + const { done, value } = await reader.read() + if (done || !value) { + break + } + buffer += decoder.decode(value, { stream: true }) + buffer = this.flushEvents(buffer, workspaceId, streamId) + if (buffer.length > MAX_EVENT_BUFFER_CHARACTERS) { + throw new Error("Instance event exceeded the stream buffer limit") + } } - buffer += decoder.decode(value, { stream: true }) - buffer = this.flushEvents(buffer, workspaceId) + } finally { + await reader.cancel().catch(() => undefined) + reader.releaseLock() } } - private flushEvents(buffer: string, workspaceId: string) { - let separatorIndex = buffer.indexOf("\n\n") + private flushEvents(buffer: string, workspaceId: string, streamId: string) { + let separator = /\r\n\r\n|\r\r|\n\n/.exec(buffer) - while (separatorIndex >= 0) { + while (separator) { + const separatorIndex = separator.index const chunk = buffer.slice(0, separatorIndex) - buffer = buffer.slice(separatorIndex + 2) - this.processChunk(chunk, workspaceId) - separatorIndex = buffer.indexOf("\n\n") + buffer = buffer.slice(separatorIndex + separator[0].length) + if (chunk.length > MAX_EVENT_BUFFER_CHARACTERS) throw new Error("Instance event exceeded the stream buffer limit") + this.processChunk(chunk, workspaceId, streamId) + separator = /\r\n\r\n|\r\r|\n\n/.exec(buffer) } return buffer } - private processChunk(chunk: string, workspaceId: string) { - const lines = chunk.split(/\r?\n/) + private processChunk(chunk: string, workspaceId: string, streamId: string) { + const lines = chunk.split(/\r\n|\r|\n/) const dataLines: string[] = [] for (const line of lines) { @@ -194,15 +212,15 @@ export class InstanceEventBridge { if (this.options.logger.isLevelEnabled("trace")) { this.options.logger.trace({ workspaceId, event }, "Instance SSE event payload") } - this.options.eventBus.publish({ type: "instance.event", instanceId: workspaceId, event }) + this.options.eventBus.publish({ type: "instance.event", instanceId: workspaceId, streamId, event }) } catch (error) { this.options.logger.warn({ workspaceId, chunk: payload, err: error }, "Failed to parse instance SSE payload") } } - private publishStatus(instanceId: string, status: InstanceStreamStatus, reason?: string) { - this.options.logger.debug({ instanceId, status, reason }, "Instance SSE status updated") - this.options.eventBus.publish({ type: "instance.eventStatus", instanceId, status, reason }) + private publishStatus(instanceId: string, streamId: string, status: InstanceStreamStatus, reason?: string) { + this.options.logger.debug({ instanceId, streamId, status, reason }, "Instance SSE status updated") + this.options.eventBus.publish({ type: "instance.eventStatus", instanceId, streamId, status, reason }) } private delay(duration: number, signal: AbortSignal) { diff --git a/packages/tauri-app/Cargo.lock b/packages/tauri-app/Cargo.lock index 9e3819644..904679971 100644 --- a/packages/tauri-app/Cargo.lock +++ b/packages/tauri-app/Cargo.lock @@ -519,6 +519,7 @@ dependencies = [ "tauri-plugin-notification", "tauri-plugin-opener", "tempfile", + "tokio", "url", "uuid", "webkit2gtk", @@ -4691,9 +4692,21 @@ dependencies = [ "mio", "pin-project-lite", "socket2", + "tokio-macros", "windows-sys 0.61.2", ] +[[package]] +name = "tokio-macros" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c55a2eff8b69ce66c84f85e1da1c233edc36ceb85a2058d11b0d6a3c7e7569c" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "tokio-rustls" version = "0.26.4" diff --git a/packages/tauri-app/src-tauri/Cargo.toml b/packages/tauri-app/src-tauri/Cargo.toml index d99219a8c..c8c223a25 100644 --- a/packages/tauri-app/src-tauri/Cargo.toml +++ b/packages/tauri-app/src-tauri/Cargo.toml @@ -9,6 +9,7 @@ tauri-build = { version = "2.5.2", features = [] } [dependencies] tauri = { version = "2.5.2", features = [ "devtools"] } +tokio = { version = "1", features = ["macros", "sync"] } serde = { version = "1", features = ["derive"] } serde_json = "1" serde_yaml = "0.9" diff --git a/packages/tauri-app/src-tauri/src/desktop_event_transport.rs b/packages/tauri-app/src-tauri/src/desktop_event_transport.rs index 372ca37c6..361203383 100644 --- a/packages/tauri-app/src-tauri/src/desktop_event_transport.rs +++ b/packages/tauri-app/src-tauri/src/desktop_event_transport.rs @@ -1,9 +1,8 @@ use parking_lot::Mutex; -use reqwest::blocking::{Client, Response}; use reqwest::StatusCode; +use reqwest::{Client, Response}; use serde::{Deserialize, Serialize}; use serde_json::Value; -use std::io::{BufRead, BufReader}; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::mpsc::{self, RecvTimeoutError, SyncSender}; use std::sync::Arc; @@ -23,12 +22,23 @@ const EVENT_STATUS_NAME: &str = "desktop:event-stream-status"; const FLUSH_INTERVAL_MS: u64 = 16; const DELTA_STREAM_WINDOW_MS: u64 = 48; const MAX_BATCH_EVENTS: usize = 256; +const MAX_BATCH_BYTES: usize = 2 * 1024 * 1024; const DEFAULT_RECONNECT_INITIAL_DELAY_MS: u64 = 1_000; const DEFAULT_RECONNECT_MAX_DELAY_MS: u64 = 10_000; const DEFAULT_RECONNECT_MULTIPLIER: f64 = 2.0; const STREAM_CONNECT_TIMEOUT_MS: u64 = 5_000; const STREAM_TCP_KEEPALIVE_MS: u64 = 30_000; -const STREAM_STALL_TIMEOUT_MS: u64 = 30_000; +const STREAM_READ_TIMEOUT_MS: u64 = 30_000; +const STREAM_STALL_TIMEOUT_MS: u64 = 35_000; +const SSE_READ_BUFFER_BYTES: usize = 8 * 1024; +const SERVER_MAX_EVENT_CHARACTERS: usize = 16 * 1024 * 1024; +const MAX_UTF8_BYTES_PER_CHARACTER: usize = 4; +const MAX_WORKSPACE_EVENT_ENVELOPE_BYTES: usize = 64 * 1024; +const MAX_SSE_LINE_BYTES: usize = + SERVER_MAX_EVENT_CHARACTERS * MAX_UTF8_BYTES_PER_CHARACTER + MAX_WORKSPACE_EVENT_ENVELOPE_BYTES; +const MAX_SSE_FRAME_BYTES: usize = MAX_SSE_LINE_BYTES + MAX_WORKSPACE_EVENT_ENVELOPE_BYTES; +const MAX_COALESCED_DELTA_BYTES: usize = 1024 * 1024; +const READER_CHANNEL_CAPACITY: usize = 1; #[derive(Clone, Debug, PartialEq, Eq)] pub struct DesktopEventStreamConfig { @@ -44,6 +54,13 @@ pub struct DesktopEventStreamConfig { #[serde(default, rename_all = "camelCase")] pub struct DesktopEventsStartRequest { pub reconnect: Option, + pub logical_start_epoch: u64, +} + +#[derive(Clone, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct DesktopEventsStartReservation { + pub logical_start_epoch: u64, } #[derive(Clone, Debug, Default, Deserialize)] @@ -60,6 +77,7 @@ pub struct DesktopEventReconnectPolicy { pub struct DesktopEventsStartResult { pub started: bool, pub generation: Option, + pub lease: Option, pub reason: Option, } @@ -159,11 +177,14 @@ struct DesktopEventTransportStats { struct DesktopEventTransportState { stop: Option>, config: Option, + lease: Option, + latest_reserved_start_epoch: u64, } pub struct DesktopEventTransportManager { state: Arc>, generation: Arc, + lease: AtomicU64, } enum ReaderMessage { @@ -178,6 +199,8 @@ enum PendingEntry { key: String, scope: String, event: Value, + serialized_bytes: usize, + delta_bytes: usize, started_at: Instant, }, Status { @@ -198,12 +221,14 @@ enum EventDeliveryPolicy { Passthrough, } +#[derive(Debug)] enum OpenStreamErrorKind { Unauthorized, Http, Transport, } +#[derive(Debug)] struct OpenStreamError { kind: OpenStreamErrorKind, message: String, @@ -213,6 +238,7 @@ struct OpenStreamError { #[derive(Default)] struct PendingBatch { events: Vec, + estimated_bytes: usize, } impl DesktopEventTransportManager { @@ -221,29 +247,52 @@ impl DesktopEventTransportManager { state: Arc::new(Mutex::new(DesktopEventTransportState { stop: None, config: None, + lease: None, + latest_reserved_start_epoch: 0, })), generation: Arc::new(AtomicU64::new(0)), + lease: AtomicU64::new(0), } } + pub fn reserve_start(&self) -> Result { + let mut state = self.state.lock(); + let logical_start_epoch = state + .latest_reserved_start_epoch + .checked_add(1) + .ok_or_else(|| "desktop event start epoch exhausted".to_string())?; + state.latest_reserved_start_epoch = logical_start_epoch; + Ok(DesktopEventsStartReservation { + logical_start_epoch, + }) + } + pub fn start( &self, app: AppHandle, stream_config: Option, - request: Option, + request: DesktopEventsStartRequest, ) -> DesktopEventsStartResult { let Some(stream_config) = stream_config else { return DesktopEventsStartResult { started: false, generation: None, + lease: None, reason: Some("desktop event stream unavailable".to_string()), }; }; - let request = request.unwrap_or_default(); let transport_config = DesktopEventTransportConfig::new(stream_config, &request); let mut state = self.state.lock(); + let Some(lease) = self.claim_start_lease(&mut state, request.logical_start_epoch) else { + return DesktopEventsStartResult { + started: false, + generation: None, + lease: None, + reason: Some("stale logical desktop event start".to_string()), + }; + }; if state .config .as_ref() @@ -254,6 +303,7 @@ impl DesktopEventTransportManager { return DesktopEventsStartResult { started: true, generation: Some(self.generation.load(Ordering::SeqCst)), + lease: Some(lease), reason: None, }; } @@ -278,17 +328,45 @@ impl DesktopEventTransportManager { DesktopEventsStartResult { started: true, generation: Some(generation), + lease: Some(lease), reason: None, } } + fn claim_start_lease( + &self, + state: &mut DesktopEventTransportState, + logical_start_epoch: u64, + ) -> Option { + if logical_start_epoch != state.latest_reserved_start_epoch { + return None; + } + + let lease = self.lease.fetch_add(1, Ordering::SeqCst) + 1; + state.lease = Some(lease); + Some(lease) + } + pub fn stop(&self) { + self.stop_current(None); + } + + pub fn stop_lease(&self, lease: u64) -> bool { + self.stop_current(Some(lease)) + } + + fn stop_current(&self, expected_lease: Option) -> bool { let mut state = self.state.lock(); + if expected_lease.is_some_and(|lease| state.lease != Some(lease)) { + return false; + } if let Some(stop) = state.stop.take() { stop.store(true, Ordering::SeqCst); } state.config = None; + state.lease = None; self.generation.fetch_add(1, Ordering::SeqCst); + true } } @@ -316,11 +394,16 @@ fn coalesced_payload_event<'a>(event: &'a Value) -> &'a Value { } } -fn coalesced_instance_id(event: &Value) -> &str { - event +fn coalesced_instance_id(event: &Value) -> String { + let instance_id = event .get("instanceId") .and_then(Value::as_str) - .unwrap_or_default() + .unwrap_or_default(); + let stream_id = event + .get("streamId") + .and_then(Value::as_str) + .unwrap_or_default(); + format!("{}@{}", instance_id, stream_id) } fn snapshot_key(event: &Value) -> Option { @@ -362,7 +445,7 @@ fn snapshot_key(event: &Value) -> Option { instance_id, session_id, message_id )) } - "session.updated" | "session.status" => { + "session.updated" => { let session_id = props .get("info") .and_then(|info| info.get("id")) @@ -440,21 +523,59 @@ fn snapshot_superseded_delta_scope(event: &Value) -> Option { )) } -fn append_delta(target: &mut Value, event: &Value) { - let next_delta = coalesced_payload_event(event) +fn delta_payload(event: &Value) -> &str { + coalesced_payload_event(event) .get("properties") .and_then(|value| value.get("delta")) .and_then(Value::as_str) - .unwrap_or_default(); + .unwrap_or_default() +} + +fn append_delta(target: &mut Value, next_delta: &str, delta_bytes: &mut usize) -> bool { + let Some(combined_len) = delta_bytes.checked_add(next_delta.len()) else { + return false; + }; + if combined_len > MAX_COALESCED_DELTA_BYTES { + return false; + } - if let Some(existing_delta) = coalesced_payload_event_mut(target) + let Some(Value::String(existing_delta)) = coalesced_payload_event_mut(target) .and_then(|event| event.get_mut("properties")) .and_then(Value::as_object_mut) .and_then(|props| props.get_mut("delta")) - { - let combined = existing_delta.as_str().unwrap_or_default().to_string() + next_delta; - *existing_delta = Value::String(combined); + else { + return false; + }; + + existing_delta.push_str(next_delta); + *delta_bytes = combined_len; + true +} + +fn serialized_json_bytes(value: &T) -> usize { + struct Counter(usize); + + impl std::io::Write for Counter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0 = self.0.saturating_add(bytes.len()); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } } + + let mut counter = Counter(0); + serde_json::to_writer(&mut counter, value).map_or(usize::MAX, |_| counter.0) +} + +fn serialized_value_bytes(value: &Value) -> usize { + serialized_json_bytes(value) +} + +fn serialized_string_content_bytes(value: &str) -> usize { + serialized_json_bytes(value).saturating_sub(2) } fn coalesced_payload_event_mut(event: &mut Value) -> Option<&mut serde_json::Map> { @@ -467,8 +588,14 @@ fn coalesced_payload_event_mut(event: &mut Value) -> Option<&mut serde_json::Map fn status_key(event: &Value) -> Option { match event.get("type")?.as_str()? { - "instance.eventStatus" => Some(coalesced_instance_id(event).to_string()), - "session.status" => snapshot_key(event), + "instance.eventStatus" => Some(format!( + "{}:{}", + coalesced_instance_id(event), + event + .get("status") + .and_then(Value::as_str) + .unwrap_or_default() + )), _ => None, } } diff --git a/packages/tauri-app/src-tauri/src/desktop_event_transport/assembler.rs b/packages/tauri-app/src-tauri/src/desktop_event_transport/assembler.rs index f91bcb760..d5756c3ae 100644 --- a/packages/tauri-app/src-tauri/src/desktop_event_transport/assembler.rs +++ b/packages/tauri-app/src-tauri/src/desktop_event_transport/assembler.rs @@ -5,53 +5,83 @@ impl PendingBatch { match classify_event(&event) { EventDeliveryPolicy::CoalesceDelta(key) => { let Some(scope) = delta_scope(&event) else { + let event_bytes = serialized_value_bytes(&event); self.events.push(PendingEntry::Event(event)); + self.estimated_bytes = self.estimated_bytes.saturating_add(event_bytes); return; }; + let mut appended_bytes = None; if let Some(PendingEntry::Delta { key: existing_key, event: existing_event, + serialized_bytes, + delta_bytes, .. }) = self.events.last_mut() { if existing_key == &key { - append_delta(existing_event, &event); - stats.delta_coalesces = stats.delta_coalesces.saturating_add(1); - return; + let next_delta = delta_payload(&event); + let next_bytes = serialized_string_content_bytes(next_delta); + if append_delta(existing_event, next_delta, delta_bytes) { + *serialized_bytes = serialized_bytes.saturating_add(next_bytes); + appended_bytes = Some(next_bytes); + stats.delta_coalesces = stats.delta_coalesces.saturating_add(1); + } } } + if let Some(appended_bytes) = appended_bytes { + self.estimated_bytes = self.estimated_bytes.saturating_add(appended_bytes); + return; + } + let event_bytes = serialized_value_bytes(&event); + let delta_bytes = delta_payload(&event).len(); self.events.push(PendingEntry::Delta { key, scope, event, + serialized_bytes: event_bytes, + delta_bytes, started_at: Instant::now(), }); + self.estimated_bytes = self.estimated_bytes.saturating_add(event_bytes); } EventDeliveryPolicy::CoalesceStatus(key) => { + let event_bytes = serialized_value_bytes(&event); if let Some(PendingEntry::Status { key: existing_key, event: existing_event, }) = self.events.last_mut() { if existing_key == &key { + let old_bytes = serialized_value_bytes(existing_event); *existing_event = event; + self.estimated_bytes = self + .estimated_bytes + .saturating_sub(old_bytes) + .saturating_add(event_bytes); stats.status_coalesces = stats.status_coalesces.saturating_add(1); return; } } self.events.push(PendingEntry::Status { key, event }); + self.estimated_bytes = self.estimated_bytes.saturating_add(event_bytes); } EventDeliveryPolicy::CoalesceSnapshot(key) => { + let event_bytes = serialized_value_bytes(&event); if let Some(part_scope) = snapshot_superseded_delta_scope(&event) { let mut dropped = 0_u64; while matches!( self.events.last(), Some(PendingEntry::Delta { scope, .. }) if scope == &part_scope ) { - self.events.pop(); + if let Some(entry) = self.events.pop() { + self.estimated_bytes = self + .estimated_bytes + .saturating_sub(pending_entry_bytes(&entry)); + } dropped = dropped.saturating_add(1); } if dropped > 0 { @@ -66,21 +96,30 @@ impl PendingBatch { }) = self.events.last_mut() { if existing_key == &key { + let old_bytes = serialized_value_bytes(existing_event); *existing_event = event; + self.estimated_bytes = self + .estimated_bytes + .saturating_sub(old_bytes) + .saturating_add(event_bytes); stats.snapshot_coalesces = stats.snapshot_coalesces.saturating_add(1); return; } } self.events.push(PendingEntry::Snapshot { key, event }); + self.estimated_bytes = self.estimated_bytes.saturating_add(event_bytes); } EventDeliveryPolicy::Passthrough => { + let event_bytes = serialized_value_bytes(&event); self.events.push(PendingEntry::Event(event)); + self.estimated_bytes = self.estimated_bytes.saturating_add(event_bytes); } } } pub(super) fn take_events(&mut self) -> Vec { + self.estimated_bytes = 0; let pending = std::mem::take(&mut self.events); pending .into_iter() @@ -101,6 +140,10 @@ impl PendingBatch { self.events.len() } + pub(super) fn pending_bytes(&self) -> usize { + self.estimated_bytes + } + pub(super) fn should_hold_single_delta(&self, now: Instant) -> bool { matches!( self.events.as_slice(), @@ -110,3 +153,14 @@ impl PendingBatch { ) } } + +fn pending_entry_bytes(entry: &PendingEntry) -> usize { + match entry { + PendingEntry::Delta { + serialized_bytes, .. + } => *serialized_bytes, + PendingEntry::Status { event, .. } + | PendingEntry::Snapshot { event, .. } + | PendingEntry::Event(event) => serialized_value_bytes(event), + } +} diff --git a/packages/tauri-app/src-tauri/src/desktop_event_transport/stream.rs b/packages/tauri-app/src-tauri/src/desktop_event_transport/stream.rs index 33737f9c8..e94de7894 100644 --- a/packages/tauri-app/src-tauri/src/desktop_event_transport/stream.rs +++ b/packages/tauri-app/src-tauri/src/desktop_event_transport/stream.rs @@ -1,17 +1,17 @@ use super::*; -use reqwest::blocking::RequestBuilder; +use reqwest::RequestBuilder; pub(super) fn build_stream_client() -> Result { + build_stream_client_with_read_timeout(Duration::from_millis(STREAM_READ_TIMEOUT_MS)) +} + +fn build_stream_client_with_read_timeout( + read_timeout: Duration, +) -> Result { Client::builder() .connect_timeout(Duration::from_millis(STREAM_CONNECT_TIMEOUT_MS)) + .read_timeout(read_timeout) .tcp_keepalive(Duration::from_millis(STREAM_TCP_KEEPALIVE_MS)) - // Note: reqwest's blocking client doesn't expose a per-read timeout. - // The global `.timeout()` would kill the entire SSE stream, so we - // rely on: - // 1. tcp_keepalive to detect dead connections (OS will RST after - // several unacked probes, typically ~2 min). - // 2. Consumer-side stall detection (STREAM_STALL_TIMEOUT_MS). - // 3. Reader thread breaking on channel send error (consumer dropped). .build() .map_err(|error: reqwest::Error| OpenStreamError { kind: OpenStreamErrorKind::Transport, @@ -36,11 +36,14 @@ pub(super) fn open_stream( config, ); - let response = request.send().map_err(|error| OpenStreamError { - kind: OpenStreamErrorKind::Transport, - message: error.to_string(), - status_code: None, - })?; + let response = + tauri::async_runtime::block_on(async { request.send().await }).map_err(|error| { + OpenStreamError { + kind: OpenStreamErrorKind::Transport, + message: error.to_string(), + status_code: None, + } + })?; if response.status().is_success() { return Ok(response); @@ -152,53 +155,164 @@ fn read_session_cookie_from_webview( } pub(super) fn read_sse( - response: Response, + mut response: Response, tx: SyncSender, stop: Arc, generation_atomic: Arc, generation: u64, + mut cancel: tokio::sync::oneshot::Receiver<()>, ) { - let mut reader = BufReader::new(response); - let mut line = String::new(); - let mut event_name: Option = None; - let mut data_lines: Vec = Vec::new(); - - loop { - if stop.load(Ordering::SeqCst) || !generation_matches(&generation_atomic, generation) { - let _ = tx.send(ReaderMessage::End(Some("stopped".to_string()))); - return; - } + let mut buffer = [0_u8; SSE_READ_BUFFER_BYTES]; + let mut decoder = SseDecoder::new(MAX_SSE_LINE_BYTES, MAX_SSE_FRAME_BYTES); - line.clear(); - match reader.read_line(&mut line) { - Ok(0) => { - let _ = flush_sse_frame(&tx, &event_name, &data_lines); - let _ = tx.send(ReaderMessage::End(Some("stream closed".to_string()))); + tauri::async_runtime::block_on(async move { + loop { + if stop.load(Ordering::SeqCst) || !generation_matches(&generation_atomic, generation) { + let _ = tx.send(ReaderMessage::End(Some("stopped".to_string()))); return; } - Ok(_) => { - if tx.send(ReaderMessage::Activity).is_err() { - return; // consumer dropped — stop reading + + let next_chunk = tokio::select! { + _ = &mut cancel => { + let _ = tx.send(ReaderMessage::End(Some("stopped".to_string()))); + return; } - let trimmed = line.trim_end_matches(['\r', '\n']); - if handle_sse_line(trimmed, &mut event_name, &mut data_lines) { - if flush_sse_frame(&tx, &event_name, &data_lines).is_err() { - return; + chunk = response.chunk() => chunk, + }; + + match next_chunk { + Ok(None) => { + decoder.discard_frame(); + let _ = tx.send(ReaderMessage::End(Some("stream closed".to_string()))); + return; + } + Ok(Some(chunk)) => { + if tx.send(ReaderMessage::Activity).is_err() { + return; // consumer dropped - stop reading + } + for bytes in chunk.chunks(buffer.len()) { + buffer[..bytes.len()].copy_from_slice(bytes); + if let Err(error) = decoder.push(&buffer[..bytes.len()], &tx) { + let _ = tx.send(ReaderMessage::End(Some(error))); + return; + } } - event_name = None; - data_lines.clear(); + } + Err(error) => { + decoder.discard_frame(); + let _ = tx.send(ReaderMessage::End(Some(error.to_string()))); + return; + } + } + } + }); +} + +struct SseDecoder { + line: Vec, + event_name: Option, + data_lines: Vec, + frame_bytes: usize, + max_line_bytes: usize, + max_frame_bytes: usize, + skip_lf: bool, +} + +impl SseDecoder { + fn new(max_line_bytes: usize, max_frame_bytes: usize) -> Self { + Self { + line: Vec::with_capacity(max_line_bytes.min(SSE_READ_BUFFER_BYTES)), + event_name: None, + data_lines: Vec::new(), + frame_bytes: 0, + max_line_bytes, + max_frame_bytes, + skip_lf: false, + } + } + + fn push(&mut self, bytes: &[u8], tx: &SyncSender) -> Result<(), String> { + for &byte in bytes { + if self.skip_lf { + self.skip_lf = false; + if byte == b'\n' { continue; } } - Err(error) => { - let _ = flush_sse_frame(&tx, &event_name, &data_lines); - let _ = tx.send(ReaderMessage::End(Some(error.to_string()))); - return; + + match byte { + b'\r' => { + self.finish_line(tx)?; + self.skip_lf = true; + } + b'\n' => self.finish_line(tx)?, + _ => { + checked_sse_size(self.line.len(), 1, self.max_line_bytes, "SSE line")?; + self.line.push(byte); + } } } + + Ok(()) + } + + fn discard_frame(&mut self) { + self.line.clear(); + self.event_name = None; + self.data_lines.clear(); + self.frame_bytes = 0; + self.skip_lf = false; + } + + fn finish_line(&mut self, tx: &SyncSender) -> Result<(), String> { + if self.line.is_empty() { + return self.flush_frame(tx); + } + + let line_bytes = self + .line + .len() + .checked_add(1) + .ok_or_else(|| "SSE frame size overflow".to_string())?; + self.frame_bytes = checked_sse_size( + self.frame_bytes, + line_bytes, + self.max_frame_bytes, + "SSE frame", + )?; + + let line = std::str::from_utf8(&self.line) + .map_err(|error| format!("invalid UTF-8 in SSE stream: {error}"))?; + handle_sse_line(line, &mut self.event_name, &mut self.data_lines); + self.line.clear(); + Ok(()) + } + + fn flush_frame(&mut self, tx: &SyncSender) -> Result<(), String> { + flush_sse_frame(tx, &self.event_name, &self.data_lines) + .map_err(|_| "desktop event consumer dropped".to_string())?; + self.event_name = None; + self.data_lines.clear(); + self.frame_bytes = 0; + Ok(()) } } +fn checked_sse_size( + current: usize, + additional: usize, + maximum: usize, + label: &str, +) -> Result { + let next = current + .checked_add(additional) + .ok_or_else(|| format!("{label} size overflow"))?; + if next > maximum { + return Err(format!("{label} exceeded {maximum} bytes")); + } + Ok(next) +} + fn handle_sse_line( trimmed: &str, event_name: &mut Option, @@ -256,6 +370,176 @@ fn parse_sse_payload(lines: &[String]) -> Option { #[cfg(test)] mod tests { use super::*; + use std::io::{Read, Write}; + use std::net::TcpListener; + + fn decode_single_event(input: &[u8]) -> Value { + let (tx, rx) = mpsc::sync_channel(1); + let mut decoder = SseDecoder::new(256, 1024); + decoder.push(input, &tx).expect("stream should decode"); + decoder.discard_frame(); + + match rx.recv().expect("event should be emitted") { + ReaderMessage::Event(payload) => payload, + _ => panic!("expected event frame"), + } + } + + #[test] + fn decodes_lf_crlf_and_cr_only_streams() { + for input in [ + &b"data: {\"ending\":\"lf\"}\n\n"[..], + &b"data: {\"ending\":\"crlf\"}\r\n\r\n"[..], + &b"data: {\"ending\":\"cr\"}\r\r"[..], + ] { + assert!(decode_single_event(input).get("ending").is_some()); + } + } + + #[test] + fn rejects_oversized_line_when_peer_withholds_lf() { + let (tx, _rx) = mpsc::sync_channel(1); + let mut decoder = SseDecoder::new(8, 64); + + let error = decoder + .push(b"data: 123", &tx) + .expect_err("ninth byte must exceed the line limit"); + + assert_eq!(error, "SSE line exceeded 8 bytes"); + } + + #[test] + fn rejects_oversized_frame_across_bounded_lines() { + let (tx, _rx) = mpsc::sync_channel(1); + let mut decoder = SseDecoder::new(16, 15); + + let error = decoder + .push(b"data: a\ndata: b\n", &tx) + .expect_err("second line must exceed the frame limit"); + + assert_eq!(error, "SSE frame exceeded 15 bytes"); + } + + #[test] + fn server_maximum_wire_frame_fits_and_next_byte_is_rejected_without_allocation() { + let maximum_server_event_bytes = SERVER_MAX_EVENT_CHARACTERS + .checked_mul(MAX_UTF8_BYTES_PER_CHARACTER) + .and_then(|bytes| bytes.checked_add(MAX_WORKSPACE_EVENT_ENVELOPE_BYTES)) + .expect("configured SSE line limit should fit usize"); + + assert_eq!(maximum_server_event_bytes, MAX_SSE_LINE_BYTES); + assert_eq!( + checked_sse_size( + 0, + maximum_server_event_bytes, + MAX_SSE_LINE_BYTES, + "SSE line" + ), + Ok(MAX_SSE_LINE_BYTES) + ); + assert_eq!( + checked_sse_size(MAX_SSE_LINE_BYTES, 1, MAX_SSE_LINE_BYTES, "SSE line"), + Err(format!("SSE line exceeded {MAX_SSE_LINE_BYTES} bytes")) + ); + assert!(MAX_SSE_LINE_BYTES + 1 <= MAX_SSE_FRAME_BYTES); + assert_eq!( + checked_sse_size(0, MAX_SSE_FRAME_BYTES, MAX_SSE_FRAME_BYTES, "SSE frame"), + Ok(MAX_SSE_FRAME_BYTES) + ); + assert_eq!( + checked_sse_size(MAX_SSE_FRAME_BYTES, 1, MAX_SSE_FRAME_BYTES, "SSE frame"), + Err(format!("SSE frame exceeded {MAX_SSE_FRAME_BYTES} bytes")) + ); + } + + #[test] + fn discards_unterminated_frame_at_eof() { + let (tx, rx) = mpsc::sync_channel(1); + let mut decoder = SseDecoder::new(256, 1024); + decoder + .push(b"data: {\"incomplete\":true}\n", &tx) + .expect("line should decode"); + + decoder.discard_frame(); + + assert!(matches!(rx.try_recv(), Err(mpsc::TryRecvError::Empty))); + } + + #[test] + fn read_timeout_ends_a_live_but_silent_response() { + let listener = TcpListener::bind("127.0.0.1:0").expect("listener should bind"); + let address = listener + .local_addr() + .expect("listener should have an address"); + let (release_tx, release_rx) = mpsc::channel(); + let server = thread::spawn(move || { + let (mut socket, _) = listener.accept().expect("client should connect"); + let mut request = [0_u8; 1024]; + socket.read(&mut request).expect("request should arrive"); + socket + .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n") + .expect("headers should send"); + socket.flush().expect("headers should flush"); + release_rx.recv().expect("test should release server"); + }); + let client = build_stream_client_with_read_timeout(Duration::from_millis(20)) + .expect("client should build"); + let mut response = tauri::async_runtime::block_on(async { + client.get(format!("http://{address}/events")).send().await + }) + .expect("response headers should arrive"); + + let error = tauri::async_runtime::block_on(async { response.chunk().await }) + .expect_err("silent body should time out"); + + assert!(error.is_timeout(), "unexpected error: {error}"); + release_tx.send(()).expect("server should release"); + server.join().expect("server should stop"); + } + + #[test] + fn cancellation_ends_a_silent_reader_promptly() { + let listener = TcpListener::bind("127.0.0.1:0").expect("listener should bind"); + let address = listener + .local_addr() + .expect("listener should have an address"); + let (release_tx, release_rx) = mpsc::channel(); + let server = thread::spawn(move || { + let (mut socket, _) = listener.accept().expect("client should connect"); + let mut request = [0_u8; 1024]; + socket.read(&mut request).expect("request should arrive"); + socket + .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n") + .expect("headers should send"); + socket.flush().expect("headers should flush"); + release_rx.recv().expect("test should release server"); + }); + let client = build_stream_client_with_read_timeout(Duration::from_secs(5)) + .expect("client should build"); + let response = tauri::async_runtime::block_on(async { + client.get(format!("http://{address}/events")).send().await + }) + .expect("response headers should arrive"); + let (tx, rx) = mpsc::sync_channel(READER_CHANNEL_CAPACITY); + let stop = Arc::new(AtomicBool::new(false)); + let generation = Arc::new(AtomicU64::new(1)); + let (cancel_tx, cancel_rx) = tokio::sync::oneshot::channel(); + let started_at = Instant::now(); + let reader = thread::spawn(move || { + read_sse(response, tx, stop, generation, 1, cancel_rx); + }); + + cancel_tx.send(()).expect("reader should be running"); + reader.join().expect("reader should stop"); + + assert!(started_at.elapsed() < Duration::from_secs(1)); + assert!(matches!( + rx.recv_timeout(Duration::from_secs(1)), + Ok(ReaderMessage::End(Some(reason))) if reason == "stopped" + )); + release_tx.send(()).expect("server should release"); + server.join().expect("server should stop"); + } #[test] fn named_ping_event_is_routed_to_ping_channel() { diff --git a/packages/tauri-app/src-tauri/src/desktop_event_transport/tests.rs b/packages/tauri-app/src-tauri/src/desktop_event_transport/tests.rs index f2440a201..162bc7cb1 100644 --- a/packages/tauri-app/src-tauri/src/desktop_event_transport/tests.rs +++ b/packages/tauri-app/src-tauri/src/desktop_event_transport/tests.rs @@ -113,7 +113,76 @@ fn coalesces_message_part_delta_events() { } #[test] -fn last_write_wins_for_status_events() { +fn splits_coalesced_delta_run_before_size_overflow_without_losing_content() { + let mut pending = PendingBatch::default(); + let mut stats = fresh_stats(); + let first = "a".repeat(MAX_COALESCED_DELTA_BYTES / 2); + let second = "b".repeat(MAX_COALESCED_DELTA_BYTES - first.len()); + pending.push(delta_event(&first), &mut stats); + pending.push(delta_event(&second), &mut stats); + + let PendingEntry::Delta { event, .. } = &pending.events[0] else { + panic!("expected coalesced delta"); + }; + assert_eq!( + event["event"]["properties"]["delta"].as_str().map(str::len), + Some(MAX_COALESCED_DELTA_BYTES) + ); + pending.push(delta_event("b"), &mut stats); + + let events = pending.take_events(); + assert_eq!(events.len(), 2); + assert_eq!( + events[0]["event"]["properties"]["delta"] + .as_str() + .map(str::len), + Some(MAX_COALESCED_DELTA_BYTES) + ); + assert_eq!(events[1]["event"]["properties"]["delta"], "b"); + assert_eq!(stats.delta_coalesces, 1); +} + +#[test] +fn tracks_cumulative_batch_bytes_for_flush_budget() { + let mut pending = PendingBatch::default(); + let mut stats = fresh_stats(); + let chunk = "x".repeat(MAX_BATCH_BYTES / 2); + + pending.push(delta_event_for("part-1", &chunk), &mut stats); + pending.push(delta_event_for("part-2", &chunk), &mut stats); + + assert!(pending.pending_bytes() >= MAX_BATCH_BYTES); +} + +#[test] +fn tracks_exact_bytes_when_appending_escaped_deltas() { + let mut pending = PendingBatch::default(); + let mut stats = fresh_stats(); + pending.push(delta_event("first\\\n"), &mut stats); + pending.push(delta_event("\"second\""), &mut stats); + + let expected = delta_event("first\\\n\"second\""); + assert_eq!(pending.pending_bytes(), serialized_value_bytes(&expected)); + assert_eq!(pending.take_events(), vec![expected]); +} + +#[test] +fn does_not_coalesce_events_across_instance_streams() { + let mut pending = PendingBatch::default(); + let mut stats = fresh_stats(); + let mut old = delta_event("old"); + old["streamId"] = Value::String("old-stream".to_string()); + let mut new = delta_event("new"); + new["streamId"] = Value::String("new-stream".to_string()); + + pending.push(old, &mut stats); + pending.push(new, &mut stats); + + assert_eq!(pending.take_events().len(), 2); +} + +#[test] +fn preserves_connecting_before_connected_status() { let mut pending = PendingBatch::default(); let mut stats = fresh_stats(); pending.push( @@ -134,8 +203,36 @@ fn last_write_wins_for_status_events() { ); let events = pending.take_events(); - assert_eq!(events.len(), 1); - assert_eq!(events[0]["status"].as_str(), Some("connected")); + assert_eq!(events.len(), 2); + assert_eq!(events[0]["status"].as_str(), Some("connecting")); + assert_eq!(events[1]["status"].as_str(), Some("connected")); +} + +#[test] +fn preserves_working_before_idle_session_status() { + let mut pending = PendingBatch::default(); + let mut stats = fresh_stats(); + for status in ["busy", "idle"] { + pending.push( + json!({ + "type": "instance.event", + "instanceId": "inst-1", + "event": { + "type": "session.status", + "properties": { + "sessionID": "sess-1", + "status": { "type": status } + } + } + }), + &mut stats, + ); + } + + let events = pending.take_events(); + assert_eq!(events.len(), 2); + assert_eq!(events[0]["event"]["properties"]["status"]["type"], "busy"); + assert_eq!(events[1]["event"]["properties"]["status"]["type"], "idle"); } #[test] @@ -300,8 +397,11 @@ fn holds_single_delta_within_stream_window() { key: "delta-key".to_string(), scope: "delta-scope".to_string(), event: delta_event("Hello"), + serialized_bytes: serialized_value_bytes(&delta_event("Hello")), + delta_bytes: "Hello".len(), started_at: Instant::now(), }], + ..PendingBatch::default() }; assert!(pending.should_hold_single_delta(Instant::now())); @@ -315,8 +415,11 @@ fn flushes_single_delta_after_stream_window() { key: "delta-key".to_string(), scope: "delta-scope".to_string(), event: delta_event("Hello"), + serialized_bytes: serialized_value_bytes(&delta_event("Hello")), + delta_bytes: "Hello".len(), started_at, }], + ..PendingBatch::default() }; assert!(!pending.should_hold_single_delta(Instant::now())); @@ -372,3 +475,57 @@ fn equivalent_transport_start_detects_material_stream_changes() { assert!(!first.is_equivalent_start(&second)); } + +#[test] +fn only_latest_lease_can_stop_a_reused_stream_generation() { + let manager = DesktopEventTransportManager::new(); + let current_stop = Arc::new(AtomicBool::new(false)); + manager.generation.store(1, Ordering::SeqCst); + { + let mut state = manager.state.lock(); + state.stop = Some(current_stop.clone()); + state.lease = Some(2); + } + + assert!(!manager.stop_lease(1)); + assert!(!current_stop.load(Ordering::SeqCst)); + assert_eq!(manager.generation.load(Ordering::SeqCst), 1); + assert!(manager.state.lock().stop.is_some()); + + assert!(manager.stop_lease(2)); + assert!(current_stop.load(Ordering::SeqCst)); + assert_eq!(manager.generation.load(Ordering::SeqCst), 2); + assert!(manager.state.lock().stop.is_none()); +} + +#[test] +fn start_reservations_remain_monotonic_across_renderer_reload() { + let manager = DesktopEventTransportManager::new(); + + let before_reload = manager.reserve_start().unwrap().logical_start_epoch; + let after_reload = manager.reserve_start().unwrap().logical_start_epoch; + + assert_eq!(before_reload, 1); + assert_eq!(after_reload, 2); +} + +#[test] +fn newer_reservation_rejects_an_older_start_arriving_out_of_order() { + let manager = DesktopEventTransportManager::new(); + let current_stop = Arc::new(AtomicBool::new(false)); + let older_epoch = manager.reserve_start().unwrap().logical_start_epoch; + let newer_epoch = manager.reserve_start().unwrap().logical_start_epoch; + let (current_lease, stale_lease) = { + let mut state = manager.state.lock(); + state.stop = Some(current_stop.clone()); + let current_lease = manager.claim_start_lease(&mut state, newer_epoch).unwrap(); + let stale_lease = manager.claim_start_lease(&mut state, older_epoch); + assert_eq!(state.lease, Some(current_lease)); + (current_lease, stale_lease) + }; + + assert_eq!(stale_lease, None); + assert!(!current_stop.load(Ordering::SeqCst)); + assert!(manager.stop_lease(current_lease)); + assert!(current_stop.load(Ordering::SeqCst)); +} diff --git a/packages/tauri-app/src-tauri/src/desktop_event_transport/transport.rs b/packages/tauri-app/src-tauri/src/desktop_event_transport/transport.rs index 5f0ed2314..703051602 100644 --- a/packages/tauri-app/src-tauri/src/desktop_event_transport/transport.rs +++ b/packages/tauri-app/src-tauri/src/desktop_event_transport/transport.rs @@ -19,7 +19,8 @@ fn send_connection_pong( )) .json(&body); - let _ = attach_session_cookie(request, app, config).send(); + let request = attach_session_cookie(request, app, config); + let _ = tauri::async_runtime::block_on(async { request.send().await }); } pub(super) fn run_transport_loop( @@ -224,16 +225,18 @@ fn consume_stream( stop: Arc, stats: &mut DesktopEventTransportStats, ) -> Option { - let (tx, rx) = mpsc::sync_channel::(4096); + let (tx, rx) = mpsc::sync_channel::(READER_CHANNEL_CAPACITY); + let (reader_cancel_tx, reader_cancel_rx) = tokio::sync::oneshot::channel(); let reader_stop = stop.clone(); let reader_generation_atomic = generation_atomic.clone(); - thread::spawn(move || { + let reader = thread::spawn(move || { read_sse( response, tx, reader_stop, reader_generation_atomic, generation, + reader_cancel_rx, ) }); @@ -241,9 +244,9 @@ fn consume_stream( let mut sequence = 0_u64; let mut last_reader_activity = Instant::now(); - loop { + let disconnect_reason = loop { if stop.load(Ordering::SeqCst) || !generation_matches(generation_atomic, generation) { - return Some("stopped".to_string()); + break Some("stopped".to_string()); } match rx.recv_timeout(Duration::from_millis(FLUSH_INTERVAL_MS)) { @@ -259,7 +262,9 @@ fn consume_stream( stats.raw_events = stats.raw_events.saturating_add(1); pending.push(event, stats); - if pending.pending_len() >= MAX_BATCH_EVENTS { + if pending.pending_len() >= MAX_BATCH_EVENTS + || pending.pending_bytes() >= MAX_BATCH_BYTES + { emit_pending_batch( app, generation, @@ -281,7 +286,7 @@ fn consume_stream( stats, ); } - return reason; + break reason; } Err(RecvTimeoutError::Timeout) => { if last_reader_activity.elapsed() >= Duration::from_millis(STREAM_STALL_TIMEOUT_MS) @@ -297,7 +302,7 @@ fn consume_stream( stats, ); } - return Some("stream stalled".to_string()); + break Some("stream stalled".to_string()); } if !pending.is_empty() { @@ -325,10 +330,15 @@ fn consume_stream( stats, ); } - return Some("reader disconnected".to_string()); + break Some("reader disconnected".to_string()); } } - } + }; + + drop(rx); + let _ = reader_cancel_tx.send(()); + let _ = reader.join(); + disconnect_reason } fn emit_pending_batch( diff --git a/packages/tauri-app/src-tauri/src/main.rs b/packages/tauri-app/src-tauri/src/main.rs index 4dc142612..27f4bee45 100644 --- a/packages/tauri-app/src-tauri/src/main.rs +++ b/packages/tauri-app/src-tauri/src/main.rs @@ -13,7 +13,8 @@ mod windows_update; use cli_manager::{CliProcessManager, CliStatus}; use desktop_event_transport::{ - DesktopEventTransportManager, DesktopEventsStartRequest, DesktopEventsStartResult, + DesktopEventTransportManager, DesktopEventsStartRequest, DesktopEventsStartReservation, + DesktopEventsStartResult, }; use keepawake::KeepAwake; use serde::Deserialize; @@ -147,19 +148,26 @@ fn cli_restart(app: AppHandle, state: tauri::State) -> Result, +) -> Result { + state.desktop_events.reserve_start() +} + #[tauri::command] fn desktop_events_start( app: AppHandle, state: tauri::State, - request: Option, + request: DesktopEventsStartRequest, ) -> DesktopEventsStartResult { let config = state.manager.desktop_event_stream_config(); state.desktop_events.start(app, config, request) } #[tauri::command] -fn desktop_events_stop(state: tauri::State) { - state.desktop_events.stop(); +fn desktop_events_stop(state: tauri::State, lease: u64) { + state.desktop_events.stop_lease(lease); } #[tauri::command] @@ -660,6 +668,7 @@ fn main() { .invoke_handler(tauri::generate_handler![ cli_get_status, cli_restart, + desktop_events_reserve_start, desktop_events_start, desktop_events_stop, wake_lock_start, diff --git a/packages/ui/src/components/instance/instance-shell2.tsx b/packages/ui/src/components/instance/instance-shell2.tsx index 004d85d1e..1ad02aef0 100644 --- a/packages/ui/src/components/instance/instance-shell2.tsx +++ b/packages/ui/src/components/instance/instance-shell2.tsx @@ -609,6 +609,7 @@ const InstanceShell2: Component = (props) => { instanceId: () => props.instance.id, instanceSessions: allInstanceSessions, activeSessionId: activeSessionIdForInstance, + visible: () => Boolean(props.isActiveInstance), }) const showEmbeddedSidebarToggle = createMemo(() => !leftPinned() && !leftOpen()) @@ -992,7 +993,11 @@ const InstanceShell2: Component = (props) => { await submit(session.id) } catch (error) { focusCreatedSessionPrompt(session.id) - throw error + const recoveryError = new Error(error instanceof Error ? error.message : String(error)) + ;(recoveryError as any).cause = error + ;(recoveryError as any).promptRecoverySessionId = session.id + if ((error as any)?.suppressPromptRecovery === true) (recoveryError as any).suppressPromptRecovery = true + throw recoveryError } } diff --git a/packages/ui/src/components/instance/shell/right-panel/useGitChanges.ts b/packages/ui/src/components/instance/shell/right-panel/useGitChanges.ts index de1899893..e4ad77633 100644 --- a/packages/ui/src/components/instance/shell/right-panel/useGitChanges.ts +++ b/packages/ui/src/components/instance/shell/right-panel/useGitChanges.ts @@ -7,7 +7,7 @@ import { getRootClient } from "../../../../stores/opencode-client" import { getOpenCodeWorkspaceIdForWorktree } from "../../../../stores/opencode-workspaces" import { requestData } from "../../../../lib/opencode-api" import { serverApi } from "../../../../lib/api-client" -import { serverEvents } from "../../../../lib/server-events" +import { sseManager } from "../../../../lib/sse-manager" import { showToastNotification } from "../../../../lib/notifications" import { adaptSdkGitStatusEntries, buildGitChangeListItems } from "./git-changes-model" @@ -427,10 +427,9 @@ export function useGitChanges(options: UseGitChangesOptions) { createEffect(() => { if (options.rightPanelTab() !== "git-changes") return - const unsubscribe = serverEvents.on("instance.event", (event) => { - if (event.type !== "instance.event") return - if (event.instanceId !== options.instanceId) return - const eventType = (event.event as { type?: unknown } | undefined)?.type + const unsubscribe = sseManager.onAcceptedEvent((instanceId, event) => { + if (instanceId !== options.instanceId) return + const eventType = (event as { type?: unknown }).type if (eventType !== "session.updated") return void passiveRefreshGitStatus({ forceReloadSelectedDiff: true }) }) diff --git a/packages/ui/src/components/instance/shell/useSessionCache.ts b/packages/ui/src/components/instance/shell/useSessionCache.ts index 35ffe5c78..20987632c 100644 --- a/packages/ui/src/components/instance/shell/useSessionCache.ts +++ b/packages/ui/src/components/instance/shell/useSessionCache.ts @@ -1,10 +1,5 @@ -import { createEffect, createSignal, type Accessor } from "solid-js" -import { messageStoreBus } from "../../../stores/message-v2/bus" -import { clearSessionRenderCache } from "../../message-block" -import { getLogger } from "../../../lib/logger" -import { invalidateSessionMessageLoad } from "../../../stores/session-state" - -const log = getLogger("session") +import { createEffect, createSignal, onCleanup, type Accessor } from "solid-js" +import { setVisibleSessionMemory } from "../../../stores/session-memory" const SESSION_CACHE_LIMIT = 5 @@ -12,6 +7,7 @@ type SessionCacheOptions = { instanceId: Accessor instanceSessions: Accessor> activeSessionId: Accessor + visible: Accessor } type SessionCacheState = { @@ -20,55 +16,14 @@ type SessionCacheState = { export function useSessionCache(options: SessionCacheOptions): SessionCacheState { const [cachedSessionIds, setCachedSessionIds] = createSignal([]) - const [pendingEvictions, setPendingEvictions] = createSignal([]) - - const evictSession = (sessionId: string) => { - if (!sessionId) return - const instanceId = options.instanceId() - log.info("Evicting cached session", { instanceId, sessionId }) - const store = messageStoreBus.getInstance(instanceId) - invalidateSessionMessageLoad(instanceId, sessionId) - store?.clearSession(sessionId, { preserveScroll: true, notify: false }) - clearSessionRenderCache(instanceId, sessionId) - } - - const scheduleEvictions = (ids: string[]) => { - if (!ids.length) return - setPendingEvictions((current) => { - const existing = new Set(current) - const next = [...current] - ids.forEach((id) => { - if (!existing.has(id)) { - next.push(id) - existing.add(id) - } - }) - return next - }) - } - - createEffect(() => { - const pending = pendingEvictions() - if (!pending.length) return - const cached = new Set(cachedSessionIds()) - const remaining: string[] = [] - pending.forEach((id) => { - if (cached.has(id)) { - remaining.push(id) - } else { - evictSession(id) - } - }) - if (remaining.length !== pending.length) { - setPendingEvictions(remaining) - } - }) createEffect(() => { const instanceSessions = options.instanceSessions() const activeId = options.activeSessionId() + const visible = options.visible() setCachedSessionIds((current) => { + if (!visible) return [] const next = current.filter((id) => id !== "info" && instanceSessions.has(id)) const touch = (id: string | null) => { @@ -84,17 +39,19 @@ export function useSessionCache(options: SessionCacheOptions): SessionCacheState touch(activeId) - const trimmed = next.length > SESSION_CACHE_LIMIT ? next.slice(0, SESSION_CACHE_LIMIT) : next - - const trimmedSet = new Set(trimmed) - const removed = current.filter((id) => !trimmedSet.has(id)) - if (removed.length) { - scheduleEvictions(removed) - } - return trimmed + return next.length > SESSION_CACHE_LIMIT ? next.slice(0, SESSION_CACHE_LIMIT) : next }) }) + createEffect(() => { + const instanceId = options.instanceId() + const mountedSessionIds = cachedSessionIds() + for (let index = mountedSessionIds.length - 1; index >= 0; index -= 1) { + setVisibleSessionMemory(instanceId, mountedSessionIds[index]!, true) + } + onCleanup(() => mountedSessionIds.forEach((sessionId) => setVisibleSessionMemory(instanceId, sessionId, false))) + }) + return { cachedSessionIds, } diff --git a/packages/ui/src/components/message-block.tsx b/packages/ui/src/components/message-block.tsx index cc3691ac5..e86ccadd0 100644 --- a/packages/ui/src/components/message-block.tsx +++ b/packages/ui/src/components/message-block.tsx @@ -18,6 +18,7 @@ import { useSpeech } from "../lib/hooks/use-speech" import { createFollowScroll } from "../lib/follow-scroll" import { inferReasoningDurationMs } from "../lib/message-timing" import type { SessionSearchMatch } from "../lib/session-search" +import { iterateTextSearchOccurrences } from "../lib/session-search-matches" import ActionOverflowMenu, { type ActionOverflowMenuItem } from "./action-overflow-menu" import { copyToClipboard } from "../lib/clipboard" import SpeechActionButton from "./speech-action-button" @@ -188,6 +189,7 @@ function clearInstanceCaches(instanceId: string) { } messageStoreBus.onInstanceDestroyed(clearInstanceCaches) +messageStoreBus.onSessionCleared(clearSessionRenderCache) function removeSearchMarks(root: HTMLElement) { const marks = Array.from(root.querySelectorAll("mark.session-search-match")) @@ -207,8 +209,8 @@ function getPartIdForSearchContainer(container: HTMLElement): string | undefined function applySearchMarks(root: HTMLElement, query: string, activeMatch?: SessionSearchMatch | null, scrollActive = false) { removeSearchMarks(root) - const normalizedQuery = query.trim().toLocaleLowerCase() - if (!normalizedQuery) return + const needle = query.trim() + if (!needle) return const containers = Array.from(root.querySelectorAll(".message-text, .tool-call, .message-reasoning-text")) let occurrenceInActivePart = 0 @@ -222,7 +224,7 @@ function applySearchMarks(root: HTMLElement, query: string, activeMatch?: Sessio const parent = node.parentElement if (!parent) return NodeFilter.FILTER_REJECT if (parent.closest("button, input, textarea, select, mark.session-search-match")) return NodeFilter.FILTER_REJECT - if (!node.nodeValue || !node.nodeValue.toLocaleLowerCase().includes(normalizedQuery)) return NodeFilter.FILTER_REJECT + if (!node.nodeValue || iterateTextSearchOccurrences(node.nodeValue, needle).next().done) return NodeFilter.FILTER_REJECT return NodeFilter.FILTER_ACCEPT }, }) @@ -234,25 +236,22 @@ function applySearchMarks(root: HTMLElement, query: string, activeMatch?: Sessio for (const textNode of textNodes) { const original = textNode.nodeValue ?? "" - const lower = original.toLocaleLowerCase() const fragment = document.createDocumentFragment() let cursor = 0 - while (cursor < original.length) { - const index = lower.indexOf(normalizedQuery, cursor) - if (index === -1) break - if (index > cursor) { - fragment.appendChild(document.createTextNode(original.slice(cursor, index))) + for (const occurrence of iterateTextSearchOccurrences(original, needle)) { + if (occurrence.start > cursor) { + fragment.appendChild(document.createTextNode(original.slice(cursor, occurrence.start))) } const mark = document.createElement("mark") const isActive = Boolean(canContainActiveMatch && activeMatch && occurrenceInActivePart === activeMatch.occurrence) mark.className = isActive ? "session-search-match session-search-match-active" : "session-search-match" - mark.textContent = original.slice(index, index + normalizedQuery.length) + mark.textContent = original.slice(occurrence.start, occurrence.end) fragment.appendChild(mark) if (canContainActiveMatch) { if (isActive) activeMark = mark occurrenceInActivePart += 1 } - cursor = index + normalizedQuery.length + cursor = occurrence.end } if (cursor < original.length) { fragment.appendChild(document.createTextNode(original.slice(cursor))) diff --git a/packages/ui/src/components/message-section.tsx b/packages/ui/src/components/message-section.tsx index f80302338..6de8bb6b1 100644 --- a/packages/ui/src/components/message-section.tsx +++ b/packages/ui/src/components/message-section.tsx @@ -7,7 +7,7 @@ import MessageBlock from "./message-block" import { getMessageAnchorId } from "./message-anchors" import MessageTimeline, { buildTimelineSegments, type TimelineSegment } from "./message-timeline" import VirtualFollowList, { type VirtualExplicitBottomPinIntent, type VirtualFollowListApi, type VirtualFollowListState, type VirtualFollowScrollSnapshot } from "./virtual-follow-list" -import { isScrollRestoreGenerationCurrent, isSnapshotAutoFollowing } from "./virtual-follow-behavior" +import { canRestoreMessageScroll, isScrollRestoreGenerationCurrent, isSnapshotAutoFollowing, shouldCancelPendingMessageScrollRestore } from "./virtual-follow-behavior" import { useConfig } from "../stores/preferences" import { getSessionInfo } from "../stores/sessions" import { messageStoreBus } from "../stores/message-v2/bus" @@ -22,10 +22,12 @@ import { partHasRenderableText } from "../types/message" import { buildRecordDisplayData } from "../stores/message-v2/record-display-cache" import { getPartCharCount } from "../lib/token-utils" import { getMessageSelectionActionPosition } from "../lib/message-selection-position" -import { buildSessionSearchMatches } from "../lib/session-search" -import type { SessionSearchMatch } from "../lib/session-search" +import { createSessionSearchPager, findLastSessionSearchPage, retainSessionSearchPage, SESSION_SEARCH_PAGE_SIZE } from "../lib/session-search" +import type { SessionSearchMatch, SessionSearchPager, SessionSearchResult } from "../lib/session-search" import { resolveThinkingExpansionDefault, resolveToolVisibility } from "./tool-call/tool-registry" +import { resolveToolRenderer } from "./tool-call/renderers" import { collectToolDeletionCompanionPartIds, executeBulkDeletionPlan } from "./tool-deletion-companions" +import { getLogger } from "../lib/logger" const MESSAGE_SCROLL_CACHE_SCOPE = "message-stream" const QUOTE_SELECTION_MAX_LENGTH = 2000 @@ -33,11 +35,13 @@ const STREAMING_TEXT_HOLD_TOP_THRESHOLD_PX = 8 const SEARCH_DEBOUNCE_MS = 250 const SEARCH_MIN_CHARS = 3 const OPEN_SESSION_SEARCH_EVENT = "codenomad:open-session-search" +const log = getLogger("session") export interface MessageSectionProps { instanceId: string sessionId: string loading?: boolean + loadComplete?: boolean loadError?: string | null emptyStateVariant?: "messages" | "no-session" onRevert?: (messageId: string) => void @@ -161,7 +165,23 @@ export default function MessageSection(props: MessageSectionProps) { const [searchedQuery, setSearchedQuery] = createSignal("") const [isSearchPending, setIsSearchPending] = createSignal(false) const [searchMatches, setSearchMatches] = createSignal([]) + const [searchTotalMatches, setSearchTotalMatches] = createSignal(0) + const [searchPageOffset, setSearchPageOffset] = createSignal(0) + const [searchHasMore, setSearchHasMore] = createSignal(false) + const [searchPageIndex, setSearchPageIndex] = createSignal(0) const [activeSearchIndex, setActiveSearchIndex] = createSignal(0) + let searchGeneration = 0 + let searchAbortController: AbortController | null = null + let searchPager: SessionSearchPager | null = null + let searchPages = new Map() + let searchPagerNextPageIndex = 0 + let searchLastPageIndex: number | null = null + let searchKnownTotalMatches: number | null = null + let searchRunQuery = "" + let searchRunIncludeThinking = false + let searchRefreshTimeout: number | undefined + let searchRefreshRequested = false + let lastCompletedSearchRevision = -1 let deleteMenuRef: HTMLDivElement | undefined let deleteMenuButtonRef: HTMLButtonElement | undefined let searchInputRef: HTMLInputElement | undefined @@ -721,6 +741,12 @@ export default function MessageSection(props: MessageSectionProps) { lastGoodScrollSnapshots.set(sessionId, snapshot) } + function cancelPendingScrollRestoreFromUser() { + if (!shouldCancelPendingMessageScrollRestore(didRestoreScroll(), restoringScrollSnapshot)) return + scrollRestoreGeneration += 1 + setDidRestoreScroll(true) + } + createEffect( on( () => props.sessionId, @@ -884,11 +910,18 @@ export default function MessageSection(props: MessageSectionProps) { const api = listApi() if (!element || !api) return if (!isActive()) return - if (props.loading) return - if (visibleMessageIds().length === 0) return + const visibleIds = visibleMessageIds() + if (visibleIds.length === 0) return if (didRestoreScroll()) return const snapshot = store().getScrollSnapshot(props.sessionId, MESSAGE_SCROLL_CACHE_SCOPE) + const failedLoadAnchorAvailable = !snapshot?.anchorKey || visibleIds.includes(snapshot.anchorKey) + if (!canRestoreMessageScroll( + Boolean(props.loading), + props.loadComplete, + Boolean(props.loadError), + failedLoadAnchorAvailable, + )) return if (!snapshot) { api.setAutoScroll(true) api.scrollToBottom({ immediate: true }) @@ -945,19 +978,164 @@ export default function MessageSection(props: MessageSectionProps) { } function closeSearch() { + cancelSearchRun() setIsSearchOpen(false) setSearchQuery("") setDebouncedSearchQuery("") setSearchedQuery("") setIsSearchPending(false) setSearchMatches([]) + setSearchTotalMatches(0) + setSearchPageOffset(0) + setSearchHasMore(false) + setSearchPageIndex(0) setActiveSearchIndex(0) } + function cancelSearchRun() { + searchGeneration += 1 + searchAbortController?.abort() + searchAbortController = null + searchPager = null + searchPages.clear() + searchPagerNextPageIndex = 0 + searchLastPageIndex = null + searchKnownTotalMatches = null + searchRunQuery = "" + if (searchRefreshTimeout !== undefined) window.clearTimeout(searchRefreshTimeout) + searchRefreshTimeout = undefined + searchRefreshRequested = false + } + + function scheduleSearchRefresh() { + if (searchRefreshTimeout !== undefined) window.clearTimeout(searchRefreshTimeout) + searchRefreshTimeout = window.setTimeout(() => { + searchRefreshTimeout = undefined + if (isSearchPending()) { + searchRefreshRequested = true + return + } + const query = debouncedSearchQuery() + if (query.trim().length < SEARCH_MIN_CHARS) return + void startSearch(query, Boolean(preferences().showThinkingBlocks)) + }, SEARCH_DEBOUNCE_MS) + } + + async function startSearch(query: string, includeThinking: boolean) { + cancelSearchRun() + const generation = searchGeneration + const controller = new AbortController() + searchAbortController = controller + searchRunQuery = query + searchRunIncludeThinking = includeThinking + resetSearchPager() + searchPages.clear() + await loadSearchPage(query, generation, 0, 0) + } + + function resetSearchPager() { + const controller = searchAbortController + if (!controller) return + searchPager = createSessionSearchPager({ + store: store(), + sessionId: props.sessionId, + query: searchRunQuery, + includeThinking: searchRunIncludeThinking, + resolveToolSearchText: (context) => resolveToolRenderer(context.toolName).getSearchText?.(context), + limit: SESSION_SEARCH_PAGE_SIZE, + signal: controller.signal, + }) + searchPagerNextPageIndex = 0 + } + + async function readSearchPage(pageIndex: number): Promise { + const cached = searchPages.get(pageIndex) + if (cached) { + retainSessionSearchPage(searchPages, pageIndex, cached) + return cached + } + if (pageIndex < searchPagerNextPageIndex) resetSearchPager() + if (!searchPager) throw new Error("Session search pager unavailable") + + while (searchPagerNextPageIndex <= pageIndex) { + const currentIndex = searchPagerNextPageIndex + const result = await searchPager.nextPage() + retainSessionSearchPage(searchPages, currentIndex, result) + searchPagerNextPageIndex += 1 + if (result.totalMatches !== null) searchKnownTotalMatches = result.totalMatches + if (!result.hasMore) searchLastPageIndex = currentIndex + if (!result.hasMore && currentIndex < pageIndex) throw new Error("Session search page unavailable") + } + return searchPages.get(pageIndex)! + } + + async function loadSearchPage(query: string, generation: number, requestedPage: number | "last", targetIndex: number) { + const startedRevision = store().getSessionRevision(props.sessionId) + const controller = searchAbortController + if (!controller || !searchPager) return + setIsSearchPending(true) + + try { + const resolved = requestedPage === "last" + ? searchLastPageIndex === null + ? await findLastSessionSearchPage(readSearchPage) + : { pageIndex: searchLastPageIndex, result: await readSearchPage(searchLastPageIndex) } + : { pageIndex: requestedPage, result: await readSearchPage(requestedPage) } + if (generation !== searchGeneration || controller.signal.aborted) return + const { pageIndex, result } = resolved + setSearchMatches(result.matches) + setSearchTotalMatches(result.totalMatches ?? searchKnownTotalMatches) + setSearchPageOffset(result.offset) + setSearchHasMore(result.hasMore) + setSearchPageIndex(pageIndex) + setSearchedQuery(query) + setActiveSearchIndex(Math.min(Math.max(targetIndex, 0), Math.max(result.matches.length - 1, 0))) + lastCompletedSearchRevision = startedRevision + } catch (error) { + if (generation !== searchGeneration || controller.signal.aborted || (error as Error)?.name === "AbortError") return + log.error("Failed to search session", error) + lastCompletedSearchRevision = store().getSessionRevision(props.sessionId) + searchRefreshRequested = false + setSearchMatches([]) + setSearchTotalMatches(0) + setSearchPageOffset(0) + setSearchHasMore(false) + setSearchPageIndex(0) + setSearchedQuery(query) + setActiveSearchIndex(0) + } finally { + if (generation === searchGeneration) { + setIsSearchPending(false) + if (searchRefreshRequested || store().getSessionRevision(props.sessionId) !== lastCompletedSearchRevision) { + searchRefreshRequested = false + scheduleSearchRefresh() + } + } + } + } + function moveSearchMatch(direction: 1 | -1) { - const count = searchMatches().length - if (count === 0) return - setActiveSearchIndex((index) => (index + direction + count) % count) + const matches = searchMatches() + if (matches.length === 0 || isSearchPending()) return + const target = activeSearchIndex() + direction + if (target >= 0 && target < matches.length) { + setActiveSearchIndex(target) + return + } + const page = searchPageIndex() + if (direction === 1 && searchHasMore()) { + void loadSearchPage(searchedQuery(), searchGeneration, page + 1, 0) + return + } + if (direction === -1 && page > 0) { + void loadSearchPage(searchedQuery(), searchGeneration, page - 1, Number.MAX_SAFE_INTEGER) + return + } + if (direction === 1) { + void loadSearchPage(searchedQuery(), searchGeneration, 0, 0) + return + } + void loadSearchPage(searchedQuery(), searchGeneration, "last", Number.MAX_SAFE_INTEGER) } function isSelectionWithinStream(range: Range | null) { @@ -1064,7 +1242,7 @@ export default function MessageSection(props: MessageSectionProps) { // to prevent O(n) per-element reactive subscriptions. The effect // only needs to re-run when `messageIds` (memo) changes. untrack(() => { - if (loading) { + if (loading && ids.length === 0) { handleClearTimelineSelection() previousTimelineIds = [] setTimelineSegments([]) @@ -1131,6 +1309,14 @@ export default function MessageSection(props: MessageSectionProps) { } } + const prefixAdded = previousTimelineIds.length > 0 && ids.length > previousTimelineIds.length && + previousTimelineIds.every((id, index) => ids[ids.length - previousTimelineIds.length + index] === id) + if (prefixAdded) { + seedTimeline() + previousTimelineIds = [...ids] + return + } + const newIds: string[] = [] ids.forEach((id) => { if (!seenTimelineMessageIds.has(id)) { @@ -1270,12 +1456,18 @@ export default function MessageSection(props: MessageSectionProps) { createEffect(() => { const query = searchQuery() + cancelSearchRun() + lastCompletedSearchRevision = -1 if (query.trim().length < SEARCH_MIN_CHARS) { setDebouncedSearchQuery("") setActiveSearchIndex(0) setSearchedQuery("") setIsSearchPending(false) setSearchMatches([]) + setSearchTotalMatches(0) + setSearchPageOffset(0) + setSearchHasMore(false) + setSearchPageIndex(0) return } setIsSearchPending(true) @@ -1286,7 +1478,6 @@ export default function MessageSection(props: MessageSectionProps) { }) createEffect(() => { - sessionRevision() const query = debouncedSearchQuery() const includeThinking = Boolean(preferences().showThinkingBlocks) if (query.trim().length < SEARCH_MIN_CHARS) { @@ -1295,18 +1486,19 @@ export default function MessageSection(props: MessageSectionProps) { setIsSearchPending(true) const frame = requestAnimationFrame(() => { - const matches = buildSessionSearchMatches({ - store: store(), - sessionId: props.sessionId, - query, - includeThinking, - }) - setSearchMatches(matches) - setSearchedQuery(query) - setActiveSearchIndex(0) - setIsSearchPending(false) + void startSearch(query, includeThinking) + }) + onCleanup(() => { + cancelAnimationFrame(frame) + cancelSearchRun() }) - onCleanup(() => cancelAnimationFrame(frame)) + }) + + createEffect(() => { + const revision = sessionRevision() + const query = debouncedSearchQuery() + if (query.trim().length < SEARCH_MIN_CHARS || lastCompletedSearchRevision < 0) return + if (revision !== lastCompletedSearchRevision) scheduleSearchRefresh() }) createEffect(() => { @@ -1449,6 +1641,7 @@ export default function MessageSection(props: MessageSectionProps) { clearQuoteSelection() persistMessageScrollSnapshot() }} + onUserScrollIntent={cancelPendingScrollRestoreFromUser} onMouseUp={() => handleStreamMouseUp()} onClick={(e) => { if (selectedTimelineIds().size === 0) return @@ -1521,7 +1714,10 @@ export default function MessageSection(props: MessageSectionProps) { + ) } @@ -488,6 +517,7 @@ function ToolCallDetails(props: { active={props.isPermissionActive} submitting={permissionSubmitting} error={permissionError} + onApprovalBlockedChange={setPermissionApprovalBlocked} renderDiff={renderDiffContent} fallbackSessionId={() => props.sessionId} onRespond={(permission, sessionId, response, message) => void handlePermissionResponse(permission, response, message)} @@ -522,6 +552,18 @@ function ToolCallDetails(props: { await copyToClipboard(text) } + const copyToolInput = async (event: MouseEvent) => { + event.preventDefault() + event.stopPropagation() + const input = props.toolInput() + if (!input) return + try { + await copyToClipboard(stringifyYaml(input)) + } catch { + await copyToClipboard(JSON.stringify(input, null, 2)) + } + } + const outputWrapTitle = () => props.outputWrapEnabled() ? props.t("toolCall.diff.disableWordWrap") @@ -535,6 +577,7 @@ function ToolCallDetails(props: { copyText?: () => string | null | undefined copyTitle?: () => string copyAriaLabel?: () => string + onCopy?: (event: MouseEvent) => void actions?: () => JSXElement wrapToggle?: () => boolean | undefined }) => ( @@ -551,18 +594,16 @@ function ToolCallDetails(props: { {(actions) => {actions()}} - - {(copyText) => ( - - )} + + @@ -632,7 +673,14 @@ function ToolCallDetails(props: {
+ + {(actions) =>
{actions()}
} +
+ {renderToolOutputBody()} + + )} >
@@ -644,6 +692,7 @@ function ToolCallDetails(props: { expanded: props.inputSectionExpanded, onToggle: props.toggleInputSection, copyText: () => toolInputDisplay()?.copyText, + onCopy: toolInputDisplay()?.copyText === null ? (event) => void copyToolInput(event) : undefined, copyTitle: () => props.t("toolCall.io.copyInputTitle"), copyAriaLabel: () => props.t("toolCall.io.copyInputAriaLabel"), }) @@ -668,6 +717,7 @@ function ToolCallDetails(props: { expanded: props.outputSectionExpanded, onToggle: props.toggleOutputSection, copyText: () => outputChrome().copyText, + onCopy: outputChrome().getCopyText ? (event) => void copyIoText(event, resolveOutputCopyText()) : undefined, copyTitle: () => props.t("toolCall.io.copyOutputTitle"), copyAriaLabel: () => props.t("toolCall.io.copyOutputAriaLabel"), actions: () => outputChrome().actions, @@ -824,10 +874,9 @@ export default function ToolCall(props: ToolCallProps) { if (override !== undefined) return override return diagnosticsDefaultExpanded() } - const diagnosticsEntries = createMemo(() => { + const diagnosticsView = createMemo(() => { const state = toolState() - if (!state) return [] - return extractDiagnostics(state) + return extractDiagnosticsView(state) }) const toggleInputSection = () => { @@ -887,6 +936,7 @@ export default function ToolCall(props: ToolCallProps) { toolName, instanceId: props.instanceId, sessionId: props.sessionId, + visibilitySessionId: props.visibilitySessionId ?? props.sessionId, t, messageVersion: () => props.messageVersion, partVersion: () => props.partVersion, @@ -938,6 +988,7 @@ export default function ToolCall(props: ToolCallProps) { } const toolTypeLabel = createMemo(() => toolName()) + const renderedToolTypeLabel = createMemo(() => limitToolTitleForRender(toolTypeLabel())) const headerTitleDetail = createMemo(() => { const rawTitle = renderToolTitle().trim() @@ -959,9 +1010,10 @@ export default function ToolCall(props: ToolCallProps) { const detail = headerTitleDetail() return [typeLabel, detail].filter(Boolean).join(" ") }) + const renderedHeaderTitleDetail = createMemo(() => limitToolTitleForRender(headerTitleDetail())) - const headerCopyText = createMemo(() => headerOutputChrome().copyText || "") - const canCopyHeaderOutput = () => headerCopyText().length > 0 + const headerCopyText = () => headerOutputChrome().copyText || headerOutputChrome().getCopyText?.() || "" + const canCopyHeaderOutput = () => Boolean(headerOutputChrome().copyText || headerOutputChrome().getCopyText) const canToggleOutputWrap = () => Boolean(headerOutputChrome().wrapToggle) const outputWrapTitle = () => outputWrapEnabled() @@ -1073,8 +1125,8 @@ export default function ToolCall(props: ToolCallProps) { > - {toolTypeLabel()} - + {renderedToolTypeLabel()} + {(detail) => {detail()}} @@ -1145,6 +1197,7 @@ export default function ToolCall(props: ToolCallProps) { toolCallIdentifier={toolCallIdentifier} instanceId={props.instanceId} sessionId={props.sessionId} + visibilitySessionId={props.visibilitySessionId ?? props.sessionId} messageId={props.messageId} messageVersion={props.messageVersion} partVersion={props.partVersion} @@ -1173,17 +1226,17 @@ export default function ToolCall(props: ToolCallProps) { /> - + {renderDiagnosticsSection( t, - diagnosticsEntries(), + diagnosticsView(), diagnosticsExpanded(), () => setDiagnosticsOverride((prev) => { const current = prev === undefined ? diagnosticsDefaultExpanded() : prev return !current }), - diagnosticFileName(diagnosticsEntries()), + diagnosticFileName(diagnosticsView().entries), )}
diff --git a/packages/ui/src/components/tool-call/diagnostic-selection.ts b/packages/ui/src/components/tool-call/diagnostic-selection.ts new file mode 100644 index 000000000..304563099 --- /dev/null +++ b/packages/ui/src/components/tool-call/diagnostic-selection.ts @@ -0,0 +1,12 @@ +export function selectSeverityBounded(values: readonly T[], rank: (value: T) => number | undefined, limit: number): T[] { + const buckets: T[][] = [[], [], []] + const scanLimit = Math.max(limit, limit * 100) + for (let index = 0; index < values.length && index < scanLimit; index += 1) { + const value = values[index] + const valueRank = rank(value) + if (valueRank === undefined) continue + const bucket = buckets[Math.max(0, Math.min(2, valueRank))]! + if (bucket.length < limit) bucket.push(value) + } + return buckets.flat().slice(0, limit) +} diff --git a/packages/ui/src/components/tool-call/diagnostics-section.tsx b/packages/ui/src/components/tool-call/diagnostics-section.tsx index b4b9057f7..e29b5cf7f 100644 --- a/packages/ui/src/components/tool-call/diagnostics-section.tsx +++ b/packages/ui/src/components/tool-call/diagnostics-section.tsx @@ -1,14 +1,40 @@ import { For, Show } from "solid-js" -import type { DiagnosticEntry } from "./diagnostics" +import { Copy } from "lucide-solid" +import { hasDiagnosticMessages, type DiagnosticsMap, type DiagnosticsView } from "./diagnostics" +import { copyToClipboard } from "../../lib/clipboard" +import { formatUnknownForCopy } from "./utils" + +export function DiagnosticsPayloadAccess(props: { + diagnostics: DiagnosticsMap + truncated: boolean + t: (key: string, params?: Record) => string +}) { + return ( +
+ + {props.t(props.truncated ? "toolCall.output.truncated" : "toolCall.diagnostics.title")} + + +
+ ) +} export function renderDiagnosticsSection( t: (key: string, params?: Record) => string, - entries: DiagnosticEntry[], + view: DiagnosticsView, expanded: boolean, toggle: () => void, fileLabel: string, ) { - if (entries.length === 0) return null + if (!hasDiagnosticMessages(view.diagnostics)) return null return (
+
- + {(entry) => (
diff --git a/packages/ui/src/components/tool-call/diagnostics.test.ts b/packages/ui/src/components/tool-call/diagnostics.test.ts new file mode 100644 index 000000000..b6d36ef08 --- /dev/null +++ b/packages/ui/src/components/tool-call/diagnostics.test.ts @@ -0,0 +1,44 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { selectSeverityBounded } from "./diagnostic-selection.ts" +import { buildDiagnosticView, extractDiagnosticsView } from "./diagnostics.ts" +import { formatUnknownForCopy } from "./utils.ts" + +test("diagnostic bounds retain errors that follow informational entries", () => { + const diagnostics = [ + ...Array.from({ length: 100 }, (_, index) => ({ message: `info ${index}`, severity: 3 })), + { message: "the error", severity: 1 }, + ] + const entries = selectSeverityBounded(diagnostics, (entry) => entry.severity === 1 ? 0 : 2, 100) + assert.equal(entries.length, 100) + assert.equal(entries[0]?.message, "the error") +}) + +test("bounded diagnostics expose truncation and retain the complete copy payload", () => { + const longMessage = `${"x".repeat(2_000)}COPY_TAIL` + const diagnostics = { + "src/file.ts": [ + { message: longMessage, severity: 1 }, + ...Array.from({ length: 100 }, (_, index) => ({ message: `warning ${index}`, severity: 2 })), + ], + } + const view = buildDiagnosticView(diagnostics, ["src/file.ts"]) + + assert.equal(view.entries.length, 100) + assert.equal(view.entries[0]?.message.includes("COPY_TAIL"), false) + assert.equal(view.truncated, true) + assert.equal(formatUnknownForCopy(view.diagnostics)?.text.includes("COPY_TAIL"), true) +}) + +test("unmatched diagnostic paths still expose the complete payload", () => { + const view = extractDiagnosticsView({ + status: "completed", + input: { filePath: "src/requested.ts" }, + metadata: { diagnostics: { "src/reported.ts": [{ message: "SEARCH_MATCH" }] } }, + output: "", + } as any) + + assert.equal(view.entries.length, 0) + assert.equal(view.truncated, true) + assert.equal(formatUnknownForCopy(view.diagnostics)?.text.includes("SEARCH_MATCH"), true) +}) diff --git a/packages/ui/src/components/tool-call/diagnostics.ts b/packages/ui/src/components/tool-call/diagnostics.ts index 8651df697..d91248f86 100644 --- a/packages/ui/src/components/tool-call/diagnostics.ts +++ b/packages/ui/src/components/tool-call/diagnostics.ts @@ -1,6 +1,7 @@ import type { ToolState } from "@opencode-ai/sdk/v2" import { getRelativePath, isToolStateCompleted, isToolStateError, isToolStateRunning } from "./utils" import { tGlobal } from "../../lib/i18n" +import { selectSeverityBounded } from "./diagnostic-selection" interface LspRangePosition { line?: number @@ -26,12 +27,28 @@ export interface DiagnosticEntry { label: string icon: string message: string + messageTruncated: boolean filePath: string displayPath: string line: number column: number } +export interface DiagnosticsView { + diagnostics: DiagnosticsMap + entries: DiagnosticEntry[] + key?: string + truncated: boolean +} + +function diagnosticListHasMessages(list: unknown): boolean { + return Array.isArray(list) && list.some((entry) => typeof entry?.message === "string") +} + +export function hasDiagnosticMessages(diagnostics: DiagnosticsMap): boolean { + return Object.values(diagnostics).some(diagnosticListHasMessages) +} + export function normalizeDiagnosticPath(path: string) { return path.replace(/\\/g, "/") } @@ -48,26 +65,33 @@ function getSeverityMeta(tone: DiagnosticEntry["tone"]) { return { label: tGlobal("toolCall.diagnostics.severity.info.short"), icon: "i", rank: 2 } } -export function extractDiagnostics(state: ToolState | undefined): DiagnosticEntry[] { - if (!state) return [] +export function extractDiagnosticsView(state: ToolState | undefined): DiagnosticsView { + if (!state) return buildDiagnosticView({}, []) const supportsMetadata = isToolStateRunning(state) || isToolStateCompleted(state) || isToolStateError(state) - if (!supportsMetadata) return [] + if (!supportsMetadata) return buildDiagnosticView({}, []) const metadata = (state.metadata || {}) as Record const input = (state.input || {}) as Record const diagnosticsMap = metadata?.diagnostics as DiagnosticsMap | undefined - if (!diagnosticsMap) return [] + if (!diagnosticsMap) return buildDiagnosticView({}, []) - return buildDiagnosticEntries(diagnosticsMap, [input.filePath, metadata.filePath, metadata.filepath, input.path].map((value) => + const view = buildDiagnosticView(diagnosticsMap, [input.filePath, metadata.filePath, metadata.filepath, input.path].map((value) => typeof value === "string" ? value : undefined, )) + if ((!view.key && hasDiagnosticMessages(diagnosticsMap)) + || (view.key && Object.entries(diagnosticsMap).some(([key, list]) => key !== view.key && diagnosticListHasMessages(list)))) { + return { ...view, truncated: true } + } + return view } -export function resolveDiagnosticsKey(diagnostics: DiagnosticsMap, preferredPaths: Array): string | undefined { - if (Object.keys(diagnostics).length === 0) return undefined +export function extractDiagnostics(state: ToolState | undefined): DiagnosticEntry[] { + return extractDiagnosticsView(state).entries +} +export function resolveDiagnosticsKey(diagnostics: DiagnosticsMap, preferredPaths: Array): string | undefined { const normalizedPreferred = preferredPaths - .filter((value): value is string => typeof value === "string" && value.length > 0) + .filter((value): value is string => typeof value === "string" && value.length > 0 && value.length <= 4_096) .map((value) => normalizeDiagnosticPath(value)) if (normalizedPreferred.length === 0) return undefined @@ -76,7 +100,16 @@ export function resolveDiagnosticsKey(diagnostics: DiagnosticsMap, preferredPath if (diagnostics[preferred]) return preferred } - const keys = Object.keys(diagnostics) + const keys: string[] = [] + let scannedKeys = 0 + for (const key in diagnostics) { + if (!Object.prototype.hasOwnProperty.call(diagnostics, key)) continue + scannedKeys += 1 + if (scannedKeys > 10_000) break + if (key.length > 4_096) continue + keys.push(key) + } + if (keys.length === 0) return undefined for (const preferred of normalizedPreferred) { const direct = keys.find((key) => normalizeDiagnosticPath(key) === preferred) @@ -94,29 +127,39 @@ export function resolveDiagnosticsKey(diagnostics: DiagnosticsMap, preferredPath return undefined } -export function buildDiagnosticEntries(diagnostics: DiagnosticsMap, preferredPaths: Array): DiagnosticEntry[] { +export function buildDiagnosticView(diagnostics: DiagnosticsMap, preferredPaths: Array): DiagnosticsView { const key = resolveDiagnosticsKey(diagnostics, preferredPaths) - if (!key) return [] + if (!key) return { diagnostics, entries: [], truncated: false } const list = diagnostics[key] - if (!Array.isArray(list) || list.length === 0) return [] + if (!Array.isArray(list) || list.length === 0) return { diagnostics, entries: [], key, truncated: false } - const entries: DiagnosticEntry[] = [] + const limit = 100 const normalizedPath = normalizeDiagnosticPath(key) - for (let index = 0; index < list.length; index++) { - const diagnostic = list[index] + const selected = selectSeverityBounded( + list, + (diagnostic) => diagnostic && typeof diagnostic.message === "string" + ? getSeverityMeta(determineSeverityTone(diagnostic.severity)).rank + : undefined, + limit, + ) + const entries: DiagnosticEntry[] = [] + for (const diagnostic of selected) { + const index = entries.length if (!diagnostic || typeof diagnostic.message !== "string") continue const tone = determineSeverityTone(typeof diagnostic.severity === "number" ? diagnostic.severity : undefined) const severityMeta = getSeverityMeta(tone) const line = typeof diagnostic.range?.start?.line === "number" ? diagnostic.range.start.line + 1 : 0 const column = typeof diagnostic.range?.start?.character === "number" ? diagnostic.range.start.character + 1 : 0 + const messageTruncated = diagnostic.message.length > 2_000 entries.push({ - id: `${normalizedPath}-${index}-${diagnostic.message}`, + id: String(index), severity: severityMeta.rank, tone, label: severityMeta.label, icon: severityMeta.icon, - message: diagnostic.message, + message: diagnostic.message.slice(0, 2_000), + messageTruncated, filePath: normalizedPath, displayPath: getRelativePath(normalizedPath), line, @@ -124,7 +167,16 @@ export function buildDiagnosticEntries(diagnostics: DiagnosticsMap, preferredPat }) } - return entries.sort((a, b) => a.severity - b.severity) + return { + diagnostics, + entries, + key, + truncated: list.length > entries.length || entries.some((entry) => entry.messageTruncated), + } +} + +export function buildDiagnosticEntries(diagnostics: DiagnosticsMap, preferredPaths: Array): DiagnosticEntry[] { + return buildDiagnosticView(diagnostics, preferredPaths).entries } export function diagnosticFileName(entries: DiagnosticEntry[]) { diff --git a/packages/ui/src/components/tool-call/diff-render.tsx b/packages/ui/src/components/tool-call/diff-render.tsx index 563fe9d6e..22d14fefb 100644 --- a/packages/ui/src/components/tool-call/diff-render.tsx +++ b/packages/ui/src/components/tool-call/diff-render.tsx @@ -5,7 +5,7 @@ import { AlignJustify, Copy, Split, WrapText } from "lucide-solid" import type { RenderCache } from "../../types/message" import type { DiffViewMode } from "../../stores/preferences" import type { DiffPayload, DiffRenderOptions, ToolScrollHelpers } from "./types" -import { getRelativePath } from "./utils" +import { getRelativePath, limitToolOutputForRender, TOOL_OUTPUT_RENDER_CHARACTER_LIMIT } from "./utils" import { getCacheEntry } from "../../lib/global-cache" import { copyToClipboard } from "../../lib/clipboard" @@ -65,6 +65,8 @@ export function createDiffContentRenderer(params: { } function renderDiffContent(payload: DiffPayload, options?: DiffRenderOptions): JSXElement | null { + const renderedDiffText = limitToolOutputForRender(payload.diffText) + const diffWasTruncated = payload.diffText.length > TOOL_OUTPUT_RENDER_CHARACTER_LIMIT const relativePath = payload.filePath ? getRelativePath(payload.filePath) : "" const toolbarLabel = options?.label || (relativePath ? params.t("toolCall.diff.label.withPath", { path: relativePath }) @@ -100,7 +102,7 @@ export function createDiffContentRenderer(params: { const cached = getCacheEntry(cacheEntryParams) if ( cached - && cached.text === payload.diffText + && cached.text === renderedDiffText && cached.theme === themeKey && cached.mode === currentMode() && cached.wrap === currentWrap() @@ -127,6 +129,9 @@ export function createDiffContentRenderer(params: { ? params.t("toolCall.diff.disableWordWrap") : params.t("toolCall.diff.enableWordWrap") const copyPatchTitle = () => params.t("toolCall.diff.copyPatch") + const copyFullDiff = async () => { + if (await copyToClipboard(payload.diffText)) options?.onFullDiffAccess?.() + } const handleDiffRendered = () => { params.handleScrollRendered() @@ -146,7 +151,7 @@ export function createDiffContentRenderer(params: {
- {cachedHtml() ? ( + {diffWasTruncated ? ( +
{renderedDiffText}
+ ) : cachedHtml() ? ( ) : ( - {payload.diffText}}> + {renderedDiffText}}> error: Accessor onRespond: (permission: PermissionRequest, sessionId: string, response: PermissionResponse, message?: string) => void | Promise + onApprovalBlockedChange?: (blocked: boolean) => void renderDiff: (payload: DiffPayload, options?: DiffRenderOptions) => JSXElement | null fallbackSessionId: Accessor } @@ -22,10 +23,12 @@ export type PermissionToolBlockProps = { export function PermissionToolBlock(props: PermissionToolBlockProps) { const { t } = useI18n() const [rejectReason, setRejectReason] = createSignal("") + const [fullDiffAccessed, setFullDiffAccessed] = createSignal(false) createEffect(() => { props.permission()?.id setRejectReason("") + setFullDiffAccessed(false) }) const diffPayload = () => { @@ -57,6 +60,10 @@ export function PermissionToolBlock(props: PermissionToolBlockProps) { respond("reject", rejectReason().trim() || undefined) } + const approvalBlocked = () => isPermissionDiffTooLarge(diffPayload()?.diffText) && !fullDiffAccessed() + + createEffect(() => props.onApprovalBlockedChange?.(approvalBlocked())) + return ( {(permission) => ( @@ -77,6 +84,7 @@ export function PermissionToolBlock(props: PermissionToolBlockProps) { {props.renderDiff(payload(), { variant: "permission-diff", disableScrollTracking: true, + onFullDiffAccess: () => setFullDiffAccessed(true), label: payload().filePath ? t("toolCall.permission.requestedDiff.withPath", { path: getRelativePath(payload().filePath || "") }) : t("toolCall.permission.requestedDiff.label"), @@ -84,6 +92,9 @@ export function PermissionToolBlock(props: PermissionToolBlockProps) {
)}
+ +
{t("toolCall.output.truncated")}
+

{t("toolCall.permission.queuedText")}

@@ -105,7 +116,7 @@ export function PermissionToolBlock(props: PermissionToolBlockProps) { + ), + } + } return { language: "diff", copyText: diffs.join("\n"), suppressInnerHeader: false } } @@ -84,19 +106,32 @@ export const applyPatchRenderer: ToolRenderer = { if (!state || state.status === "pending") return null const payload = readToolStatePayload(state) - const files = createMemo(() => { + const allFiles = createMemo(() => { const list = (payload.metadata as any).files return Array.isArray(list) ? (list as ApplyPatchFile[]) : [] }) + const files = createMemo(() => allFiles().slice(0, APPLY_PATCH_FILE_RENDER_LIMIT)) const diagnosticsMap = createMemo(() => { const value = (payload.metadata as any).diagnostics return value && typeof value === "object" ? (value as DiagnosticsMap) : {} }) + const diagnosticViews = createMemo(() => files().map((file) => + buildDiagnosticView(diagnosticsMap(), [file.filePath, file.relativePath]), + )) + const hasDiagnostics = createMemo(() => hasDiagnosticMessages(diagnosticsMap())) + const diagnosticsTruncated = createMemo(() => { + const views = diagnosticViews() + const renderedKeys = new Set(views.map((view) => view.key).filter(Boolean)) + return views.some((view) => view.truncated) + || Object.entries(diagnosticsMap()).some(([key, list]) => + !renderedKeys.has(key) && Array.isArray(list) && list.some((entry) => typeof entry?.message === "string"), + ) + }) if (files().length === 0) { const fallback = isToolStateCompleted(state) && typeof state.output === "string" ? state.output : null if (!fallback) return null - return renderMarkdown({ content: fallback, size: "large", disableHighlight: state.status === "running" }) + return renderMarkdown({ content: limitToolOutputForRender(fallback), size: "large", disableHighlight: state.status === "running" }) } return ( @@ -106,7 +141,7 @@ export const applyPatchRenderer: ToolRenderer = { const labelBase = file.relativePath || file.filePath || t("toolCall.applyPatch.fileFallback", { number: index() + 1 }) const diffText = typeof file.diff === "string" ? file.diff : typeof file.patch === "string" ? file.patch : "" const filePath = typeof file.filePath === "string" ? file.filePath : file.relativePath - const entries = createMemo(() => buildDiagnosticEntries(diagnosticsMap(), [file.filePath, file.relativePath])) + const entries = createMemo(() => diagnosticViews()[index()]?.entries ?? []) return (
@@ -124,6 +159,16 @@ export const applyPatchRenderer: ToolRenderer = { ) }} + APPLY_PATCH_FILE_RENDER_LIMIT}> +
{t("toolCall.output.truncated")}
+
+ +
+
+ +
+
+
) }, diff --git a/packages/ui/src/components/tool-call/renderers/bash.tsx b/packages/ui/src/components/tool-call/renderers/bash.tsx index 36275163d..d5451bc2c 100644 --- a/packages/ui/src/components/tool-call/renderers/bash.tsx +++ b/packages/ui/src/components/tool-call/renderers/bash.tsx @@ -1,7 +1,7 @@ import { Show, createEffect, createMemo, onCleanup, type Accessor } from "solid-js" import type { ToolState } from "@opencode-ai/sdk/v2" import type { ToolRenderer, ToolScrollHelpers } from "../types" -import { ensureMarkdownContent, formatUnknown, getToolName, isToolStateCompleted, isToolStateError, isToolStateRunning, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, formatUnknownForCopy, formatUnknownForRender, getToolName, isToolStateCompleted, isToolStateError, isToolStateRunning, limitToolOutputForRender, readToolStatePayload } from "../utils" import { tGlobal } from "../../../lib/i18n" import { createStableAnsiStreamUpdater } from "../ansi-render" import { ansiToHtml, hasAnsi } from "../../../lib/ansi" @@ -99,7 +99,7 @@ function getBashCopyText(state: ToolState | undefined): string { const { input, metadata } = readToolStatePayload(state) const command = typeof input.command === "string" && input.command.length > 0 ? `$ ${input.command}` : "" - const outputResult = formatUnknown( + const outputResult = formatUnknownForCopy( isToolStateCompleted(state) ? state.output : (isToolStateRunning(state) || isToolStateError(state)) && metadata.output @@ -122,8 +122,8 @@ function BashToolBody(props: { if (!current || current.status === "pending") return "" const { input, metadata } = readToolStatePayload(current) - const command = typeof input.command === "string" && input.command.length > 0 ? `$ ${input.command}` : "" - const outputResult = formatUnknown( + const command = typeof input.command === "string" && input.command.length > 0 ? limitToolOutputForRender(`$ ${input.command}`) : "" + const outputResult = formatUnknownForRender( isToolStateCompleted(current) ? current.output : (isToolStateRunning(current) || isToolStateError(current)) && metadata.output @@ -132,10 +132,11 @@ function BashToolBody(props: { ) return [command, outputResult?.text].filter(Boolean).join("\n") }) + const renderedContent = createMemo(() => limitToolOutputForRender(joinedContent())) const finalMarkdown = createMemo(() => { const current = state() - const content = joinedContent() + const content = renderedContent() if (!current || current.status === "pending" || current.status === "running" || content.length === 0) { return null } @@ -147,7 +148,7 @@ function BashToolBody(props: { const finalAnsiHtml = createMemo(() => { const current = state() - const content = joinedContent() + const content = renderedContent() if (!current || current.status === "pending" || current.status === "running" || content.length === 0) { return null } @@ -169,7 +170,7 @@ function BashToolBody(props: { } > - + ) @@ -196,9 +197,9 @@ export const bashRenderer: ToolRenderer = { return `${baseTitle} · ${tGlobal("toolCall.renderer.bash.title.timeout", { timeout: timeoutLabel })}` }, getOutputChrome({ toolState }) { - const text = getBashCopyText(toolState()) - if (!text) return undefined - return { language: "bash", copyText: text, wrapToggle: true, suppressInnerHeader: true } + const state = toolState() + if (!state || state.status === "pending") return undefined + return { language: "bash", getCopyText: () => getBashCopyText(state), wrapToggle: true, suppressInnerHeader: true } }, renderBody({ toolState, renderMarkdown, scrollHelpers, onContentRendered }) { return diff --git a/packages/ui/src/components/tool-call/renderers/default.tsx b/packages/ui/src/components/tool-call/renderers/default.tsx index f19682675..e0835f606 100644 --- a/packages/ui/src/components/tool-call/renderers/default.tsx +++ b/packages/ui/src/components/tool-call/renderers/default.tsx @@ -1,5 +1,5 @@ import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, formatUnknown, isToolStateCompleted, isToolStateError, isToolStateRunning, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, formatUnknownForCopy, formatUnknownForRender, isToolStateCompleted, isToolStateError, isToolStateRunning, readToolStatePayload } from "../utils" import { getDefaultToolSearchText } from "../search-text" export const defaultRenderer: ToolRenderer = { @@ -16,12 +16,12 @@ export const defaultRenderer: ToolRenderer = { ? metadata.output : metadata.diff ?? metadata.preview ?? input.content - const result = formatUnknown(primaryOutput) - if (!result) return undefined + if (primaryOutput === undefined || primaryOutput === null) return undefined + const rendered = formatUnknownForRender(primaryOutput) return { - language: result.language ?? "text", - copyText: result.text, + language: rendered?.language ?? "text", + getCopyText: () => formatUnknownForCopy(primaryOutput)?.text ?? null, wrapToggle: true, suppressInnerHeader: true, } @@ -37,7 +37,7 @@ export const defaultRenderer: ToolRenderer = { ? metadata.output : metadata.diff ?? metadata.preview ?? input.content - const result = formatUnknown(primaryOutput) + const result = formatUnknownForRender(primaryOutput) if (!result) return null const content = ensureMarkdownContent(result.text, result.language, true) diff --git a/packages/ui/src/components/tool-call/renderers/edit.tsx b/packages/ui/src/components/tool-call/renderers/edit.tsx index b4f84fb1d..bde7cb018 100644 --- a/packages/ui/src/components/tool-call/renderers/edit.tsx +++ b/packages/ui/src/components/tool-call/renderers/edit.tsx @@ -1,5 +1,5 @@ import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, extractDiffPayload, getRelativePath, getToolName, isToolStateCompleted, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, extractDiffPayload, getRelativePath, getToolName, isToolStateCompleted, limitToolOutputForRender, readToolStatePayload } from "../utils" import { tGlobal } from "../../../lib/i18n" import { getDiffToolSearchText } from "../search-text" @@ -43,7 +43,8 @@ export const editRenderer: ToolRenderer = { const { metadata } = readToolStatePayload(state) const diffText = typeof metadata.diff === "string" ? metadata.diff : null const fallback = isToolStateCompleted(state) && typeof state.output === "string" ? state.output : null - const content = ensureMarkdownContent(diffText || fallback, "diff", true) + const value = diffText || fallback + const content = ensureMarkdownContent(value ? limitToolOutputForRender(value) : value, "diff", true) if (!content) return null return renderMarkdown({ content, size: "large", disableHighlight: state.status === "running" }) diff --git a/packages/ui/src/components/tool-call/renderers/patch.tsx b/packages/ui/src/components/tool-call/renderers/patch.tsx index 356bba5a6..a190dfdfa 100644 --- a/packages/ui/src/components/tool-call/renderers/patch.tsx +++ b/packages/ui/src/components/tool-call/renderers/patch.tsx @@ -1,5 +1,5 @@ import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, extractDiffPayload, getRelativePath, getToolName, isToolStateCompleted, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, extractDiffPayload, getRelativePath, getToolName, isToolStateCompleted, limitToolOutputForRender, readToolStatePayload } from "../utils" import { tGlobal } from "../../../lib/i18n" import { getDiffToolSearchText } from "../search-text" @@ -43,7 +43,8 @@ export const patchRenderer: ToolRenderer = { const { metadata } = readToolStatePayload(state) const diffText = typeof metadata.diff === "string" ? metadata.diff : null const fallback = isToolStateCompleted(state) && typeof state.output === "string" ? state.output : null - const content = ensureMarkdownContent(diffText || fallback, "diff", true) + const value = diffText || fallback + const content = ensureMarkdownContent(value ? limitToolOutputForRender(value) : value, "diff", true) if (!content) return null return renderMarkdown({ content, size: "large", disableHighlight: state.status === "running" }) diff --git a/packages/ui/src/components/tool-call/renderers/read.tsx b/packages/ui/src/components/tool-call/renderers/read.tsx index ff37e2e85..e4b73be66 100644 --- a/packages/ui/src/components/tool-call/renderers/read.tsx +++ b/packages/ui/src/components/tool-call/renderers/read.tsx @@ -1,5 +1,5 @@ import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, getRelativePath, getToolName, inferLanguageFromPath, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, getRelativePath, getToolName, inferLanguageFromPath, limitToolOutputForRender, readToolStatePayload } from "../utils" import { tGlobal } from "../../../lib/i18n" import { getReadToolSearchText } from "../search-text" @@ -56,7 +56,7 @@ export const readRenderer: ToolRenderer = { const { metadata, input } = readToolStatePayload(state) const preview = typeof metadata.preview === "string" ? metadata.preview : null const language = inferLanguageFromPath(getReadPath(input)) - const content = ensureMarkdownContent(preview, language, true) + const content = ensureMarkdownContent(preview ? limitToolOutputForRender(preview) : preview, language, true) if (!content) return null return renderMarkdown({ content, disableHighlight: state.status === "running" }) }, diff --git a/packages/ui/src/components/tool-call/renderers/skill.tsx b/packages/ui/src/components/tool-call/renderers/skill.tsx index adaa2cc66..af6810537 100644 --- a/packages/ui/src/components/tool-call/renderers/skill.tsx +++ b/packages/ui/src/components/tool-call/renderers/skill.tsx @@ -1,5 +1,5 @@ import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, formatUnknown, getToolName } from "../utils" +import { ensureMarkdownContent, formatUnknownForCopy, formatUnknownForRender, getToolName } from "../utils" import { getDefaultToolSearchText } from "../search-text" export const skillRenderer: ToolRenderer = { @@ -12,15 +12,14 @@ export const skillRenderer: ToolRenderer = { const state = toolState() if (!state || state.status !== "completed") return undefined - const output = formatUnknown(state.output)?.text ?? null - if (!output) return undefined - return { copyText: output, suppressInnerHeader: true } + if (state.output === undefined || state.output === null) return undefined + return { getCopyText: () => formatUnknownForCopy(state.output)?.text ?? null, suppressInnerHeader: true } }, renderBody({ toolState, renderMarkdown }) { const state = toolState() if (!state || state.status !== "completed") return null - const output = formatUnknown(state.output)?.text ?? null + const output = formatUnknownForRender(state.output)?.text ?? null const content = ensureMarkdownContent(output, undefined, false) if (!content) return null return
{renderMarkdown({ content })}
diff --git a/packages/ui/src/components/tool-call/renderers/task-summary.ts b/packages/ui/src/components/tool-call/renderers/task-summary.ts new file mode 100644 index 000000000..7eafa5f2d --- /dev/null +++ b/packages/ui/src/components/tool-call/renderers/task-summary.ts @@ -0,0 +1,14 @@ +export const TASK_STEP_RENDER_LIMIT = 200 + +export function getLegacyTaskSummary(summary: unknown) { + const entries = Array.isArray(summary) ? summary : [] + return { + entries, + renderedEntries: entries.slice(-TASK_STEP_RENDER_LIMIT), + truncated: entries.length > TASK_STEP_RENDER_LIMIT, + } +} + +export function stringifyLegacyTaskSummary(summary: unknown): string { + return JSON.stringify(getLegacyTaskSummary(summary).entries, null, 2) +} diff --git a/packages/ui/src/components/tool-call/renderers/task-transcript-load.test.ts b/packages/ui/src/components/tool-call/renderers/task-transcript-load.test.ts new file mode 100644 index 000000000..c5f55ae01 --- /dev/null +++ b/packages/ui/src/components/tool-call/renderers/task-transcript-load.test.ts @@ -0,0 +1,38 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { createRoot, createSignal } from "solid-js" +import { useTaskTranscriptLoad } from "./task-transcript-load.ts" + +test("a remounted task retries a cancelled transcript load after its loading marker clears", async () => { + const [loading, setLoading] = createSignal(false) + const [loaded, setLoaded] = createSignal(false) + const requests: Array<{ signal: AbortSignal; resolve: () => void }> = [] + const load = (_sessionId: string, signal: AbortSignal) => new Promise((resolve) => { + requests.push({ signal, resolve }) + setLoading(true) + }) + const mount = () => createRoot((dispose) => { + useTaskTranscriptLoad({ sessionId: () => "child", enabled: () => true, loaded, loading, load }) + return dispose + }) + + const disposeFirst = mount() + assert.equal(requests.length, 1) + assert.equal(requests[0].signal.aborted, false, "setting the owned loading marker must not cancel its request") + + disposeFirst() + assert.equal(requests[0].signal.aborted, true) + + const disposeRemount = mount() + assert.equal(requests.length, 1, "the remount joins the still-marked request") + + setLoading(false) + assert.equal(requests.length, 2, "the remount retries when cancellation clears the marker") + + setLoaded(true) + setLoading(false) + requests[0].resolve() + requests[1].resolve() + await Promise.resolve() + disposeRemount() +}) diff --git a/packages/ui/src/components/tool-call/renderers/task-transcript-load.ts b/packages/ui/src/components/tool-call/renderers/task-transcript-load.ts new file mode 100644 index 000000000..6ef861e21 --- /dev/null +++ b/packages/ui/src/components/tool-call/renderers/task-transcript-load.ts @@ -0,0 +1,34 @@ +import { createEffect, onCleanup, type Accessor } from "solid-js" + +interface TaskTranscriptLoadOptions { + sessionId: Accessor + enabled: Accessor + loaded: Accessor + loading: Accessor + load: (sessionId: string, signal: AbortSignal) => Promise +} + +export function useTaskTranscriptLoad(options: TaskTranscriptLoadOptions): void { + let owned: { sessionId: string; controller: AbortController } | undefined + + createEffect(() => { + const sessionId = options.sessionId() + const enabled = options.enabled() + const loaded = options.loaded() + const loading = options.loading() + + if (owned && (!enabled || owned.sessionId !== sessionId)) { + owned.controller.abort() + owned = undefined + } + if (!sessionId || !enabled || loaded || loading || owned) return + + const request = { sessionId, controller: new AbortController() } + owned = request + void options.load(sessionId, request.controller.signal).catch(() => undefined).finally(() => { + if (owned === request) owned = undefined + }) + }) + + onCleanup(() => owned?.controller.abort()) +} diff --git a/packages/ui/src/components/tool-call/renderers/task.tsx b/packages/ui/src/components/tool-call/renderers/task.tsx index fce9b3781..40ed369f9 100644 --- a/packages/ui/src/components/tool-call/renderers/task.tsx +++ b/packages/ui/src/components/tool-call/renderers/task.tsx @@ -1,11 +1,17 @@ -import { For, Index, Show, createEffect, createMemo, createSignal, untrack } from "solid-js" +import { For, Index, Show, createEffect, createMemo, createSignal, onCleanup, untrack } from "solid-js" import type { ToolState } from "@opencode-ai/sdk/v2" import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, getDefaultToolAction, getToolIcon, getToolName, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, getDefaultToolAction, getToolIcon, getToolName, limitToolOutputForRender, limitToolTitleForRender, readToolStatePayload } from "../utils" import { messageStoreBus } from "../../../stores/message-v2/bus" +import { activeInstanceId } from "../../../stores/instances" import { loadMessages } from "../../../stores/session-api" -import { loading, messagesLoaded } from "../../../stores/session-state" +import { activeSessionId, loading, messagesLoaded } from "../../../stores/session-state" +import { setVisibleSessionMemory } from "../../../stores/session-memory" import { getTaskToolSearchText } from "../search-text" +import { Copy } from "lucide-solid" +import { copyToClipboard } from "../../../lib/clipboard" +import { getLegacyTaskSummary, stringifyLegacyTaskSummary, TASK_STEP_RENDER_LIMIT } from "./task-summary" +import { useTaskTranscriptLoad } from "./task-transcript-load" interface TaskSummaryItem { id: string @@ -38,6 +44,7 @@ function TaskToolCallRow(props: { toolKey: string store: ReturnType sessionId: string + visibilitySessionId: string renderToolCall: NonNullable }) { const parts = createMemo(() => splitToolKey(props.toolKey)) @@ -74,6 +81,7 @@ function TaskToolCallRow(props: { messageVersion: messageVersion(), partVersion: partVersion(), sessionId: props.sessionId, + visibilitySessionId: props.visibilitySessionId, forceCollapsed: true, }) }) @@ -169,9 +177,31 @@ export const taskRenderer: ToolRenderer = { const { input } = readToolStatePayload(state) return describeTaskTitle(input) }, - renderBody({ toolState, instanceId, renderToolCall, messageVersion, partVersion, scrollHelpers, renderMarkdown, t, onContentRendered }) { + getOutputChrome({ toolState, t }) { + const state = toolState() + if (!state) return undefined + const { input } = readToolStatePayload(state) + const prompt = typeof input.prompt === "string" ? input.prompt : "" + const output = state && "output" in state && typeof state.output === "string" ? state.output : null + if (prompt.length > 10_000) { + return { + actions: ( + + ), + } + } + return output ? { copyText: output } : undefined + }, + renderBody({ toolState, instanceId, visibilitySessionId, renderToolCall, messageVersion, partVersion, scrollHelpers, renderMarkdown, t, onContentRendered }) { const store = messageStoreBus.getOrCreate(instanceId) - const [requestedChildLoad, setRequestedChildLoad] = createSignal(false) const childSessionId = createMemo(() => { const state = toolState() @@ -185,6 +215,11 @@ export const taskRenderer: ToolRenderer = { return loadedForInstance?.has(id) ?? false }) + const childSessionResident = createMemo(() => { + const id = childSessionId() + return Boolean(id && store.getSessionMessageIds(id).length > 0) + }) + const childSessionLoading = createMemo(() => { const id = childSessionId() if (!id) return false @@ -192,17 +227,23 @@ export const taskRenderer: ToolRenderer = { return loadingSet?.has(id) ?? false }) + useTaskTranscriptLoad({ + sessionId: childSessionId, + enabled: () => activeInstanceId() === instanceId && activeSessionId().get(instanceId) === visibilitySessionId, + loaded: childSessionLoaded, + loading: childSessionLoading, + load: (id, signal) => loadMessages(instanceId, id, { signal, force: true }), + }) + createEffect(() => { const id = childSessionId() - if (!id) return - if (requestedChildLoad()) return - if (childSessionLoaded()) return - if (childSessionLoading()) return - setRequestedChildLoad(true) - void loadMessages(instanceId, id) + if (!id || activeInstanceId() !== instanceId || activeSessionId().get(instanceId) !== visibilitySessionId) return + setVisibleSessionMemory(instanceId, id, true) + onCleanup(() => setVisibleSessionMemory(instanceId, id, false)) }) const [childToolKeys, setChildToolKeys] = createSignal([]) + const [childToolsTruncated, setChildToolsTruncated] = createSignal(false) let indexedSessionId = "" let indexedMessageCount = 0 @@ -215,20 +256,23 @@ export const taskRenderer: ToolRenderer = { indexedMessageTail = "" indexedPartCounts.clear() setChildToolKeys([]) + setChildToolsTruncated(false) } - function scanMessageToolParts(messageId: string, startIndex: number) { + function scanMessageToolParts(messageId: string, startIndex: number, limit: number) { const record = store.getMessage(messageId) if (!record) return [] as string[] const partIds = record.partIds const keys: string[] = [] - for (let idx = startIndex; idx < partIds.length; idx += 1) { + const oldestScannedIndex = Math.max(startIndex, partIds.length - 1_000) + if (oldestScannedIndex > startIndex) setChildToolsTruncated(true) + for (let idx = partIds.length - 1; idx >= oldestScannedIndex && keys.length < limit; idx -= 1) { const partId = partIds[idx] const entry = record.parts?.[partId] const data = entry?.data if (!data || (data as any).type !== "tool") continue - keys.push(`${messageId}::${partId}`) + keys.unshift(`${messageId}::${partId}`) } indexedPartCounts.set(messageId, partIds.length) return keys @@ -241,15 +285,18 @@ export const taskRenderer: ToolRenderer = { indexedPartCounts.clear() const nextKeys: string[] = [] - for (const messageId of messageIds) { - nextKeys.push(...scanMessageToolParts(messageId, 0)) + const oldestScannedIndex = Math.max(0, messageIds.length - 1_000) + for (let index = messageIds.length - 1; index >= oldestScannedIndex && nextKeys.length < TASK_STEP_RENDER_LIMIT; index -= 1) { + const keys = scanMessageToolParts(messageIds[index], 0, TASK_STEP_RENDER_LIMIT - nextKeys.length) + for (let keyIndex = keys.length - 1; keyIndex >= 0; keyIndex -= 1) nextKeys.unshift(keys[keyIndex]) } - setChildToolKeys(nextKeys) + setChildToolsTruncated(childToolsTruncated() || messageIds.length > 1_000 || nextKeys.length === TASK_STEP_RENDER_LIMIT) + setChildToolKeys(nextKeys.slice(-TASK_STEP_RENDER_LIMIT)) } createEffect(() => { const id = childSessionId() - const loaded = childSessionLoaded() + const loaded = childSessionLoaded() || childSessionResident() if (!id || !loaded) { if (indexedSessionId) { @@ -286,9 +333,11 @@ export const taskRenderer: ToolRenderer = { const appendedKeys: string[] = [] // Scan any new messages appended since last index. - for (let idx = indexedMessageCount; idx < messageIds.length; idx += 1) { + for (let idx = Math.max(indexedMessageCount, messageIds.length - 1_000); idx < messageIds.length; idx += 1) { const messageId = messageIds[idx] - appendedKeys.push(...scanMessageToolParts(messageId, 0)) + const keys = scanMessageToolParts(messageId, 0, TASK_STEP_RENDER_LIMIT) + for (const key of keys) appendedKeys.push(key) + if (appendedKeys.length > TASK_STEP_RENDER_LIMIT) appendedKeys.splice(0, appendedKeys.length - TASK_STEP_RENDER_LIMIT) } // Scan a small window of recent messages for newly appended parts. @@ -302,15 +351,22 @@ export const taskRenderer: ToolRenderer = { const record = store.getMessage(messageId) const nextPartCount = record?.partIds.length ?? 0 if (nextPartCount > previousPartCount) { - appendedKeys.push(...scanMessageToolParts(messageId, previousPartCount)) + const keys = scanMessageToolParts(messageId, previousPartCount, TASK_STEP_RENDER_LIMIT) + for (const key of keys) appendedKeys.push(key) + if (appendedKeys.length > TASK_STEP_RENDER_LIMIT) appendedKeys.splice(0, appendedKeys.length - TASK_STEP_RENDER_LIMIT) } } indexedMessageCount = messageIds.length indexedMessageTail = messageIds[messageIds.length - 1] ?? "" + if (indexedPartCounts.size > 1_000) { + const retainedIds = new Set(messageIds.slice(-1_000)) + for (const messageId of indexedPartCounts.keys()) if (!retainedIds.has(messageId)) indexedPartCounts.delete(messageId) + } if (appendedKeys.length > 0) { - setChildToolKeys((prev) => [...prev, ...appendedKeys]) + if (childToolKeys().length + appendedKeys.length > TASK_STEP_RENDER_LIMIT) setChildToolsTruncated(true) + setChildToolKeys((prev) => [...prev, ...appendedKeys].slice(-TASK_STEP_RENDER_LIMIT)) } }) }) @@ -319,14 +375,14 @@ export const taskRenderer: ToolRenderer = { if (!state) return null const { input } = readToolStatePayload(state) const prompt = typeof input.prompt === "string" ? input.prompt : null - return ensureMarkdownContent(prompt, undefined, false) + return ensureMarkdownContent(prompt ? limitToolOutputForRender(prompt) : prompt, undefined, false) }) const outputContent = createMemo(() => { const state = toolState() if (!state) return null const output = typeof (state as { output?: unknown }).output === "string" ? ((state as { output?: string }).output as string) : null - return ensureMarkdownContent(output, undefined, false) + return ensureMarkdownContent(output ? limitToolOutputForRender(output) : output, undefined, false) }) const agentLabel = createMemo(() => { @@ -358,21 +414,22 @@ export const taskRenderer: ToolRenderer = { return null }) - const legacyItems = createMemo(() => { + const legacySummary = createMemo(() => { // Track the reactive change points so we only recompute when the part/message changes messageVersion?.() partVersion?.() const state = toolState() - if (!state) return [] + if (!state) return getLegacyTaskSummary(undefined) + const { metadata } = readToolStatePayload(state) + return getLegacyTaskSummary((metadata as any).summary) + }) + const legacyItems = createMemo(() => { // Prefer deriving steps from the child session when loaded. if (childSessionLoaded()) return [] - const { metadata } = readToolStatePayload(state) - const summary = Array.isArray((metadata as any).summary) ? ((metadata as any).summary as any[]) : [] - - return summary.map((entry, index) => { + return legacySummary().renderedEntries.map((entry, index) => { const tool = typeof entry?.tool === "string" ? (entry.tool as string) : "unknown" const stateValue = typeof entry?.state === "object" ? (entry.state as ToolState) : undefined const metadataFromEntry = typeof entry?.metadata === "object" && entry.metadata ? entry.metadata : {} @@ -419,11 +476,27 @@ export const taskRenderer: ToolRenderer = {
{t("toolCall.task.sections.steps")} -
+ +
{t("toolCall.output.truncated")}
+
0} fallback={ @@ -438,8 +511,8 @@ export const taskRenderer: ToolRenderer = { {(item) => { const icon = getToolIcon(item.tool) - const description = describeToolTitle(item) - const toolLabel = getToolName(item.tool) + const description = limitToolTitleForRender(describeToolTitle(item)) + const toolLabel = limitToolTitleForRender(getToolName(item.tool)) const status = normalizeStatus(item.status ?? item.state?.status) const statusIcon = summarizeStatusIcon(status) const statusKey = summarizeStatusLabel(status) @@ -483,6 +556,7 @@ export const taskRenderer: ToolRenderer = { toolKey={key()} store={store} sessionId={childSessionId()} + visibilitySessionId={visibilitySessionId} renderToolCall={render()} /> )} diff --git a/packages/ui/src/components/tool-call/renderers/webfetch.tsx b/packages/ui/src/components/tool-call/renderers/webfetch.tsx index 59aa5645a..88684bea3 100644 --- a/packages/ui/src/components/tool-call/renderers/webfetch.tsx +++ b/packages/ui/src/components/tool-call/renderers/webfetch.tsx @@ -1,5 +1,5 @@ import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, formatUnknown, getToolName, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, formatUnknownForCopy, formatUnknownForRender, getToolName, readToolStatePayload } from "../utils" import { tGlobal } from "../../../lib/i18n" import { getWebfetchToolSearchText } from "../search-text" @@ -21,16 +21,13 @@ export const webfetchRenderer: ToolRenderer = { if (!state || state.status === "pending") return undefined const { metadata } = readToolStatePayload(state) - const result = formatUnknown( - state.status === "completed" - ? state.output - : metadata.output, - ) - if (!result) return undefined + const output = state.status === "completed" ? state.output : metadata.output + if (output === undefined || output === null) return undefined + const rendered = formatUnknownForRender(output) return { - language: result.language ?? "text", - copyText: result.text, + language: rendered?.language ?? "text", + getCopyText: () => formatUnknownForCopy(output)?.text ?? null, wrapToggle: true, suppressInnerHeader: true, } @@ -40,7 +37,7 @@ export const webfetchRenderer: ToolRenderer = { if (!state || state.status === "pending") return null const { metadata } = readToolStatePayload(state) - const result = formatUnknown( + const result = formatUnknownForRender( state.status === "completed" ? state.output : metadata.output, diff --git a/packages/ui/src/components/tool-call/renderers/write.tsx b/packages/ui/src/components/tool-call/renderers/write.tsx index 9a7f3d540..f72c0438f 100644 --- a/packages/ui/src/components/tool-call/renderers/write.tsx +++ b/packages/ui/src/components/tool-call/renderers/write.tsx @@ -1,5 +1,5 @@ import type { ToolRenderer } from "../types" -import { ensureMarkdownContent, getRelativePath, getToolName, inferLanguageFromPath, readToolStatePayload } from "../utils" +import { ensureMarkdownContent, getRelativePath, getToolName, inferLanguageFromPath, limitToolOutputForRender, readToolStatePayload } from "../utils" import { tGlobal } from "../../../lib/i18n" import { getWriteToolSearchText } from "../search-text" @@ -35,7 +35,7 @@ export const writeRenderer: ToolRenderer = { const { metadata, input } = readToolStatePayload(state) const contentValue = typeof input.content === "string" ? input.content : metadata.content const filePath = typeof input.filePath === "string" ? input.filePath : undefined - const content = ensureMarkdownContent(contentValue ?? null, inferLanguageFromPath(filePath), true) + const content = ensureMarkdownContent(typeof contentValue === "string" ? limitToolOutputForRender(contentValue) : null, inferLanguageFromPath(filePath), true) if (!content) return null return renderMarkdown({ content, size: "large", disableHighlight: state.status === "running" }) }, diff --git a/packages/ui/src/components/tool-call/search-text.test.ts b/packages/ui/src/components/tool-call/search-text.test.ts new file mode 100644 index 000000000..5576959e9 --- /dev/null +++ b/packages/ui/src/components/tool-call/search-text.test.ts @@ -0,0 +1,487 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { + buildSessionSearchMatches, + createSessionSearchPager, + findLastSessionSearchPage, + retainSessionSearchPage, +} from "../../lib/session-search.ts" +import { getApplyPatchToolSearchText, getDefaultToolSearchText, getDiffToolSearchText, getReadToolSearchText, getTaskToolSearchText, getWriteToolSearchText } from "./search-text.ts" +import { getLegacyTaskSummary, stringifyLegacyTaskSummary, TASK_STEP_RENDER_LIMIT } from "./renderers/task-summary.ts" + +function textRecord(id: string, text: string) { + const partId = `${id}-part` + return { + id, + sessionId: "session", + role: "user", + status: "complete", + createdAt: 0, + updatedAt: 0, + revision: 1, + partIds: [partId], + parts: { [partId]: { id: partId, revision: 1, data: { id: partId, type: "text", text } } }, + } +} + +function searchStore(messageIds: string[], records: Map>) { + return { + getSessionMessageIds: () => messageIds, + getMessage: (messageId: string) => records.get(messageId), + getMessageInfo: () => undefined, + } as any +} + +async function collect(values: AsyncIterable): Promise { + const result: string[] = [] + for await (const value of values) result.push(value) + return result +} + +test("task search keeps visible output ahead of an oversized prompt", async () => { + const values = await collect(getTaskToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "task", state: {} } as any, + toolName: "task", + toolState: { status: "completed", input: { prompt: "x".repeat(20_000) }, output: "UNIQUE_RESULT" } as any, + })) + assert.equal(values.join("\n").includes("UNIQUE_RESULT"), true) +}) + +test("tool search retains output beyond the render preview limit", async () => { + const values = await collect(getTaskToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "task", state: {} } as any, + toolName: "task", + toolState: { status: "completed", input: {}, output: `${"x".repeat(20_000)}TAIL_NEEDLE` } as any, + })) + assert.equal(values.join("\n").includes("TAIL_NEEDLE"), true) +}) + +test("apply_patch search includes the exact rendered patch input", async () => { + const input = { patchText: "*** Begin Patch\n*** Add File: exact-needle.txt\n+needle\n*** End Patch" } + const values = await collect(getApplyPatchToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "apply_patch", state: {} } as any, + toolName: "apply_patch", + toolState: { status: "completed", input, output: "" } as any, + })) + + assert.equal(values.includes(JSON.stringify(input, null, 2)), true) +}) + +test("task search includes the exact rendered model metadata", async () => { + const values = await collect(getTaskToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "task", state: {} } as any, + toolName: "task", + toolState: { + status: "completed", + input: {}, + metadata: { model: { providerID: "exact-provider", modelID: "exact-model" } }, + output: "", + } as any, + })) + + assert.equal(values.includes("exact-provider/exact-model"), true) +}) + +test("edit and patch search include complete diagnostics", async () => { + for (const toolName of ["edit", "patch"]) { + const values = await collect(getDiffToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: toolName, state: {} } as any, + toolName, + toolState: { + status: "completed", + input: { filePath: "src/file.ts" }, + metadata: { diagnostics: { "src/file.ts": [{ message: `${"x".repeat(2_000)}DIAGNOSTIC_TAIL` }] } }, + output: "", + } as any, + })) + assert.equal(values.join("\n").includes("DIAGNOSTIC_TAIL"), true) + } +}) + +test("write search includes complete diagnostics", async () => { + const values = await collect(getWriteToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "write", state: {} } as any, + toolName: "write", + toolState: { + status: "completed", + input: { filePath: "src/file.ts", content: "content" }, + metadata: { diagnostics: { "src/file.ts": [{ message: `${"x".repeat(2_000)}DIAGNOSTIC_TAIL` }] } }, + output: "", + } as any, + })) + assert.equal(values.join("\n").includes("DIAGNOSTIC_TAIL"), true) +}) + +test("task output tail matches remain exposed through a bounded visible preview", async () => { + const record = textRecord("tool-tail", "") + const partId = record.partIds[0] + ;(record.parts as any)[partId] = { + id: partId, + revision: 1, + data: { id: partId, type: "tool", tool: "task", state: { status: "completed", input: {}, output: `${"x".repeat(20_000)}TAIL_NEEDLE` } }, + } as any + const result = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: "TAIL_NEEDLE", + includeThinking: false, + resolveToolSearchText: getTaskToolSearchText, + }) + + assert.equal(result.matches[0]?.preview.includes("TAIL_NEEDLE"), true) + assert.ok((result.matches[0]?.preview.length ?? 0) <= 120) +}) + +test("legacy task summary search matches remain available from the bounded view's copy source", async () => { + const summary = [ + { title: "OMITTED_SEARCH_MATCH" }, + ...Array.from({ length: TASK_STEP_RENDER_LIMIT }, (_, index) => ({ title: `step-${index}` })), + ] + const view = getLegacyTaskSummary(summary) + const searchValues = await collect(getTaskToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "task", state: {} } as any, + toolName: "task", + toolState: { status: "completed", input: {}, metadata: { summary } } as any, + })) + + assert.equal(view.renderedEntries.length, TASK_STEP_RENDER_LIMIT) + assert.equal(view.renderedEntries.some((entry) => entry.title === "OMITTED_SEARCH_MATCH"), false) + assert.equal(searchValues.join("\n").includes("OMITTED_SEARCH_MATCH"), true) + assert.equal(stringifyLegacyTaskSummary(view.entries).includes("OMITTED_SEARCH_MATCH"), true) +}) + +test("case-fold expansion preserves original string offsets", async () => { + const record = textRecord("unicode", "İİxy") + const result = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: "xy", + includeThinking: false, + }) + + assert.deepEqual(result.matches.map(({ start, end }) => ({ start, end })), [{ start: 2, end: 4 }]) +}) + +test("small structured output remains searchable as rendered JSON", async () => { + const values = await collect(getTaskToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "task", state: {} } as any, + toolName: "task", + toolState: { status: "completed", input: {}, output: { status: "failed", detail: null } } as any, + })) + assert.equal(values.some((value) => value.includes('"status": "failed"')), true) + assert.equal(values.some((value) => value.includes('"detail": null')), true) +}) + +test("complete-session search includes messages older than the former corpus window", async () => { + const messageIds = ["old", ...Array.from({ length: 10_000 }, (_, index) => `new-${index}`)] + const { matches } = await buildSessionSearchMatches({ + store: searchStore(messageIds, new Map([["old", textRecord("old", "distant needle")]])), + sessionId: "session", + query: "needle", + includeThinking: false, + }) + assert.equal(matches[0]?.messageId, "old") +}) + +test("complete-session search includes every part", async () => { + const record = textRecord("message", "") + record.partIds = Array.from({ length: 20_001 }, (_, index) => `part-${index}`) + record.parts = Object.fromEntries(record.partIds.map((id, index) => [ + id, + { id, revision: 1, data: { id, type: "text", text: index === 20_000 ? "final needle" : "" } }, + ])) + const { matches } = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: "needle", + includeThinking: false, + }) + assert.equal(matches[0]?.partId, "part-20000") +}) + +test("complete-session search includes characters beyond former per-part and total limits", async () => { + const record = textRecord("large", `${"x".repeat(5_000_001)}needle`) + const { matches } = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: "needle", + includeThinking: false, + }) + assert.equal(matches[0]?.start, 5_000_001) +}) + +test("chunked search finds a match spanning a chunk boundary", async () => { + const record = textRecord("boundary", `${"x".repeat(65_534)}needle`) + const { matches, totalMatches } = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: "needle", + includeThinking: false, + }) + assert.equal(totalMatches, 1) + assert.equal(matches[0]?.start, 65_534) +}) + +test("chunked search carries non-overlap offsets across chunk boundaries", async () => { + const record = textRecord("non-overlap-boundary", "a".repeat(65_550)) + const { matches } = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: "aaa", + includeThinking: false, + limit: 30_000, + }) + const starts = matches.map((match) => match.start) + const boundaryIndex = starts.indexOf(65_535) + + assert.deepEqual(starts.slice(boundaryIndex, boundaryIndex + 3), [65_535, 65_538, 65_541]) + assert.equal(starts.every((start, index) => start === index * 3), true) +}) + +test("complete-session search does not shorten the query", async () => { + const query = "q".repeat(1_001) + const record = textRecord("query", query) + const { matches } = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query, + includeThinking: false, + }) + assert.equal(matches[0]?.end, query.length) +}) + +test("oversized literal matches preserve offsets and retain only a bounded preview", async () => { + const query = "Ab".repeat(50_000) + const record = textRecord("oversized-preview", `İİ${query.toLowerCase()}:suffix`) + const { matches } = await buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query, + includeThinking: false, + }) + + assert.deepEqual(matches.map(({ start, end }) => ({ start, end })), [{ start: 2, end: 2 + query.length }]) + assert.ok((matches[0]?.preview.length ?? 0) < 300) +}) + +test("oversized literal misses stay responsive and cancel at the scanner checkpoint", async () => { + const record = textRecord("oversized-miss", "a".repeat(200_000)) + const controller = new AbortController() + const startedAt = Date.now() + let yields = 0 + + await assert.rejects( + buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: `${"a".repeat(100_000)}b`, + includeThinking: false, + signal: controller.signal, + yieldControl: async () => { + yields += 1 + controller.abort() + }, + }), + (error: any) => error?.name === "AbortError", + ) + + assert.ok(yields > 0) + assert.ok(Date.now() - startedAt < 2_000) +}) + +test("paged search keeps matches after the first thousand reachable", async () => { + const record = textRecord("many", `${"needle ".repeat(1_001)}tail needle`) + const store = searchStore([record.id], new Map([[record.id, record]])) + const pager = createSessionSearchPager({ + store, + sessionId: "session", + query: "needle", + includeThinking: false, + limit: 250, + }) + const first = await pager.nextPage() + let later = first + while (later.hasMore) later = await pager.nextPage() + + assert.equal(first.totalMatches, null) + assert.equal(first.matches.length, 250) + assert.equal(later.totalMatches, 1_002) + assert.equal(later.offset, 1_000) + assert.equal(later.matches.at(-1)?.occurrence, 1_001) +}) + +test("session search yields and honors stale-generation cancellation", async () => { + const records = new Map>() + const messageIds = Array.from({ length: 200 }, (_, index) => { + const id = `message-${index}` + records.set(id, textRecord(id, "searchable text")) + return id + }) + const controller = new AbortController() + let yields = 0 + + await assert.rejects( + buildSessionSearchMatches({ + store: searchStore(messageIds, records), + sessionId: "session", + query: "absent", + includeThinking: false, + signal: controller.signal, + yieldControl: async () => { + yields += 1 + controller.abort() + }, + }), + (error: any) => error?.name === "AbortError", + ) + assert.ok(yields > 0) +}) + +test("tool search handles deeply nested and cyclic structured output", async () => { + let output: any = "DEEPEST_NEEDLE" + for (let depth = 0; depth < 20_000; depth += 1) output = [output] + const cycle: any = { output } + cycle.self = cycle + + const values = await collect(getTaskToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "task", state: {} } as any, + toolName: "task", + toolState: { status: "completed", input: {}, output: cycle } as any, + })) + assert.equal(values.includes("DEEPEST_NEEDLE"), true) +}) + +test("tool search includes text rendered in specialized titles", async () => { + const base = { toolCall: { type: "tool", id: "tool", tool: "grep", state: {} } as any } + const searchValues = await collect(getDefaultToolSearchText({ + ...base, + toolName: "grep", + toolState: { status: "completed", input: { pattern: "release-[0-9]+" }, output: "" } as any, + })) + const taskValues = await collect(getTaskToolSearchText({ + ...base, + toolName: "task", + toolState: { status: "completed", input: { description: "audit reconnect" }, output: "" } as any, + })) + const readValues = await collect(getReadToolSearchText({ + ...base, + toolName: "read", + toolState: { status: "completed", input: { path: "file.ts", offset: 40, limit: 20 } } as any, + })) + + assert.ok(searchValues.includes("release-[0-9]+")) + assert.ok(taskValues.includes("audit reconnect")) + assert.ok(readValues.some((value) => value.includes("40"))) + assert.ok(readValues.some((value) => value.includes("20"))) +}) + +test("running default tools search the same fallback content they render", async () => { + for (const fallback of [ + { diff: "DIFF_NEEDLE" }, + { preview: "PREVIEW_NEEDLE" }, + ]) { + const values = await collect(getDefaultToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "grep", state: {} } as any, + toolName: "grep", + toolState: { status: "running", input: { content: "INPUT_NEEDLE" }, metadata: fallback } as any, + })) + assert.equal(values.includes(Object.values(fallback)[0]!), true) + } + + const values = await collect(getDefaultToolSearchText({ + toolCall: { type: "tool", id: "tool", tool: "invalid", state: {} } as any, + toolName: "invalid", + toolState: { status: "error", input: { content: "INPUT_NEEDLE" }, metadata: { output: "" } } as any, + })) + assert.equal(values.includes("INPUT_NEEDLE"), true) +}) + +test("pager resumes extraction without revisiting earlier transcript content", async () => { + const records = new Map([["message", textRecord("message", "needle ".repeat(600))]]) + let recordReads = 0 + const store = searchStore(["message"], records) + const originalGetMessage = store.getMessage + store.getMessage = (messageId: string) => { + recordReads += 1 + return originalGetMessage(messageId) + } + const pager = createSessionSearchPager({ store, sessionId: "session", query: "needle", includeThinking: false }) + + await pager.nextPage() + await pager.nextPage() + + assert.equal(recordReads, 1) +}) + +test("retained search pages stay bounded and last-page discovery reaches the final match", async () => { + const records = new Map([["message", textRecord("message", "needle ".repeat(9))]]) + const pager = createSessionSearchPager({ + store: searchStore(["message"], records), + sessionId: "session", + query: "needle", + includeThinking: false, + limit: 2, + }) + const retained = new Map>>() + let nextPageIndex = 0 + const last = await findLastSessionSearchPage(async (pageIndex) => { + assert.equal(pageIndex, nextPageIndex) + const page = await pager.nextPage() + retainSessionSearchPage(retained, pageIndex, page, 3) + nextPageIndex += 1 + return page + }) + + assert.equal(last.pageIndex, 4) + assert.equal(last.result.matches.at(-1)?.occurrence, 8) + assert.deepEqual([...retained.keys()], [2, 3, 4]) +}) + +test("a refreshed pager clamps a shrunken result set to its first valid page", async () => { + const records = new Map([["message", textRecord("message", "needle ".repeat(300))]]) + const store = searchStore(["message"], records) + const first = await createSessionSearchPager({ store, sessionId: "session", query: "needle", includeThinking: false }).nextPage() + assert.equal(first.hasMore, true) + + records.set("message", textRecord("message", "needle ".repeat(100))) + const refreshed = await createSessionSearchPager({ store, sessionId: "session", query: "needle", includeThinking: false }).nextPage() + assert.equal(refreshed.offset, 0) + assert.equal(refreshed.matches.length, 100) + assert.equal(refreshed.totalMatches, 100) +}) + +test("tool extraction yields and can be cancelled inside a large payload", async () => { + const record = textRecord("tool-message", "") + const partId = "tool-part" + record.partIds = [partId] + record.parts = { + [partId]: { + id: partId, + revision: 1, + data: { + id: partId, + type: "tool", + tool: "task", + state: { status: "completed", input: {}, output: Array.from({ length: 20_000 }, (_, index) => `value-${index}`) }, + }, + }, + } as any + const controller = new AbortController() + let yields = 0 + + await assert.rejects( + buildSessionSearchMatches({ + store: searchStore([record.id], new Map([[record.id, record]])), + sessionId: "session", + query: "absent", + includeThinking: false, + signal: controller.signal, + yieldControl: async () => { + yields += 1 + controller.abort() + }, + }), + (error: any) => error?.name === "AbortError", + ) + assert.ok(yields > 0) +}) diff --git a/packages/ui/src/components/tool-call/search-text.ts b/packages/ui/src/components/tool-call/search-text.ts index b50cf4ef2..2bfddca64 100644 --- a/packages/ui/src/components/tool-call/search-text.ts +++ b/packages/ui/src/components/tool-call/search-text.ts @@ -1,184 +1,224 @@ import type { ToolSearchTextContext } from "./types" import { - formatUnknown, isToolStateCompleted, isToolStateError, isToolStateRunning, readToolStatePayload, } from "./utils" +import { exceedsRetainedByteLimit } from "../../lib/session-memory-budget" +import { tGlobal } from "../../lib/i18n" type QuestionOption = { label?: unknown; description?: unknown } -type QuestionPrompt = { header?: unknown; question?: unknown; options?: unknown; multiple?: unknown; answer?: unknown } +type QuestionPrompt = { header?: unknown; question?: unknown; options?: unknown; answer?: unknown } -function appendString(values: string[], value: unknown) { - if (typeof value === "string" && value.trim().length > 0) values.push(value) +function* strings(...values: unknown[]): Generator { + for (const value of values) if (typeof value === "string" && value.length > 0) yield value } -function appendFormatted(values: string[], value: unknown) { - const result = formatUnknown(value) - if (result?.text.trim()) values.push(result.text) +function* objectSearchValues(record: Record): Generator { + for (const key in record) { + if (!Object.prototype.hasOwnProperty.call(record, key)) continue + yield key + try { + yield record[key] + } catch { + // Ignore hostile getters while retaining the rest of the payload. + } + } +} + +async function* formatted(value: unknown, context: ToolSearchTextContext): AsyncGenerator { + let canRenderJson = false + try { + canRenderJson = value !== null && typeof value === "object" && !exceedsRetainedByteLimit(value, 10_000) + } catch { + // Fall through to guarded traversal. + } + if (canRenderJson) { + try { + const rendered = JSON.stringify(value, null, 2) + if (rendered) { + yield rendered + return + } + } catch { + // Cyclic values use guarded traversal. + } + } + + type Pending = { value: unknown } | { iterator: Iterator } + const pending: Pending[] = [{ value }] + const seen = new WeakSet() + let units = 0 + while (pending.length > 0) { + const item = pending.pop()! + if ("iterator" in item) { + let next: IteratorResult + try { + next = item.iterator.next() + } catch { + continue + } + if (!next.done) pending.push(item, { value: next.value }) + continue + } + const current = item.value + if (current === null) yield "null" + else if (typeof current === "string") yield current + else if (typeof current === "number" || typeof current === "boolean" || typeof current === "bigint") yield String(current) + else if (current && typeof current === "object" && !seen.has(current)) { + seen.add(current) + pending.push({ iterator: Array.isArray(current) ? current.values() : objectSearchValues(current as Record) }) + } + units += 1 + if (units % 64 === 0) await context.checkpoint?.() + } } -function appendBaseToolText(values: string[], context: ToolSearchTextContext) { +function* base(context: ToolSearchTextContext): Generator { const { metadata } = readToolStatePayload(context.toolState) - appendString(values, context.toolName) - appendString(values, metadata.title) - appendString(values, metadata.description) - appendString(values, context.toolState && "title" in context.toolState ? (context.toolState as any).title : undefined) + yield* strings( + context.toolName, + metadata.title, + metadata.description, + context.toolState && "title" in context.toolState ? (context.toolState as any).title : undefined, + ) } -function appendToolErrorText(values: string[], context: ToolSearchTextContext) { - appendString(values, context.toolState && "message" in context.toolState ? (context.toolState as any).message : undefined) - appendString(values, context.toolState && "error" in context.toolState ? (context.toolState as any).error : undefined) +function* errors(context: ToolSearchTextContext): Generator { + yield* strings( + context.toolState && "message" in context.toolState ? (context.toolState as any).message : undefined, + context.toolState && "error" in context.toolState ? (context.toolState as any).error : undefined, + ) } -export function getDefaultToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getDefaultToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const state = context.toolState const { input, metadata, output } = readToolStatePayload(state) - appendBaseToolText(values, context) - - const primaryOutput = state && isToolStateCompleted(state) - ? output - : state && (isToolStateRunning(state) || isToolStateError(state)) && metadata.output - ? metadata.output - : metadata.diff ?? metadata.preview ?? input.content - - appendString(values, typeof input.command === "string" ? `$ ${input.command}` : undefined) - appendString(values, input.filePath) - appendString(values, input.path) - appendFormatted(values, primaryOutput) - appendToolErrorText(values, context) - return values + yield* base(context) + yield* strings( + typeof input.command === "string" ? `$ ${input.command}` : undefined, + input.description, + input.pattern, + input.filePath, + input.path, + ) + yield* formatted( + state && isToolStateCompleted(state) + ? output + : state && (isToolStateRunning(state) || isToolStateError(state)) + ? metadata.output || (metadata.diff ?? metadata.preview ?? input.content) + : metadata.diff ?? metadata.preview ?? input.content, + context, + ) + yield* errors(context) } -export function getBashToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getBashToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const state = context.toolState const { input, metadata, output } = readToolStatePayload(state) - appendBaseToolText(values, context) - appendString(values, typeof input.command === "string" && input.command.length > 0 ? `$ ${input.command}` : undefined) - appendFormatted( - values, + yield* base(context) + yield* strings( + typeof input.command === "string" && input.command.length > 0 ? `$ ${input.command}` : undefined, + input.description, + typeof input.timeout === "number" ? String(input.timeout) : undefined, + ) + yield* formatted( state && isToolStateCompleted(state) ? output - : state && (isToolStateRunning(state) || isToolStateError(state)) - ? metadata.output - : undefined, + : state && (isToolStateRunning(state) || isToolStateError(state)) ? metadata.output : undefined, + context, ) - appendToolErrorText(values, context) - return values + yield* errors(context) } -export function getReadToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getReadToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const { input, metadata } = readToolStatePayload(context.toolState) - appendBaseToolText(values, context) - appendString(values, input.filePath) - appendString(values, input.path) - appendString(values, input.name) - appendString(values, metadata.preview) - appendToolErrorText(values, context) - return values + yield* base(context) + yield* strings(input.filePath, input.path, input.name, metadata.preview) + if (typeof input.offset === "number") yield tGlobal("toolCall.renderer.read.detail.offset", { offset: input.offset }) + if (typeof input.limit === "number") yield tGlobal("toolCall.renderer.read.detail.limit", { limit: input.limit }) + yield* errors(context) } -export function getWriteToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getWriteToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const { input, metadata } = readToolStatePayload(context.toolState) - appendBaseToolText(values, context) - appendString(values, input.filePath) - appendString(values, typeof input.content === "string" ? input.content : metadata.content) - appendToolErrorText(values, context) - return values + yield* base(context) + yield* strings(input.filePath, typeof input.content === "string" ? input.content : metadata.content) + yield* formatted((metadata as any).diagnostics, context) + yield* errors(context) } -export function getDiffToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getDiffToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const { input, metadata, output } = readToolStatePayload(context.toolState) - appendBaseToolText(values, context) - appendString(values, input.filePath) - appendString(values, input.path) - appendString(values, metadata.diff) - appendFormatted(values, output) - appendFormatted(values, metadata.output) - appendToolErrorText(values, context) - return values + yield* base(context) + yield* strings(input.filePath, input.path, metadata.diff) + yield* formatted(output, context) + yield* formatted(metadata.output, context) + yield* formatted((metadata as any).diagnostics, context) + yield* errors(context) } -export function getApplyPatchToolSearchText(context: ToolSearchTextContext): string[] { - const values = getDiffToolSearchText(context) - const { metadata, output } = readToolStatePayload(context.toolState) +export async function* getApplyPatchToolSearchText(context: ToolSearchTextContext): AsyncGenerator { + yield* getDiffToolSearchText(context) + const { input, metadata } = readToolStatePayload(context.toolState) + yield* formatted(input, context) const files = Array.isArray((metadata as any).files) ? ((metadata as any).files as any[]) : [] - - for (const file of files) { - appendString(values, file?.filePath) - appendString(values, file?.relativePath) - appendString(values, file?.diff) - appendString(values, file?.patch) + for (let index = 0; index < files.length; index += 1) { + const file = files[index] + yield* strings(file?.filePath, file?.relativePath, file?.diff, file?.patch) + if (index % 64 === 63) await context.checkpoint?.() } - - appendFormatted(values, (metadata as any).diagnostics) - appendFormatted(values, output) - return values } -export function getWebfetchToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getWebfetchToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const state = context.toolState const { input, metadata, output } = readToolStatePayload(state) - appendBaseToolText(values, context) - appendString(values, input.url) - appendFormatted(values, state && isToolStateCompleted(state) ? output : metadata.output) - appendToolErrorText(values, context) - return values + yield* base(context) + yield* strings(input.url) + yield* formatted(state && isToolStateCompleted(state) ? output : metadata.output, context) + yield* errors(context) } -export function getTaskToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getTaskToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const { input, metadata, output } = readToolStatePayload(context.toolState) - appendBaseToolText(values, context) - appendString(values, input.prompt) - appendString(values, input.subagent_type) - appendFormatted(values, output) - appendFormatted(values, metadata.summary) - appendToolErrorText(values, context) - return values + const model = (metadata as any).model + const providerId = model && typeof model === "object" && typeof model.providerID === "string" ? model.providerID : undefined + const modelId = model && typeof model === "object" && typeof model.modelID === "string" ? model.modelID : undefined + yield* base(context) + yield* errors(context) + yield* formatted(output, context) + yield* formatted(metadata.summary, context) + yield* strings(input.description, input.subagent_type, input.prompt) + yield* strings(providerId && modelId ? `${providerId}/${modelId}` : providerId ?? modelId) } -export function getTodoToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getTodoToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const { metadata } = readToolStatePayload(context.toolState) const todos = Array.isArray((metadata as any).todos) ? ((metadata as any).todos as any[]) : [] - appendBaseToolText(values, context) - - for (const todo of todos) { - appendString(values, todo?.content) - appendString(values, todo?.status) + yield* base(context) + for (let index = 0; index < todos.length; index += 1) { + yield* strings(todos[index]?.content, todos[index]?.status) + if (index % 64 === 63) await context.checkpoint?.() } - - appendToolErrorText(values, context) - return values + yield* errors(context) } -export function getQuestionToolSearchText(context: ToolSearchTextContext): string[] { - const values: string[] = [] +export async function* getQuestionToolSearchText(context: ToolSearchTextContext): AsyncGenerator { const { input, metadata } = readToolStatePayload(context.toolState) const questions = Array.isArray(input.questions) ? (input.questions as QuestionPrompt[]) : [] const answers = Array.isArray((metadata as any).answers) ? ((metadata as any).answers as unknown[]) : [] - appendBaseToolText(values, context) - + yield* base(context) for (const question of questions) { - appendString(values, question.header) - appendString(values, question.question) + yield* strings(question.header, question.question) const options = Array.isArray(question.options) ? (question.options as QuestionOption[]) : [] - for (const option of options) { - appendString(values, option.label) - appendString(values, option.description) + for (let index = 0; index < options.length; index += 1) { + yield* strings(options[index].label, options[index].description) + if (index % 64 === 63) await context.checkpoint?.() } - appendFormatted(values, question.answer) + yield* formatted(question.answer, context) + await context.checkpoint?.() } - - appendFormatted(values, answers) - appendToolErrorText(values, context) - return values + yield* formatted(answers, context) + yield* errors(context) } diff --git a/packages/ui/src/components/tool-call/types.ts b/packages/ui/src/components/tool-call/types.ts index 496337877..61123577b 100644 --- a/packages/ui/src/components/tool-call/types.ts +++ b/packages/ui/src/components/tool-call/types.ts @@ -37,6 +37,7 @@ export interface DiffRenderOptions { variant?: string disableScrollTracking?: boolean label?: string + onFullDiffAccess?: () => void /** * Optional cache key suffix to avoid collisions when rendering multiple diffs * within the same tool call (e.g. apply_patch). @@ -57,6 +58,7 @@ export interface ToolRendererContext { toolName: Accessor instanceId: string sessionId: string + visibilitySessionId: string t: (key: string, params?: Record) => string messageVersion?: Accessor partVersion?: Accessor @@ -73,6 +75,7 @@ export interface ToolRendererContext { messageVersion?: number partVersion?: number sessionId: string + visibilitySessionId?: string forceCollapsed?: boolean }) => JSXElement | null outputWrapEnabled?: Accessor @@ -84,6 +87,7 @@ export interface ToolSearchTextContext { toolCall: ToolCallPart toolState: ToolState | undefined toolName: string + checkpoint?: () => Promise } export interface ToolRenderer { @@ -95,7 +99,7 @@ export interface ToolRenderer { * Text that is visible or directly revealable through this renderer. Keep this * in sync with custom renderBody output when adding specialized tool UIs. */ - getSearchText?(context: ToolSearchTextContext): string[] + getSearchText?(context: ToolSearchTextContext): AsyncIterable | Promise | string[] renderBody(context: ToolRendererContext): JSXElement | null } @@ -103,6 +107,7 @@ export interface ToolOutputChrome { title?: string language?: string copyText?: string | null + getCopyText?: () => string | null actions?: JSXElement wrapToggle?: boolean suppressInnerHeader?: boolean diff --git a/packages/ui/src/components/tool-call/utils.test.ts b/packages/ui/src/components/tool-call/utils.test.ts new file mode 100644 index 000000000..7f1b5efc2 --- /dev/null +++ b/packages/ui/src/components/tool-call/utils.test.ts @@ -0,0 +1,36 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { formatUnknownForCopy, limitToolOutputForRender, limitToolTitleForRender, TOOL_OUTPUT_RENDER_CHARACTER_LIMIT, TOOL_TITLE_RENDER_CHARACTER_LIMIT } from "./utils.ts" + +test("tool output rendering keeps a bounded prefix without exposing fenced tail content", () => { + const text = `HEAD${"x".repeat(20_000)}TAIL` + const rendered = limitToolOutputForRender(text) + assert.ok(rendered.length < text.length) + assert.ok(rendered.length < TOOL_OUTPUT_RENDER_CHARACTER_LIMIT + 100) + assert.ok(rendered.startsWith("HEAD")) + assert.equal(rendered.includes("TAIL"), false) +}) + +test("tool output truncation cannot expose HTML from the tail of a fenced block", () => { + const rendered = limitToolOutputForRender(`\`\`\`text\n${"x".repeat(20_000)}\n\`\`\`\n`) + assert.equal(rendered.includes(" { + const title = `title ${"x".repeat(20_000)}` + const rendered = limitToolTitleForRender(title) + assert.equal(rendered.length, TOOL_TITLE_RENDER_CHARACTER_LIMIT) + assert.ok(rendered.endsWith("...")) + assert.equal(rendered.includes("truncated"), false) +}) + +test("oversized structured output remains available for an explicit copy", () => { + const value = { output: "x".repeat(20_000) } + assert.equal(formatUnknownForCopy(value)?.text, JSON.stringify(value, null, 2)) +}) + +test("explicit copy fails safely for pathologically nested arrays", () => { + let value: unknown = "leaf" + for (let depth = 0; depth < 20_000; depth += 1) value = [value] + assert.doesNotThrow(() => formatUnknownForCopy(value)) +}) diff --git a/packages/ui/src/components/tool-call/utils.ts b/packages/ui/src/components/tool-call/utils.ts index 1ea469b5e..c8db80931 100644 --- a/packages/ui/src/components/tool-call/utils.ts +++ b/packages/ui/src/components/tool-call/utils.ts @@ -4,6 +4,7 @@ import type { ToolState } from "@opencode-ai/sdk/v2" import type { DiffPayload } from "./types" import { getLogger } from "../../lib/logger" import { tGlobal } from "../../lib/i18n" +import { exceedsRetainedByteLimit } from "../../lib/session-memory-budget" const log = getLogger("session") @@ -12,6 +13,18 @@ export type ToolStateCompleted = import("@opencode-ai/sdk/v2").ToolStateComplete export type ToolStateError = import("@opencode-ai/sdk/v2").ToolStateError export const diffCapableTools = new Set(["edit", "patch"]) +export const TOOL_OUTPUT_RENDER_CHARACTER_LIMIT = 10_000 +export const TOOL_TITLE_RENDER_CHARACTER_LIMIT = 384 + +export function limitToolOutputForRender(text: string): string { + if (text.length <= TOOL_OUTPUT_RENDER_CHARACTER_LIMIT) return text + return `${text.slice(0, TOOL_OUTPUT_RENDER_CHARACTER_LIMIT)}\n\n${tGlobal("toolCall.output.truncated")}` +} + +export function limitToolTitleForRender(text: string): string { + if (text.length <= TOOL_TITLE_RENDER_CHARACTER_LIMIT) return text + return `${text.slice(0, TOOL_TITLE_RENDER_CHARACTER_LIMIT - 3)}...` +} export function isToolStateRunning(state: ToolState): state is ToolStateRunning { return state.status === "running" @@ -150,6 +163,23 @@ export function formatUnknown(value: unknown): { text: string; language?: string return null } +export function formatUnknownForRender(value: unknown): { text: string; language?: string } | null { + if (typeof value !== "string" && exceedsRetainedByteLimit(value, TOOL_OUTPUT_RENDER_CHARACTER_LIMIT)) { + return { text: tGlobal("toolCall.output.tooLarge") } + } + const result = formatUnknown(value) + return result ? { ...result, text: limitToolOutputForRender(result.text) } : null +} + +export function formatUnknownForCopy(value: unknown): { text: string; language?: string } | null { + try { + return formatUnknown(value) + } catch (error) { + log.error("Failed to format tool call output for copy", error) + return { text: tGlobal("toolCall.output.tooLarge") } + } +} + export function inferLanguageFromPath(path?: string): string | undefined { return getLanguageFromPath(path || "") } @@ -236,13 +266,14 @@ export function buildToolSpeechText(options: { }): string { const sections: string[] = [] - if (options.title.trim()) { - sections.push(options.title.trim()) + const title = limitToolOutputForRender(options.title).trim() + if (title) { + sections.push(title) } const { input, output } = readToolStatePayload(options.state) - const formattedInput = formatUnknown(input) - const formattedOutput = formatUnknown(output) + const formattedInput = formatUnknownForRender(input) + const formattedOutput = formatUnknownForRender(output) if (formattedInput?.text?.trim()) { sections.push(`${options.t("toolCall.io.input")}:\n${formattedInput.text.trim()}`) @@ -252,13 +283,14 @@ export function buildToolSpeechText(options: { sections.push(`${options.t("toolCall.io.output")}:\n${formattedOutput.text.trim()}`) } - if (options.state?.status === "error" && options.state.error?.trim()) { - sections.push(`${options.t("toolCall.error.label")} ${options.state.error.trim()}`) + const error = options.state?.status === "error" ? limitToolOutputForRender(options.state.error ?? "").trim() : "" + if (error) { + sections.push(`${options.t("toolCall.error.label")} ${error}`) } if (sections.length === 1 && options.state?.status === "pending") { sections.push(options.t("toolCall.pending.waitingToRun")) } - return sections.join("\n\n").trim() + return limitToolOutputForRender(sections.join("\n\n").trim()) } diff --git a/packages/ui/src/components/virtual-follow-behavior.test.ts b/packages/ui/src/components/virtual-follow-behavior.test.ts index 8c4ba2626..5c0ea8c2e 100644 --- a/packages/ui/src/components/virtual-follow-behavior.test.ts +++ b/packages/ui/src/components/virtual-follow-behavior.test.ts @@ -5,6 +5,7 @@ import { ANCHOR_RESTORE_MAX_FRAMES, ANCHOR_RESTORE_STABLE_FRAMES, AnchorRestoreStabilizer, + canRestoreMessageScroll, BOTTOM_FOLLOW_EPSILON_PX, ScrollRestoreTokenGuard, VirtualScrollController, @@ -12,6 +13,7 @@ import { isAutoFollowing, isScrollRestoreGenerationCurrent, isSnapshotAutoFollowing, + shouldCancelPendingMessageScrollRestore, resolveAutoPinHoldElement, restoreFollowModeFromSnapshot, selectTopViewportAnchor, @@ -335,4 +337,20 @@ describe("virtual follow behavior", () => { assert.equal(isScrollRestoreGenerationCurrent("session-a", 3, "session-b", 4), false) assert.equal(isScrollRestoreGenerationCurrent("session-a", 3, "session-a", 4), false) }) + + it("waits for message hydration before restoring while preserving resident defaults", () => { + assert.equal(canRestoreMessageScroll(false, false), false) + assert.equal(canRestoreMessageScroll(true, false), false) + assert.equal(canRestoreMessageScroll(false, true), true) + assert.equal(canRestoreMessageScroll(false, false, true), true) + assert.equal(canRestoreMessageScroll(false, false, true, false), false) + assert.equal(canRestoreMessageScroll(true, false, true), false) + assert.equal(canRestoreMessageScroll(false), true) + }) + + it("cancels only a pending delayed restore for user scroll intent", () => { + assert.equal(shouldCancelPendingMessageScrollRestore(false, false), true) + assert.equal(shouldCancelPendingMessageScrollRestore(true, false), false) + assert.equal(shouldCancelPendingMessageScrollRestore(false, true), false) + }) }) diff --git a/packages/ui/src/components/virtual-follow-behavior.ts b/packages/ui/src/components/virtual-follow-behavior.ts index 3d604f44d..f162d5e7b 100644 --- a/packages/ui/src/components/virtual-follow-behavior.ts +++ b/packages/ui/src/components/virtual-follow-behavior.ts @@ -104,6 +104,19 @@ export function isScrollRestoreGenerationCurrent( return startedSessionId === currentSessionId && startedGeneration === currentGeneration } +export function canRestoreMessageScroll( + loading: boolean, + loadComplete?: boolean, + loadFailed = false, + failedLoadAnchorAvailable = true, +) { + return !loading && (loadComplete !== false || (loadFailed && failedLoadAnchorAvailable)) +} + +export function shouldCancelPendingMessageScrollRestore(didRestore: boolean, restoring: boolean): boolean { + return !didRestore && !restoring +} + export type FollowEffect = | { type: "none" } | { type: "scroll-top"; immediate: boolean } diff --git a/packages/ui/src/components/virtual-follow-list.tsx b/packages/ui/src/components/virtual-follow-list.tsx index 8a7f17ad4..abcee4171 100644 --- a/packages/ui/src/components/virtual-follow-list.tsx +++ b/packages/ui/src/components/virtual-follow-list.tsx @@ -83,6 +83,7 @@ export interface VirtualFollowListProps { onScrollElementChange?: (element: HTMLDivElement | undefined) => void onShellElementChange?: (element: HTMLDivElement | undefined) => void onScroll?: () => void + onUserScrollIntent?: () => void onExplicitBottomPinCancelled?: () => void onMouseUp?: (event: MouseEvent) => void onClick?: (event: MouseEvent) => void @@ -175,6 +176,7 @@ export default function VirtualFollowList(props: VirtualFollowListProps) { } function markUserScrollIntent(direction: "up" | "down" | null) { + props.onUserScrollIntent?.() cancelActiveScrollRestore() scrollController.setUserIntent(direction, performance.now() + USER_SCROLL_INTENT_WINDOW_MS) if (direction === "up") { @@ -613,11 +615,13 @@ export default function VirtualFollowList(props: VirtualFollowListProps) { if ((event.target as HTMLElement | null)?.closest(INTERACTIVE_KEY_TARGET_SELECTOR)) return if (event.key === "End") { event.preventDefault() + props.onUserScrollIntent?.() scrollToBottom(true) return } if (event.key === "Home") { event.preventDefault() + props.onUserScrollIntent?.() scrollToTop(true) return } diff --git a/packages/ui/src/lib/api-client.ts b/packages/ui/src/lib/api-client.ts index e94fd39b1..0a3937d88 100644 --- a/packages/ui/src/lib/api-client.ts +++ b/packages/ui/src/lib/api-client.ts @@ -205,8 +205,8 @@ export const serverApi = { return request(`/api/usage/${encodeURIComponent(providerId)}${query ? `?${query}` : ""}`) }, - fetchWorktrees(id: string): Promise { - return request(`/api/workspaces/${encodeURIComponent(id)}/worktrees`) + fetchWorktrees(id: string, signal?: AbortSignal): Promise { + return request(`/api/workspaces/${encodeURIComponent(id)}/worktrees`, { signal }) }, createWorktree(id: string, payload: WorktreeCreateRequest): Promise<{ slug: string; directory: string; branch?: string }> { diff --git a/packages/ui/src/lib/event-transport-contract.ts b/packages/ui/src/lib/event-transport-contract.ts index e4d91629c..a441802e5 100644 --- a/packages/ui/src/lib/event-transport-contract.ts +++ b/packages/ui/src/lib/event-transport-contract.ts @@ -9,6 +9,14 @@ export interface DesktopEventTransportStartOptions { reconnect?: Partial } +export interface DesktopEventsStartRequest extends DesktopEventTransportStartOptions { + logicalStartEpoch: number +} + +export interface DesktopEventsStartReservation { + logicalStartEpoch: number +} + export type DesktopEventTransportState = | "connecting" | "connected" @@ -41,6 +49,7 @@ export interface DesktopEventTransportStatusPayload { export interface DesktopEventsStartResult { started: boolean generation?: number + lease?: number reason?: string } diff --git a/packages/ui/src/lib/global-cache.test.ts b/packages/ui/src/lib/global-cache.test.ts new file mode 100644 index 000000000..61fc0fcb7 --- /dev/null +++ b/packages/ui/src/lib/global-cache.test.ts @@ -0,0 +1,67 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { clearCacheForInstance, getCacheEntry, setCacheEntry } from "./global-cache.ts" + +test("global render cache rejects a single oversized derived value", () => { + const params = { instanceId: "instance", sessionId: "session", scope: "markdown", cacheId: "entry", version: "1" } + try { + setCacheEntry(params, "x".repeat(3 * 1024 * 1024)) + assert.equal(getCacheEntry(params), undefined) + } finally { + clearCacheForInstance("instance") + } +}) + +test("global render cache counts raw buffers and zero-byte entries", () => { + const oversized = { instanceId: "buffer", sessionId: "session", scope: "raw", cacheId: "entry", version: "1" } + try { + setCacheEntry(oversized, new ArrayBuffer(5 * 1024 * 1024)) + assert.equal(getCacheEntry(oversized), undefined) + + for (let index = 0; index <= 4_096; index += 1) { + setCacheEntry({ instanceId: "entries", sessionId: String(index), scope: "raw", cacheId: "entry", version: "1" }, null) + } + assert.equal(getCacheEntry({ instanceId: "entries", sessionId: "0", scope: "raw", cacheId: "entry", version: "1" }), undefined) + assert.equal(getCacheEntry({ instanceId: "entries", sessionId: "4096", scope: "raw", cacheId: "entry", version: "1" }), null) + } finally { + clearCacheForInstance("buffer") + clearCacheForInstance("entries") + } +}) + +test("global render cache counts retained backing buffers and cache keys", () => { + const viewParams = { instanceId: "view", sessionId: "session", scope: "raw", cacheId: "entry", version: "1" } + const keyParams = { instanceId: "key", sessionId: "session", scope: "raw", cacheId: "x".repeat(3 * 1024 * 1024), version: "1" } + try { + setCacheEntry(viewParams, new Uint8Array(new ArrayBuffer(5 * 1024 * 1024), 0, 1)) + setCacheEntry(keyParams, null) + assert.equal(getCacheEntry(viewParams), undefined) + assert.equal(getCacheEntry(keyParams), undefined) + } finally { + clearCacheForInstance("view") + clearCacheForInstance("key") + } +}) + +test("global render cache rejects growable buffers", () => { + const BufferConstructor = ArrayBuffer as typeof ArrayBuffer & { new(length: number, options: { maxByteLength: number }): ArrayBuffer & { resize?: (length: number) => void } } + const buffer = new BufferConstructor(1, { maxByteLength: 8 * 1024 * 1024 }) + if (typeof buffer.resize !== "function") return + const params = { instanceId: "growable", sessionId: "session", scope: "raw", cacheId: "entry", version: "1" } + try { + setCacheEntry(params, buffer) + assert.equal(getCacheEntry(params), undefined) + } finally { + clearCacheForInstance("growable") + } +}) + +test("global render cache does not trust spoofed buffer tags", () => { + const params = { instanceId: "spoof", sessionId: "session", scope: "raw", cacheId: "entry", version: "1" } + try { + setCacheEntry(params, { [Symbol.toStringTag]: "ArrayBuffer", payload: "x".repeat(3 * 1024 * 1024) }) + assert.equal(getCacheEntry(params), undefined) + } finally { + clearCacheForInstance("spoof") + } +}) diff --git a/packages/ui/src/lib/global-cache.ts b/packages/ui/src/lib/global-cache.ts index 1a5cb2575..c0e659d1c 100644 --- a/packages/ui/src/lib/global-cache.ts +++ b/packages/ui/src/lib/global-cache.ts @@ -1,3 +1,5 @@ +import { estimateRetainedBytes } from "./session-memory-budget" + export interface CacheEntryBaseParams { instanceId?: string sessionId?: string @@ -12,6 +14,7 @@ export interface CacheEntryParams extends CacheEntryBaseParams { type VersionedCacheEntry = { version: string value: unknown + byteSize: number } type CacheValueMap = Map @@ -19,7 +22,37 @@ type CacheScopeMap = Map type CacheSessionMap = Map const GLOBAL_KEY = "GLOBAL" +const MAX_SCOPE_CACHE_ENTRIES = 64 +const MAX_CACHE_ENTRY_BYTES = 4 * 1024 * 1024 +const MAX_GLOBAL_CACHE_BYTES = 32 * 1024 * 1024 +const MAX_GLOBAL_CACHE_ENTRIES = 4_096 const cacheStore = new Map() +let retainedBytes = 0 +let retainedEntries = 0 + +function estimateCacheBytes(value: unknown): number { + return estimateRetainedBytes(value, MAX_CACHE_ENTRY_BYTES) +} + +function estimateCacheKeyBytes(params: CacheEntryParams): number { + return [params.instanceId, params.sessionId, params.scope, params.cacheId, params.version] + .reduce((total, value) => total + (value?.length ?? 0) * 2 + 16, 0) +} + +function recalculateRetainedBytes(): void { + retainedBytes = 0 + retainedEntries = 0 + for (const sessionMap of cacheStore.values()) { + for (const scopeMap of sessionMap.values()) { + for (const valueMap of scopeMap.values()) { + for (const entry of valueMap.values()) { + retainedBytes += entry.byteSize + retainedEntries += 1 + } + } + } + } +} function resolveKey(value?: string) { return value && value.length > 0 ? value : GLOBAL_KEY @@ -89,13 +122,42 @@ export function setCacheEntry(params: CacheEntryParams, value: T | undefined) if (value === undefined) { const existingMap = getScopeValueMap(params, false) + const existing = existingMap?.get(params.cacheId) + retainedBytes -= existing?.byteSize ?? 0 + if (existing) retainedEntries -= 1 existingMap?.delete(params.cacheId) cleanupHierarchy(instanceKey, sessionKey, params.scope) return } - const scopeEntries = getScopeValueMap(params, true) - scopeEntries?.set(params.cacheId, { version: params.version, value }) + const scopeEntries = getScopeValueMap(params, false) + const existing = scopeEntries?.get(params.cacheId) + retainedBytes -= existing?.byteSize ?? 0 + if (existing) retainedEntries -= 1 + const byteSize = estimateCacheBytes(value) + estimateCacheKeyBytes(params) + if (byteSize > MAX_CACHE_ENTRY_BYTES) { + scopeEntries?.delete(params.cacheId) + cleanupHierarchy(instanceKey, sessionKey, params.scope) + return + } + if (retainedBytes + byteSize > MAX_GLOBAL_CACHE_BYTES || retainedEntries >= MAX_GLOBAL_CACHE_ENTRIES) { + cacheStore.clear() + retainedBytes = 0 + retainedEntries = 0 + } + const target = getScopeValueMap(params, true) + if (!target) return + target.delete(params.cacheId) + target.set(params.cacheId, { version: params.version, value, byteSize }) + retainedBytes += byteSize + retainedEntries += 1 + while (target.size > MAX_SCOPE_CACHE_ENTRIES) { + const oldest = target.keys().next().value + if (oldest === undefined) break + retainedBytes -= target.get(oldest)?.byteSize ?? 0 + target.delete(oldest) + retainedEntries -= 1 + } } export function getCacheEntry(params: CacheEntryParams): T | undefined { @@ -104,6 +166,8 @@ export function getCacheEntry(params: CacheEntryParams): T | undefined { if (!entry || entry.version !== params.version) { return undefined } + scopeEntries!.delete(params.cacheId) + scopeEntries!.set(params.cacheId, entry) return entry.value as T } @@ -116,6 +180,7 @@ export function clearCacheScope(params: CacheEntryBaseParams): void { if (!scopeMap) return scopeMap.delete(params.scope) cleanupHierarchy(instanceKey, sessionKey) + recalculateRetainedBytes() } export function clearCacheForSession(instanceId?: string, sessionId?: string): void { @@ -127,10 +192,12 @@ export function clearCacheForSession(instanceId?: string, sessionId?: string): v if (sessionMap.size === 0) { cacheStore.delete(instanceKey) } + recalculateRetainedBytes() } export function clearCacheForInstance(instanceId?: string): void { const instanceKey = resolveKey(instanceId) cacheStore.delete(instanceKey) + recalculateRetainedBytes() } diff --git a/packages/ui/src/lib/hooks/use-commands.ts b/packages/ui/src/lib/hooks/use-commands.ts index e11c97650..807aebeef 100644 --- a/packages/ui/src/lib/hooks/use-commands.ts +++ b/packages/ui/src/lib/hooks/use-commands.ts @@ -7,15 +7,15 @@ import type { VisibilityPreference, } from "../../stores/preferences" import { createCommandRegistry, type Command } from "../commands" -import { activeInstanceId } from "../../stores/instances" +import { activeInstanceId, isInstanceRuntimeCurrent } from "../../stores/instances" import { selectNextAppTab, selectPreviousAppTab } from "../../stores/app-tabs" -import type { ClientPart, MessageInfo } from "../../types/message" +import type { ClientPart } from "../../types/message" import { getSessions, getVisibleSessionIds, setActiveSession, setActiveSessionFromList } from "../../stores/sessions" import { showAlertDialog } from "../../stores/alerts" import type { Instance } from "../../types/instance" import type { MessageRecord } from "../../stores/message-v2/types" import { messageStoreBus } from "../../stores/message-v2/bus" -import { cleanupBlankSessions } from "../../stores/session-state" +import { cleanupBlankSessions, invalidateSessionMessageLoad } from "../../stores/session-state" import { getLogger } from "../logger" import { requestData } from "../opencode-api" import { emitSessionSidebarRequest } from "../session-sidebar-events" @@ -299,25 +299,21 @@ export function useCommands(options: UseCommandsOptions) { const store = messageStoreBus.getOrCreate(instance.id) const messageIds = store.getSessionMessageIds(sessionId) - const infoMap = new Map() - messageIds.forEach((id) => { - const info = store.getMessageInfo(id) - if (info) infoMap.set(id, info) - }) const revertState = store.getSessionRevert(sessionId) ?? session.revert let after = 0 if (revertState?.messageID) { - const revertInfo = infoMap.get(revertState.messageID) ?? store.getMessageInfo(revertState.messageID) + const revertInfo = store.getMessageInfo(revertState.messageID) after = revertInfo?.time?.created || 0 } let messageID = "" let restoredText: string | null = null - for (let i = messageIds.length - 1; i >= 0; i--) { + const firstScannedIndex = Math.max(0, messageIds.length - 10_000) + for (let i = messageIds.length - 1; i >= firstScannedIndex; i--) { const id = messageIds[i] const record = store.getMessage(id) - const info = infoMap.get(id) ?? store.getMessageInfo(id) + const info = store.getMessageInfo(id) if (record?.role === "user" && info?.time?.created) { if (after > 0 && info.time.created >= after) { continue @@ -344,6 +340,10 @@ export function useCommands(options: UseCommandsOptions) { }), "session.revert", ) + if (!isInstanceRuntimeCurrent(instance.id, instance)) return + if (store.getSessionRevert(sessionId)?.messageID !== messageID) { + invalidateSessionMessageLoad(instance.id, sessionId) + } if (!restoredText) { const fallbackRecord = store.getMessage(messageID) diff --git a/packages/ui/src/lib/i18n/messages/de/messaging.ts b/packages/ui/src/lib/i18n/messages/de/messaging.ts index 238db557f..bfbae0e2f 100644 --- a/packages/ui/src/lib/i18n/messages/de/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/de/messaging.ts @@ -197,6 +197,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "Nachricht senden", "promptInput.send.errorFallback": "Nachricht konnte nicht gesendet werden", "promptInput.send.errorTitle": "Senden fehlgeschlagen", + "promptInput.send.ambiguousTitle": "Zustellung konnte nicht bestätigt werden", + "promptInput.send.ambiguousMessage": "Der Server hat diese Nachricht möglicherweise erhalten. Sie wurde als Entwurf wiederhergestellt und wird nur erneut gesendet, wenn Sie Senden wählen.", "promptInput.conversationMode.enable.title": "Konversationsmodus aktivieren", "promptInput.conversationMode.disable.title": "Konversationsmodus deaktivieren", "promptInput.conversationMode.error.title": "Konversationswiedergabe fehlgeschlagen", diff --git a/packages/ui/src/lib/i18n/messages/de/settings.ts b/packages/ui/src/lib/i18n/messages/de/settings.ts index 49794fa81..606b2d6bf 100644 --- a/packages/ui/src/lib/i18n/messages/de/settings.ts +++ b/packages/ui/src/lib/i18n/messages/de/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "Startverhalten", "settings.appearance.startup.subtitle": "Lege fest, was dieses Gerät beim Start von CodeNomad wiederherstellt.", "settings.appearance.startup.restore.title": "Vorherigen Zustand wiederherstellen", - "settings.appearance.startup.restore.subtitle": "Arbeitsbereichs- und Sidecar-Tabs, aktive Sitzungen und nicht gesendete Nachrichten erneut öffnen sowie Scrollpositionen, Panel-Layout, Fensterposition und Zoom wiederherstellen.", + "settings.appearance.startup.restore.subtitle": "Tabs, aktive Sitzungen und nicht gesendete Nachrichten erneut öffnen und das Layout wiederherstellen.", "settings.appearance.startup.clear.title": "Gespeicherter Startzustand", - "settings.appearance.startup.clear.subtitle": "Gespeicherte Tabs, Entwürfe, Scrollpositionen und das Panel-Layout entfernen. CodeNomad- und OpenCode-Daten werden nicht gelöscht.", + "settings.appearance.startup.clear.subtitle": "Gespeicherte Tabs, Entwürfe, Scrollpositionen und das Panel-Layout entfernen. Projektdateien und OpenCode-Unterhaltungen werden nicht gelöscht.", "settings.appearance.startup.clear.action": "Gespeicherten Zustand löschen", "settings.appearance.startup.clearSuccess": "Gespeicherter Startzustand wurde gelöscht.", "settings.appearance.startup.clearError": "Der gespeicherte Startzustand konnte nicht gelöscht werden.", diff --git a/packages/ui/src/lib/i18n/messages/de/toolCall.ts b/packages/ui/src/lib/i18n/messages/de/toolCall.ts index c736663f8..670d344c6 100644 --- a/packages/ui/src/lib/i18n/messages/de/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/de/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "Verzeichnis wird aufgelistet...", "toolCall.renderer.bash.title.timeout": "Zeitüberschreitung: {timeout}", + "toolCall.output.truncated": "[Ausgabe für die Darstellung gekürzt; kopieren Sie sie für die vollständige Ausgabe]", + "toolCall.input.tooLarge": "Die Eingabe wird nicht dargestellt, da sie zu groß ist.", + "toolCall.output.tooLarge": "Die strukturierte Ausgabe wird nicht dargestellt, da sie zu groß ist.", "toolCall.renderer.read.detail.offset": "Offset: {offset}", "toolCall.renderer.read.detail.limit": "Limit: {limit}", diff --git a/packages/ui/src/lib/i18n/messages/en/messaging.ts b/packages/ui/src/lib/i18n/messages/en/messaging.ts index d337618c2..b0616a773 100644 --- a/packages/ui/src/lib/i18n/messages/en/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/en/messaging.ts @@ -197,6 +197,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "Send message", "promptInput.send.errorFallback": "Failed to send message", "promptInput.send.errorTitle": "Send failed", + "promptInput.send.ambiguousTitle": "Delivery could not be confirmed", + "promptInput.send.ambiguousMessage": "The server may have received this message. It was restored as a draft and will not be sent again unless you choose Send.", "promptInput.conversationMode.enable.title": "Enable conversation mode", "promptInput.conversationMode.disable.title": "Disable conversation mode", "promptInput.conversationMode.error.title": "Conversation playback failed", diff --git a/packages/ui/src/lib/i18n/messages/en/settings.ts b/packages/ui/src/lib/i18n/messages/en/settings.ts index 7e79f5e06..6161618c9 100644 --- a/packages/ui/src/lib/i18n/messages/en/settings.ts +++ b/packages/ui/src/lib/i18n/messages/en/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "Startup", "settings.appearance.startup.subtitle": "Choose what this device restores when CodeNomad starts.", "settings.appearance.startup.restore.title": "Restore previous state", - "settings.appearance.startup.restore.subtitle": "Reopen workspace and sidecar tabs, active sessions, and unsent messages, and restore scroll positions, panel layout, window position, and zoom.", + "settings.appearance.startup.restore.subtitle": "Reopen workspace and sidecar tabs, active sessions, and unsent messages, and restore the layout.", "settings.appearance.startup.clear.title": "Saved startup state", - "settings.appearance.startup.clear.subtitle": "Remove saved tabs, drafts, scroll positions, and panel layout. CodeNomad and OpenCode data are not deleted.", + "settings.appearance.startup.clear.subtitle": "Remove saved tabs, drafts, scroll positions, and panel layout. Project files and OpenCode conversations are not deleted.", "settings.appearance.startup.clear.action": "Clear saved state", "settings.appearance.startup.clearSuccess": "Saved startup state cleared.", "settings.appearance.startup.clearError": "Could not clear the saved startup state.", diff --git a/packages/ui/src/lib/i18n/messages/en/toolCall.ts b/packages/ui/src/lib/i18n/messages/en/toolCall.ts index 8bf7bec07..c7659c845 100644 --- a/packages/ui/src/lib/i18n/messages/en/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/en/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "Listing directory...", "toolCall.renderer.bash.title.timeout": "Timeout: {timeout}", + "toolCall.output.truncated": "[Output truncated for rendering; copy to access the full output]", + "toolCall.input.tooLarge": "Input omitted from rendering because it is too large.", + "toolCall.output.tooLarge": "Structured output omitted from rendering because it is too large.", "toolCall.renderer.read.detail.offset": "Offset: {offset}", "toolCall.renderer.read.detail.limit": "Limit: {limit}", diff --git a/packages/ui/src/lib/i18n/messages/es/messaging.ts b/packages/ui/src/lib/i18n/messages/es/messaging.ts index 43d6530b1..5e3863645 100644 --- a/packages/ui/src/lib/i18n/messages/es/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/es/messaging.ts @@ -200,6 +200,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "Enviar mensaje", "promptInput.send.errorFallback": "No se pudo enviar el mensaje", "promptInput.send.errorTitle": "Error al enviar", + "promptInput.send.ambiguousTitle": "No se pudo confirmar la entrega", + "promptInput.send.ambiguousMessage": "Es posible que el servidor haya recibido este mensaje. Se restauro como borrador y no se volvera a enviar a menos que elijas Enviar.", "promptInput.conversationMode.enable.title": "Activar modo conversacion", "promptInput.conversationMode.disable.title": "Desactivar modo conversacion", "promptInput.conversationMode.error.title": "Fallo la reproduccion de la conversacion", diff --git a/packages/ui/src/lib/i18n/messages/es/settings.ts b/packages/ui/src/lib/i18n/messages/es/settings.ts index 268c30623..06582f9d0 100644 --- a/packages/ui/src/lib/i18n/messages/es/settings.ts +++ b/packages/ui/src/lib/i18n/messages/es/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "Inicio", "settings.appearance.startup.subtitle": "Elige qué restaura este dispositivo al iniciar CodeNomad.", "settings.appearance.startup.restore.title": "Restaurar el estado anterior", - "settings.appearance.startup.restore.subtitle": "Volver a abrir las pestañas de espacios de trabajo y sidecars, las sesiones activas y los mensajes no enviados, y restaurar las posiciones de desplazamiento, el diseño de paneles, la posición de la ventana y el zoom.", + "settings.appearance.startup.restore.subtitle": "Volver a abrir pestañas, sesiones activas y mensajes no enviados, y restaurar el diseño.", "settings.appearance.startup.clear.title": "Estado de inicio guardado", - "settings.appearance.startup.clear.subtitle": "Elimina pestañas, borradores, posiciones de desplazamiento y el diseño de paneles guardados. No se eliminan datos de CodeNomad ni de OpenCode.", + "settings.appearance.startup.clear.subtitle": "Elimina pestañas, borradores, posiciones de desplazamiento y el diseño de paneles guardados. No se eliminan los archivos del proyecto ni las conversaciones de OpenCode.", "settings.appearance.startup.clear.action": "Borrar estado guardado", "settings.appearance.startup.clearSuccess": "Se borró el estado de inicio guardado.", "settings.appearance.startup.clearError": "No se pudo borrar el estado de inicio guardado.", diff --git a/packages/ui/src/lib/i18n/messages/es/toolCall.ts b/packages/ui/src/lib/i18n/messages/es/toolCall.ts index ea206440e..3cd68824e 100644 --- a/packages/ui/src/lib/i18n/messages/es/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/es/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "Listando directorio...", "toolCall.renderer.bash.title.timeout": "Tiempo de espera: {timeout}", + "toolCall.output.truncated": "[Salida truncada para la visualización; cópiala para acceder a la salida completa]", + "toolCall.input.tooLarge": "La entrada no se muestra porque es demasiado grande.", + "toolCall.output.tooLarge": "La salida estructurada no se muestra porque es demasiado grande.", "toolCall.renderer.read.detail.offset": "Desplazamiento: {offset}", "toolCall.renderer.read.detail.limit": "Límite: {limit}", diff --git a/packages/ui/src/lib/i18n/messages/fr/messaging.ts b/packages/ui/src/lib/i18n/messages/fr/messaging.ts index 5ae20972a..612231c70 100644 --- a/packages/ui/src/lib/i18n/messages/fr/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/fr/messaging.ts @@ -200,6 +200,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "Envoyer le message", "promptInput.send.errorFallback": "Impossible d'envoyer le message", "promptInput.send.errorTitle": "Échec de l'envoi", + "promptInput.send.ambiguousTitle": "La remise n'a pas pu etre confirmee", + "promptInput.send.ambiguousMessage": "Le serveur a peut-etre recu ce message. Il a ete restaure comme brouillon et ne sera renvoye que si vous choisissez Envoyer.", "promptInput.conversationMode.enable.title": "Activer le mode conversation", "promptInput.conversationMode.disable.title": "Desactiver le mode conversation", "promptInput.conversationMode.error.title": "La lecture de la conversation a echoue", diff --git a/packages/ui/src/lib/i18n/messages/fr/settings.ts b/packages/ui/src/lib/i18n/messages/fr/settings.ts index 65c924f73..4819dccc2 100644 --- a/packages/ui/src/lib/i18n/messages/fr/settings.ts +++ b/packages/ui/src/lib/i18n/messages/fr/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "Démarrage", "settings.appearance.startup.subtitle": "Choisissez ce que cet appareil restaure au démarrage de CodeNomad.", "settings.appearance.startup.restore.title": "Restaurer l’état précédent", - "settings.appearance.startup.restore.subtitle": "Rouvrir les onglets d’espaces de travail et de sidecars, les sessions actives et les messages non envoyés, et restaurer les positions de défilement, la disposition des panneaux, la position de la fenêtre et le zoom.", + "settings.appearance.startup.restore.subtitle": "Rouvrir les onglets, les sessions actives et les messages non envoyés, et restaurer la disposition.", "settings.appearance.startup.clear.title": "État de démarrage enregistré", - "settings.appearance.startup.clear.subtitle": "Supprime les onglets, brouillons, positions de défilement et la disposition des panneaux enregistrés. Les données CodeNomad et OpenCode ne sont pas supprimées.", + "settings.appearance.startup.clear.subtitle": "Supprime les onglets, brouillons, positions de défilement et la disposition des panneaux enregistrés. Les fichiers du projet et les conversations OpenCode ne sont pas supprimés.", "settings.appearance.startup.clear.action": "Effacer l’état enregistré", "settings.appearance.startup.clearSuccess": "L’état de démarrage enregistré a été effacé.", "settings.appearance.startup.clearError": "Impossible d’effacer l’état de démarrage enregistré.", diff --git a/packages/ui/src/lib/i18n/messages/fr/toolCall.ts b/packages/ui/src/lib/i18n/messages/fr/toolCall.ts index e706e641f..5ab53393f 100644 --- a/packages/ui/src/lib/i18n/messages/fr/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/fr/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "Liste du répertoire...", "toolCall.renderer.bash.title.timeout": "Délai : {timeout}", + "toolCall.output.truncated": "[Sortie tronquée pour l’affichage ; copiez-la pour accéder à la sortie complète]", + "toolCall.input.tooLarge": "Entrée omise de l’affichage car elle est trop volumineuse.", + "toolCall.output.tooLarge": "Sortie structurée omise de l’affichage car elle est trop volumineuse.", "toolCall.renderer.read.detail.offset": "Décalage : {offset}", "toolCall.renderer.read.detail.limit": "Limite : {limit}", diff --git a/packages/ui/src/lib/i18n/messages/he/messaging.ts b/packages/ui/src/lib/i18n/messages/he/messaging.ts index a342b71e2..f24b21851 100644 --- a/packages/ui/src/lib/i18n/messages/he/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/he/messaging.ts @@ -197,6 +197,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "שלח הודעה", "promptInput.send.errorFallback": "שליחת ההודעה נכשלה", "promptInput.send.errorTitle": "השליחה נכשלה", + "promptInput.send.ambiguousTitle": "לא ניתן לאשר את המסירה", + "promptInput.send.ambiguousMessage": "ייתכן שהשרת קיבל את ההודעה. היא שוחזרה כטיוטה ולא תישלח שוב אלא אם תבחרו לשלוח.", "promptInput.conversationMode.enable.title": "הפעל מצב שיחה", "promptInput.conversationMode.disable.title": "כבה מצב שיחה", "promptInput.conversationMode.error.title": "ניגון השיחה נכשל", diff --git a/packages/ui/src/lib/i18n/messages/he/settings.ts b/packages/ui/src/lib/i18n/messages/he/settings.ts index a7f9c7fb6..f66c82637 100644 --- a/packages/ui/src/lib/i18n/messages/he/settings.ts +++ b/packages/ui/src/lib/i18n/messages/he/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "הפעלה", "settings.appearance.startup.subtitle": "בחר מה ישוחזר במכשיר זה בעת הפעלת CodeNomad.", "settings.appearance.startup.restore.title": "שחזור המצב הקודם", - "settings.appearance.startup.restore.subtitle": "פתיחה מחדש של כרטיסיות סביבות עבודה ו-sidecar, הפעלות פעילות והודעות שלא נשלחו, ושחזור מיקומי גלילה, פריסת חלוניות, מיקום החלון ורמת התקריב.", + "settings.appearance.startup.restore.subtitle": "פתיחה מחדש של כרטיסיות, הפעלות פעילות והודעות שלא נשלחו, ושחזור הפריסה.", "settings.appearance.startup.clear.title": "מצב הפעלה שמור", - "settings.appearance.startup.clear.subtitle": "הסרת כרטיסיות, טיוטות, מיקומי גלילה ופריסת חלוניות שנשמרו. נתוני CodeNomad ו-OpenCode לא יימחקו.", + "settings.appearance.startup.clear.subtitle": "הסרת כרטיסיות, טיוטות, מיקומי גלילה ופריסת חלוניות שנשמרו. קובצי הפרויקט ושיחות OpenCode לא יימחקו.", "settings.appearance.startup.clear.action": "נקה מצב שמור", "settings.appearance.startup.clearSuccess": "מצב ההפעלה השמור נוקה.", "settings.appearance.startup.clearError": "לא ניתן לנקות את מצב ההפעלה השמור.", diff --git a/packages/ui/src/lib/i18n/messages/he/toolCall.ts b/packages/ui/src/lib/i18n/messages/he/toolCall.ts index 678092bc6..67e98a582 100644 --- a/packages/ui/src/lib/i18n/messages/he/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/he/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "מפרט ספרייה...", "toolCall.renderer.bash.title.timeout": "פסק זמן: {timeout}", + "toolCall.output.truncated": "[הפלט קוצר לצורך תצוגה; יש להעתיק כדי לגשת לפלט המלא]", + "toolCall.input.tooLarge": "הקלט לא מוצג מכיוון שהוא גדול מדי.", + "toolCall.output.tooLarge": "הפלט המובנה לא מוצג מכיוון שהוא גדול מדי.", "toolCall.renderer.read.detail.offset": "היסט: {offset}", "toolCall.renderer.read.detail.limit": "מגבלה: {limit}", diff --git a/packages/ui/src/lib/i18n/messages/ja/messaging.ts b/packages/ui/src/lib/i18n/messages/ja/messaging.ts index 587ddc177..236172091 100644 --- a/packages/ui/src/lib/i18n/messages/ja/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/ja/messaging.ts @@ -200,6 +200,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "メッセージを送信", "promptInput.send.errorFallback": "メッセージの送信に失敗しました", "promptInput.send.errorTitle": "送信に失敗", + "promptInput.send.ambiguousTitle": "配信を確認できませんでした", + "promptInput.send.ambiguousMessage": "サーバーがこのメッセージを受信した可能性があります。下書きとして復元され、送信を選択しない限り再送信されません。", "promptInput.conversationMode.enable.title": "会話モードを有効化", "promptInput.conversationMode.disable.title": "会話モードを無効化", "promptInput.conversationMode.error.title": "会話の読み上げに失敗しました", diff --git a/packages/ui/src/lib/i18n/messages/ja/settings.ts b/packages/ui/src/lib/i18n/messages/ja/settings.ts index 6b266a108..f0cf76184 100644 --- a/packages/ui/src/lib/i18n/messages/ja/settings.ts +++ b/packages/ui/src/lib/i18n/messages/ja/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "起動", "settings.appearance.startup.subtitle": "CodeNomad の起動時にこのデバイスで復元する内容を選択します。", "settings.appearance.startup.restore.title": "前回の状態を復元", - "settings.appearance.startup.restore.subtitle": "ワークスペースとサイドカーのタブ、アクティブなセッション、未送信メッセージを再度開き、スクロール位置、パネル配置、ウィンドウ位置、ズームを復元します。", + "settings.appearance.startup.restore.subtitle": "タブ、アクティブなセッション、未送信メッセージを再度開き、配置を復元します。", "settings.appearance.startup.clear.title": "保存された起動状態", - "settings.appearance.startup.clear.subtitle": "保存されたタブ、下書き、スクロール位置、パネル配置を削除します。CodeNomad と OpenCode のデータは削除されません。", + "settings.appearance.startup.clear.subtitle": "保存されたタブ、下書き、スクロール位置、パネル配置を削除します。プロジェクトファイルと OpenCode の会話は削除されません。", "settings.appearance.startup.clear.action": "保存状態を消去", "settings.appearance.startup.clearSuccess": "保存された起動状態を消去しました。", "settings.appearance.startup.clearError": "保存された起動状態を消去できませんでした。", diff --git a/packages/ui/src/lib/i18n/messages/ja/toolCall.ts b/packages/ui/src/lib/i18n/messages/ja/toolCall.ts index 681d67f72..67624feb6 100644 --- a/packages/ui/src/lib/i18n/messages/ja/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/ja/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "ディレクトリ一覧を取得中...", "toolCall.renderer.bash.title.timeout": "タイムアウト: {timeout}", + "toolCall.output.truncated": "[表示用に出力を省略しました。完全な出力にアクセスするにはコピーしてください]", + "toolCall.input.tooLarge": "入力が大きすぎるため表示を省略しました。", + "toolCall.output.tooLarge": "構造化出力が大きすぎるため表示を省略しました。", "toolCall.renderer.read.detail.offset": "オフセット: {offset}", "toolCall.renderer.read.detail.limit": "上限: {limit}", diff --git a/packages/ui/src/lib/i18n/messages/ne/messaging.ts b/packages/ui/src/lib/i18n/messages/ne/messaging.ts index 5ba82e028..dd3b643d9 100644 --- a/packages/ui/src/lib/i18n/messages/ne/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/ne/messaging.ts @@ -197,6 +197,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "सन्देश पठाउनुहोस्", "promptInput.send.errorFallback": "सन्देश पठाउन असफल भयो", "promptInput.send.errorTitle": "पठाउन असफल", + "promptInput.send.ambiguousTitle": "डेलिभरी पुष्टि गर्न सकिएन", + "promptInput.send.ambiguousMessage": "सर्भरले यो सन्देश प्राप्त गरेको हुन सक्छ। यसलाई ड्राफ्टका रूपमा पुनर्स्थापित गरिएको छ र तपाईंले पठाउनुहोस् नचुनेसम्म फेरि पठाइने छैन।", "promptInput.conversationMode.enable.title": "कुराकानी मोड सक्षम गर्नुहोस्", "promptInput.conversationMode.disable.title": "कुराकानी मोड अक्षम गर्नुहोस्", "promptInput.conversationMode.error.title": "कुराकानी वाचन असफल भयो", diff --git a/packages/ui/src/lib/i18n/messages/ne/settings.ts b/packages/ui/src/lib/i18n/messages/ne/settings.ts index cbfbb0953..d9ed7f4a1 100644 --- a/packages/ui/src/lib/i18n/messages/ne/settings.ts +++ b/packages/ui/src/lib/i18n/messages/ne/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "सुरुआत", "settings.appearance.startup.subtitle": "CodeNomad सुरु हुँदा यस यन्त्रमा के पुनर्स्थापना गर्ने छान्नुहोस्।", "settings.appearance.startup.restore.title": "अघिल्लो अवस्था पुनर्स्थापना गर्नुहोस्", - "settings.appearance.startup.restore.subtitle": "कार्यस्थान र साइडकार ट्याब, सक्रिय सत्र र नपठाइएका सन्देशहरू फेरि खोल्नुहोस्, र स्क्रोल स्थिति, प्यानल लेआउट, विन्डोको स्थान र जुम पुनर्स्थापना गर्नुहोस्।", + "settings.appearance.startup.restore.subtitle": "ट्याब, सक्रिय सत्र र नपठाइएका सन्देशहरू फेरि खोल्नुहोस् र लेआउट पुनर्स्थापना गर्नुहोस्।", "settings.appearance.startup.clear.title": "सुरक्षित सुरुआत अवस्था", - "settings.appearance.startup.clear.subtitle": "सुरक्षित ट्याब, मस्यौदा, स्क्रोल स्थिति र प्यानल लेआउट हटाउनुहोस्। CodeNomad र OpenCode का डेटा मेटिँदैनन्।", + "settings.appearance.startup.clear.subtitle": "सुरक्षित ट्याब, मस्यौदा, स्क्रोल स्थिति र प्यानल लेआउट हटाउनुहोस्। परियोजना फाइलहरू र OpenCode संवादहरू मेटिँदैनन्।", "settings.appearance.startup.clear.action": "सुरक्षित अवस्था खाली गर्नुहोस्", "settings.appearance.startup.clearSuccess": "सुरक्षित सुरुआत अवस्था खाली गरियो।", "settings.appearance.startup.clearError": "सुरक्षित सुरुआत अवस्था खाली गर्न सकिएन।", diff --git a/packages/ui/src/lib/i18n/messages/ne/toolCall.ts b/packages/ui/src/lib/i18n/messages/ne/toolCall.ts index ad2b2a8c7..51356f933 100644 --- a/packages/ui/src/lib/i18n/messages/ne/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/ne/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "डाइरेक्टरी सूचीबद्ध गर्दै...", "toolCall.renderer.bash.title.timeout": "समय समाप्त: {timeout}", + "toolCall.output.truncated": "[प्रदर्शनका लागि आउटपुट छोट्याइएको छ; पूर्ण आउटपुटका लागि प्रतिलिपि गर्नुहोस्]", + "toolCall.input.tooLarge": "इनपुट धेरै ठूलो भएकाले प्रदर्शन गरिएको छैन।", + "toolCall.output.tooLarge": "संरचित आउटपुट धेरै ठूलो भएकाले प्रदर्शन गरिएको छैन।", "toolCall.renderer.read.detail.offset": "अफसेट: {offset}", "toolCall.renderer.read.detail.limit": "सीमा: {limit}", diff --git a/packages/ui/src/lib/i18n/messages/ru/messaging.ts b/packages/ui/src/lib/i18n/messages/ru/messaging.ts index 02aaa1c3c..83b736c46 100644 --- a/packages/ui/src/lib/i18n/messages/ru/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/ru/messaging.ts @@ -200,6 +200,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "Отправить сообщение", "promptInput.send.errorFallback": "Не удалось отправить сообщение", "promptInput.send.errorTitle": "Не удалось отправить", + "promptInput.send.ambiguousTitle": "Не удалось подтвердить доставку", + "promptInput.send.ambiguousMessage": "Возможно, сервер получил это сообщение. Оно восстановлено как черновик и не будет отправлено повторно, пока вы не выберете Отправить.", "promptInput.conversationMode.enable.title": "Включить режим разговора", "promptInput.conversationMode.disable.title": "Выключить режим разговора", "promptInput.conversationMode.error.title": "Сбой озвучивания разговора", diff --git a/packages/ui/src/lib/i18n/messages/ru/settings.ts b/packages/ui/src/lib/i18n/messages/ru/settings.ts index d364ab5a6..2bcce10d6 100644 --- a/packages/ui/src/lib/i18n/messages/ru/settings.ts +++ b/packages/ui/src/lib/i18n/messages/ru/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "Запуск", "settings.appearance.startup.subtitle": "Выберите, что восстанавливать на этом устройстве при запуске CodeNomad.", "settings.appearance.startup.restore.title": "Восстанавливать предыдущее состояние", - "settings.appearance.startup.restore.subtitle": "Повторно открывать вкладки рабочих пространств и sidecar, активные сессии и неотправленные сообщения, а также восстанавливать позиции прокрутки, расположение панелей, положение окна и масштаб.", + "settings.appearance.startup.restore.subtitle": "Повторно открывать вкладки, активные сессии и неотправленные сообщения и восстанавливать макет.", "settings.appearance.startup.clear.title": "Сохранённое состояние запуска", - "settings.appearance.startup.clear.subtitle": "Удалить сохранённые вкладки, черновики, позиции прокрутки и расположение панелей. Данные CodeNomad и OpenCode не удаляются.", + "settings.appearance.startup.clear.subtitle": "Удалить сохранённые вкладки, черновики, позиции прокрутки и расположение панелей. Файлы проекта и разговоры OpenCode не удаляются.", "settings.appearance.startup.clear.action": "Очистить сохранённое состояние", "settings.appearance.startup.clearSuccess": "Сохранённое состояние запуска очищено.", "settings.appearance.startup.clearError": "Не удалось очистить сохранённое состояние запуска.", diff --git a/packages/ui/src/lib/i18n/messages/ru/toolCall.ts b/packages/ui/src/lib/i18n/messages/ru/toolCall.ts index 7f18981ff..d84160e15 100644 --- a/packages/ui/src/lib/i18n/messages/ru/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/ru/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "Просмотр каталога…", "toolCall.renderer.bash.title.timeout": "Таймаут: {timeout}", + "toolCall.output.truncated": "[Вывод сокращён для отображения; скопируйте его для доступа к полному выводу]", + "toolCall.input.tooLarge": "Ввод не отображается, поскольку он слишком большой.", + "toolCall.output.tooLarge": "Структурированный вывод не отображается, поскольку он слишком большой.", "toolCall.renderer.read.detail.offset": "Смещение: {offset}", "toolCall.renderer.read.detail.limit": "Лимит: {limit}", diff --git a/packages/ui/src/lib/i18n/messages/zh-Hans/messaging.ts b/packages/ui/src/lib/i18n/messages/zh-Hans/messaging.ts index bb36a3814..02cacfacd 100644 --- a/packages/ui/src/lib/i18n/messages/zh-Hans/messaging.ts +++ b/packages/ui/src/lib/i18n/messages/zh-Hans/messaging.ts @@ -199,6 +199,8 @@ export const messagingMessages = { "promptInput.send.ariaLabel": "发送消息", "promptInput.send.errorFallback": "发送消息失败", "promptInput.send.errorTitle": "发送失败", + "promptInput.send.ambiguousTitle": "无法确认消息是否送达", + "promptInput.send.ambiguousMessage": "服务器可能已收到此消息。消息已恢复为草稿,除非您选择发送,否则不会再次发送。", "promptInput.conversationMode.enable.title": "开启对话模式", "promptInput.conversationMode.disable.title": "关闭对话模式", "promptInput.conversationMode.error.title": "对话播报失败", diff --git a/packages/ui/src/lib/i18n/messages/zh-Hans/settings.ts b/packages/ui/src/lib/i18n/messages/zh-Hans/settings.ts index 1eed86ce0..2a35480b5 100644 --- a/packages/ui/src/lib/i18n/messages/zh-Hans/settings.ts +++ b/packages/ui/src/lib/i18n/messages/zh-Hans/settings.ts @@ -129,9 +129,9 @@ export const settingsMessages = { "settings.appearance.startup.title": "启动", "settings.appearance.startup.subtitle": "选择 CodeNomad 启动时在此设备上恢复的内容。", "settings.appearance.startup.restore.title": "恢复上次状态", - "settings.appearance.startup.restore.subtitle": "重新打开工作区和边栏标签页、活动会话和未发送的消息,并恢复滚动位置、面板布局、窗口位置和缩放级别。", + "settings.appearance.startup.restore.subtitle": "重新打开标签页、活动会话和未发送的消息,并恢复布局。", "settings.appearance.startup.clear.title": "已保存的启动状态", - "settings.appearance.startup.clear.subtitle": "移除已保存的标签页、草稿、滚动位置和面板布局。不会删除 CodeNomad 或 OpenCode 数据。", + "settings.appearance.startup.clear.subtitle": "移除已保存的标签页、草稿、滚动位置和面板布局。不会删除项目文件或 OpenCode 对话。", "settings.appearance.startup.clear.action": "清除已保存状态", "settings.appearance.startup.clearSuccess": "已清除保存的启动状态。", "settings.appearance.startup.clearError": "无法清除保存的启动状态。", diff --git a/packages/ui/src/lib/i18n/messages/zh-Hans/toolCall.ts b/packages/ui/src/lib/i18n/messages/zh-Hans/toolCall.ts index 07641548e..00621c696 100644 --- a/packages/ui/src/lib/i18n/messages/zh-Hans/toolCall.ts +++ b/packages/ui/src/lib/i18n/messages/zh-Hans/toolCall.ts @@ -56,6 +56,9 @@ export const toolCallMessages = { "toolCall.renderer.action.listingDirectory": "正在列出目录...", "toolCall.renderer.bash.title.timeout": "超时:{timeout}", + "toolCall.output.truncated": "[输出已截断以便显示;复制即可访问完整输出]", + "toolCall.input.tooLarge": "输入内容过大,已省略显示。", + "toolCall.output.tooLarge": "结构化输出过大,已省略显示。", "toolCall.renderer.read.detail.offset": "偏移:{offset}", "toolCall.renderer.read.detail.limit": "限制:{limit}", diff --git a/packages/ui/src/lib/native/desktop-events-reconnect.test.ts b/packages/ui/src/lib/native/desktop-events-reconnect.test.ts new file mode 100644 index 000000000..00164a5c3 --- /dev/null +++ b/packages/ui/src/lib/native/desktop-events-reconnect.test.ts @@ -0,0 +1,154 @@ +import assert from "node:assert/strict" +import test from "node:test" + +import { connectTauriWorkspaceEvents } from "./desktop-events.ts" + +test("native reconnect opens once per connected transition", async () => { + let statusHandler: ((event: { payload: any }) => void) | undefined + const bridge = { + invoke: async (command: string) => { + if (command === "desktop_events_reserve_start") return { logicalStartEpoch: 1 } + if (command === "desktop_events_start") return { started: true, generation: 7, lease: 1 } + return undefined + }, + listen: async (eventName: string, handler: (event: { payload: any }) => void) => { + if (eventName === "desktop:event-stream-status") statusHandler = handler + return () => {} + }, + } as any + let opens = 0 + const connection = await connectTauriWorkspaceEvents({ + onBatch: () => {}, + onOpen: () => { opens += 1 }, + }, { reconnect: {} }, bridge) + + assert.ok(statusHandler) + statusHandler({ payload: { generation: 7, state: "connected" } }) + statusHandler({ payload: { generation: 7, state: "connected" } }) + assert.equal(opens, 1) + + statusHandler({ payload: { generation: 7, state: "disconnected" } }) + statusHandler({ payload: { generation: 7, state: "connected" } }) + assert.equal(opens, 2) + connection.disconnect() +}) + +test("native startup replay preserves status and batch order", async () => { + let resolveStart!: (value: { started: true; generation: number; lease: number }) => void + const start = new Promise<{ started: true; generation: number; lease: number }>((resolve) => { resolveStart = resolve }) + const handlers = new Map void>() + const bridge = { + invoke: (command: string) => command === "desktop_events_reserve_start" + ? Promise.resolve({ logicalStartEpoch: 1 }) + : start, + listen: async (eventName: string, handler: (event: { payload: any }) => void) => { + handlers.set(eventName, handler) + return () => {} + }, + } as any + const events: string[] = [] + const pending = connectTauriWorkspaceEvents({ + onBatch: () => { events.push("batch") }, + onOpen: () => { events.push("open") }, + onStatus: (status) => { events.push(status) }, + }, { reconnect: {} }, bridge) + await Promise.resolve() + await Promise.resolve() + + handlers.get("desktop:event-stream-status")?.({ payload: { generation: 7, state: "connected" } }) + handlers.get("desktop:event-batch")?.({ payload: { generation: 7, sequence: 1, emittedAt: 0, events: [{}] } }) + handlers.get("desktop:event-stream-status")?.({ payload: { generation: 7, state: "disconnected" } }) + resolveStart({ started: true, generation: 7, lease: 1 }) + const connection = await pending + + assert.deepEqual(events, ["connected", "open", "batch", "disconnected"]) + connection.disconnect() +}) + +test("a later native reservation rejects an older start delayed by listener setup", async () => { + let releaseOlderListen!: () => void + const olderListenBlocked = new Promise((resolve) => { + releaseOlderListen = resolve + }) + const stops: number[] = [] + const arrivedEpochs: number[] = [] + let listenCalls = 0 + let latestEpoch = 0 + const bridge = { + invoke: async (command: string, args?: { lease?: number; request?: { logicalStartEpoch: number } }) => { + if (command === "desktop_events_reserve_start") { + latestEpoch += 1 + return { logicalStartEpoch: latestEpoch } + } + if (command === "desktop_events_stop") { + stops.push(args?.lease ?? -1) + return + } + const epoch = args?.request?.logicalStartEpoch ?? 0 + arrivedEpochs.push(epoch) + if (epoch !== latestEpoch) return { started: false, reason: "stale logical desktop event start" } + return { started: true, generation: 1, lease: 1 } + }, + listen: async () => { + listenCalls += 1 + if (listenCalls === 1) await olderListenBlocked + return () => {} + }, + } as any + + const olderPending = connectTauriWorkspaceEvents({ onBatch: () => {} }, { reconnect: {} }, bridge) + await Promise.resolve() + await Promise.resolve() + assert.equal(arrivedEpochs.length, 0) + + const current = await connectTauriWorkspaceEvents({ onBatch: () => {} }, { reconnect: {} }, bridge) + releaseOlderListen() + await assert.rejects(olderPending, /stale logical desktop event start/) + + assert.equal(arrivedEpochs.length, 2) + assert.ok(arrivedEpochs[0] > arrivedEpochs[1]) + assert.deepEqual(stops, []) + current.disconnect() + assert.deepEqual(stops, [1]) +}) + +test("a reloaded webview module reserves a newer native start epoch", async () => { + type DesktopEventsModule = typeof import("./desktop-events.ts") + let moduleId = 0 + const loadModule = () => import(`./desktop-events.ts?reload=${moduleId++}`) as Promise + const firstModule = await loadModule() + const reloadedModule = await loadModule() + const reservations: number[] = [] + const starts: number[] = [] + let latestEpoch = 40 + let activeLease: number | undefined + const bridge = { + invoke: async (command: string, args?: { lease?: number; request?: { logicalStartEpoch: number } }) => { + if (command === "desktop_events_reserve_start") { + latestEpoch += 1 + reservations.push(latestEpoch) + return { logicalStartEpoch: latestEpoch } + } + if (command === "desktop_events_start") { + const epoch = args?.request?.logicalStartEpoch ?? 0 + starts.push(epoch) + if (epoch !== latestEpoch) return { started: false, reason: "stale logical desktop event start" } + activeLease = epoch + return { started: true, generation: epoch, lease: epoch } + } + if (command === "desktop_events_stop" && args?.lease === activeLease) activeLease = undefined + return undefined + }, + listen: async () => () => {}, + } as any + + const first = await firstModule.connectTauriWorkspaceEvents({ onBatch: () => {} }, { reconnect: {} }, bridge) + const reloaded = await reloadedModule.connectTauriWorkspaceEvents({ onBatch: () => {} }, { reconnect: {} }, bridge) + + assert.deepEqual(reservations, [41, 42]) + assert.deepEqual(starts, [41, 42]) + first.disconnect() + assert.equal(activeLease, 42) + reloaded.disconnect() + assert.equal(activeLease, undefined) +}) diff --git a/packages/ui/src/lib/native/desktop-events.test.ts b/packages/ui/src/lib/native/desktop-events.test.ts index c6d0e3882..4bb8c53e9 100644 --- a/packages/ui/src/lib/native/desktop-events.test.ts +++ b/packages/ui/src/lib/native/desktop-events.test.ts @@ -44,8 +44,11 @@ describe("connectTauriWorkspaceEvents", () => { const unlistened: string[] = [] const bridge = { invoke: async (command: string) => { + if (command === "desktop_events_reserve_start") { + return { logicalStartEpoch: 1 } + } if (command === "desktop_events_start") { - return { started: true, generation: 1 } + return { started: true, generation: 1, lease: 1 } } if (command === "desktop_events_stop") { return undefined diff --git a/packages/ui/src/lib/native/desktop-events.ts b/packages/ui/src/lib/native/desktop-events.ts index 2ee51d4cd..f17fa3992 100644 --- a/packages/ui/src/lib/native/desktop-events.ts +++ b/packages/ui/src/lib/native/desktop-events.ts @@ -2,6 +2,8 @@ import { invoke } from "@tauri-apps/api/core" import { listen } from "@tauri-apps/api/event" import type { WorkspaceEventPayload } from "../../../../server/src/api-types" import type { + DesktopEventsStartRequest, + DesktopEventsStartReservation, DesktopEventsStartResult, DesktopEventTransportStartOptions, DesktopEventTransportState, @@ -55,20 +57,31 @@ export async function connectTauriWorkspaceEvents( options: DesktopEventTransportStartOptions, bridge: DesktopEventTransportBridge = defaultDesktopEventTransportBridge, ): Promise { + const reservation = await bridge.invoke("desktop_events_reserve_start") + if (!Number.isSafeInteger(reservation?.logicalStartEpoch) || reservation.logicalStartEpoch < 1) { + throw new Error("desktop event transport did not return a valid start reservation") + } + const request: DesktopEventsStartRequest = { + ...options, + logicalStartEpoch: reservation.logicalStartEpoch, + } let closed = false - let opened = false + let connected = false + let lease!: number let expectedGeneration: number | null = null const notifyTerminalError = createTerminalErrorNotifier(callbacks) - const pendingBatches: WorkspaceEventBatchPayload[] = [] - const pendingStatuses: DesktopEventTransportStatusPayload[] = [] + const pendingEvents: Array< + | { type: "batch"; payload: WorkspaceEventBatchPayload } + | { type: "status"; payload: DesktopEventTransportStatusPayload } + > = [] const matchesGeneration = (generation: number) => expectedGeneration === generation const handleBatchPayload = (payload: WorkspaceEventBatchPayload) => { if (!payload || !matchesGeneration(payload.generation)) return - if (!opened) { - opened = true + if (!connected) { + connected = true callbacks.onStatus?.("connected") callbacks.onOpen?.() } @@ -86,9 +99,13 @@ export async function connectTauriWorkspaceEvents( callbacks.onStatus?.(mapDesktopEventTransportStatus(payload.state)) - if (payload.state === "connected" && !opened) { - opened = true - callbacks.onOpen?.() + if (payload.state === "connected") { + if (!connected) { + connected = true + callbacks.onOpen?.() + } + } else { + connected = false } if (payload.state === "unauthorized") { @@ -126,11 +143,9 @@ export async function connectTauriWorkspaceEvents( const flushPending = () => { if (expectedGeneration === null) return - for (const payload of pendingStatuses.splice(0, pendingStatuses.length)) { - handleStatusPayload(payload) - } - for (const payload of pendingBatches.splice(0, pendingBatches.length)) { - handleBatchPayload(payload) + for (const event of pendingEvents.splice(0, pendingEvents.length)) { + if (event.type === "status") handleStatusPayload(event.payload) + else handleBatchPayload(event.payload) } } @@ -139,7 +154,7 @@ export async function connectTauriWorkspaceEvents( const payload = event.payload if (!payload) return if (expectedGeneration === null) { - pendingBatches.push(payload) + pendingEvents.push({ type: "batch", payload }) return } handleBatchPayload(payload) @@ -150,18 +165,25 @@ export async function connectTauriWorkspaceEvents( const payload = event.payload if (!payload) return if (expectedGeneration === null) { - pendingStatuses.push(payload) + pendingEvents.push({ type: "status", payload }) return } handleStatusPayload(payload) }) try { - const result = await bridge.invoke("desktop_events_start", { request: options }) + const result = await bridge.invoke("desktop_events_start", { request }) if (!result?.started) { throw new Error(result?.reason ?? "desktop event transport unavailable") } - expectedGeneration = result.generation ?? null + if (result.generation === undefined) { + throw new Error("desktop event transport did not return a generation") + } + if (result.lease === undefined) { + throw new Error("desktop event transport did not return a lease") + } + lease = result.lease + expectedGeneration = result.generation flushPending() } catch (error) { unlistenBatch() @@ -178,7 +200,7 @@ export async function connectTauriWorkspaceEvents( closed = true unlistenBatch() unlistenStatus() - void bridge.invoke("desktop_events_stop").catch((error) => { + void bridge.invoke("desktop_events_stop", { lease }).catch((error) => { log.warn("Failed to stop native desktop event transport", error) }) }, diff --git a/packages/ui/src/lib/opencode-api.test.ts b/packages/ui/src/lib/opencode-api.test.ts index 0897e435a..d85596e82 100644 --- a/packages/ui/src/lib/opencode-api.test.ts +++ b/packages/ui/src/lib/opencode-api.test.ts @@ -1,7 +1,7 @@ import assert from "node:assert/strict" import { describe, it } from "node:test" -import { getOpencodeErrorMessage } from "./opencode-api.ts" +import { getOpencodeErrorMessage, isDeliveryAmbiguousError } from "./opencode-api.ts" describe("getOpencodeErrorMessage", () => { it("uses the detailed message from a nested SDK error", () => { @@ -38,3 +38,27 @@ describe("getOpencodeErrorMessage", () => { assert.equal(getOpencodeErrorMessage(first, "Unable to load sessions"), "Unable to load sessions") }) }) + +describe("isDeliveryAmbiguousError", () => { + it("recognizes standard transport codes and status zero", () => { + assert.equal(isDeliveryAmbiguousError({ code: "ECONNRESET", message: "socket hang up" }), true) + assert.equal(isDeliveryAmbiguousError({ code: "ETIMEDOUT", message: "timed out" }), true) + assert.equal(isDeliveryAmbiguousError({ status: 0, cause: new TypeError("Failed to fetch") }), true) + }) + + it("recognizes common browser fetch failure messages", () => { + for (const message of ["Network request failed", "Load failed", "fetch failed"]) { + assert.equal(isDeliveryAmbiguousError(new TypeError(message)), true) + } + }) + + it("keeps definite HTTP failures replayable", () => { + assert.equal(isDeliveryAmbiguousError({ response: { status: 400 }, error: { message: "Bad command" } }), false) + assert.equal(isDeliveryAmbiguousError({ response: { status: 429 }, error: { message: "Rate limited" } }), false) + }) + + it("treats unknown and post-success parsing failures as ambiguous", () => { + assert.equal(isDeliveryAmbiguousError(new SyntaxError("Unexpected end of JSON input")), true) + assert.equal(isDeliveryAmbiguousError({ response: { status: 200 }, error: new TypeError("terminated") }), true) + }) +}) diff --git a/packages/ui/src/lib/opencode-api.ts b/packages/ui/src/lib/opencode-api.ts index 70bbd73d8..32a473220 100644 --- a/packages/ui/src/lib/opencode-api.ts +++ b/packages/ui/src/lib/opencode-api.ts @@ -33,14 +33,35 @@ export function getOpencodeErrorMessage(error: unknown, fallback: string): strin return extract(error) ?? fallback } +export function isDeliveryAmbiguousError(error: unknown): boolean { + const seen = new Set() + const pending = [error] + while (pending.length > 0) { + const value = pending.pop() + if (!value || typeof value !== "object" || seen.has(value)) continue + seen.add(value) + const candidate = value as any + const status = candidate.status ?? candidate.statusCode ?? candidate.response?.status + if (typeof status === "number") { + if (status >= 400 && status < 500 && status !== 408) return false + } + pending.push(candidate.cause, candidate.error, candidate.response) + } + // Once dispatch has started, unknown failures are ambiguous unless the server + // returned a definitive client rejection. + return true +} + type RequestResultLike = | { data: T error?: undefined + response?: { status?: number } } | { data?: undefined error: unknown + response?: { status?: number } } export async function requestData( @@ -52,7 +73,10 @@ export async function requestData( throw new OpencodeApiError(`${label} returned no result`) } if ((result as any).error) { - throw new OpencodeApiError(`${label} failed`, { cause: (result as any).error }) + const response = (result as any).response + throw new OpencodeApiError(`${label} failed`, { + cause: response ? { error: (result as any).error, response } : (result as any).error, + }) } return (result as any).data as T } diff --git a/packages/ui/src/lib/retry-utils.ts b/packages/ui/src/lib/retry-utils.ts index 54350105e..52359e272 100644 --- a/packages/ui/src/lib/retry-utils.ts +++ b/packages/ui/src/lib/retry-utils.ts @@ -1,3 +1,19 @@ +const RETRYABLE_TRANSPORT_CODES = new Set([ + "ECONNABORTED", + "ECONNREFUSED", + "ECONNRESET", + "EHOSTUNREACH", + "ENETDOWN", + "ENETUNREACH", + "ENOTFOUND", + "EPIPE", + "ETIMEDOUT", + "UND_ERR_BODY_TIMEOUT", + "UND_ERR_CONNECT_TIMEOUT", + "UND_ERR_HEADERS_TIMEOUT", + "UND_ERR_SOCKET", +]) + interface RetryOptions { maxAttempts?: number initialDelayMs?: number @@ -29,14 +45,22 @@ export async function retryWithBackoff( try { if (timeoutMs) { const controller = new AbortController() - const timer = setTimeout(() => controller.abort(), timeoutMs) + let timer: ReturnType | undefined try { - const result = await fn(controller.signal) - clearTimeout(timer) - return result - } catch (error) { - clearTimeout(timer) - throw error + const operation = fn(controller.signal) + return await Promise.race([ + operation, + new Promise((_, reject) => { + timer = setTimeout(() => { + const error = new Error(`Operation timed out after ${timeoutMs}ms`) + error.name = "TimeoutError" + controller.abort(error) + reject(error) + }, timeoutMs) + }), + ]) + } finally { + if (timer) clearTimeout(timer) } } @@ -58,9 +82,8 @@ export async function retryWithBackoff( } export function isRetryableError(error: Error): boolean { - if (error.name === "AbortError" || error.name === "TimeoutError") return true - if (error.message.includes("Failed to fetch")) return true - if (error.message.includes("NetworkError")) return true - if (error.message.includes("timeout")) return true - return false + const code = (error as Error & { code?: unknown }).code + if (typeof code === "string" && RETRYABLE_TRANSPORT_CODES.has(code.toUpperCase())) return true + if (error.name === "AbortError" || error.name === "NetworkError" || error.name === "TimeoutError") return true + return /failed to fetch|fetch failed|network request failed|networkerror|load failed|timed?\s*out|timeout|socket hang up|connection reset/i.test(error.message) } diff --git a/packages/ui/src/lib/server-events.ts b/packages/ui/src/lib/server-events.ts index bb747f1a0..db209404f 100644 --- a/packages/ui/src/lib/server-events.ts +++ b/packages/ui/src/lib/server-events.ts @@ -48,7 +48,10 @@ class ServerEvents { try { const connection = await connectWorkspaceEvents({ - onBatch: (events) => this.dispatchBatch(events), + onBatch: (events) => { + if (generation !== this.connectGeneration) return + this.dispatchBatch(events) + }, onError: () => { if (generation !== this.connectGeneration) { return @@ -70,6 +73,7 @@ class ServerEvents { this.openHandlers.forEach((handler) => handler()) }, onPing: (payload) => { + if (generation !== this.connectGeneration) return const identity = getClientIdentity() const pongPayload = { ...identity, pingTs: payload.ts } @@ -179,6 +183,7 @@ class ServerEvents { restart(reason = "manual restart"): void { this.retryDelay = RETRY_BASE_DELAY this.clearReconnectTimer() + this.emitTransportStatus("disconnected") if (this.connection) { this.connection.disconnect() diff --git a/packages/ui/src/lib/session-memory-budget.test.ts b/packages/ui/src/lib/session-memory-budget.test.ts new file mode 100644 index 000000000..d30a61ad9 --- /dev/null +++ b/packages/ui/src/lib/session-memory-budget.test.ts @@ -0,0 +1,63 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { + estimateRetainedBytes, + estimateRetainedValuesIncrementally, + exceedsRetainedByteLimit, + selectSessionMemoryEvictions, +} from "./session-memory-budget.ts" + +test("session memory eviction applies one byte budget across workspaces and subagents", () => { + const entries = Array.from({ length: 5 }, (_, workspace) => [ + { key: `${workspace}:parent`, byteSize: 8, lastTouched: workspace * 2 + 1, protected: workspace === 4 }, + { key: `${workspace}:subagent`, byteSize: 8, lastTouched: workspace * 2 + 2, protected: false }, + ]).flat() + + const evictions = selectSessionMemoryEvictions(entries, 40) + assert.deepEqual(evictions, ["0:parent", "0:subagent", "1:parent", "1:subagent", "2:parent"]) + assert.equal(evictions.includes("4:parent"), false) +}) + +test("session memory eviction favors old large sessions and permits protected overage", () => { + assert.deepEqual(selectSessionMemoryEvictions([ + { key: "visible", byteSize: 12, lastTouched: 3, protected: true }, + { key: "old-large", byteSize: 10, lastTouched: 1, protected: false }, + { key: "new-small", byteSize: 2, lastTouched: 2, protected: false }, + ], 12), ["old-large", "new-small"]) + assert.deepEqual(selectSessionMemoryEvictions([{ key: "visible", byteSize: 20, lastTouched: 1, protected: true }], 10), []) +}) + +test("retained byte estimates handle cycles without allocating serialized copies", () => { + const value: { text: string; self?: unknown } = { text: "hello" } + value.self = value + assert.ok(estimateRetainedBytes(value) >= 26) + assert.equal(exceedsRetainedByteLimit(Array.from({ length: 1_000 }, () => ({})), 10_000), true) +}) + +test("incremental retained byte estimates yield without changing the estimate", async () => { + const values = Array.from({ length: 20 }, (_, index) => ({ index, text: `message-${index}` })) + let yields = 0 + const measured = await estimateRetainedValuesIncrementally(values, { + yieldEvery: 5, + yieldControl: async () => { yields += 1 }, + }) + + assert.equal(measured, values.reduce((total, value) => total + estimateRetainedBytes(value), 0)) + assert.ok(yields > 1) +}) + +test("incremental retained byte estimates stop at a cancellation checkpoint", async () => { + const controller = new AbortController() + let yields = 0 + await assert.rejects(estimateRetainedValuesIncrementally([ + Array.from({ length: 1_000 }, (_, index) => ({ index })), + ], { + signal: controller.signal, + yieldEvery: 10, + yieldControl: async () => { + yields += 1 + controller.abort(new Error("measurement replaced")) + }, + }), /measurement replaced/) + assert.equal(yields, 1) +}) diff --git a/packages/ui/src/lib/session-memory-budget.ts b/packages/ui/src/lib/session-memory-budget.ts new file mode 100644 index 000000000..97f0171b7 --- /dev/null +++ b/packages/ui/src/lib/session-memory-budget.ts @@ -0,0 +1,153 @@ +export const MAX_HOT_SESSION_MESSAGE_BYTES = 64 * 1024 * 1024 + +export interface SessionMemoryEntry { + key: string + byteSize: number + lastTouched: number + protected: boolean +} + +const arrayBufferByteLength = Object.getOwnPropertyDescriptor(ArrayBuffer.prototype, "byteLength")?.get +const arrayBufferResizable = Object.getOwnPropertyDescriptor(ArrayBuffer.prototype, "resizable")?.get +const sharedBufferPrototype = typeof SharedArrayBuffer === "undefined" ? undefined : SharedArrayBuffer.prototype +const sharedBufferByteLength = sharedBufferPrototype && Object.getOwnPropertyDescriptor(sharedBufferPrototype, "byteLength")?.get +const sharedBufferGrowable = sharedBufferPrototype && Object.getOwnPropertyDescriptor(sharedBufferPrototype, "growable")?.get + +function readBufferSize(value: object): { byteLength: number; growable: boolean } | undefined { + try { + if (arrayBufferByteLength) { + const byteLength = arrayBufferByteLength.call(value) as number + return { byteLength, growable: Boolean(arrayBufferResizable?.call(value)) } + } + } catch { + // Not an ArrayBuffer. + } + try { + if (sharedBufferByteLength) { + const byteLength = sharedBufferByteLength.call(value) as number + return { byteLength, growable: Boolean(sharedBufferGrowable?.call(value)) } + } + } catch { + // Not a SharedArrayBuffer. + } + return undefined +} + +function measureRetainedBytes(value: unknown, limit: number): number { + const seen = new WeakSet() + const pending: unknown[] = [value] + let total = 0 + while (pending.length > 0 && total <= limit) { + const current = pending.pop() + if (typeof current === "string") total += current.length * 2 + 16 + else if (typeof current === "number" || typeof current === "bigint") total += 8 + else if (typeof current === "boolean") total += 4 + else if (current && typeof current === "object") { + const buffer = readBufferSize(current) + if (buffer) { + total += buffer.growable ? limit + 1 : buffer.byteLength + continue + } + if (ArrayBuffer.isView(current)) { + const backing = readBufferSize(current.buffer) + total += backing?.growable ? limit + 1 : backing?.byteLength ?? current.byteLength + continue + } + if (seen.has(current)) continue + seen.add(current) + total += Array.isArray(current) ? 24 + current.length * 8 : 32 + for (const key in current) { + if (!Object.prototype.hasOwnProperty.call(current, key)) continue + total += key.length * 2 + 8 + if (total > limit) break + pending.push((current as Record)[key]) + } + } + } + return total +} + +export function estimateRetainedBytes(value: unknown, limit = Number.POSITIVE_INFINITY): number { + return measureRetainedBytes(value, limit) +} + +export async function estimateRetainedValuesIncrementally( + values: Iterable, + options: { + signal?: AbortSignal + yieldEvery?: number + yieldControl?: () => Promise + } = {}, +): Promise { + const yieldEvery = Math.max(1, options.yieldEvery ?? 250) + const yieldControl = options.yieldControl ?? (() => new Promise((resolve) => setTimeout(resolve, 0))) + let processed = 0 + let total = 0 + + const checkpoint = (): Promise | undefined => { + if (options.signal?.aborted) throw options.signal.reason ?? new Error("Session memory measurement aborted") + processed += 1 + if (processed % yieldEvery !== 0) return + return yieldControl().then(() => { + if (options.signal?.aborted) throw options.signal.reason ?? new Error("Session memory measurement aborted") + }) + } + + for (const value of values) { + const seen = new WeakSet() + const pending: unknown[] = [value] + while (pending.length > 0) { + const pause = checkpoint() + if (pause) await pause + const current = pending.pop() + if (typeof current === "string") total += current.length * 2 + 16 + else if (typeof current === "number" || typeof current === "bigint") total += 8 + else if (typeof current === "boolean") total += 4 + else if (current && typeof current === "object") { + const buffer = readBufferSize(current) + if (buffer) { + total += buffer.growable ? Number.POSITIVE_INFINITY : buffer.byteLength + continue + } + if (ArrayBuffer.isView(current)) { + const backing = readBufferSize(current.buffer) + total += backing?.growable ? Number.POSITIVE_INFINITY : backing?.byteLength ?? current.byteLength + continue + } + if (seen.has(current)) continue + seen.add(current) + total += Array.isArray(current) ? 24 + current.length * 8 : 32 + for (const key in current) { + if (!Object.prototype.hasOwnProperty.call(current, key)) continue + const propertyPause = checkpoint() + if (propertyPause) await propertyPause + total += key.length * 2 + 8 + pending.push((current as Record)[key]) + } + } + } + } + return total +} + +export function exceedsRetainedByteLimit(value: unknown, limit: number): boolean { + return measureRetainedBytes(value, limit) > limit +} + +export function selectSessionMemoryEvictions( + entries: readonly SessionMemoryEntry[], + byteLimit = MAX_HOT_SESSION_MESSAGE_BYTES, +): string[] { + let total = entries.reduce((sum, entry) => sum + Math.max(0, entry.byteSize), 0) + if (total <= byteLimit) return [] + const candidates = entries + .filter((entry) => !entry.protected) + .sort((left, right) => left.lastTouched - right.lastTouched || right.byteSize - left.byteSize || left.key.localeCompare(right.key)) + const evictions: string[] = [] + for (const entry of candidates) { + if (total <= byteLimit) break + evictions.push(entry.key) + total -= Math.max(0, entry.byteSize) + } + return evictions +} diff --git a/packages/ui/src/lib/session-search-matches.ts b/packages/ui/src/lib/session-search-matches.ts new file mode 100644 index 000000000..673d3e377 --- /dev/null +++ b/packages/ui/src/lib/session-search-matches.ts @@ -0,0 +1,78 @@ +export interface TextSearchOccurrence { + start: number + end: number + occurrence: number + preview: string +} + +const PREVIEW_RADIUS = 56 +const LITERAL_REGEX_CHUNK_CHARACTERS = 4_096 + +function escapeLiteralPattern(value: string): string { + return value.replace(/[.*+?^${}()|[\]\\]/g, "\\$&") +} + +export function makeTextSearchPreview(text: string, start: number, end: number): string { + const previewStart = Math.max(0, start - PREVIEW_RADIUS) + const previewEnd = Math.min(text.length, end + PREVIEW_RADIUS) + const value = end - start <= PREVIEW_RADIUS * 2 + ? text.slice(previewStart, previewEnd) + : `${text.slice(previewStart, start + PREVIEW_RADIUS)}...${text.slice(end - PREVIEW_RADIUS, previewEnd)}` + return `${previewStart > 0 ? "..." : ""}${value.replace(/\s+/g, " ").trim()}${previewEnd < text.length ? "..." : ""}` +} + +export function* iterateTextSearchOccurrences(text: string, query: string, from = 0): Generator<{ start: number; end: number }> { + if (!query) return + if (query.length <= LITERAL_REGEX_CHUNK_CHARACTERS) { + const pattern = new RegExp(escapeLiteralPattern(query), "giu") + pattern.lastIndex = Math.max(0, from) + for (const match of text.matchAll(pattern)) { + const start = match.index + if (start === undefined) continue + yield { start, end: start + match[0].length } + } + return + } + + const foldedText = text.toLowerCase() + const foldedQuery = query.toLowerCase() + let foldedFrom = text.slice(0, Math.max(0, from)).toLowerCase().length + let originalOffset = 0 + let foldedOffset = 0 + const originalBoundary = (target: number): number | undefined => { + if (target < foldedOffset) return undefined + while (foldedOffset < target && originalOffset < text.length) { + const width = text.codePointAt(originalOffset)! > 0xffff ? 2 : 1 + foldedOffset += text.slice(originalOffset, originalOffset + width).toLowerCase().length + originalOffset += width + } + return foldedOffset === target ? originalOffset : undefined + } + + while (true) { + const foldedStart = foldedText.indexOf(foldedQuery, foldedFrom) + if (foldedStart < 0) return + const start = originalBoundary(foldedStart) + const end = originalBoundary(foldedStart + foldedQuery.length) + if (start !== undefined && end !== undefined && text.slice(start, end).toLowerCase() === foldedQuery) { + yield { start, end } + foldedFrom = foldedStart + foldedQuery.length + } else { + foldedFrom = foldedStart + 1 + } + } +} + +export function findTextSearchOccurrences(text: string, query: string, limit: number): TextSearchOccurrence[] { + const matches: TextSearchOccurrence[] = [] + for (const { start, end } of iterateTextSearchOccurrences(text, query)) { + if (matches.length >= limit) break + matches.push({ + start, + end, + occurrence: matches.length, + preview: makeTextSearchPreview(text, start, end), + }) + } + return matches +} diff --git a/packages/ui/src/lib/session-search.test.ts b/packages/ui/src/lib/session-search.test.ts new file mode 100644 index 000000000..ae134fc33 --- /dev/null +++ b/packages/ui/src/lib/session-search.test.ts @@ -0,0 +1,43 @@ +import assert from "node:assert/strict" +import { readFile } from "node:fs/promises" +import test from "node:test" +import { findTextSearchOccurrences } from "./session-search-matches.ts" + +test("session search bounds matches from repetitive output", () => { + assert.equal(findTextSearchOccurrences("match ".repeat(2_000), "match", 1_000).length, 1_000) +}) + +test("text occurrence search is case-insensitive on the original string", () => { + assert.equal(findTextSearchOccurrences("Needle", "needle", 1)[0]?.start, 0) +}) + +test("text occurrence search retains original offsets around expanding case folds", () => { + assert.deepEqual(findTextSearchOccurrences("İİxy", "xy", 1)[0], { + start: 2, + end: 4, + occurrence: 0, + preview: "İİxy", + }) +}) + +test("text occurrence search retains original offsets around contextual case folds", () => { + assert.deepEqual(findTextSearchOccurrences("I\u0307abc", "ABC", 1)[0], { + start: 2, + end: 5, + occurrence: 0, + preview: "I\u0307abc", + }) +}) + +test("oversized literal search preserves the complete case-insensitive query without one oversized regex", () => { + const query = "Ab".repeat(50_000) + const prefix = "prefix:" + const result = findTextSearchOccurrences(`${prefix}${query.toLowerCase()}:suffix`, query, 1)[0] + assert.equal(result?.start, prefix.length) + assert.equal(result?.end, prefix.length + query.length) +}) + +test("message search errors settle the current revision before refresh scheduling", async () => { + const source = await readFile(new URL("../components/message-section.tsx", import.meta.url), "utf8") + assert.match(source, /catch \(error\)[\s\S]*lastCompletedSearchRevision = store\(\)\.getSessionRevision\(props\.sessionId\)[\s\S]*searchRefreshRequested = false[\s\S]*finally/) +}) diff --git a/packages/ui/src/lib/session-search.ts b/packages/ui/src/lib/session-search.ts index cd894780f..db674ead7 100644 --- a/packages/ui/src/lib/session-search.ts +++ b/packages/ui/src/lib/session-search.ts @@ -2,8 +2,9 @@ import type { ClientPart, MessageInfo } from "../types/message" import { isHiddenSyntheticTextPart } from "../types/message" import type { InstanceMessageStore } from "../stores/message-v2/instance-store" import type { MessageRecord, MessageRole } from "../stores/message-v2/types" -import { resolveToolRenderer } from "../components/tool-call/renderers" import { getDefaultToolSearchText } from "../components/tool-call/search-text" +import type { ToolSearchTextContext } from "../components/tool-call/types" +import { iterateTextSearchOccurrences, makeTextSearchPreview } from "./session-search-matches" export interface SessionSearchMatch { id: string @@ -17,176 +18,279 @@ export interface SessionSearchMatch { preview: string } -interface SearchablePartText { - partId?: string - partType?: string - text: string -} +type ToolSearchTextSource = AsyncIterable | Promise | string[] | undefined export interface BuildSessionSearchMatchesOptions { store: InstanceMessageStore sessionId: string query: string includeThinking: boolean + resolveToolSearchText?: (context: ToolSearchTextContext) => ToolSearchTextSource + limit?: number + signal?: AbortSignal + yieldControl?: () => Promise } -const PREVIEW_RADIUS = 56 +export interface SessionSearchResult { + matches: SessionSearchMatch[] + totalMatches: number | null + offset: number + hasMore: boolean +} -function normalizeSearchValue(value: string): string { - return value.toLocaleLowerCase() +export interface SessionSearchPager { + nextPage(): Promise } -function segmentToText(segment: unknown): string { - if (typeof segment === "string") return segment - if (Array.isArray(segment)) return segment.map((entry) => segmentToText(entry)).filter(Boolean).join("\n") - if (!segment || typeof segment !== "object") return "" +export const SESSION_SEARCH_PAGE_SIZE = 250 +export const SESSION_SEARCH_RETAINED_PAGE_LIMIT = 3 +const SEARCH_TEXT_CHUNK_CHARACTERS = 64 * 1024 +const OVERSIZED_LITERAL_QUERY_CHARACTERS = 4_096 +const SEARCH_YIELD_INTERVAL_MS = 8 +const SEARCH_CHECKPOINT_UNITS = 64 - const candidate = segment as { text?: unknown; value?: unknown; content?: unknown[] } - const parts: string[] = [] - if (typeof candidate.text === "string") parts.push(candidate.text) - if (typeof candidate.value === "string") parts.push(candidate.value) - if (Array.isArray(candidate.content)) { - parts.push(candidate.content.map((entry) => segmentToText(entry)).filter(Boolean).join("\n")) - } - return parts.filter(Boolean).join("\n") +function yieldToEventLoop(): Promise { + const schedulerYield = (globalThis as any).scheduler?.yield + if (typeof schedulerYield === "function") return schedulerYield.call((globalThis as any).scheduler) + const setImmediate = (globalThis as any).setImmediate + if (typeof setImmediate === "function") return new Promise((resolve) => setImmediate(resolve)) + return new Promise((resolve) => setTimeout(resolve, 0)) } -function extractReasoningText(part: ClientPart): string { - const text = segmentToText((part as any).text) - const content = Array.isArray((part as any).content) - ? (part as any).content.map((entry: unknown) => segmentToText(entry)).filter(Boolean).join("\n") - : "" - return [text, content].filter(Boolean).join("\n") +function abortSearch(): never { + const error = new Error("Session search cancelled") + error.name = "AbortError" + throw error } -function extractGenericPartText(part: ClientPart): string { - const candidate = part as Record - const values = [ - candidate.text, - candidate.content, - candidate.value, - candidate.title, - candidate.name, - candidate.filename, - candidate.message, - ] - return values.map((value) => segmentToText(value)).filter(Boolean).join("\n") -} - -function extractToolText(part: Extract): string { - const toolName = typeof part.tool === "string" ? part.tool : "" - const context = { toolCall: part, toolState: (part as any).state, toolName } - const renderer = resolveToolRenderer(toolName) - const values = renderer.getSearchText?.(context) ?? getDefaultToolSearchText(context) - return values.filter((value) => value.trim().length > 0).join("\n") +async function* segmentTexts(segment: unknown, checkpoint: () => Promise): AsyncGenerator { + type Pending = { value: unknown } | { array: unknown[]; index: number } + const pending: Pending[] = [{ value: segment }] + const seen = new WeakSet() + while (pending.length > 0) { + const item = pending.pop()! + if ("array" in item) { + if (item.index < item.array.length) pending.push({ array: item.array, index: item.index + 1 }, { value: item.array[item.index] }) + continue + } + const current = item.value + if (typeof current === "string" && current.length > 0) yield current + else if (current && typeof current === "object" && !seen.has(current)) { + seen.add(current) + if (Array.isArray(current)) pending.push({ array: current, index: 0 }) + else { + const candidate = current as { text?: unknown; value?: unknown; content?: unknown } + if (candidate.content !== undefined) pending.push({ value: candidate.content }) + if (candidate.value !== undefined) pending.push({ value: candidate.value }) + if (candidate.text !== undefined) pending.push({ value: candidate.text }) + } + } + await checkpoint() + } } -function extractMessageInfoText(info: MessageInfo | undefined): string { - if (!info || info.role !== "assistant" || !info.error) return "" - const error = info.error as { data?: { message?: unknown }; message?: unknown; name?: unknown } - const values = [error.data?.message, error.message, error.name] - return values.filter((value): value is string => typeof value === "string" && value.trim().length > 0).join("\n") +async function* toolTexts( + part: Extract, + checkpoint: () => Promise, + resolver?: BuildSessionSearchMatchesOptions["resolveToolSearchText"], +): AsyncGenerator { + const toolName = typeof part.tool === "string" ? part.tool : "" + const context = { toolCall: part, toolState: (part as any).state, toolName, checkpoint } + const source = resolver?.(context) ?? getDefaultToolSearchText(context) + const resolved = await source + if (resolved && Symbol.asyncIterator in Object(resolved)) { + for await (const text of resolved as AsyncIterable) if (text.length > 0) yield text + } else if (Array.isArray(resolved)) { + for (const text of resolved) if (text.length > 0) yield text + } } -function extractSearchablePartText(part: ClientPart, includeThinking: boolean): SearchablePartText | null { - if (!part || typeof part !== "object") return null - if (isHiddenSyntheticTextPart(part)) return null - - const partId = typeof (part as any).id === "string" ? (part as any).id : undefined - const partType = typeof (part as any).type === "string" ? (part as any).type : undefined - +async function* partTexts( + part: ClientPart, + includeThinking: boolean, + checkpoint: () => Promise, + resolver?: BuildSessionSearchMatchesOptions["resolveToolSearchText"], +): AsyncGenerator { + if (!part || typeof part !== "object" || isHiddenSyntheticTextPart(part)) return if (part.type === "text") { - const text = typeof (part as any).text === "string" ? (part as any).text : segmentToText((part as any).text) - return text.trim().length > 0 ? { partId, partType, text } : null + yield* segmentTexts((part as any).text, checkpoint) + return } - if (part.type === "reasoning") { - if (!includeThinking) return null - const text = extractReasoningText(part) - return text.trim().length > 0 ? { partId, partType, text } : null + if (!includeThinking) return + yield* segmentTexts((part as any).text, checkpoint) + yield* segmentTexts((part as any).content, checkpoint) + return } - if (part.type === "file") { - const filename = (part as any).filename - return typeof filename === "string" && filename.trim().length > 0 ? { partId, partType, text: filename } : null + if (typeof (part as any).filename === "string") yield (part as any).filename + return } - if (part.type === "tool") { - const text = extractToolText(part) - return text.trim().length > 0 ? { partId, partType, text } : null + yield* toolTexts(part, checkpoint, resolver) + return } - if (part.type === "compaction") { - const text = (part as any).auto ? "Session auto-compacted" : "Session compacted" - return { partId, partType, text } + yield (part as any).auto ? "Session auto-compacted" : "Session compacted" + return + } + const candidate = part as Record + for (const value of [candidate.text, candidate.content, candidate.value, candidate.title, candidate.name, candidate.filename, candidate.message]) { + yield* segmentTexts(value, checkpoint) } - - const text = extractGenericPartText(part) - return text.trim().length > 0 ? { partId, partType, text } : null } -function buildPreview(text: string, start: number, end: number): string { - const from = Math.max(0, start - PREVIEW_RADIUS) - const to = Math.min(text.length, end + PREVIEW_RADIUS) - const prefix = from > 0 ? "..." : "" - const suffix = to < text.length ? "..." : "" - return `${prefix}${text.slice(from, to).replace(/\s+/g, " ").trim()}${suffix}` +async function* infoTexts(info: MessageInfo | undefined, checkpoint: () => Promise): AsyncGenerator { + if (!info || info.role !== "assistant" || !info.error) return + const error = info.error as { data?: { message?: unknown }; message?: unknown; name?: unknown } + for (const value of [error.data?.message, error.message, error.name]) yield* segmentTexts(value, checkpoint) } -function collectRecordSearchableText(store: InstanceMessageStore, record: MessageRecord, includeThinking: boolean): SearchablePartText[] { - const results: SearchablePartText[] = [] - for (const partId of record.partIds) { - const part = record.parts[partId]?.data - if (!part) continue - const text = extractSearchablePartText(part, includeThinking) - if (text) results.push(text) - } - - const infoText = extractMessageInfoText(store.getMessageInfo(record.id)) - if (infoText.trim().length > 0) { - results.push({ partType: "error", text: infoText }) +async function* scanTextSource( + messageId: string, + record: MessageRecord, + source: { partId?: string; partType?: string; texts: AsyncIterable }, + query: string, + checkpoint: (force?: boolean) => Promise, +): AsyncGenerator { + let partOffset = 0 + let occurrence = 0 + for await (const text of source.texts) { + const chunkSize = Math.max(SEARCH_TEXT_CHUNK_CHARACTERS, query.length) + const overlap = Math.max(0, query.length - 1) + let nextAllowedStart = 0 + for (let chunkStart = 0; chunkStart < text.length; chunkStart += chunkSize) { + const primaryEnd = Math.min(text.length, chunkStart + chunkSize) + const windowEnd = Math.min(text.length, primaryEnd + overlap) + const windowText = text.slice(chunkStart, windowEnd) + const windowFrom = Math.max(0, nextAllowedStart - chunkStart) + for (const match of iterateTextSearchOccurrences(windowText, query, windowFrom)) { + const start = chunkStart + match.start + if (start >= primaryEnd) break + const end = chunkStart + match.end + nextAllowedStart = end + const absoluteStart = partOffset + start + yield { + id: `${messageId}:${source.partId ?? source.partType ?? "info"}:${absoluteStart}`, + messageId, + partId: source.partId, + partType: source.partType, + role: record.role, + start: absoluteStart, + end: partOffset + end, + occurrence, + preview: makeTextSearchPreview(text, start, end), + } + occurrence += 1 + await checkpoint() + } + await checkpoint(query.length > OVERSIZED_LITERAL_QUERY_CHARACTERS) + } + partOffset += text.length + 1 } - - return results } -export function buildSessionSearchMatches(options: BuildSessionSearchMatchesOptions): SessionSearchMatch[] { +async function* scanSessionSearchMatches(options: BuildSessionSearchMatchesOptions): AsyncGenerator { const query = options.query.trim() - if (!query) return [] - - const needle = normalizeSearchValue(query) - const matches: SessionSearchMatch[] = [] - const messageIds = options.store.getSessionMessageIds(options.sessionId) + if (!query) return + const yieldControl = options.yieldControl ?? yieldToEventLoop + let checkpointUnits = 0 + let lastYieldAt = Date.now() + const checkpoint = async (force = false) => { + if (options.signal?.aborted) abortSearch() + checkpointUnits += 1 + if (!force && checkpointUnits < SEARCH_CHECKPOINT_UNITS && Date.now() - lastYieldAt < SEARCH_YIELD_INTERVAL_MS) return + checkpointUnits = 0 + await yieldControl() + lastYieldAt = Date.now() + if (options.signal?.aborted) abortSearch() + } - for (const messageId of messageIds) { + for (const messageId of options.store.getSessionMessageIds(options.sessionId)) { + if (options.signal?.aborted) abortSearch() const record = options.store.getMessage(messageId) if (!record) continue - const searchableParts = collectRecordSearchableText(options.store, record, options.includeThinking) - - for (const searchable of searchableParts) { - const haystack = normalizeSearchValue(searchable.text) - let from = 0 - let occurrence = 0 - while (from < haystack.length) { - const index = haystack.indexOf(needle, from) - if (index === -1) break - const end = index + query.length - matches.push({ - id: `${messageId}:${searchable.partId ?? searchable.partType ?? "info"}:${index}`, - messageId, - partId: searchable.partId, - partType: searchable.partType, - role: record.role, - start: index, - end, - occurrence, - preview: buildPreview(searchable.text, index, end), - }) - occurrence += 1 - from = end > index ? end : index + 1 - } + for (const partId of record.partIds) { + const part = record.parts[partId]?.data + if (!part) continue + yield* scanTextSource( + messageId, + record, + { partId, partType: part.type, texts: partTexts(part, options.includeThinking, checkpoint, options.resolveToolSearchText) }, + query, + checkpoint, + ) } + yield* scanTextSource( + messageId, + record, + { partType: "error", texts: infoTexts(options.store.getMessageInfo(record.id), checkpoint) }, + query, + checkpoint, + ) + await checkpoint() } +} + +export function createSessionSearchPager(options: BuildSessionSearchMatchesOptions): SessionSearchPager { + const iterator = scanSessionSearchMatches(options)[Symbol.asyncIterator]() + const limit = Math.max(1, Math.trunc(options.limit ?? SESSION_SEARCH_PAGE_SIZE)) + let lookahead: SessionSearchMatch | undefined + let offset = 0 + let done = false + + return { + async nextPage() { + const pageOffset = offset + const matches: SessionSearchMatch[] = [] + if (lookahead) { + matches.push(lookahead) + lookahead = undefined + } + while (matches.length < limit && !done) { + const next = await iterator.next() + done = Boolean(next.done) + if (!next.done) matches.push(next.value) + } + if (!done) { + const next = await iterator.next() + done = Boolean(next.done) + if (!next.done) lookahead = next.value + } + offset += matches.length + return { + matches, + offset: pageOffset, + hasMore: Boolean(lookahead), + totalMatches: done ? offset : null, + } + }, + } +} + +export function retainSessionSearchPage( + pages: Map, + pageIndex: number, + result: SessionSearchResult, + limit = SESSION_SEARCH_RETAINED_PAGE_LIMIT, +): void { + pages.delete(pageIndex) + pages.set(pageIndex, result) + while (pages.size > limit) pages.delete(pages.keys().next().value!) +} + +export async function findLastSessionSearchPage( + loadPage: (pageIndex: number) => Promise, +): Promise<{ pageIndex: number; result: SessionSearchResult }> { + let pageIndex = 0 + let result = await loadPage(pageIndex) + while (result.hasMore) { + pageIndex += 1 + result = await loadPage(pageIndex) + } + return { pageIndex, result } +} - return matches +export async function buildSessionSearchMatches(options: BuildSessionSearchMatchesOptions): Promise { + return createSessionSearchPager(options).nextPage() } diff --git a/packages/ui/src/lib/sse-manager-generation.test.ts b/packages/ui/src/lib/sse-manager-generation.test.ts new file mode 100644 index 000000000..7224ce51c --- /dev/null +++ b/packages/ui/src/lib/sse-manager-generation.test.ts @@ -0,0 +1,11 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { acceptInstanceStreamId } from "./sse-manager.ts" + +test("instance stream ids reject events from a replaced runtime", () => { + const streams = new Map() + assert.equal(acceptInstanceStreamId(streams, "instance", "old"), true) + assert.equal(acceptInstanceStreamId(streams, "instance", "new"), false) + assert.equal(acceptInstanceStreamId(streams, "instance", "new", true), true) + assert.equal(acceptInstanceStreamId(streams, "instance", "old"), false) +}) diff --git a/packages/ui/src/lib/sse-manager.ts b/packages/ui/src/lib/sse-manager.ts index eb788067e..a3bcbbbb5 100644 --- a/packages/ui/src/lib/sse-manager.ts +++ b/packages/ui/src/lib/sse-manager.ts @@ -125,13 +125,30 @@ type SSEEvent = const [connectionStatus, setConnectionStatus] = createSignal>(new Map()) const [transportStatus, setTransportStatus] = createSignal("connecting") +export function acceptInstanceStreamId( + streams: Map, + instanceId: string, + streamId: string | undefined, + replace = false, +): boolean { + if (streamId === undefined) return true + const current = streams.get(instanceId) + if (current === undefined || replace) streams.set(instanceId, streamId) + return current === undefined || replace || current === streamId +} + class SSEManager { + private readonly streamIds = new Map() + private readonly acceptedEventHandlers = new Set<(instanceId: string, event: SSEEvent | InstanceStreamEvent) => void>() + constructor() { log.info("sseManager initialized: listening for SSE disconnect and reconnect") serverEvents.on("instance.eventStatus", (event) => { const payload = event as InstanceStatusPayload + if (!acceptInstanceStreamId(this.streamIds, payload.instanceId, payload.streamId, payload.status === "connecting")) return this.updateConnectionStatus(payload.instanceId, payload.status) + if (payload.status === "connected") this.onConnectionRestored?.(payload.instanceId) if (payload.status === "disconnected") { if (payload.reason === "workspace stopped") { return @@ -143,12 +160,14 @@ class SSEManager { serverEvents.on("instance.event", (event) => { const payload = event as InstanceEventPayload - this.updateConnectionStatus(payload.instanceId, "connected") + if (!acceptInstanceStreamId(this.streamIds, payload.instanceId, payload.streamId)) return + if (payload.streamId) this.updateConnectionStatus(payload.instanceId, "connected") this.handleEvent(payload.instanceId, payload.event as SSEEvent) }) serverEvents.onTransportStatus((status) => { log.info("SSE transport status changed", { status }) + if (status !== "connected") this.streamIds.clear() setTransportStatus(status) }) } @@ -166,7 +185,8 @@ class SSEManager { log.warn("Dropping malformed event", event) return } - + if (this.shouldHandleEvent && !this.shouldHandleEvent(instanceId)) return + for (const handler of this.acceptedEventHandlers) handler(instanceId, event) log.info("Received event", { type: event.type, event }) switch (event.type) { @@ -293,7 +313,13 @@ class SSEManager { onInstanceDisposed?: (instanceId: string, event: ServerInstanceDisposedEvent) => void onWorktreeReady?: (instanceId: string, event: WorktreeReadyEvent) => void | Promise onConnectionLost?: (instanceId: string, reason: string) => void | Promise + onConnectionRestored?: (instanceId: string) => void + shouldHandleEvent?: (instanceId: string) => boolean + onAcceptedEvent(handler: (instanceId: string, event: SSEEvent | InstanceStreamEvent) => void): () => void { + this.acceptedEventHandlers.add(handler) + return () => this.acceptedEventHandlers.delete(handler) + } getStatus(instanceId: string): ConnectionStatus | null { return deriveDisplayConnectionStatus(connectionStatus().get(instanceId) ?? null, transportStatus()) } diff --git a/packages/ui/src/lib/trailing-resync.test.ts b/packages/ui/src/lib/trailing-resync.test.ts index 7e4a5d167..4a37a43d4 100644 --- a/packages/ui/src/lib/trailing-resync.test.ts +++ b/packages/ui/src/lib/trailing-resync.test.ts @@ -1,6 +1,7 @@ import assert from "node:assert/strict" import { it } from "node:test" import { TrailingResyncCoordinator, waitForSettledPrerequisite } from "./trailing-resync.ts" +import { retryWithBackoff } from "./retry-utils.ts" function deferred() { let resolve!: () => void, reject!: (error: unknown) => void const promise = new Promise((resolvePromise, rejectPromise) => { @@ -49,6 +50,28 @@ it("does not lose a request queued as the previous pass settles", async () => { firstPass.resolve(); await first; await boundaryRequest assert.equal(calls, 2) }) +it("coalesces arbitrarily many reconnects into one trailing pass", async () => { + const passes = [deferred(), deferred()] + const { coordinator, calls } = coordinatorFor(passes) + const active = coordinator.request("workspace-1") + const queued = Array.from({ length: 10_000 }, () => coordinator.request("workspace-1")) + assert.equal(new Set(queued).size, 1) + assert.equal(queued[0], active) + await Promise.resolve() + assert.equal(calls(), 1) + passes[0]!.resolve() + await turn() + assert.equal(calls(), 2) + passes[1]!.resolve() + await Promise.all(queued) + assert.equal(calls(), 2) +}) it("continues recovery after a prerequisite rejects", async () => { await assert.doesNotReject(waitForSettledPrerequisite(Promise.reject(new Error("initial hydration failed")))) }) +it("rejects a timed-out request even when it ignores abort", async () => { + await assert.rejects( + retryWithBackoff(() => new Promise(() => undefined), { maxAttempts: 1, timeoutMs: 10 }), + (error: any) => error?.name === "TimeoutError", + ) +}) diff --git a/packages/ui/src/lib/trailing-resync.ts b/packages/ui/src/lib/trailing-resync.ts index 3a52e80f3..39e07fb12 100644 --- a/packages/ui/src/lib/trailing-resync.ts +++ b/packages/ui/src/lib/trailing-resync.ts @@ -1,5 +1,5 @@ export class TrailingResyncCoordinator { - private readonly tails = new Map>() + private readonly active = new Map }>() constructor( private readonly run: (key: string) => Promise, @@ -7,19 +7,27 @@ export class TrailingResyncCoordinator { ) {} request(key: string): Promise { - const previous = this.tails.get(key) ?? Promise.resolve() - const task = previous.then(async () => { - try { - await this.run(key) - } catch (error) { - this.onError(key, error) - } - }) - const tracked = task.finally(() => { - if (this.tails.get(key) === tracked) this.tails.delete(key) + const existing = this.active.get(key) + if (existing) { + existing.dirty = true + return existing.promise + } + + const state = { dirty: false, promise: undefined as unknown as Promise } + this.active.set(key, state) + state.promise = (async () => { + do { + state.dirty = false + try { + await this.run(key) + } catch (error) { + this.onError(key, error) + } + } while (state.dirty) + })().finally(() => { + if (this.active.get(key) === state) this.active.delete(key) }) - this.tails.set(key, tracked) - return tracked + return state.promise } } diff --git a/packages/ui/src/stores/commands.test.ts b/packages/ui/src/stores/commands.test.ts new file mode 100644 index 000000000..47751e525 --- /dev/null +++ b/packages/ui/src/stores/commands.test.ts @@ -0,0 +1,15 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { clearCommands, fetchCommands, getCommands } from "./commands.ts" + +test("command hydration cannot commit after runtime replacement", async () => { + let resolve!: (value: unknown) => void + const response = new Promise((done) => { resolve = done }) + let current = true + const hydration = fetchCommands("instance", { command: { list: () => response } } as any, () => current) + current = false + resolve({ data: [{ name: "stale" }] }) + await hydration + assert.deepEqual(getCommands("instance"), []) + clearCommands("instance") +}) diff --git a/packages/ui/src/stores/commands.ts b/packages/ui/src/stores/commands.ts index a9f0ed3c3..2ed0035aa 100644 --- a/packages/ui/src/stores/commands.ts +++ b/packages/ui/src/stores/commands.ts @@ -5,8 +5,17 @@ import { requestData } from "../lib/opencode-api" const [commandMap, setCommandMap] = createSignal>(new Map()) -export async function fetchCommands(instanceId: string, client: OpencodeClient): Promise { - const commands = await requestData(client.command.list(), "command.list").catch(() => []) +export async function fetchCommands( + instanceId: string, + client: OpencodeClient, + isCurrent: () => boolean = () => true, + signal?: AbortSignal, +): Promise { + const commands = await requestData( + (client.command.list as any)(undefined, signal ? { signal } : undefined), + "command.list", + ).catch(() => []) + if (!isCurrent()) return setCommandMap((prev) => { const next = new Map(prev) next.set(instanceId, commands) diff --git a/packages/ui/src/stores/delta-buffer.test.ts b/packages/ui/src/stores/delta-buffer.test.ts index 09072783e..8b627dfd6 100644 --- a/packages/ui/src/stores/delta-buffer.test.ts +++ b/packages/ui/src/stores/delta-buffer.test.ts @@ -6,8 +6,10 @@ import { clearPendingDeltasForPart, enqueueDelta, flushPendingDeltasForMessage, + hasPendingDeltasForMessage, resetDeltaBufferForTests, setFlushCallback, + setRecoveryCallback, } from "./delta-buffer.ts" type DeltaBatch = Array<{ instanceId: string; messageId: string; partId: string; field: string; delta: string }> @@ -99,4 +101,43 @@ describe("delta buffer", () => { ], ]) }) + + it("drops oversized deltas and requests bounded authoritative recovery", async () => { + const recoveries: unknown[] = [] + const flushed: DeltaBatch[] = [] + setRecoveryCallback((pending) => recoveries.push(pending)) + setFlushCallback((batch) => flushed.push(batch)) + + enqueueDelta("instance-1", "message-1", "part-1", "text", "x".repeat(300_000), "session-1") + await delay(75) + + assert.deepEqual(flushed, []) + assert.deepEqual(recoveries, [{ + instanceId: "instance-1", sessionId: "session-1", messageId: "message-1", partId: "part-1", field: "text", + }]) + }) + + it("reports buffered message deltas without consuming them", async () => { + const flushed: DeltaBatch[] = [] + setFlushCallback((batch) => flushed.push(batch)) + enqueueDelta("instance-1", "message-1", "part-1", "text", "pending") + + assert.equal(hasPendingDeltasForMessage("instance-1", "message-1"), true) + assert.equal(hasPendingDeltasForMessage("instance-1", "message-2"), false) + await delay(75) + assert.equal(flushed[0]?.[0]?.delta, "pending") + }) + + it("retains a delta until the flush callback has attempted application", async () => { + let pendingDuringApplication = false + setFlushCallback(() => { + pendingDuringApplication = hasPendingDeltasForMessage("instance-1", "message-1") + }) + enqueueDelta("instance-1", "message-1", "part-1", "text", "orphan", "session-1") + + await delay(75) + + assert.equal(pendingDuringApplication, true) + assert.equal(hasPendingDeltasForMessage("instance-1", "message-1"), false) + }) }) diff --git a/packages/ui/src/stores/delta-buffer.ts b/packages/ui/src/stores/delta-buffer.ts index 3307257a4..42a39435d 100644 --- a/packages/ui/src/stores/delta-buffer.ts +++ b/packages/ui/src/stores/delta-buffer.ts @@ -8,20 +8,52 @@ */ const DELTA_FLUSH_INTERVAL = 50 +const MAX_PENDING_DELTA_CHARACTERS = 64_000 +const MAX_PENDING_DELTA_ENTRIES = 64 -const pendingDeltas = new Map() +type PendingDelta = { instanceId: string; sessionId?: string; messageId: string; partId: string; field: string; delta: string } +type DeltaRecoveryRequest = Omit & { delta?: string } +const pendingDeltas = new Map() let deltaFlushTimer: ReturnType | null = null -export function enqueueDelta(instanceId: string, messageId: string, partId: string, field: string, delta: string) { +export function enqueueDelta(instanceId: string, messageId: string, partId: string, field: string, delta: string, sessionId?: string) { const key = `${instanceId}:${messageId}:${partId}:${field}` const existing = pendingDeltas.get(key) const accumulated = existing ? existing.delta + delta : delta - pendingDeltas.set(key, { instanceId, messageId, partId, field, delta: accumulated }) + const resolvedSessionId = sessionId ?? existing?.sessionId + if (accumulated.length > MAX_PENDING_DELTA_CHARACTERS || (!existing && pendingDeltas.size >= MAX_PENDING_DELTA_ENTRIES)) { + pendingDeltas.delete(key) + recoveryCallback?.({ instanceId, ...(resolvedSessionId ? { sessionId: resolvedSessionId } : {}), messageId, partId, field }) + return + } + pendingDeltas.set(key, { instanceId, ...(resolvedSessionId ? { sessionId: resolvedSessionId } : {}), messageId, partId, field, delta: accumulated }) if (deltaFlushTimer === null) { deltaFlushTimer = setTimeout(flushDeltas, DELTA_FLUSH_INTERVAL) } } +export function requestDeltaRecovery(pending: DeltaRecoveryRequest): void { + recoveryCallback?.(pending) +} + +export function clearPendingDeltasForSession(instanceId: string, sessionId: string): void { + for (const [key, pending] of pendingDeltas) if (pending.instanceId === instanceId && pending.sessionId === sessionId) pendingDeltas.delete(key) +} + +export function clearPendingDeltasForInstance(instanceId: string): void { + for (const [key, pending] of pendingDeltas) if (pending.instanceId === instanceId) pendingDeltas.delete(key) +} + +export function clearPendingDeltasForMessage(instanceId: string, messageId: string): boolean { + let cleared = false + for (const [key, pending] of pendingDeltas) { + if (pending.instanceId !== instanceId || pending.messageId !== messageId) continue + pendingDeltas.delete(key) + cleared = true + } + return cleared +} + export function clearPendingDeltasForPart(instanceId: string, messageId: string, partId: string) { const keysToDelete: string[] = [] for (const key of pendingDeltas.keys()) { @@ -37,7 +69,7 @@ export function clearPendingDeltasForPart(instanceId: string, messageId: string, export function flushPendingDeltasForMessage( instanceId: string, messageId: string, - applyDelta: (instanceId: string, delta: { messageId: string; partId: string; field: string; delta: string }) => void + applyDelta: (instanceId: string, delta: { messageId: string; partId: string; field: string; delta: string }) => boolean | void ): void { const prefix = `${instanceId}:${messageId}:` const keysToFlush: string[] = [] @@ -49,24 +81,47 @@ export function flushPendingDeltasForMessage( for (const key of keysToFlush) { const pending = pendingDeltas.get(key) if (pending) { - pendingDeltas.delete(key) - applyDelta(instanceId, { + const applied = applyDelta(instanceId, { messageId: pending.messageId, partId: pending.partId, field: pending.field, delta: pending.delta, }) + if (applied !== false) pendingDeltas.delete(key) } } } export function setFlushCallback( - callback: (batch: Array<{ instanceId: string; messageId: string; partId: string; field: string; delta: string }>) => void + callback: (batch: PendingDelta[]) => void ) { // Store callback for flushDeltas to use flushCallback = callback } +export function hasPendingDeltasForMessage(instanceId: string, messageId: string): boolean { + const prefix = `${instanceId}:${messageId}:` + for (const key of pendingDeltas.keys()) if (key.startsWith(prefix)) return true + return false +} + +export function getPendingDeltasForMessage( + instanceId: string, + messageId: string, +): Array<{ partId: string; field: string; delta: string }> { + const result: Array<{ partId: string; field: string; delta: string }> = [] + for (const pending of pendingDeltas.values()) { + if (pending.instanceId === instanceId && pending.messageId === messageId) { + result.push({ partId: pending.partId, field: pending.field, delta: pending.delta }) + } + } + return result +} + +export function setRecoveryCallback(callback: (pending: DeltaRecoveryRequest) => void) { + recoveryCallback = callback +} + export function resetDeltaBufferForTests() { pendingDeltas.clear() if (deltaFlushTimer !== null) { @@ -74,16 +129,21 @@ export function resetDeltaBufferForTests() { deltaFlushTimer = null } flushCallback = null + recoveryCallback = null } -let flushCallback: ((batch: Array<{ instanceId: string; messageId: string; partId: string; field: string; delta: string }>) => void) | null = null +let flushCallback: ((batch: PendingDelta[]) => void) | null = null +let recoveryCallback: ((pending: DeltaRecoveryRequest) => void) | null = null function flushDeltas() { deltaFlushTimer = null if (pendingDeltas.size === 0) return const batch = Array.from(pendingDeltas.values()) - pendingDeltas.clear() if (flushCallback) { flushCallback(batch) } + for (const pending of batch) { + const key = `${pending.instanceId}:${pending.messageId}:${pending.partId}:${pending.field}` + if (pendingDeltas.get(key) === pending) pendingDeltas.delete(key) + } } diff --git a/packages/ui/src/stores/instances-reconnect-resync.test.ts b/packages/ui/src/stores/instances-reconnect-resync.test.ts new file mode 100644 index 000000000..29a996cf6 --- /dev/null +++ b/packages/ui/src/stores/instances-reconnect-resync.test.ts @@ -0,0 +1,947 @@ +import assert from "node:assert/strict" +import { describe, it } from "node:test" + +import { sdkManager } from "../lib/sdk-manager.ts" +import { sseManager } from "../lib/sse-manager.ts" +import { serverApi } from "../lib/api-client.ts" +import { serverEvents } from "../lib/server-events.ts" +import { + addInstance, + addPermissionToQueue, + addQuestionToQueue, + getPermissionQueue, + getQuestionQueue, + hasAnsweredQuestion, + markQuestionAnswered, + removeInstance, + updateInstance, +} from "./instances.ts" +import { + handlePermissionReplied, + handlePermissionUpdated, + handleQuestionAnswered, + handleQuestionAsked, +} from "./session-events.ts" +import { messageStoreBus } from "./message-v2/bus.ts" +import { agents, getSessionDraftPrompt, sessions, setActiveSession, setMessagesLoaded, setSessionDraftPrompt, setSessions } from "./session-state.ts" +import { reloadOpenCodeWorkspaces } from "./opencode-workspaces.ts" +import { fetchSessions } from "./session-api.ts" +import { setVisibleSessionMemory } from "./session-memory.ts" +import { reloadWorktrees } from "./worktrees.ts" + +function deferred() { + let resolve!: (value: T) => void + const promise = new Promise((done) => { resolve = done }) + return { promise, resolve } +} + +async function waitFor(check: () => boolean, timeoutMs = 1_000): Promise { + const deadline = Date.now() + timeoutMs + while (!check()) { + if (Date.now() >= deadline) throw new Error("Timed out waiting for reconnect resync") + await new Promise((resolve) => setTimeout(resolve, 5)) + } +} + +type PendingResponse = Promise | (() => Promise) +type PendingMessages = (input: any, options?: { signal?: AbortSignal }) => Promise + +function resolvePending(pending: PendingResponse | undefined, fallback: any): Promise { + if (typeof pending === "function") return pending() + return pending ?? Promise.resolve(fallback) +} + +function setup(instanceId: string, pending: { + sessions?: PendingResponse + permissions?: PendingResponse + questions?: PendingResponse + legacyPermissions?: PendingResponse + legacyQuestions?: PendingResponse + messages?: PendingMessages +} = {}) { + let sessionLists = 0 + let permissionLists = 0 + let questionLists = 0 + let legacyPermissionLists = 0 + let legacyQuestionLists = 0 + let messageLists = 0 + const messageSessionIds: string[] = [] + const permissionLocations: any[] = [] + const questionLocations: any[] = [] + const client = { + session: { + list: () => { sessionLists += 1; return resolvePending(pending.sessions, { data: [] }) }, + status: async () => ({ data: {} }), + messages: async (input: any, options?: { signal?: AbortSignal }) => { + messageLists += 1 + messageSessionIds.push(input.sessionID) + return pending.messages?.(input, options) ?? { data: [] } + }, + }, + permission: { list: () => { legacyPermissionLists += 1; return resolvePending(pending.legacyPermissions, { data: [] }) } }, + question: { list: () => { legacyQuestionLists += 1; return resolvePending(pending.legacyQuestions, { data: [] }) } }, + v2: { + permission: { request: { list: (input: any) => { permissionLists += 1; permissionLocations.push(input?.location); return resolvePending(pending.permissions, { data: { data: [] } }) } } }, + question: { request: { list: (input: any) => { questionLists += 1; questionLocations.push(input?.location); return resolvePending(pending.questions, { data: { data: [] } }) } } }, + }, + } as any + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, client) + addInstance({ id: instanceId, folder: "/work", port: 0, pid: 1, proxyPath: "", status: "ready", client }) + return { + client, + sessionLists: () => sessionLists, + permissionLists: () => permissionLists, + questionLists: () => questionLists, + legacyPermissionLists: () => legacyPermissionLists, + legacyQuestionLists: () => legacyQuestionLists, + messageLists: () => messageLists, + messageSessionIds, + permissionLocations, + questionLocations, + cleanup() { + removeInstance(instanceId, { authoritative: false }) + sdkManager.destroyClientsForInstance(instanceId) + }, + } +} + +describe("reconnect interruption resync", () => { + it("does not commit a session refresh after its status request is aborted", async () => { + const instanceId = "reconnect-aborted-status", sessionId = "session" + const status = deferred() + const harness = setup(instanceId) + let statusParameters: unknown + let statusSignal: AbortSignal | undefined + harness.client.session.status = (parameters: unknown, options?: { signal?: AbortSignal }) => { + statusParameters = parameters + statusSignal = options?.signal + return status.promise + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: "Resident", + status: "idle", + model: { providerId: "", modelId: "" }, + } as any]]))) + setSessionDraftPrompt(instanceId, sessionId, "unsent draft") + const controller = new AbortController() + + try { + const refresh = fetchSessions(instanceId, { authoritativeDeletes: true, signal: controller.signal }) + await waitFor(() => statusSignal !== undefined) + assert.equal(statusParameters, undefined) + assert.equal(statusSignal, controller.signal) + controller.abort(new Error("refresh timed out")) + status.resolve({ data: {} }) + await refresh + + assert.equal(sessions().get(instanceId)?.has(sessionId), true) + assert.equal(getSessionDraftPrompt(instanceId, sessionId), "unsent draft") + } finally { + harness.cleanup() + } + }) + + it("does not commit a session list that settles after abort", async () => { + const instanceId = "reconnect-aborted-list", sessionId = "session" + const list = deferred() + const harness = setup(instanceId) + let listSignal: AbortSignal | undefined + harness.client.session.list = (_parameters: unknown, options?: { signal?: AbortSignal }) => { + listSignal = options?.signal + return list.promise + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: "Resident", + status: "idle", + model: { providerId: "", modelId: "" }, + } as any]]))) + setSessionDraftPrompt(instanceId, sessionId, "unsent draft") + const controller = new AbortController() + + try { + const refresh = fetchSessions(instanceId, { authoritativeDeletes: true, signal: controller.signal }) + await waitFor(() => listSignal !== undefined) + controller.abort(new Error("refresh timed out")) + list.resolve({ data: [] }) + await refresh + + assert.equal(sessions().get(instanceId)?.has(sessionId), true) + assert.equal(getSessionDraftPrompt(instanceId, sessionId), "unsent draft") + } finally { + harness.cleanup() + } + }) + + it("keeps resident state when a reconnect session refresh is empty", async () => { + const instanceId = "reconnect-merge-only", sessionId = "session" + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [ + { slug: "root", directory: "/work", kind: "root" }, + { slug: "feature", directory: "/feature", kind: "worktree" }, + ], + }) + const harness = setup(instanceId) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: "Resident", + status: "idle", + model: { providerId: "", modelId: "" }, + } as any]]))) + setSessionDraftPrompt(instanceId, sessionId, "unsent draft") + messageStoreBus.getOrCreate(instanceId).upsertMessage({ + id: "message", + sessionId, + role: "assistant", + status: "complete", + parts: [{ id: "part", type: "text", text: "resident" }] as any, + }) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.sessionLists() === 1) + await new Promise((resolve) => setTimeout(resolve, 20)) + + assert.equal(sessions().get(instanceId)?.has(sessionId), true) + assert.equal(getSessionDraftPrompt(instanceId, sessionId), "unsent draft") + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), ["message"]) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("deletes missing sessions only after complete root topology and project scope", async () => { + const instanceId = "reconnect-authoritative-sessions", sessionId = "deleted" + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + const harness = setup(instanceId) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: "Deleted", + status: "idle", + model: { providerId: "", modelId: "" }, + } as any]]))) + setSessionDraftPrompt(instanceId, sessionId, "discard with deletion") + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.sessionLists() === 1) + await waitFor(() => !sessions().get(instanceId)?.has(sessionId)) + assert.equal(getSessionDraftPrompt(instanceId, sessionId), "") + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("reloads selected and visible sessions whose previous message loads were empty", async () => { + const instanceId = "reconnect-selected-empty", sessionId = "selected-empty", visibleSessionId = "visible-empty" + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + const harness = setup(instanceId, { + sessions: Promise.resolve({ data: [ + { id: sessionId, title: "Selected", directory: "/work", time: { created: 1 } }, + { id: visibleSessionId, title: "Visible", directory: "/work", time: { created: 1 } }, + ] }), + }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([sessionId, visibleSessionId].map((id) => [id, { + id, + instanceId, + parentId: null, + title: id, + status: "idle", + model: { providerId: "", modelId: "" }, + } as any])))) + setActiveSession(instanceId, sessionId) + setVisibleSessionMemory(instanceId, visibleSessionId, true) + setMessagesLoaded((prev) => new Map(prev).set(instanceId, new Set([sessionId, visibleSessionId]))) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.messageLists() === 2) + assert.deepEqual(new Set(harness.messageSessionIds), new Set([sessionId, visibleSessionId])) + } finally { + setVisibleSessionMemory(instanceId, visibleSessionId, false) + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("aborts a stalled reconnect message reload", async () => { + const instanceId = "reconnect-message-timeout", sessionId = "selected" + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + let messageSignal: AbortSignal | undefined + const harness = setup(instanceId, { + sessions: Promise.resolve({ data: [{ id: sessionId, title: "Selected", directory: "/work", time: { created: 1 } }] }), + messages: (_input, options) => { + messageSignal = options?.signal + return new Promise((_resolve, reject) => { + messageSignal?.addEventListener("abort", () => reject(messageSignal?.reason), { once: true }) + }) + }, + }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: sessionId, + status: "idle", + model: { providerId: "", modelId: "" }, + } as any]]))) + setActiveSession(instanceId, sessionId) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => Boolean(messageSignal)) + await Promise.race([ + new Promise((resolve) => messageSignal?.addEventListener("abort", () => resolve(), { once: true })), + new Promise((_resolve, reject) => setTimeout(() => reject(new Error("reconnect message reload did not abort")), 6_000)), + ]) + assert.equal(messageSignal?.aborted, true) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("continues interruption resync when the session refresh fails", async () => { + const instanceId = "reconnect-session-list-failure" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const harness = setup(instanceId, { + sessions: async () => { throw new Error("session list unavailable") }, + permissions: Promise.resolve({ data: { data: [permission] } }), + }) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => getPermissionQueue(instanceId).length === 1) + assert.equal(getPermissionQueue(instanceId)[0]?.id, permission.id) + } finally { + harness.cleanup() + } + }) + + it("does not apply a stale revert when reconnect metadata refresh fails", async () => { + const instanceId = "reconnect-stale-revert", sessionId = "session" + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + const staleRevert = { messageID: "stale-anchor" } + const messageIds = ["before", "stale-anchor", "current-anchor", "after"] + const harness = setup(instanceId, { + sessions: async () => { throw new Error("session list unavailable") }, + messages: async () => ({ + data: messageIds.map((id) => ({ + info: { + id, sessionID: sessionId, role: "assistant", agent: "build", + providerID: "provider", modelID: "model", time: { created: 1, completed: 2 }, + }, + parts: [{ id: `${id}-part`, type: "text", text: id }], + })), + }), + }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: sessionId, + status: "idle", + model: { providerId: "", modelId: "" }, + revert: staleRevert, + } as any]]))) + messageStoreBus.getOrCreate(instanceId).setSessionRevert(sessionId, staleRevert) + setActiveSession(instanceId, sessionId) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.messageLists() === 1, 5_000) + await waitFor(() => messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId).length === messageIds.length, 5_000) + + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), messageIds) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionRevert(sessionId), staleRevert) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("does not apply a stale revert when a capped reconnect list omits the active session", async () => { + const instanceId = "reconnect-capped-stale-revert", sessionId = "session" + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + const staleRevert = { messageID: "stale-anchor" } + const messageIds = ["before", "stale-anchor", "current-anchor", "after"] + const harness = setup(instanceId, { + sessions: Promise.resolve({ + data: Array.from({ length: 10_000 }, (_, index) => ({ + id: `other-${index}`, + title: `Other ${index}`, + directory: "/work", + metadata: { owner: "test" }, + time: { created: 1 }, + })), + }), + messages: async () => ({ + data: messageIds.map((id) => ({ + info: { + id, sessionID: sessionId, role: "assistant", agent: "build", + providerID: "provider", modelID: "model", time: { created: 1, completed: 2 }, + }, + parts: [{ id: `${id}-part`, type: "text", text: id }], + })), + }), + }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: sessionId, + status: "idle", + model: { providerId: "", modelId: "" }, + revert: staleRevert, + } as any]]))) + messageStoreBus.getOrCreate(instanceId).setSessionRevert(sessionId, staleRevert) + setActiveSession(instanceId, sessionId) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.messageLists() === 1, 5_000) + await waitFor(() => messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId).length === messageIds.length, 5_000) + + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), messageIds) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionRevert(sessionId), staleRevert) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("reconciles missed permissions and questions through their authoritative lists", async () => { + const instanceId = "reconnect-interruptions" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const harness = setup(instanceId, { + permissions: Promise.resolve({ data: { data: [permission] } }), + questions: Promise.resolve({ data: { data: [question] } }), + }) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => getPermissionQueue(instanceId).length === 1 && getQuestionQueue(instanceId).length === 1) + assert.equal(getPermissionQueue(instanceId)[0]?.id, permission.id) + assert.equal(getQuestionQueue(instanceId)[0]?.id, question.id) + } finally { + harness.cleanup() + } + }) + + it("rejects stale list results after newer accepted replies", async () => { + const instanceId = "reconnect-interruption-fences" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const permissions = deferred() + const questions = deferred() + const harness = setup(instanceId, { permissions: permissions.promise, questions: questions.promise }) + + try { + addPermissionToQueue(instanceId, permission, "v2") + addQuestionToQueue(instanceId, question, "v2") + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.permissionLists() === 1 && harness.questionLists() === 1) + + handlePermissionReplied(instanceId, { type: "permission.replied", properties: { requestID: permission.id } } as any) + handleQuestionAnswered(instanceId, { type: "question.replied", properties: { requestID: question.id } } as any) + permissions.resolve({ data: { data: [permission] } }) + questions.resolve({ data: { data: [question] } }) + await new Promise((resolve) => setTimeout(resolve, 20)) + + assert.deepEqual(getPermissionQueue(instanceId), []) + assert.deepEqual(getQuestionQueue(instanceId), []) + } finally { + harness.cleanup() + } + }) + + it("does not delete asked events that arrive while reconnect lists are in flight", async () => { + const instanceId = "reconnect-newer-asked-events" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const permissions = deferred() + const questions = deferred() + const harness = setup(instanceId, { permissions: permissions.promise, questions: questions.promise }) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.permissionLists() === 1 && harness.questionLists() === 1) + handlePermissionUpdated(instanceId, { type: "permission.v2.asked", properties: permission } as any) + handleQuestionAsked(instanceId, { type: "question.v2.asked", properties: question } as any) + permissions.resolve({ data: { data: [] } }) + questions.resolve({ data: { data: [] } }) + await new Promise((resolve) => setTimeout(resolve, 20)) + + assert.equal(getPermissionQueue(instanceId)[0]?.id, permission.id) + assert.equal(getQuestionQueue(instanceId)[0]?.id, question.id) + } finally { + harness.cleanup() + } + }) + + it("keeps local queues and accepts V2 additions when legacy snapshots fail", async () => { + const instanceId = "reconnect-incomplete-legacy" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const v2Permission = { id: "v2-permission", sessionID: "session", permission: "write", patterns: [] } as any + const v2Question = { id: "v2-question", sessionID: "session", questions: [] } as any + const harness = setup(instanceId, { + legacyPermissions: () => Promise.reject(new Error("permission list unavailable")), + legacyQuestions: () => Promise.reject(new Error("question list unavailable")), + permissions: Promise.resolve({ data: { data: [v2Permission] } }), + questions: Promise.resolve({ data: { data: [v2Question] } }), + }) + + try { + addPermissionToQueue(instanceId, permission, "legacy") + addQuestionToQueue(instanceId, question, "legacy") + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.legacyPermissionLists() === 3 && harness.legacyQuestionLists() === 3) + + assert.deepEqual(getPermissionQueue(instanceId).map((entry) => entry.id), [permission.id, v2Permission.id]) + assert.deepEqual(getQuestionQueue(instanceId).map((entry) => entry.id), [question.id, v2Question.id]) + assert.equal(harness.permissionLists(), 3) + assert.equal(harness.questionLists(), 3) + } finally { + harness.cleanup() + } + }) + + it("retries a transient V2 snapshot and restores the missed request", async () => { + const instanceId = "reconnect-v2-retry" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + let attempts = 0 + const harness = setup(instanceId, { + permissions: () => { + attempts += 1 + if (attempts === 1) return Promise.reject(new Error("temporary V2 failure")) + return Promise.resolve({ data: { data: [permission] } }) + }, + }) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => getPermissionQueue(instanceId).length === 1) + assert.equal(harness.permissionLists(), 2) + } finally { + harness.cleanup() + } + }) + + it("ignores permission and question snapshots that settle after their retry timeout", async () => { + const instanceId = "reconnect-late-timeout" + const permission = { id: "stale-permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "stale-question", sessionID: "session", questions: [] } as any + const firstPermissions = deferred() + const firstQuestions = deferred() + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + const harness = setup(instanceId) + let permissionCalls = 0 + let questionCalls = 0 + let firstPermissionParameters: unknown + let firstQuestionParameters: unknown + let firstPermissionSignal: AbortSignal | undefined + let firstQuestionSignal: AbortSignal | undefined + harness.client.permission.list = (parameters: unknown, options?: { signal?: AbortSignal }) => { + permissionCalls += 1 + if (permissionCalls === 1) { + firstPermissionParameters = parameters + firstPermissionSignal = options?.signal + } + return permissionCalls === 1 ? firstPermissions.promise : Promise.resolve({ data: [] }) + } + harness.client.question.list = (parameters: unknown, options?: { signal?: AbortSignal }) => { + questionCalls += 1 + if (questionCalls === 1) { + firstQuestionParameters = parameters + firstQuestionSignal = options?.signal + } + return questionCalls === 1 ? firstQuestions.promise : Promise.resolve({ data: [] }) + } + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => permissionCalls === 2 && questionCalls === 2, 7_000) + assert.equal(firstPermissionParameters, undefined) + assert.equal(firstQuestionParameters, undefined) + assert.equal(firstPermissionSignal?.aborted, true) + assert.equal(firstQuestionSignal?.aborted, true) + firstPermissions.resolve({ data: [permission] }) + firstQuestions.resolve({ data: [question] }) + await new Promise((resolve) => setTimeout(resolve, 50)) + + assert.deepEqual(getPermissionQueue(instanceId), []) + assert.deepEqual(getQuestionQueue(instanceId), []) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("keeps local queues when V2 snapshots exhaust their retries", async () => { + const instanceId = "reconnect-incomplete-v2" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const legacyPermission = { id: "legacy-permission", sessionID: "session", permission: "write", patterns: [] } as any + const legacyQuestion = { id: "legacy-question", sessionID: "session", questions: [] } as any + const harness = setup(instanceId, { + legacyPermissions: Promise.resolve({ data: [legacyPermission] }), + legacyQuestions: Promise.resolve({ data: [legacyQuestion] }), + permissions: () => Promise.reject(new Error("permission V2 unavailable")), + questions: () => Promise.reject(new Error("question V2 unavailable")), + }) + + try { + addPermissionToQueue(instanceId, permission, "v2") + addQuestionToQueue(instanceId, question, "v2") + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.permissionLists() === 3 && harness.questionLists() === 3) + + assert.deepEqual(getPermissionQueue(instanceId).map((entry) => entry.id), [permission.id, legacyPermission.id]) + assert.deepEqual(getQuestionQueue(instanceId).map((entry) => entry.id), [question.id, legacyQuestion.id]) + } finally { + harness.cleanup() + } + }) + + it("retains reply tombstones through an empty reconnect snapshot", async () => { + const instanceId = "reconnect-delayed-stale-events" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const harness = setup(instanceId) + + try { + addPermissionToQueue(instanceId, permission, "v2") + addQuestionToQueue(instanceId, question, "v2") + handlePermissionReplied(instanceId, { type: "permission.replied", properties: { requestID: permission.id } } as any) + handleQuestionAnswered(instanceId, { type: "question.replied", properties: { requestID: question.id } } as any) + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.permissionLists() === 1 && harness.questionLists() === 1) + await new Promise((resolve) => setTimeout(resolve, 20)) + + handlePermissionUpdated(instanceId, { type: "permission.v2.asked", properties: permission } as any) + handleQuestionAsked(instanceId, { type: "question.v2.asked", properties: question } as any) + assert.deepEqual(getPermissionQueue(instanceId), []) + assert.deepEqual(getQuestionQueue(instanceId), []) + } finally { + harness.cleanup() + } + }) + + it("tombstones authoritative absence before delayed ask events", async () => { + const instanceId = "reconnect-authoritative-absence" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + const harness = setup(instanceId) + + try { + addPermissionToQueue(instanceId, permission, "v2") + addQuestionToQueue(instanceId, question, "v2") + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.permissionLists() === 1 && harness.questionLists() === 1) + await waitFor(() => getPermissionQueue(instanceId).length === 0 && getQuestionQueue(instanceId).length === 0) + + handlePermissionUpdated(instanceId, { type: "permission.v2.asked", properties: permission } as any) + handleQuestionAsked(instanceId, { type: "question.v2.asked", properties: question } as any) + assert.deepEqual(getPermissionQueue(instanceId), []) + assert.deepEqual(getQuestionQueue(instanceId), []) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("refreshes worktree topology before reconnect interruption lists", async () => { + const instanceId = "reconnect-worktree-refresh" + const originalFetchWorktrees = serverApi.fetchWorktrees + let worktreeLists = 0 + serverApi.fetchWorktrees = async () => { + worktreeLists += 1 + return { isGitRepo: true, worktrees: [{ slug: "root", directory: "/work", kind: "root" }] } + } + const harness = setup(instanceId) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => worktreeLists === 1 && harness.permissionLists() === 1) + assert.equal(worktreeLists, 1) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("keeps interruption queues when reconnect worktree freshness fails", async () => { + const instanceId = "reconnect-stale-worktrees" + const permission = { id: "permission", sessionID: "session", permission: "read", patterns: [] } as any + const question = { id: "question", sessionID: "session", questions: [] } as any + const originalFetchWorktrees = serverApi.fetchWorktrees + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + const harness = setup(instanceId) + + try { + assert.equal(await reloadWorktrees(instanceId), true) + addPermissionToQueue(instanceId, permission, "v2") + addQuestionToQueue(instanceId, question, "v2") + serverApi.fetchWorktrees = async () => { throw new Error("worktree list unavailable") } + + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.permissionLists() === 1 && harness.questionLists() === 1) + assert.deepEqual(getPermissionQueue(instanceId).map(({ id }) => id), [permission.id]) + assert.deepEqual(getQuestionQueue(instanceId).map(({ id }) => id), [question.id]) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("force-refreshes OpenCode workspace ids before reconnect lists", async () => { + const instanceId = "reconnect-workspace-id-refresh" + const originalFetchWorktrees = serverApi.fetchWorktrees + let workspaceId = "workspace-old" + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [ + { slug: "root", directory: "/work", kind: "root" }, + { slug: "feature", directory: "/feature", kind: "worktree" }, + ], + }) + const harness = setup(instanceId) + harness.client.experimental = { + workspace: { + syncList: async () => ({ data: [] }), + list: async () => ({ data: [{ id: workspaceId, directory: "/feature" }] }), + }, + } + + try { + await reloadWorktrees(instanceId) + await reloadOpenCodeWorkspaces(instanceId) + workspaceId = "workspace-new" + + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.permissionLocations.some((location) => location?.workspace === workspaceId)) + assert.equal(harness.permissionLocations.some((location) => location?.workspace === "workspace-old"), false) + assert.equal(harness.questionLocations.some((location) => location?.workspace === workspaceId), true) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + harness.cleanup() + } + }) + + it("bounds answered-question tombstones while retaining ordinary delayed-event fences", () => { + const instanceId = "bounded-question-tombstones" + try { + markQuestionAnswered(instanceId, "recent", 1_000) + assert.equal(hasAnsweredQuestion(instanceId, "recent", 2_000), true) + assert.equal(hasAnsweredQuestion(instanceId, "recent", Number.MAX_SAFE_INTEGER), false) + + for (let index = 0; index <= 4_096; index += 1) { + markQuestionAnswered(instanceId, `question-${index}`, 10_000 + index) + } + assert.equal(hasAnsweredQuestion(instanceId, "question-0", 20_000), false) + assert.equal(hasAnsweredQuestion(instanceId, "question-4096", 20_000), true) + } finally { + removeInstance(instanceId, { authoritative: false }) + } + }) + + it("does not carry a reconnect pass into a replacement runtime", async () => { + const instanceId = "reconnect-runtime-fence" + const sessions = deferred() + const harness = setup(instanceId, { sessions: sessions.promise }) + + try { + sseManager.onConnectionRestored?.(instanceId) + await waitFor(() => harness.sessionLists() === 1) + updateInstance(instanceId, { client: { session: {} } as any }) + sessions.resolve({ data: [] }) + await new Promise((resolve) => setTimeout(resolve, 20)) + assert.equal(harness.permissionLists(), 0) + assert.equal(harness.questionLists(), 0) + } finally { + harness.cleanup() + } + }) + + it("lets healthy initial session hydration commit before reconnect refresh", async () => { + const instanceId = "reconnect-during-initial-hydration", sessionId = "restored-idle" + const originalFetchWorktrees = serverApi.fetchWorktrees + const originalReadWorktreeMap = serverApi.readWorktreeMap + const initialSessions = deferred() + const reconnectSessions = deferred() + let sessionLists = 0 + const apiSession = { id: sessionId, title: "Restored", directory: "/work", time: { created: 1, updated: 1 } } + const client = { + session: { + list: () => ++sessionLists === 1 ? initialSessions.promise : reconnectSessions.promise, + status: async () => ({ data: {} }), + }, + app: { agents: async () => ({ data: [] }) }, + config: { providers: async () => ({ data: { providers: [], default: {} } }) }, + command: { list: async () => ({ data: [] }) }, + experimental: { workspace: { + syncList: async () => ({ data: [] }), + list: async () => ({ data: [] }), + } }, + permission: { list: async () => ({ data: [] }) }, + question: { list: async () => ({ data: [] }) }, + v2: { + permission: { request: { list: async () => ({ data: { data: [] } }) } }, + question: { request: { list: async () => ({ data: { data: [] } }) } }, + }, + } as any + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + serverApi.readWorktreeMap = async () => null as any + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, client) + + try { + ;(serverEvents as any).dispatch({ + type: "workspace.started", + workspace: { + id: instanceId, path: "/work", status: "ready", pid: 1, port: 1, + proxyPath: `/workspaces/${instanceId}/instance`, binaryId: "test", binaryLabel: "Test", + createdAt: new Date(0).toISOString(), updatedAt: new Date(0).toISOString(), + }, + }) + await waitFor(() => sessionLists === 1) + sseManager.onConnectionRestored?.(instanceId) + initialSessions.resolve({ data: [apiSession] }) + await waitFor(() => sessionLists === 2) + + assert.equal(sessions().get(instanceId)?.has(sessionId), true) + reconnectSessions.resolve({ data: [apiSession] }) + } finally { + reconnectSessions.resolve({ data: [apiSession] }) + serverApi.fetchWorktrees = originalFetchWorktrees + serverApi.readWorktreeMap = originalReadWorktreeMap + removeInstance(instanceId, { authoritative: false }) + sdkManager.destroyClientsForInstance(instanceId) + } + }) + + it("times out stalled initial hydration and runs the queued trailing reconnect", { timeout: 15_000 }, async () => { + const instanceId = "reconnect-stalled-initial-hydration" + const originalFetchWorktrees = serverApi.fetchWorktrees + const originalReadWorktreeMap = serverApi.readWorktreeMap + let sessionLists = 0 + let permissionLists = 0 + const agentResponse = deferred() + let agentSignal: AbortSignal | undefined + const client = { + session: { + list: async () => { + sessionLists += 1 + return { data: [] } + }, + status: async () => ({ data: {} }), + }, + app: { agents: (_parameters: unknown, options?: { signal?: AbortSignal }) => { + agentSignal = options?.signal + return agentResponse.promise + } }, + experimental: { workspace: { + syncList: async () => ({ data: [] }), + list: async () => ({ data: [] }), + } }, + permission: { list: async () => ({ data: [] }) }, + question: { list: async () => ({ data: [] }) }, + v2: { + permission: { request: { list: async () => { + permissionLists += 1 + return { data: { data: [] } } + } } }, + question: { request: { list: async () => ({ data: { data: [] } }) } }, + }, + } as any + serverApi.fetchWorktrees = async () => ({ + isGitRepo: true, + worktrees: [{ slug: "root", directory: "/work", kind: "root" }], + }) + serverApi.readWorktreeMap = async () => null as any + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, client) + + try { + ;(serverEvents as any).dispatch({ + type: "workspace.started", + workspace: { + id: instanceId, + path: "/work", + status: "ready", + pid: 1, + port: 1, + proxyPath: `/workspaces/${instanceId}/instance`, + binaryId: "test", + binaryLabel: "Test", + createdAt: new Date(0).toISOString(), + updatedAt: new Date(0).toISOString(), + }, + }) + await waitFor(() => sessionLists === 1 && agentSignal !== undefined) + + sseManager.onConnectionRestored?.(instanceId) + sseManager.onConnectionRestored?.(instanceId) + + await waitFor(() => permissionLists === 2, 12_000) + assert.equal(agentSignal?.aborted, true) + assert.ok(sessionLists >= 3) + agentResponse.resolve({ data: [{ name: "stale-agent", mode: "primary" }] }) + await new Promise((resolve) => setImmediate(resolve)) + assert.equal(agents().get(instanceId)?.some((agent) => agent.name === "stale-agent") ?? false, false) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + serverApi.readWorktreeMap = originalReadWorktreeMap + removeInstance(instanceId, { authoritative: false }) + sdkManager.destroyClientsForInstance(instanceId) + } + }) +}) diff --git a/packages/ui/src/stores/instances.ts b/packages/ui/src/stores/instances.ts index 04a40f3dc..b13648f9b 100644 --- a/packages/ui/src/stores/instances.ts +++ b/packages/ui/src/stores/instances.ts @@ -16,6 +16,7 @@ import { fetchSessions, fetchAgents, fetchProviders, + loadMessages, clearInstanceDraftPrompts, clearSessionListRequestState, clearInstanceDeletedSessionAuthority, @@ -31,17 +32,24 @@ import { reloadWorktrees, } from "./worktrees" import { getRootClient } from "./opencode-client" -import { clearOpenCodeWorkspaceCache, getOpenCodeWorkspaceIdForSession, getOpenCodeWorkspaceIdForWorktree, syncOpenCodeWorkspaces } from "./opencode-workspaces" +import { clearOpenCodeWorkspaceCache, getOpenCodeWorkspaceIdForSession, getOpenCodeWorkspaceIdForWorktree, reloadOpenCodeWorkspaces, syncOpenCodeWorkspaces } from "./opencode-workspaces" import { fetchCommands, clearCommands } from "./commands" import { serverSettings } from "./preferences" import { reconcileSessionPendingState, + activeSessionId, + invalidateSessionMessageLoad, + invalidateOwnedSessionMessageLoads, + getSessionRoot, + purgeInstanceSessionState, sessions, setSessionPendingPermission, setSessionPendingQuestion, + resetInstanceSessionRequestState, } from "./session-state" import { setHasInstances } from "./ui" import { messageStoreBus } from "./message-v2/bus" +import { clearPendingDeltasForInstance } from "./delta-buffer" import { upsertPermissionV2, removePermissionV2, upsertQuestionV2, removeQuestionV2 } from "./message-v2/bridge" import { clearRepliedPermissions, @@ -62,7 +70,7 @@ import { getLogger } from "../lib/logger" import { mergeInstanceMetadata, clearInstanceMetadata } from "./instance-metadata" import { showWorkspaceLaunchError } from "./launch-errors" import { activeSidecarToken } from "./sidecars" -import { buildV2RequestLocations, type V2Location } from "./request-locations" +import { buildV2RequestLocations, type V2RequestLocationSnapshot } from "./request-locations" import { showToastNotification } from "../lib/notifications" import { tGlobal } from "../lib/i18n" import { appSessionRestoreGateActive } from "./app-session-restore-gate" @@ -73,6 +81,7 @@ import { getAbortReason } from "./app-session-restore-timeout" import { AbortCreatedWorkspaceCleanup } from "./abort-created-workspace-cleanup" import { TrailingResyncCoordinator, waitForSettledPrerequisite } from "../lib/trailing-resync" import { retryWithBackoff } from "../lib/retry-utils" +import { getVisibleSessionMemoryIds } from "./session-memory" import { cancelRestoreCreation } from "./restore-creation-cancellation" import { RestoreWorkspaceCommitGates, type RestoreWorkspaceCommitGate, type RestoreWorkspaceTerminal, @@ -80,6 +89,22 @@ import { import { WorkspaceListReconciliationFence } from "./workspace-list-reconciliation-fence" const log = getLogger("api") +const RECONNECT_LIST_TIMEOUT_MS = 5_000 + +async function waitForReconnectPrerequisite(prerequisite: Promise | undefined, signal?: AbortSignal): Promise { + if (!signal) return waitForSettledPrerequisite(prerequisite) + let onAbort: (() => void) | undefined + const aborted = new Promise((_resolve, reject) => { + onAbort = () => reject(signal.reason) + if (signal.aborted) onAbort() + else signal.addEventListener("abort", onAbort, { once: true }) + }) + try { + await Promise.race([waitForSettledPrerequisite(prerequisite), aborted]) + } finally { + if (onAbort) signal.removeEventListener("abort", onAbort) + } +} setPermissionAutoAcceptFamilyRootResolver((instanceId, sessionId) => { const instanceSessions = sessions().get(instanceId) @@ -113,8 +138,21 @@ serverEvents.on("yolo.autoAccepted", (event) => { }) const [instances, setInstances] = createSignal>(new Map()) +let workspaceListInitialized = false +sseManager.shouldHandleEvent = (instanceId) => !workspaceListInitialized || instances().has(instanceId) -const [activeInstanceId, setActiveInstanceId] = createSignal(null) +const [activeInstanceId, writeActiveInstanceId] = createSignal(null) + +function setActiveInstanceId(instanceId: string | null): void { + const previousInstanceId = activeInstanceId() + if (previousInstanceId && previousInstanceId !== instanceId) { + const previousSessionId = activeSessionId().get(previousInstanceId) + if (previousSessionId) { + invalidateOwnedSessionMessageLoads(previousInstanceId, getSessionRoot(previousInstanceId, previousSessionId)?.id ?? previousSessionId) + } + } + writeActiveInstanceId(instanceId) +} const [instanceLogs, setInstanceLogs] = createSignal>(new Map()) const [logStreamingState, setLogStreamingState] = createSignal>(new Map()) @@ -124,10 +162,16 @@ const [activePermissionId, setActivePermissionId] = createSignal>(new Map()) const [activeQuestionId, setActiveQuestionId] = createSignal>(new Map()) +type AnsweredQuestion = { answeredAt: number; missingPasses: number } +const answeredQuestionIdsByInstance = new Map>() +const ANSWERED_QUESTION_TOMBSTONE_TTL_MS = 10 * 60 * 1_000 +const MAX_ANSWERED_QUESTION_TOMBSTONES = 4_096 class InterruptionRegistry { private readonly enqueuedAt = new Map() + private readonly mutationVersions = new Map() private readonly sources = new Map>() private readonly sessionCounts = new Map>() + private mutationVersion = 0 constructor(private readonly defaultSource: S) {} @@ -143,6 +187,18 @@ class InterruptionRegistry { return this.enqueuedAt.get(requestId) ?? Date.now() } + markMutation(requestId: string): void { + this.mutationVersions.set(requestId, ++this.mutationVersion) + } + + snapshotMutationVersion(): number { + return this.mutationVersion + } + + wasMutatedAfter(requestId: string, version: number): boolean { + return (this.mutationVersions.get(requestId) ?? 0) > version + } + setSource(instanceId: string, requestId: string, source: S): void { const sources = this.sources.get(instanceId) ?? new Map() sources.set(requestId, source) @@ -155,6 +211,7 @@ class InterruptionRegistry { remove(instanceId: string, requestId: string): void { this.enqueuedAt.delete(requestId) + this.mutationVersions.delete(requestId) const sources = this.sources.get(instanceId) sources?.delete(requestId) if (sources?.size === 0) this.sources.delete(instanceId) @@ -180,7 +237,10 @@ class InterruptionRegistry { } clear(instanceId: string, requests: readonly T[], clearPending: (sessionId: string) => void): void { - requests.forEach(({ id }) => this.enqueuedAt.delete(id)) + requests.forEach(({ id }) => { + this.enqueuedAt.delete(id) + this.mutationVersions.delete(id) + }) this.sources.delete(instanceId) for (const sessionId of this.sessionCounts.get(instanceId)?.keys() ?? []) clearPending(sessionId) this.sessionCounts.delete(instanceId) @@ -194,19 +254,23 @@ type InterruptionKind = "permission" | "question" type ActiveInterruption = { kind: InterruptionKind; id: string } | null -async function getV2RequestLocations(instanceId: string): Promise { +async function getV2RequestLocations(instanceId: string, topologyComplete = true): Promise { const instance = instances().get(instanceId) const worktrees = getWorktrees(instanceId) const workspaceBySlug = new Map() for (const worktree of worktrees) { if (!worktree.slug || worktree.slug === "root") continue - const workspace = await getOpenCodeWorkspaceIdForWorktree(instanceId, worktree.slug) + const workspace = await getOpenCodeWorkspaceIdForWorktree(instanceId, worktree.slug).catch((error) => { + log.warn("Failed to resolve OpenCode workspace for pending requests", { instanceId, slug: worktree.slug, error }) + return null + }) if (!workspace) continue workspaceBySlug.set(worktree.slug, workspace) } - return buildV2RequestLocations(instance?.folder, worktrees, workspaceBySlug) + const snapshot = buildV2RequestLocations(instance?.folder, worktrees, workspaceBySlug) + return { ...snapshot, complete: topologyComplete && snapshot.complete } } const [activeInterruption, setActiveInterruption] = createSignal>(new Map()) @@ -227,9 +291,22 @@ const MAX_LOG_ENTRIES = 1000 const pendingDisposeRequests = new Map>() const pendingRehydrations = new Map>() -const initialHydrations = new Map>() +type InitialHydration = { promise: Promise; controller: AbortController; authority: number } +const initialHydrations = new Map() const initialSessionHydrations = new Map>() const initialWorkspaceMetadataHydrations = new Map>() +const hydrationAuthorities = new Map() +let hydrationAuthoritySequence = 0 + +function claimHydrationAuthority(instanceId: string): number { + const authority = ++hydrationAuthoritySequence + hydrationAuthorities.set(instanceId, authority) + return authority +} + +function hasHydrationAuthority(instanceId: string, authority: number): boolean { + return hydrationAuthorities.get(instanceId) === authority +} type RestoreWorkspaceDescriptor = WorkspaceDescriptor & { reused?: boolean } const workspaceListReconciliationFence = new WorkspaceListReconciliationFence() @@ -254,10 +331,81 @@ const restoreCreationCommitGates = new RestoreWorkspaceCommitGates { - await waitForSettledPrerequisite(initialHydrations.get(instanceId)) const instance = instances().get(instanceId) if (!instance?.client || instance.status !== "ready") return - await fetchSessions(instanceId, { reset: false }) + const initialHydration = initialHydrations.get(instanceId) + await retryWithBackoff( + (signal) => waitForReconnectPrerequisite(initialHydration?.promise, signal), + { maxAttempts: 1, timeoutMs: RECONNECT_LIST_TIMEOUT_MS }, + ).catch((error) => { + if (initialHydration && initialHydrations.get(instanceId) === initialHydration) { + initialHydration.controller.abort(error) + initialHydrations.delete(instanceId) + initialSessionHydrations.delete(instanceId) + initialWorkspaceMetadataHydrations.delete(instanceId) + } + log.warn("Initial hydration did not settle before reconnect resync", { instanceId, error }) + }) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return + const reconnectAuthority = claimHydrationAuthority(instanceId) + const hasCommitAuthority = () => hasHydrationAuthority(instanceId, reconnectAuthority) && + isInstanceRuntimeCurrent(instanceId, instance) + if (!hasCommitAuthority()) return + const worktreeTopologyComplete = await retryWithBackoff((signal) => reloadWorktrees(instanceId, signal), { + maxAttempts: 1, + timeoutMs: RECONNECT_LIST_TIMEOUT_MS, + }).catch((error) => { + log.warn("Failed to refresh worktrees after instance connection", { instanceId, error }) + return false + }) + if (!hasCommitAuthority()) return + await reloadOpenCodeWorkspaces(instanceId) + if (!hasCommitAuthority()) return + const sessionScopeComplete = worktreeTopologyComplete && (await getV2RequestLocations(instanceId, true)).complete + const refreshedSessionIds = await retryWithBackoff((signal) => fetchSessions(instanceId, { + reset: false, + authoritativeDeletes: sessionScopeComplete, + signal, + hasCommitAuthority, + }), { maxAttempts: 1, timeoutMs: RECONNECT_LIST_TIMEOUT_MS }).catch((error) => { + log.warn("Failed to refresh sessions after instance connection", { instanceId, error }) + return new Set() + }) + if (!hasCommitAuthority()) return + const interruptionSyncs = await Promise.allSettled([ + retryWithBackoff((signal) => syncPendingPermissions(instanceId, instance, signal, worktreeTopologyComplete), { + maxAttempts: 3, + initialDelayMs: 100, + maxDelayMs: 200, + timeoutMs: RECONNECT_LIST_TIMEOUT_MS, + }), + retryWithBackoff((signal) => syncPendingQuestions(instanceId, instance, signal, worktreeTopologyComplete), { + maxAttempts: 3, + initialDelayMs: 100, + maxDelayMs: 200, + timeoutMs: RECONNECT_LIST_TIMEOUT_MS, + }), + ]) + for (const result of interruptionSyncs) { + if (result.status === "rejected") log.warn("Failed to resync pending requests after instance connection", { instanceId, error: result.reason }) + } + if (!hasCommitAuthority()) return + const residentSessionIds = messageStoreBus.getInstance(instanceId)?.getResidentSessionIds() ?? [] + const activeId = activeSessionId().get(instanceId) + const visibleSessionIds = getVisibleSessionMemoryIds(instanceId) + const reloadSessionIds = new Set(visibleSessionIds) + if (activeId) reloadSessionIds.add(activeId) + const knownSessions = sessions().get(instanceId) + for (const sessionId of new Set([...residentSessionIds, ...reloadSessionIds])) { + if (knownSessions?.has(sessionId)) invalidateSessionMessageLoad(instanceId, sessionId) + } + await Promise.all(Array.from(reloadSessionIds) + .filter((sessionId) => knownSessions?.has(sessionId)) + .map((sessionId) => loadMessages(instanceId, sessionId, { + force: true, + timeoutMs: RECONNECT_LIST_TIMEOUT_MS, + applySessionRevert: refreshedSessionIds.has(sessionId), + }))) }, (instanceId, error) => { log.warn("Failed to resync sessions after instance connection", { instanceId, error }) @@ -268,13 +416,12 @@ function resyncConnectedInstance(instanceId: string): void { void connectionResyncs.request(instanceId) } -serverEvents.on("instance.eventStatus", (event) => { - if (event.type !== "instance.eventStatus" || event.status !== "connected") return - if (disconnectedInstance()?.id === event.instanceId) { +sseManager.onConnectionRestored = (instanceId) => { + if (disconnectedInstance()?.id === instanceId) { setDisconnectedInstance(null) } - resyncConnectedInstance(event.instanceId) -}) + resyncConnectedInstance(instanceId) +} function createRestoreCreationRequestId(): string { if (typeof globalThis.crypto?.randomUUID === "function") { @@ -351,7 +498,8 @@ function ensureActiveInstanceSelected(): void { function upsertWorkspace(descriptor: WorkspaceDescriptor, projectName?: string) { const mapped = workspaceDescriptorToInstance(descriptor, projectName) if (instances().has(descriptor.id)) { - updateInstance(descriptor.id, mapped) + const { client: _client, port: _port, proxyPath: _proxyPath, ...metadata } = mapped + updateInstance(descriptor.id, metadata) } else { addInstance(mapped) } @@ -376,8 +524,11 @@ function attachClient(descriptor: WorkspaceDescriptor) { const nextPort = descriptor.port ?? instance.port const nextProxyPath = descriptor.proxyPath + const runtimeChanged = Boolean(instance.client && ( + (descriptor.pid && instance.attachedPid !== descriptor.pid) || instance.proxyPath !== nextProxyPath + )) - if (instance.client && instance.proxyPath === nextProxyPath) { + if (!runtimeChanged && instance.client && instance.proxyPath === nextProxyPath) { if (nextPort && instance.port !== nextPort) { updateInstance(descriptor.id, { port: nextPort }) } @@ -385,6 +536,10 @@ function attachClient(descriptor: WorkspaceDescriptor) { } if (instance.client) { + initialHydrations.get(descriptor.id)?.controller.abort(new Error("Instance runtime replaced")) + clearOpenCodeWorkspaceCache(descriptor.id) + clearPendingDeltasForInstance(descriptor.id) + if (runtimeChanged) clearReloadableInstanceState(descriptor.id) sdkManager.destroyClientsForInstance(descriptor.id) } @@ -393,18 +548,24 @@ function attachClient(descriptor: WorkspaceDescriptor) { client, port: nextPort ?? 0, proxyPath: nextProxyPath, + attachedPid: descriptor.pid ?? instance.pid, status: "ready", }) sseManager.seedStatusIfMissing(descriptor.id, "connecting") - const sessionHydration = startInstanceSessionHydration(descriptor.id) + const controller = new AbortController() + const authority = claimHydrationAuthority(descriptor.id) + const hasCommitAuthority = () => hasHydrationAuthority(descriptor.id, authority) + const sessionHydration = startInstanceSessionHydration(descriptor.id, false, controller.signal, hasCommitAuthority) initialSessionHydrations.set(descriptor.id, sessionHydration.sessions) initialWorkspaceMetadataHydrations.set(descriptor.id, sessionHydration.workspaceMetadata) const hydration = hydrateInstanceData(descriptor.id, { propagateErrors: true, sessionHydration: sessionHydration.sessions, workspaceMetadataHydration: sessionHydration.workspaceMetadata, + signal: controller.signal, + hasCommitAuthority, }) - initialHydrations.set(descriptor.id, hydration) + initialHydrations.set(descriptor.id, { promise: hydration, controller, authority }) if (sseManager.getStatuses().get(descriptor.id) === "connected") { resyncConnectedInstance(descriptor.id) } @@ -442,7 +603,7 @@ function waitForInstanceReady(instanceId: string): Promise { async function waitForInstanceInitialHydration(instanceId: string): Promise { await waitForInstanceReady(instanceId) - await initialHydrations.get(instanceId) + await initialHydrations.get(instanceId)?.promise } async function waitForInstanceInitialSessionHydration(instanceId: string): Promise { @@ -466,49 +627,65 @@ function releaseInstanceResources(instanceId: string) { sseManager.seedStatus(instanceId, "disconnected") } -async function syncPendingPermissions(instanceId: string): Promise { - const instance = instances().get(instanceId) +async function syncPendingPermissions(instanceId: string, expectedInstance?: Instance, signal?: AbortSignal, topologyComplete = true): Promise { + const instance = expectedInstance ?? instances().get(instanceId) if (!instance?.client) return + const isCurrent = () => isInstanceRuntimeCurrent(instanceId, instance) && !signal?.aborted + if (!isCurrent()) return try { const syncStartedAt = Date.now() + const mutationVersion = permissionRegistry.snapshotMutationVersion() const remote: Array<{ request: PermissionRequest; source: PermissionSource }> = [] - const legacyRemote = await requestData( - instance.client.permission.list(), - "permission.list", - ).catch((error) => { - log.warn("Failed to list legacy pending permissions", { instanceId, error }) - return [] - }) - for (const permission of legacyRemote) { - permissionRegistry.setSource(instanceId, permission.id, "legacy") - remote.push({ request: permission, source: "legacy" }) + let requestFailed = false + try { + const legacyRemote = await requestData( + (instance.client.permission.list as any)(undefined, signal ? { signal } : undefined), + "permission.list", + ) + if (!isCurrent()) return + for (const permission of legacyRemote) remote.push({ request: permission, source: "legacy" }) + } catch (error) { + if (!isCurrent()) return + requestFailed = true + log.warn("Failed to sync legacy pending permissions", { instanceId, error }) } - for (const location of await getV2RequestLocations(instanceId)) { - const response = await requestData<{ location?: unknown; data: PermissionRequest[] }>( - instance.client.v2.permission.request.list({ location }), - "v2.permission.request.list", - ) - log.info("v2.permission.request.list", { instanceId, location, resolvedLocation: response.location }) - for (const permission of response.data) { - permissionRegistry.setSource(instanceId, permission.id, "v2") - remote.push({ request: permission, source: "v2" }) + const locationSnapshot = await getV2RequestLocations(instanceId, topologyComplete) + if (!isCurrent()) return + let complete = locationSnapshot.complete && !requestFailed + for (const location of locationSnapshot.locations) { + try { + const response = await requestData<{ location?: unknown; data: PermissionRequest[] }>( + (instance.client.v2.permission.request.list as any)({ location }, signal ? { signal } : undefined), + "v2.permission.request.list", + ) + if (!isCurrent()) return + log.info("v2.permission.request.list", { instanceId, location, resolvedLocation: response.location }) + for (const permission of response.data) remote.push({ request: permission, source: "v2" }) + } catch (error) { + if (!isCurrent()) return + complete = false + requestFailed = true + log.warn("Failed to sync a permission request location", { instanceId, location, error }) } } + if (!isCurrent()) return const remotePendingIds = new Set(remote.map((item) => item.request.id)) - pruneRepliedPermissions(instanceId, remotePendingIds, syncStartedAt) + if (complete) pruneRepliedPermissions(instanceId, remotePendingIds, syncStartedAt) const pendingRemote = remote.filter((item) => !hasRepliedPermission(instanceId, item.request.id)) const remoteIds = new Set(pendingRemote.map((item) => item.request.id)) const local = getPermissionQueue(instanceId) - // Remove any stale local permissions missing from server. - for (const entry of local) { - if (!remoteIds.has(entry.id)) { - removePermissionFromQueue(instanceId, entry.id) - removePermissionV2(instanceId, entry.id) + if (complete) { + for (const entry of local) { + if (!remoteIds.has(entry.id) && !permissionRegistry.wasMutatedAfter(entry.id, mutationVersion)) { + markPermissionReplied(instanceId, entry.id) + removePermissionFromQueue(instanceId, entry.id) + removePermissionV2(instanceId, entry.id) + } } } @@ -518,65 +695,94 @@ async function syncPendingPermissions(instanceId: string): Promise { upsertPermissionV2(instanceId, queuedPermission) } reconcilePendingSessionIndicators(instanceId) + if (requestFailed) throw new Error("One or more permission request locations failed") } catch (error) { log.warn("Failed to sync pending permissions", { instanceId, error }) + throw error } } -async function syncPendingQuestions(instanceId: string): Promise { - const instance = instances().get(instanceId) +async function syncPendingQuestions(instanceId: string, expectedInstance?: Instance, signal?: AbortSignal, topologyComplete = true): Promise { + const instance = expectedInstance ?? instances().get(instanceId) if (!instance?.client) return + const isCurrent = () => isInstanceRuntimeCurrent(instanceId, instance) && !signal?.aborted + if (!isCurrent()) return try { + const syncStartedAt = Date.now() + const mutationVersion = questionRegistry.snapshotMutationVersion() const remote: Array<{ request: QuestionRequest; source: QuestionSource }> = [] - const legacyRemote = await requestData( - instance.client.question.list(), - "question.list", - ).catch((error) => { - log.warn("Failed to list legacy pending questions", { instanceId, error }) - return [] - }) - for (const request of legacyRemote) { - questionRegistry.setSource(instanceId, request.id, "legacy") - remote.push({ request, source: "legacy" }) + let requestFailed = false + try { + const legacyRemote = await requestData( + (instance.client.question.list as any)(undefined, signal ? { signal } : undefined), + "question.list", + ) + if (!isCurrent()) return + for (const request of legacyRemote) remote.push({ request, source: "legacy" }) + } catch (error) { + if (!isCurrent()) return + requestFailed = true + log.warn("Failed to sync legacy pending questions", { instanceId, error }) } - for (const location of await getV2RequestLocations(instanceId)) { - const response = await requestData<{ location?: unknown; data: QuestionRequest[] }>( - instance.client.v2.question.request.list({ location }), - "v2.question.request.list", - ) - log.info("v2.question.request.list", { instanceId, location, resolvedLocation: response.location }) - for (const request of response.data) { - questionRegistry.setSource(instanceId, request.id, "v2") - remote.push({ request, source: "v2" }) + const locationSnapshot = await getV2RequestLocations(instanceId, topologyComplete) + if (!isCurrent()) return + let complete = locationSnapshot.complete && !requestFailed + for (const location of locationSnapshot.locations) { + try { + const response = await requestData<{ location?: unknown; data: QuestionRequest[] }>( + (instance.client.v2.question.request.list as any)({ location }, signal ? { signal } : undefined), + "v2.question.request.list", + ) + if (!isCurrent()) return + log.info("v2.question.request.list", { instanceId, location, resolvedLocation: response.location }) + for (const request of response.data) remote.push({ request, source: "v2" }) + } catch (error) { + if (!isCurrent()) return + complete = false + requestFailed = true + log.warn("Failed to sync a question request location", { instanceId, location, error }) } } - const remoteIds = new Set(remote.map((item) => item.request.id)) + if (!isCurrent()) return + const remotePendingIds = new Set(remote.map((item) => item.request.id)) + if (complete) pruneAnsweredQuestions(instanceId, remotePendingIds, syncStartedAt) + const pendingRemote = remote.filter((item) => !hasAnsweredQuestion(instanceId, item.request.id)) + const remoteIds = new Set(pendingRemote.map((item) => item.request.id)) const local = getQuestionQueue(instanceId) - // Remove any stale local requests missing from server. - for (const entry of local) { - if (!remoteIds.has(entry.id)) { - removeQuestionFromQueue(instanceId, entry.id) - removeQuestionV2(instanceId, entry.id) + if (complete) { + for (const entry of local) { + if (!remoteIds.has(entry.id) && !questionRegistry.wasMutatedAfter(entry.id, mutationVersion)) { + markQuestionAnswered(instanceId, entry.id) + removeQuestionFromQueue(instanceId, entry.id) + removeQuestionV2(instanceId, entry.id) + } } } // Upsert all server-side pending questions. - for (const { request, source } of remote) { + for (const { request, source } of pendingRemote) { questionRegistry.ensureEnqueuedAt(request) addQuestionToQueue(instanceId, request, source) upsertQuestionV2(instanceId, request) } reconcilePendingSessionIndicators(instanceId) + if (requestFailed) throw new Error("One or more question request locations failed") } catch (error) { log.warn("Failed to sync pending questions", { instanceId, error }) + throw error } } -function startInstanceSessionHydration(instanceId: string, force = false): { +function startInstanceSessionHydration( + instanceId: string, + force = false, + signal?: AbortSignal, + hasCommitAuthority: () => boolean = () => true, +): { sessions: Promise workspaceMetadata: Promise } { @@ -584,15 +790,20 @@ function startInstanceSessionHydration(instanceId: string, force = false): { ? reloadWorktreeMap(instanceId) : ensureWorktreeMapLoaded(instanceId) const worktreeHydration = force - ? reloadWorktrees(instanceId) - : ensureWorktreesLoaded(instanceId) + ? reloadWorktrees(instanceId, signal) + : ensureWorktreesLoaded(instanceId, signal) const sessions = Promise.all([worktreeHydration, worktreeMapHydration]).then(async () => { + signal?.throwIfAborted() + if (!hasCommitAuthority()) return resetSessionPagination(instanceId) - await fetchSessions(instanceId).catch((error) => { + await fetchSessions(instanceId, { signal, hasCommitAuthority }).catch((error) => { log.error("Failed to hydrate sessions", { instanceId, error }) }) + signal?.throwIfAborted() }) const workspaceMetadata = worktreeHydration.then(async () => { + signal?.throwIfAborted() + if (!hasCommitAuthority()) return await Promise.all([worktreeMapHydration, syncOpenCodeWorkspaces(instanceId)]) }) return { sessions, workspaceMetadata } @@ -603,24 +814,40 @@ async function hydrateInstanceData(instanceId: string, options?: { propagateErrors?: boolean sessionHydration?: Promise workspaceMetadataHydration?: Promise + signal?: AbortSignal + hasCommitAuthority?: () => boolean }) { + const hasCommitAuthority = () => !options?.signal?.aborted && (options?.hasCommitAuthority?.() ?? true) try { const hydration = options?.sessionHydration ? { sessions: options.sessionHydration, workspaceMetadata: options.workspaceMetadataHydration ?? Promise.resolve(), } - : startInstanceSessionHydration(instanceId, options?.force) + : startInstanceSessionHydration(instanceId, options?.force, options?.signal, hasCommitAuthority) await hydration.sessions + if (!hasCommitAuthority()) return await hydration.workspaceMetadata - await fetchAgents(instanceId) - await fetchProviders(instanceId) + if (!hasCommitAuthority()) return + await fetchAgents(instanceId, options?.signal, hasCommitAuthority) + if (!hasCommitAuthority()) return + await fetchProviders(instanceId, options?.signal, hasCommitAuthority) + if (!hasCommitAuthority()) return await ensureInstanceConfigLoaded(instanceId) + if (!hasCommitAuthority()) return const instance = instances().get(instanceId) if (!instance?.client) return - await fetchCommands(instanceId, instance.client) - await syncPendingPermissions(instanceId) - await syncPendingQuestions(instanceId) + await fetchCommands( + instanceId, + instance.client, + () => hasCommitAuthority() && isInstanceRuntimeCurrent(instanceId, instance), + options?.signal, + ) + if (!hasCommitAuthority()) return + await Promise.all([ + syncPendingPermissions(instanceId, instance, options?.signal), + syncPendingQuestions(instanceId, instance, options?.signal), + ]) } catch (error) { log.error("Failed to fetch initial data", error) if (options?.propagateErrors) throw error @@ -666,12 +893,15 @@ async function postInstanceDispose(instanceId: string): Promise { } function clearReloadableInstanceState(instanceId: string): void { + resetInstanceSessionRequestState(instanceId) clearCacheForInstance(instanceId) clearCommands(instanceId) clearInstanceMetadata(instanceId) messageStoreBus.clearInstanceScrollSnapshots(instanceId) clearPermissionQueue(instanceId) + clearRepliedPermissions(instanceId) clearQuestionQueue(instanceId) + answeredQuestionIdsByInstance.delete(instanceId) } async function rehydrateInstance(instanceId: string, options?: { reason?: string }): Promise { @@ -679,8 +909,8 @@ async function rehydrateInstance(instanceId: string, options?: { reason?: string return pendingRehydrations.get(instanceId) } + const instance = instances().get(instanceId) const promise = (async () => { - const instance = instances().get(instanceId) if (!instance?.client) { return } @@ -689,8 +919,9 @@ async function rehydrateInstance(instanceId: string, options?: { reason?: string clearReloadableInstanceState(instanceId) await hydrateInstanceData(instanceId, { force: true }) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return })().finally(() => { - pendingRehydrations.delete(instanceId) + if (pendingRehydrations.get(instanceId) === promise) pendingRehydrations.delete(instanceId) }) pendingRehydrations.set(instanceId, promise) @@ -702,14 +933,15 @@ async function disposeInstance(instanceId: string): Promise { return pendingDisposeRequests.get(instanceId)! } + const instance = instances().get(instanceId) const promise = (async () => { const ok = await postInstanceDispose(instanceId) - if (ok) { + if (ok && isInstanceRuntimeCurrent(instanceId, instance)) { await rehydrateInstance(instanceId, { reason: "disposed" }) } return ok })().finally(() => { - pendingDisposeRequests.delete(instanceId) + if (pendingDisposeRequests.get(instanceId) === promise) pendingDisposeRequests.delete(instanceId) }) pendingDisposeRequests.set(instanceId, promise) @@ -754,6 +986,8 @@ const initialWorkspaceLoad = (async function initializeWorkspaces(): Promise<{ e } catch (error) { log.error("Failed to load workspaces", error) return { error } + } finally { + workspaceListInitialized = true } })() let latestWorkspaceLoad = initialWorkspaceLoad @@ -906,7 +1140,7 @@ function addInstance(instance: Instance) { .length setInstances((prev) => { const next = new Map(prev) - next.set(instance.id, instance) + next.set(instance.id, { ...instance, runtimeToken: Symbol(instance.id) }) return next }) ensureLogContainer(instance.id) @@ -925,13 +1159,26 @@ function updateInstance(id: string, updates: Partial) { const next = new Map(prev) const instance = next.get(id) if (instance) { - next.set(id, { ...instance, ...updates }) + next.set(id, { + ...instance, + ...updates, + runtimeToken: (updates.client !== undefined && updates.client !== instance.client) || + (updates.pid !== undefined && updates.pid !== instance.pid) || + (updates.proxyPath !== undefined && updates.proxyPath !== instance.proxyPath) + ? Symbol(id) + : instance.runtimeToken, + }) } return next }) syncHasInstancesFlag() } +function isInstanceRuntimeCurrent(instanceId: string, captured: Instance | undefined): boolean { + const current = instances().get(instanceId) + return Boolean(current && captured && current.runtimeToken === captured.runtimeToken) +} + function removeInstance(id: string, options: { authoritative?: boolean } = {}) { const removedInstance = instances().get(id) const removedOccurrence = removedInstance @@ -977,12 +1224,17 @@ function removeInstance(id: string, options: { authoritative?: boolean } = {}) { clearPermissionQueue(id) clearRepliedPermissions(id) clearQuestionQueue(id) + answeredQuestionIdsByInstance.delete(id) clearInstanceMetadata(id) clearPermissionAutoAcceptForInstance(id) clearSyncedYoloSessionsForInstance(id) + initialHydrations.get(id)?.controller.abort(new Error(`Workspace ${id} was removed`)) initialHydrations.delete(id) initialSessionHydrations.delete(id) initialWorkspaceMetadataHydrations.delete(id) + pendingDisposeRequests.delete(id) + pendingRehydrations.delete(id) + hydrationAuthorities.delete(id) settleInstanceReadyWaiters(id, new Error(`Workspace ${id} was removed before it became ready`)) if (activeInstanceId() === id) { @@ -998,6 +1250,7 @@ function removeInstance(id: string, options: { authoritative?: boolean } = {}) { clearInstanceDeletedSessionAuthority(id) clearInstanceSessionExpansionState(id) clearInstanceSessionSelection(id) + purgeInstanceSessionState(id) if (removedInstance && removedOccurrence >= 0 && options.authoritative !== false) { publishInstanceLifecycleAuthority({ type: "removed", @@ -1396,6 +1649,7 @@ function addPermissionToQueue(instanceId: string, permission: PermissionRequest, let updated = false let previousPermission: PermissionRequest | undefined let queuedPermission = permission + permissionRegistry.markMutation(permission.id) if (source) permissionRegistry.setSource(instanceId, permission.id, source) setPermissionQueues((prev) => { @@ -1555,6 +1809,7 @@ function clearPermissionQueue(instanceId: string): void { function addQuestionToQueue(instanceId: string, request: QuestionRequest, source: QuestionSource = "v2"): void { let inserted = false + questionRegistry.markMutation(request.id) questionRegistry.setSource(instanceId, request.id, source) setQuestionQueues((prev) => { @@ -1630,6 +1885,60 @@ function clearQuestionQueue(instanceId: string): void { recomputeActiveInterruption(instanceId) } +function markQuestionAnswered(instanceId: string, requestId: string, answeredAt = Date.now()): void { + if (!requestId) return + const answered = answeredQuestionIdsByInstance.get(instanceId) ?? new Map() + pruneExpiredAnsweredQuestions(answered, answeredAt) + answered.delete(requestId) + answered.set(requestId, { answeredAt, missingPasses: 0 }) + while (answered.size > MAX_ANSWERED_QUESTION_TOMBSTONES) answered.delete(answered.keys().next().value!) + answeredQuestionIdsByInstance.set(instanceId, answered) +} + +function hasAnsweredQuestion(instanceId: string, requestId: string, now = Date.now()): boolean { + const answered = answeredQuestionIdsByInstance.get(instanceId) + if (!answered) return false + pruneExpiredAnsweredQuestions(answered, now) + if (answered.size === 0) answeredQuestionIdsByInstance.delete(instanceId) + return answered.has(requestId) +} + +function pruneAnsweredQuestions(instanceId: string, remotePendingIds: Set, syncStartedAt: number): void { + const answered = answeredQuestionIdsByInstance.get(instanceId) + if (!answered) return + pruneExpiredAnsweredQuestions(answered, syncStartedAt) + for (const [requestId, state] of answered) { + if (remotePendingIds.has(requestId)) { + state.missingPasses = 0 + continue + } + if (syncStartedAt < state.answeredAt) continue + if (state.missingPasses > 0) answered.delete(requestId) + else state.missingPasses += 1 + } + if (answered.size === 0) answeredQuestionIdsByInstance.delete(instanceId) +} + +function pruneExpiredAnsweredQuestions(answered: Map, now: number): void { + for (const [requestId, state] of answered) { + if (now - state.answeredAt < ANSWERED_QUESTION_TOMBSTONE_TTL_MS) break + answered.delete(requestId) + } +} + +function clearSessionInterruptions(instanceId: string, sessionId: string): void { + for (const permission of [...getPermissionQueue(instanceId)]) { + if (getPermissionSessionId(permission) !== sessionId) continue + removePermissionFromQueue(instanceId, permission.id) + removePermissionV2(instanceId, permission.id) + } + for (const question of [...getQuestionQueue(instanceId)]) { + if (getQuestionSessionId(question) !== sessionId) continue + removeQuestionFromQueue(instanceId, question.id) + removeQuestionV2(instanceId, question.id) + } +} + function setActivePermissionIdForInstance(instanceId: string, permissionId: string): void { setActiveInterruptionForInstance(instanceId, { kind: "permission", id: permissionId }) } @@ -1655,6 +1964,7 @@ async function sendQuestionReply( if (source === "legacy") { const workspace = sessionId ? await getOpenCodeWorkspaceIdForSession(instanceId, sessionId) : null + if (!isInstanceRuntimeCurrent(instanceId, instance)) return await requestData( client.question.reply({ requestID: requestId, @@ -1674,7 +1984,10 @@ async function sendQuestionReply( ) } + if (!isInstanceRuntimeCurrent(instanceId, instance)) return + markQuestionAnswered(instanceId, requestId) removeQuestionFromQueue(instanceId, requestId) + removeQuestionV2(instanceId, requestId) } catch (error) { log.error("Failed to send question reply", error) throw error @@ -1693,6 +2006,7 @@ async function sendQuestionReject(instanceId: string, sessionId: string, request if (source === "legacy") { const workspace = sessionId ? await getOpenCodeWorkspaceIdForSession(instanceId, sessionId) : null + if (!isInstanceRuntimeCurrent(instanceId, instance)) return await requestData( client.question.reject({ requestID: requestId, @@ -1710,7 +2024,10 @@ async function sendQuestionReject(instanceId: string, sessionId: string, request ) } + if (!isInstanceRuntimeCurrent(instanceId, instance)) return + markQuestionAnswered(instanceId, requestId) removeQuestionFromQueue(instanceId, requestId) + removeQuestionV2(instanceId, requestId) } catch (error) { log.error("Failed to send question reject", error) throw error @@ -1735,6 +2052,7 @@ async function sendPermissionResponse( if (source === "legacy") { const workspace = sessionId ? await getOpenCodeWorkspaceIdForSession(instanceId, sessionId) : null + if (!isInstanceRuntimeCurrent(instanceId, instance)) return await requestData( client.permission.reply({ requestID: requestId, @@ -1756,6 +2074,7 @@ async function sendPermissionResponse( ) } + if (!isInstanceRuntimeCurrent(instanceId, instance)) return markPermissionReplied(instanceId, requestId) // Remove from both local queues after successful response; the SSE replied event // is still accepted, but the UI no longer depends on receiving it. @@ -1838,7 +2157,9 @@ export { setActiveInstanceId, addInstance, updateInstance, + isInstanceRuntimeCurrent, removeInstance, + clearSessionInterruptions, createInstance, cancelRestoreCreationRequest, disposeRestoreCreatedInstance, @@ -1880,6 +2201,8 @@ export { getQuestionQueueLength, getQuestionEnqueuedAtForInstance, addQuestionToQueue, + hasAnsweredQuestion, + markQuestionAnswered, removeQuestionFromQueue, clearQuestionQueue, sendQuestionReply, diff --git a/packages/ui/src/stores/message-prompt-display.test.ts b/packages/ui/src/stores/message-prompt-display.test.ts index 3cbaa93c3..bd1fb2a90 100644 --- a/packages/ui/src/stores/message-prompt-display.test.ts +++ b/packages/ui/src/stores/message-prompt-display.test.ts @@ -162,6 +162,14 @@ describe("message prompt display overrides", () => { clearPromptDisplayOverridesForSession(instanceId, "session-a") + assert.equal(getPromptDisplayOverride("other-instance", "session-a", "msg-1"), undefined) + assert.deepEqual(getPromptDisplayOverride("other-instance", "session-b", "msg-2"), metadata) + assert.deepEqual( + JSON.parse(storage.getItem("codenomad:prompt-display:v3") ?? "{}"), + { "session-b:msg-2": metadata }, + ) + + resetPromptDisplayOverrideStateForTests() assert.equal(getPromptDisplayOverride("other-instance", "session-a", "msg-1"), undefined) assert.deepEqual(getPromptDisplayOverride("other-instance", "session-b", "msg-2"), metadata) diff --git a/packages/ui/src/stores/message-v2/bridge.ts b/packages/ui/src/stores/message-v2/bridge.ts index 44c3872a1..7e9345360 100644 --- a/packages/ui/src/stores/message-v2/bridge.ts +++ b/packages/ui/src/stores/message-v2/bridge.ts @@ -14,7 +14,7 @@ interface SessionMetadata { parentId?: string | null } -function resolveSessionMetadata(session?: Session | null): SessionMetadata | undefined { +function resolveSessionMetadata(session?: Session | SessionMetadata | null): SessionMetadata | undefined { if (!session) return undefined return { id: session.id, @@ -81,17 +81,19 @@ export function upsertMessageInfoV2(instanceId: string, info: MessageInfo | null return } const store = messageStoreBus.getOrCreate(instanceId) - const timeInfo = (info.time ?? {}) as { created?: number; end?: number } + const timeInfo = (info.time ?? {}) as { created?: number; end?: number; completed?: number } const createdAt = typeof timeInfo.created === "number" ? timeInfo.created : Date.now() - const endAt = typeof timeInfo.end === "number" ? timeInfo.end : undefined + const endAt = typeof timeInfo.end === "number" ? timeInfo.end : timeInfo.completed + const status = options?.status ?? "complete" store.upsertMessage({ id: info.id, sessionId: info.sessionID, role: info.role === "user" ? "user" : "assistant", - status: options?.status ?? "complete", + status, createdAt, updatedAt: endAt ?? createdAt, + isEphemeral: info.role === "user" ? false : status === "sending" || status === "streaming", bumpRevision: Boolean(options?.bumpRevision), }) store.setMessageInfo(info.id, info) @@ -111,12 +113,12 @@ export function applyPartUpdateV2(instanceId: string, part: ClientPart | null | export function applyPartDeltaV2( instanceId: string, input: { messageId: string; partId: string; field: string; delta: string }, -): void { +): boolean { if (!input?.messageId || !input.partId || !input.field || typeof input.delta !== "string") { - return + return false } const store = messageStoreBus.getOrCreate(instanceId) - store.applyPartDelta({ + return store.applyPartDelta({ messageId: input.messageId, partId: input.partId, field: input.field, diff --git a/packages/ui/src/stores/message-v2/bus.ts b/packages/ui/src/stores/message-v2/bus.ts index 86900cd9f..f20307fa2 100644 --- a/packages/ui/src/stores/message-v2/bus.ts +++ b/packages/ui/src/stores/message-v2/bus.ts @@ -1,6 +1,6 @@ import { createInstanceMessageStore } from "./instance-store" import type { InstanceMessageStore } from "./instance-store" -import { clearCacheForInstance } from "../../lib/global-cache" +import { clearCacheForInstance, clearCacheForSession } from "../../lib/global-cache" import { getLogger } from "../../lib/logger" import type { ScrollSnapshot } from "./types" @@ -16,6 +16,7 @@ class MessageStoreBus { private stores = new Map() private teardownHandlers = new Set<(instanceId: string) => void>() private sessionClearHandlers = new Set<(instanceId: string, sessionId: string) => void>() + private sessionChangeHandlers = new Set<(instanceId: string, sessionId: string) => void>() private scrollSnapshotHandlers = new Set< (instanceId: string, sessionId: string, scope: string, snapshot: ScrollSnapshot) => void >() @@ -30,6 +31,7 @@ class MessageStoreBus { store ?? createInstanceMessageStore(instanceId, { onSessionCleared: (id, sessionId) => this.notifySessionCleared(id, sessionId), + onSessionChanged: (id, sessionId) => this.notifySessionChanged(id, sessionId), onScrollSnapshotChanged: (id, sessionId, scope, snapshot) => this.notifyScrollSnapshotChanged(id, sessionId, scope, snapshot), }) @@ -52,6 +54,7 @@ class MessageStoreBus { } private notifySessionCleared(instanceId: string, sessionId: string) { + clearCacheForSession(instanceId, sessionId) for (const handler of this.sessionClearHandlers) { try { handler(instanceId, sessionId) @@ -61,6 +64,21 @@ class MessageStoreBus { } } + onSessionChanged(handler: (instanceId: string, sessionId: string) => void): () => void { + this.sessionChangeHandlers.add(handler) + return () => this.sessionChangeHandlers.delete(handler) + } + + private notifySessionChanged(instanceId: string, sessionId: string) { + for (const handler of this.sessionChangeHandlers) { + try { + handler(instanceId, sessionId) + } catch (error) { + log.error("Failed to run session change handler", error) + } + } + } + onScrollSnapshotChanged( handler: (instanceId: string, sessionId: string, scope: string, snapshot: ScrollSnapshot) => void, ): () => void { @@ -111,6 +129,10 @@ class MessageStoreBus { return this.registerInstance(instanceId) } + entries(): IterableIterator<[string, InstanceMessageStore]> { + return this.stores.entries() + } + clearInstanceScrollSnapshots(instanceId: string): void { this.stores.get(instanceId)?.clearScrollSnapshots() this.scrollSnapshotSeeds.delete(instanceId) diff --git a/packages/ui/src/stores/message-v2/instance-store.test.ts b/packages/ui/src/stores/message-v2/instance-store.test.ts index 975b4b9f6..1f5549b81 100644 --- a/packages/ui/src/stores/message-v2/instance-store.test.ts +++ b/packages/ui/src/stores/message-v2/instance-store.test.ts @@ -1,6 +1,7 @@ import assert from "node:assert/strict" import { describe, it } from "node:test" +import { estimateRetainedBytes } from "../../lib/session-memory-budget.ts" import { createInstanceMessageStore } from "./instance-store.ts" describe("message-v2 permission state", () => { @@ -41,4 +42,96 @@ describe("message-v2 permission state", () => { assert.equal(store.getPermissionState(undefined, "permission-2")?.active, true) }) + it("protects legacy pending permissions that use sessionId", () => { + const store = createInstanceMessageStore("instance-1") + store.upsertPermission({ permission: { id: "legacy", sessionId: "session-1", permission: "edit" }, enqueuedAt: 1 }) + assert.equal(store.hasSessionActiveWork("session-1"), true) + }) + + it("charges large canonical interruption payloads once to their session", async () => { + const store = createInstanceMessageStore("instance-1") + store.addOrUpdateSession({ id: "session-1" }) + store.addOrUpdateSession({ id: "session-2" }) + const baseline = await store.getSessionApproximateByteSizeIncrementally("session-1") + const permission = { + id: "permission-large", + sessionID: "session-1", + action: "edit", + resources: ["p".repeat(256 * 1024)], + } + const question = { + id: "question-large", + sessionID: "session-1", + questions: [{ header: "Confirm", question: "q".repeat(256 * 1024), options: [] }], + } + + store.upsertPermission({ permission, messageId: "message-1", partId: "part-1", enqueuedAt: 1 }) + store.upsertQuestion({ request: question as any, messageId: "message-1", partId: "part-1", enqueuedAt: 2 }) + const expected = baseline + 3 * ( + estimateRetainedBytes(store.state.permissions.queue[0]?.permission) + + estimateRetainedBytes(store.state.questions.queue[0]?.request) + ) + assert.equal(await store.getSessionApproximateByteSizeIncrementally("session-1"), expected) + assert.ok(await store.getSessionApproximateByteSizeIncrementally("session-2") < expected / 2) + + store.upsertPermission({ permission, messageId: "message-1", partId: "part-1", enqueuedAt: 1 }) + store.upsertQuestion({ request: question as any, messageId: "message-1", partId: "part-1", enqueuedAt: 2 }) + assert.equal(store.state.permissions.queue.length, 1) + assert.equal(store.state.questions.queue.length, 1) + assert.equal(await store.getSessionApproximateByteSizeIncrementally("session-1"), baseline + 3 * ( + estimateRetainedBytes(store.state.permissions.queue[0]?.permission) + + estimateRetainedBytes(store.state.questions.queue[0]?.request) + )) + }) + +}) + +describe("message-v2 authoritative hydration", () => { + const message = (id: string) => ({ + id, + sessionId: "session-1", + role: "assistant" as const, + status: "complete" as const, + parts: [{ id: `part-${id}`, type: "text", text: id, messageID: id, sessionID: "session-1" }] as any, + }) + const info = (id: string) => ({ id, sessionID: "session-1", role: "assistant", time: { created: 1 } }) as any + + it("replaces stale messages and accepts an authoritative empty session", () => { + const store = createInstanceMessageStore("instance-1") + store.hydrateMessages("session-1", [message("message-1"), message("message-2")], [info("message-1"), info("message-2")]) + + store.hydrateMessages("session-1", [message("message-2")], [info("message-2")]) + assert.equal(store.getMessage("message-1"), undefined) + assert.equal(store.getMessageInfo("message-1"), undefined) + assert.deepEqual(store.getSessionMessageIds("session-1"), ["message-2"]) + + store.hydrateMessages("session-1", [], []) + assert.equal(store.getMessage("message-2"), undefined) + assert.deepEqual(store.getSessionMessageIds("session-1"), []) + }) + + it("removes stale parts when authoritative hydration returns an empty part list", () => { + const store = createInstanceMessageStore("instance-1") + store.hydrateMessages("session-1", [message("message-1")], [info("message-1")]) + + store.hydrateMessages("session-1", [{ ...message("message-1"), parts: [] }], [info("message-1")]) + assert.deepEqual(store.getMessage("message-1")?.partIds, []) + assert.deepEqual(store.getMessage("message-1")?.parts, {}) + }) + + it("releases a directly removed message and its info version", () => { + const store = createInstanceMessageStore("instance-1") + store.hydrateMessages("session-1", [message("message-1")], [info("message-1")]) + store.removeMessage("message-1") + assert.equal("message-1" in store.state.messages, false) + assert.equal("message-1" in store.state.messageInfoVersion, false) + }) + + it("bumps authority when a revert anchor is not resident", () => { + const store = createInstanceMessageStore("instance-1") + store.hydrateMessages("session-1", [message("message-1")], [info("message-1")]) + const revision = store.getSessionRevision("session-1") + store.setSessionRevert("session-1", { messageID: "evicted-anchor" }) + assert.ok(store.getSessionRevision("session-1") > revision) + }) }) diff --git a/packages/ui/src/stores/message-v2/instance-store.ts b/packages/ui/src/stores/message-v2/instance-store.ts index 437bea3d4..40e563c6a 100644 --- a/packages/ui/src/stores/message-v2/instance-store.ts +++ b/packages/ui/src/stores/message-v2/instance-store.ts @@ -11,9 +11,13 @@ import { setPromptDisplayOverride, } from "../message-prompt-display" import type { ClientPart, MessageInfo } from "../../types/message" -import { mergePermissionRequest } from "../../types/permission" +import { getPermissionSessionId, mergePermissionRequest } from "../../types/permission" +import { getQuestionSessionId } from "../../types/question" import { clearRecordDisplayCacheForMessages } from "./record-display-cache" +import { estimateRetainedValuesIncrementally } from "../../lib/session-memory-budget" import { mergePendingRequestEntry, shouldSkipPendingRequestUpsert } from "./pending-request-dedupe" + +const DERIVED_RENDER_MEMORY_MULTIPLIER = 3 import type { InstanceMessageState, LatestTodoSnapshot, @@ -35,6 +39,7 @@ const storeLog = getLogger("session") interface MessageStoreHooks { onSessionCleared?: (instanceId: string, sessionId: string) => void + onSessionChanged?: (instanceId: string, sessionId: string) => void onScrollSnapshotChanged?: (instanceId: string, sessionId: string, scope: string, snapshot: ScrollSnapshot) => void } @@ -225,7 +230,7 @@ export interface InstanceMessageStore { delta: string bumpRevision?: boolean bumpSessionRevision: boolean - }) => void + }) => boolean removeMessage: (messageId: string, fallbackSessionId?: string) => void removeMessagePart: (messageId: string, partId: string, fallbackSessionId?: string) => void bufferPendingPart: (entry: PendingPartEntry) => void @@ -248,13 +253,18 @@ export interface InstanceMessageStore { getScrollSnapshot: (sessionId: string, scope: string) => ScrollSnapshot | undefined getSessionRevision: (sessionId: string) => number getSessionMessageIds: (sessionId: string) => string[] + getResidentSessionIds: () => string[] + getSessionApproximateByteSizeIncrementally: (sessionId: string, signal?: AbortSignal) => Promise + interruptSessionActiveMessages: (sessionId: string) => void + hasSessionActiveWork: (sessionId: string) => boolean + hasSessionPendingInput: (sessionId: string) => boolean getLastAssistantMessageId: (sessionId: string) => string | undefined // Index of the most recent message in the session that contains a compaction part. // Returns -1 if there has been no compaction. getLastCompactionMessageIndex: (sessionId: string) => number getMessage: (messageId: string) => MessageRecord | undefined getLatestTodoSnapshot: (sessionId: string) => LatestTodoSnapshot | undefined - clearSession: (sessionId: string, options?: { preserveScroll?: boolean; notify?: boolean }) => void + clearSession: (sessionId: string, options?: { preserveScroll?: boolean; preservePromptDisplay?: boolean; notify?: boolean }) => void clearScrollSnapshots: () => void clearInstance: () => void } @@ -265,6 +275,27 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt const TODO_TOOL_NAME = "todowrite" const messageInfoCache = new Map() + const activeMessageCounts = new Map() + + function isActiveMessage(record: MessageRecord | undefined): boolean { + return record?.status === "sending" || record?.status === "streaming" + } + + function updateActiveMessageCount(previous: MessageRecord | undefined, next: MessageRecord | undefined) { + if (isActiveMessage(previous)) { + const count = (activeMessageCounts.get(previous!.sessionId) ?? 1) - 1 + if (count > 0) activeMessageCounts.set(previous!.sessionId, count) + else activeMessageCounts.delete(previous!.sessionId) + } + if (isActiveMessage(next)) activeMessageCounts.set(next!.sessionId, (activeMessageCounts.get(next!.sessionId) ?? 0) + 1) + } + + function setActiveMessageCount(sessionId: string, records: Iterable) { + let count = 0 + for (const record of records) if (isActiveMessage(record)) count += 1 + if (count > 0) activeMessageCounts.set(sessionId, count) + else activeMessageCounts.delete(sessionId) + } function findLastAssistantMessageId(messageIds: readonly string[]): string | undefined { for (let index = messageIds.length - 1; index >= 0; index -= 1) { @@ -347,12 +378,77 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt function bumpSessionRevision(sessionId: string) { if (!sessionId) return setState("sessionRevisions", sessionId, (value = 0) => value + 1) + hooks?.onSessionChanged?.(instanceId, sessionId) } function getSessionRevisionValue(sessionId: string) { return state.sessionRevisions[sessionId] ?? 0 } + function getResidentSessionIds() { + const sessionIds = new Set(Object.values(state.sessions) + .filter((session) => session.messageIds.length > 0) + .map((session) => session.id)) + for (const entry of state.permissions.queue) { + const sessionId = getPermissionSessionId(entry.permission) + if (sessionId) sessionIds.add(sessionId) + } + for (const entry of state.questions.queue) { + const sessionId = getQuestionSessionId(entry.request) + if (sessionId) sessionIds.add(sessionId) + } + return [...sessionIds] + } + + async function getSessionApproximateByteSizeIncrementally(sessionId: string, signal?: AbortSignal) { + function* values(): Iterable { + const session = state.sessions[sessionId] + yield session + yield state.usage[sessionId] + for (const messageId of session?.messageIds ?? []) { + yield state.messages[messageId] + yield messageInfoCache.get(messageId) + } + for (const entry of state.permissions.queue) { + if (getPermissionSessionId(entry.permission) === sessionId) yield entry.permission + } + for (const entry of state.questions.queue) { + if (getQuestionSessionId(entry.request) === sessionId) yield entry.request + } + } + return (await estimateRetainedValuesIncrementally(values(), { signal })) * DERIVED_RENDER_MEMORY_MULTIPLIER + } + + function interruptSessionActiveMessages(sessionId: string) { + const messageIds = state.sessions[sessionId]?.messageIds ?? [] + let changed = false + setState("messages", produce((draft) => { + for (const messageId of messageIds) { + const message = draft[messageId] + if (!message || (message.status !== "sending" && message.status !== "streaming")) continue + message.status = "error" + message.isEphemeral = false + message.updatedAt = Date.now() + message.revision += 1 + changed = true + } + })) + if (changed) { + activeMessageCounts.delete(sessionId) + bumpSessionRevision(sessionId) + } + } + + function hasSessionActiveWork(sessionId: string) { + if (hasSessionPendingInput(sessionId)) return true + return (activeMessageCounts.get(sessionId) ?? 0) > 0 + } + + function hasSessionPendingInput(sessionId: string) { + return state.permissions.queue.some((entry) => getPermissionSessionId(entry.permission) === sessionId) || + state.questions.queue.some((entry) => getQuestionSessionId(entry.request) === sessionId) + } + function getLastAssistantMessageIdValue(sessionId: string) { return state.lastAssistantMessageIds[sessionId] } @@ -432,11 +528,15 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt } function hydrateMessages(sessionId: string, inputs: MessageUpsertInput[], infos?: Iterable) { - if (!Array.isArray(inputs) || inputs.length === 0) return + if (!Array.isArray(inputs)) return ensureSessionEntry(sessionId) const incomingIds = inputs.map((item) => item.id) + const incomingIdSet = new Set(incomingIds) + const staleIds = Object.values(state.messages) + .filter((record) => record.sessionId === sessionId && !incomingIdSet.has(record.id)) + .map((record) => record.id) const normalizedRecords: Record = {} const now = Date.now() @@ -471,6 +571,23 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt const nextPermissionsByMessage: Record> = { ...state.permissions.byMessage, } + const nextQuestionsByMessage: Record> = { + ...state.questions.byMessage, + } + setActiveMessageCount(sessionId, Object.values(normalizedRecords)) + + if (staleIds.length > 0) { + clearRecordDisplayCacheForMessages(instanceId, staleIds) + for (const id of staleIds) { + delete nextMessages[id] + delete nextMessageInfoVersion[id] + delete nextPendingParts[id] + delete nextPermissionsByMessage[id] + delete nextQuestionsByMessage[id] + messageInfoCache.delete(id) + clearPromptDisplayOverride(instanceId, sessionId, id) + } + } Object.entries(normalizedRecords).forEach(([id, record]) => { nextMessages[id] = record @@ -486,10 +603,11 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt } batch(() => { - setState("messages", () => nextMessages) - setState("messageInfoVersion", () => nextMessageInfoVersion) - setState("pendingParts", () => nextPendingParts) - setState("permissions", "byMessage", () => nextPermissionsByMessage) + setState("messages", reconcile(nextMessages)) + setState("messageInfoVersion", reconcile(nextMessageInfoVersion)) + setState("pendingParts", reconcile(nextPendingParts)) + setState("permissions", "byMessage", reconcile(nextPermissionsByMessage)) + setState("questions", "byMessage", reconcile(nextQuestionsByMessage)) if (usageState) { setState("usage", sessionId, usageState) @@ -502,6 +620,7 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt })) recomputeLastAssistantMessageId(sessionId, incomingIds) + clearLatestTodoSnapshot(sessionId) Object.values(normalizedRecords).forEach((record) => { maybeUpdateLatestTodoFromRecord(record) }) @@ -521,7 +640,7 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt } function normalizeParts(messageId: string, parts: ClientPart[] | undefined) { - if (!parts || parts.length === 0) { + if (parts == null) { return null } const map: MessageRecord["parts"] = {} @@ -547,8 +666,10 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt const now = Date.now() let nextRecord: MessageRecord | undefined + let previousRecord: MessageRecord | undefined setState("messages", input.id, (previous) => { + previousRecord = previous ? { ...previous } : undefined const revision = previous ? previous.revision + (shouldBump ? 1 : 0) : 0 const clientPromptDisplayMetadata = resolveClientPromptDisplayText(instanceId, input, previous) const record: MessageRecord = { @@ -570,6 +691,7 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt }) if (nextRecord) { + updateActiveMessageCount(previousRecord, nextRecord) maybeUpdateLatestTodoFromRecord(nextRecord) } @@ -695,13 +817,13 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt bumpSessionRevision?: boolean }) { if (!input?.messageId || !input.partId || !input.field || typeof input.delta !== "string") { - return + return false } const message = state.messages[input.messageId] if (!message) { // Best-effort: drop deltas for unknown messages. - return + return false } let applied = false @@ -730,6 +852,7 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt if (applied && (input.bumpSessionRevision ?? true)) { bumpSessionRevision(message.sessionId) } + return applied } function removeMessage(messageId: string, fallbackSessionId?: string) { @@ -760,19 +883,8 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt setState("sessions", sessionId, "messageIds", (ids = []) => ids.filter((id) => id !== messageId)) }) - setState("messages", (prev) => { - if (!prev[messageId]) return prev - const next = { ...prev } - delete next[messageId] - return next - }) - - setState("messageInfoVersion", (prev) => { - if (!(messageId in prev)) return prev - const next = { ...prev } - delete next[messageId] - return next - }) + setState("messages", messageId, undefined as any) + setState("messageInfoVersion", messageId, undefined as any) messageInfoCache.delete(messageId) @@ -796,6 +908,7 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt clearLatestTodoSnapshot(sessionId) } recomputeLastAssistantMessageId(sessionId) + setActiveMessageCount(sessionId, (state.sessions[sessionId]?.messageIds ?? []).flatMap((id) => state.messages[id] ? [state.messages[id]] : [])) bumpSessionRevision(sessionId) }) }) @@ -1010,9 +1123,12 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt draft.active = draft.queue[0] ?? null }), ) + const sessionId = getPermissionSessionId(entry.permission) + if (sessionId) bumpSessionRevision(sessionId) } function removePermission(permissionId: string) { + const sessionId = getPermissionSessionId(state.permissions.queue.find((item) => item.permission.id === permissionId)?.permission) setState( "permissions", produce((draft) => { @@ -1033,6 +1149,7 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt }) }), ) + if (sessionId) bumpSessionRevision(sessionId) } function getPermissionState(messageId?: string, partId?: string) { @@ -1098,9 +1215,12 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt } }), ) + const sessionId = getQuestionSessionId(entry.request) + if (sessionId) bumpSessionRevision(sessionId) } function removeQuestion(requestId: string) { + const sessionId = getQuestionSessionId(state.questions.queue.find((item) => item.request.id === requestId)?.request) setState( "questions", produce((draft) => { @@ -1121,6 +1241,7 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt }) }), ) + if (sessionId) bumpSessionRevision(sessionId) } function getQuestionState(messageId?: string, partId?: string) { @@ -1132,14 +1253,14 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt return { entry, active } } - function pruneMessagesAfterRevert(sessionId: string, revertMessageId: string) { + function pruneMessagesAfterRevert(sessionId: string, revertMessageId: string): boolean { const session = state.sessions[sessionId] - if (!session) return + if (!session) return false const stopIndex = session.messageIds.indexOf(revertMessageId) - if (stopIndex === -1) return + if (stopIndex === -1) return false const removedIds = session.messageIds.slice(stopIndex) const keptIds = session.messageIds.slice(0, stopIndex) - if (removedIds.length === 0) return + if (removedIds.length === 0) return false removedIds.forEach((messageId) => clearPromptDisplayOverride(instanceId, sessionId, messageId)) @@ -1188,16 +1309,17 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt }) recomputeLastAssistantMessageId(sessionId, keptIds) + setActiveMessageCount(sessionId, keptIds.flatMap((id) => state.messages[id] ? [state.messages[id]] : [])) bumpSessionRevision(sessionId) + return true } function setSessionRevert(sessionId: string, revert?: SessionRecord["revert"] | null) { if (!sessionId) return ensureSessionEntry(sessionId) - if (revert?.messageID) { - pruneMessagesAfterRevert(sessionId, revert.messageID) - } + const pruned = revert?.messageID ? pruneMessagesAfterRevert(sessionId, revert.messageID) : false setState("sessions", sessionId, "revert", revert ?? null) + if (!pruned) bumpSessionRevision(sessionId) } function getSessionRevert(sessionId: string) { @@ -1225,56 +1347,37 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt return state.scrollState[key] } - function clearSession(sessionId: string, options?: { preserveScroll?: boolean; notify?: boolean }) { + function clearSession(sessionId: string, options?: { preserveScroll?: boolean; preservePromptDisplay?: boolean; notify?: boolean }) { if (!sessionId) return - clearPromptDisplayOverridesForSession(instanceId, sessionId) + if (!options?.preservePromptDisplay) clearPromptDisplayOverridesForSession(instanceId, sessionId) - const messageIds = Object.values(state.messages) - .filter((record) => record.sessionId === sessionId) - .map((record) => record.id) + const messageIds = [...(state.sessions[sessionId]?.messageIds ?? [])] storeLog.info("Clearing session data", { instanceId, sessionId, messageCount: messageIds.length }) clearRecordDisplayCacheForMessages(instanceId, messageIds) batch(() => { - setState("messages", (prev) => { - const next = { ...prev } - messageIds.forEach((id) => delete next[id]) - return next - }) + const permissionQueue = state.permissions.queue.filter((entry) => getPermissionSessionId(entry.permission) !== sessionId) + const questionQueue = state.questions.queue.filter((entry) => getQuestionSessionId(entry.request) !== sessionId) + setState("permissions", "queue", permissionQueue) + setState("permissions", "active", (active) => + active && getPermissionSessionId(active.permission) === sessionId ? permissionQueue[0] ?? null : active) + setState("questions", "queue", questionQueue) + setState("questions", "active", (active) => + active && getQuestionSessionId(active.request) === sessionId ? questionQueue[0] ?? null : active) - setState("messageInfoVersion", (prev) => { - const next = { ...prev } - messageIds.forEach((id) => delete next[id]) - return next - }) + setState("messages", produce((next) => { messageIds.forEach((id) => delete next[id]) })) + + setState("messageInfoVersion", produce((next) => { messageIds.forEach((id) => delete next[id]) })) messageIds.forEach((id) => messageInfoCache.delete(id)) - setState("pendingParts", (prev) => { - const next = { ...prev } - messageIds.forEach((id) => { - if (next[id]) delete next[id] - }) - return next - }) + setState("pendingParts", produce((next) => { messageIds.forEach((id) => delete next[id]) })) - setState("permissions", "byMessage", (prev) => { - const next = { ...prev } - messageIds.forEach((id) => { - if (next[id]) delete next[id] - }) - return next - }) + setState("permissions", "byMessage", produce((next) => { messageIds.forEach((id) => delete next[id]) })) - setState("questions", "byMessage", (prev) => { - const next = { ...prev } - messageIds.forEach((id) => { - if (next[id]) delete next[id] - }) - return next - }) + setState("questions", "byMessage", produce((next) => { messageIds.forEach((id) => delete next[id]) })) setState("usage", (prev) => { const next = { ...prev } @@ -1322,14 +1425,16 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt }) clearLatestTodoSnapshot(sessionId) + activeMessageCounts.delete(sessionId) if (options?.notify !== false) hooks?.onSessionCleared?.(instanceId, sessionId) } - function clearInstance() { - clearPromptDisplayOverridesForInstance(instanceId, Object.keys(state.sessions)) - messageInfoCache.clear() + function clearInstance() { + clearPromptDisplayOverridesForInstance(instanceId, Object.keys(state.sessions)) + messageInfoCache.clear() + activeMessageCounts.clear() setState(reconcile(createInitialState(instanceId))) } @@ -1370,6 +1475,11 @@ export function createInstanceMessageStore(instanceId: string, hooks?: MessageSt getScrollSnapshot, getSessionRevision: getSessionRevisionValue, getSessionMessageIds: (sessionId: string) => state.sessions[sessionId]?.messageIds ?? [], + getResidentSessionIds, + getSessionApproximateByteSizeIncrementally, + interruptSessionActiveMessages, + hasSessionActiveWork, + hasSessionPendingInput, getLastAssistantMessageId: getLastAssistantMessageIdValue, getLastCompactionMessageIndex, getMessage: (messageId: string) => state.messages[messageId], diff --git a/packages/ui/src/stores/opencode-workspaces.test.ts b/packages/ui/src/stores/opencode-workspaces.test.ts index 789b18702..ee25fb54a 100644 --- a/packages/ui/src/stores/opencode-workspaces.test.ts +++ b/packages/ui/src/stores/opencode-workspaces.test.ts @@ -1,7 +1,32 @@ import assert from "node:assert/strict" -import { describe, it } from "node:test" +import { describe, it, mock } from "node:test" +import { sdkManager } from "../lib/sdk-manager.ts" +import { addInstance, removeInstance } from "./instances.ts" import { mapOpenCodeWorkspacesToWorktreeSlugs } from "./opencode-workspace-matching.ts" +import { + clearOpenCodeWorkspaceCache, + getOpenCodeWorkspaceIdForWorktree, + reloadOpenCodeWorkspaces, + syncOpenCodeWorkspaces, +} from "./opencode-workspaces.ts" + +function deferred() { + let resolve!: (value: T) => void + const promise = new Promise((done) => { resolve = done }) + return { promise, resolve } +} + +function setupWorkspaceInstance(instanceId: string, workspace: Record) { + const client = { experimental: { workspace } } as any + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, client) + addInstance({ id: instanceId, folder: "/work", port: 0, pid: 0, proxyPath: "", status: "ready", client }) + return () => { + clearOpenCodeWorkspaceCache(instanceId) + removeInstance(instanceId, { authoritative: false }) + sdkManager.destroyClientsForInstance(instanceId) + } +} describe("mapOpenCodeWorkspacesToWorktreeSlugs", () => { it("matches POSIX worktree directories case-sensitively", () => { @@ -58,3 +83,101 @@ describe("mapOpenCodeWorkspacesToWorktreeSlugs", () => { assert.equal(result.size, 0) }) }) + +describe("OpenCode workspace sync", () => { + it("deduplicates callers onto one shared SDK sync", async () => { + const instanceId = "shared-workspace-sync" + const gate = deferred() + let syncCalls = 0 + let listCalls = 0 + const cleanup = setupWorkspaceInstance(instanceId, { + syncList: async () => { syncCalls += 1; await gate.promise; return { data: [] } }, + list: async () => { listCalls += 1; return { data: [] } }, + }) + + try { + const requests = [ + syncOpenCodeWorkspaces(instanceId), + syncOpenCodeWorkspaces(instanceId), + syncOpenCodeWorkspaces(instanceId), + ] + while (syncCalls === 0) await new Promise((resolve) => setImmediate(resolve)) + assert.equal(syncCalls, 1) + gate.resolve() + await Promise.all(requests) + assert.equal(listCalls, 1) + } finally { + cleanup() + } + }) + + it("queues reconnect replacement without releasing or aborting shared callers", async () => { + const instanceId = "replacement-workspace-sync" + const first = deferred() + const second = deferred() + const signals: AbortSignal[] = [] + let syncCalls = 0 + const cleanup = setupWorkspaceInstance(instanceId, { + syncList: async (_parameters: unknown, options?: { signal?: AbortSignal }) => { + syncCalls += 1 + if (options?.signal) signals.push(options.signal) + await (syncCalls === 1 ? first.promise : second.promise) + return { data: [] } + }, + list: async () => ({ data: [] }), + }) + + try { + let sharedSettled = false + const shared = syncOpenCodeWorkspaces(instanceId).then(() => { sharedSettled = true }) + while (syncCalls === 0) await new Promise((resolve) => setImmediate(resolve)) + const reconnect = reloadOpenCodeWorkspaces(instanceId) + assert.equal(signals[0]?.aborted, false) + + first.resolve() + while (syncCalls < 2) await new Promise((resolve) => setImmediate(resolve)) + assert.equal(sharedSettled, false) + assert.equal(signals[0]?.aborted, false) + + second.resolve() + await Promise.all([shared, reconnect]) + assert.equal(sharedSettled, true) + } finally { + cleanup() + } + }) + + it("aborts timed-out SDK I/O and negative-caches the shared result", async () => { + const instanceId = "timed-out-workspace-sync" + let syncCalls = 0 + let sdkSignal: AbortSignal | undefined + const cleanup = setupWorkspaceInstance(instanceId, { + syncList: (_parameters: unknown, options?: { signal?: AbortSignal }) => { + syncCalls += 1 + sdkSignal = options?.signal + return new Promise((_resolve, reject) => options?.signal?.addEventListener( + "abort", + () => reject(options.signal?.reason), + { once: true }, + )) + }, + list: async () => ({ data: [] }), + }) + mock.timers.enable({ apis: ["setTimeout"] }) + + try { + const request = syncOpenCodeWorkspaces(instanceId) + while (!sdkSignal) await new Promise((resolve) => setImmediate(resolve)) + mock.timers.tick(5_000) + await request + assert.equal(sdkSignal.aborted, true) + + assert.equal(await getOpenCodeWorkspaceIdForWorktree(instanceId, "missing-one"), null) + assert.equal(await getOpenCodeWorkspaceIdForWorktree(instanceId, "missing-two"), null) + assert.equal(syncCalls, 1) + } finally { + mock.timers.reset() + cleanup() + } + }) +}) diff --git a/packages/ui/src/stores/opencode-workspaces.ts b/packages/ui/src/stores/opencode-workspaces.ts index b6829583b..9bcf8b196 100644 --- a/packages/ui/src/stores/opencode-workspaces.ts +++ b/packages/ui/src/stores/opencode-workspaces.ts @@ -5,20 +5,29 @@ import { mapOpenCodeWorkspacesToWorktreeSlugs } from "./opencode-workspace-match const WORKSPACE_SYNC_TIMEOUT_MS = 5_000 -function withWorkspaceSyncTimeout(operation: Promise): Promise { - return new Promise((resolve, reject) => { - const timer = setTimeout(() => reject(new Error("OpenCode workspace sync timed out")), WORKSPACE_SYNC_TIMEOUT_MS) - operation.then( - (value) => { - clearTimeout(timer) - resolve(value) - }, - (error) => { - clearTimeout(timer) - reject(error) - }, - ) +async function withWorkspaceSyncTimeout( + controller: AbortController, + operation: (signal: AbortSignal) => Promise, +): Promise { + let timer: ReturnType | undefined + let rejectAbort!: (reason?: unknown) => void + const aborted = new Promise((_resolve, reject) => { rejectAbort = reject }) + const onAbort = () => rejectAbort(controller.signal.reason) + if (controller.signal.aborted) onAbort() + else controller.signal.addEventListener("abort", onAbort, { once: true }) + const timeout = new Promise((_resolve, reject) => { + timer = setTimeout(() => { + const error = new Error("OpenCode workspace sync timed out") + controller.abort(error) + reject(error) + }, WORKSPACE_SYNC_TIMEOUT_MS) }) + try { + return await Promise.race([operation(controller.signal), timeout, aborted]) + } finally { + if (timer) clearTimeout(timer) + controller.signal.removeEventListener("abort", onAbort) + } } const log = getLogger("api") @@ -33,7 +42,12 @@ type OpenCodeWorkspace = { } const workspaceIdByWorktreeSlug = new Map>() -const workspaceSyncs = new Map>() +type WorkspaceSync = { + promise: Promise + controller: AbortController + controllers: Set +} +const workspaceSyncs = new Map() async function getInstance(instanceId: string) { const { instances } = await import("./instances") @@ -49,55 +63,93 @@ function getCachedOpenCodeWorkspaceIdForSession(instanceId: string, sessionId: s return getCachedOpenCodeWorkspaceIdForWorktree(instanceId, getWorktreeSlugForSession(instanceId, sessionId)) } -async function syncOpenCodeWorkspaces(instanceId: string): Promise { - if (!instanceId) return - const existing = workspaceSyncs.get(instanceId) - if (existing) return existing - - const task = (async () => { +function startWorkspaceSync(instanceId: string, previous?: WorkspaceSync, reset = false): WorkspaceSync { + let runtimeToken: symbol | undefined + const controller = new AbortController() + const isCurrent = async () => (await getInstance(instanceId))?.runtimeToken === runtimeToken + let task!: Promise + task = (async () => { + await previous?.promise + controller.signal.throwIfAborted() + if (reset) workspaceIdByWorktreeSlug.delete(instanceId) const instance = await getInstance(instanceId) if (!instance?.client || !instance.folder) return + runtimeToken = instance.runtimeToken const rootClient = getRootClient(instanceId) as any const workspaceApi = rootClient.experimental?.workspace if (!workspaceApi?.syncList || !workspaceApi?.list) { log.warn("OpenCode experimental workspace API unavailable", { instanceId }) - workspaceIdByWorktreeSlug.set(instanceId, new Map()) + if (await isCurrent()) workspaceIdByWorktreeSlug.set(instanceId, new Map()) return } - await withWorkspaceSyncTimeout(workspaceApi.syncList({ directory: instance.folder })) - const result = await withWorkspaceSyncTimeout(workspaceApi.list({ directory: instance.folder })) + const result = await withWorkspaceSyncTimeout(controller, async (signal) => { + await workspaceApi.syncList({ directory: instance.folder }, { signal }) + signal.throwIfAborted() + if (!await isCurrent()) return null + return workspaceApi.list({ directory: instance.folder }, { signal }) + }) + if (!await isCurrent()) return const workspaces = Array.isArray(result?.data) ? (result.data as OpenCodeWorkspace[]) : [] const next = mapOpenCodeWorkspacesToWorktreeSlugs(getWorktrees(instanceId), workspaces) workspaceIdByWorktreeSlug.set(instanceId, next) })() - .catch((error) => { + .catch(async (error) => { + if (controller.signal.aborted) { + if (controller.signal.reason instanceof Error && controller.signal.reason.message === "OpenCode workspace sync timed out" && await isCurrent()) { + log.warn("Failed to sync OpenCode workspaces", { instanceId, error }) + workspaceIdByWorktreeSlug.set(instanceId, new Map()) + } + return + } + if (!await isCurrent()) return log.warn("Failed to sync OpenCode workspaces", { instanceId, error }) if (!workspaceIdByWorktreeSlug.has(instanceId)) { workspaceIdByWorktreeSlug.set(instanceId, new Map()) } }) .finally(() => { - if (workspaceSyncs.get(instanceId) === task) { + if (workspaceSyncs.get(instanceId)?.promise === task) { workspaceSyncs.delete(instanceId) } }) - workspaceSyncs.set(instanceId, task) - return task + const sync = { + promise: task, + controller, + controllers: new Set([...(previous?.controllers ?? []), controller]), + } + workspaceSyncs.set(instanceId, sync) + return sync +} + +async function awaitWorkspaceSync(instanceId: string, initial: WorkspaceSync): Promise { + let current = initial + while (true) { + await current.promise + const replacement = workspaceSyncs.get(instanceId) + if (!replacement || replacement === current) return + current = replacement + } +} + +async function syncOpenCodeWorkspaces(instanceId: string): Promise { + if (!instanceId) return + const sync = workspaceSyncs.get(instanceId) ?? startWorkspaceSync(instanceId) + await awaitWorkspaceSync(instanceId, sync) } async function reloadOpenCodeWorkspaces(instanceId: string): Promise { - await workspaceSyncs.get(instanceId) - await syncOpenCodeWorkspaces(instanceId) + if (!instanceId) return + const replacement = startWorkspaceSync(instanceId, workspaceSyncs.get(instanceId), true) + await awaitWorkspaceSync(instanceId, replacement) } async function getOpenCodeWorkspaceIdForWorktree(instanceId: string, slug: string): Promise { if (!slug || slug === "root") return null - const cached = getCachedOpenCodeWorkspaceIdForWorktree(instanceId, slug) - if (cached) return cached + if (workspaceIdByWorktreeSlug.has(instanceId)) return getCachedOpenCodeWorkspaceIdForWorktree(instanceId, slug) await syncOpenCodeWorkspaces(instanceId) return getCachedOpenCodeWorkspaceIdForWorktree(instanceId, slug) } @@ -108,6 +160,7 @@ async function getOpenCodeWorkspaceIdForSession(instanceId: string, sessionId: s } function clearOpenCodeWorkspaceCache(instanceId: string): void { + for (const controller of workspaceSyncs.get(instanceId)?.controllers ?? []) controller.abort() workspaceSyncs.delete(instanceId) workspaceIdByWorktreeSlug.delete(instanceId) } diff --git a/packages/ui/src/stores/permission-replies.test.ts b/packages/ui/src/stores/permission-replies.test.ts index 1079b7207..d13b50024 100644 --- a/packages/ui/src/stores/permission-replies.test.ts +++ b/packages/ui/src/stores/permission-replies.test.ts @@ -16,7 +16,7 @@ describe("replied permission tracking", () => { markPermissionReplied(instanceId, permissionId, 1_000) pruneRepliedPermissions(instanceId, new Set(), 900) - assert.equal(hasRepliedPermission(instanceId, permissionId), true) + assert.equal(hasRepliedPermission(instanceId, permissionId, 1_100), true) clearRepliedPermissions(instanceId) }) @@ -27,18 +27,33 @@ describe("replied permission tracking", () => { markPermissionReplied(instanceId, permissionId, 1_000) pruneRepliedPermissions(instanceId, new Set([permissionId]), 1_100) - assert.equal(hasRepliedPermission(instanceId, permissionId), true) + assert.equal(hasRepliedPermission(instanceId, permissionId, 1_100), true) clearRepliedPermissions(instanceId) }) - it("clears replied ids once a newer sync observes them missing", () => { + it("clears replied ids only after two newer syncs observe them missing", () => { const instanceId = "instance-new-sync" const permissionId = "permission-1" markPermissionReplied(instanceId, permissionId, 1_000) pruneRepliedPermissions(instanceId, new Set(), 1_100) + assert.equal(hasRepliedPermission(instanceId, permissionId, 1_100), true) + pruneRepliedPermissions(instanceId, new Set(), 1_200) - assert.equal(hasRepliedPermission(instanceId, permissionId), false) + assert.equal(hasRepliedPermission(instanceId, permissionId, 1_200), false) + clearRepliedPermissions(instanceId) + }) + + it("expires and caps replied ids", () => { + const instanceId = "instance-bounded" + markPermissionReplied(instanceId, "expired", 1_000) + assert.equal(hasRepliedPermission(instanceId, "expired", Number.MAX_SAFE_INTEGER), false) + + for (let index = 0; index <= 4_096; index += 1) { + markPermissionReplied(instanceId, `permission-${index}`, 10_000 + index) + } + assert.equal(hasRepliedPermission(instanceId, "permission-0", 20_000), false) + assert.equal(hasRepliedPermission(instanceId, "permission-4096", 20_000), true) clearRepliedPermissions(instanceId) }) }) diff --git a/packages/ui/src/stores/permission-replies.ts b/packages/ui/src/stores/permission-replies.ts index c354c63b9..1ac8bf8b6 100644 --- a/packages/ui/src/stores/permission-replies.ts +++ b/packages/ui/src/stores/permission-replies.ts @@ -1,14 +1,30 @@ -const repliedPermissionIdsByInstance = new Map>() +type RepliedPermission = { repliedAt: number; missingPasses: number } +const repliedPermissionIdsByInstance = new Map>() +const REPLIED_PERMISSION_TOMBSTONE_TTL_MS = 10 * 60 * 1_000 +const MAX_REPLIED_PERMISSION_TOMBSTONES = 4_096 + +function pruneExpiredRepliedPermissions(replied: Map, now: number): void { + for (const [permissionId, state] of replied) { + if (now - state.repliedAt < REPLIED_PERMISSION_TOMBSTONE_TTL_MS) break + replied.delete(permissionId) + } +} function pruneRepliedPermissions(instanceId: string, remotePendingIds: Set, syncStartedAt: number): void { const replied = repliedPermissionIdsByInstance.get(instanceId) if (!replied) return - for (const [permissionId, repliedAt] of replied) { - // Only a sync started after the local reply can prove the server no longer - // considers this permission pending. - if (!remotePendingIds.has(permissionId) && syncStartedAt >= repliedAt) { - replied.delete(permissionId) + pruneExpiredRepliedPermissions(replied, syncStartedAt) + for (const [permissionId, state] of replied) { + if (remotePendingIds.has(permissionId)) { + state.missingPasses = 0 + continue } + // Only a sync started after the local reply can prove the server no longer + // considers this permission pending. Keep it through one complete reconnect + // snapshot so delayed events from that stream remain fenced. + if (syncStartedAt < state.repliedAt) continue + if (state.missingPasses > 0) replied.delete(permissionId) + else state.missingPasses += 1 } if (replied.size === 0) { repliedPermissionIdsByInstance.delete(instanceId) @@ -19,15 +35,20 @@ function markPermissionReplied(instanceId: string, permissionId: string, replied if (!permissionId) return let replied = repliedPermissionIdsByInstance.get(instanceId) if (!replied) { - replied = new Map() + replied = new Map() repliedPermissionIdsByInstance.set(instanceId, replied) } - replied.set(permissionId, repliedAt) + pruneExpiredRepliedPermissions(replied, repliedAt) + replied.delete(permissionId) + replied.set(permissionId, { repliedAt, missingPasses: 0 }) + while (replied.size > MAX_REPLIED_PERMISSION_TOMBSTONES) replied.delete(replied.keys().next().value!) } -function hasRepliedPermission(instanceId: string, permissionId: string): boolean { +function hasRepliedPermission(instanceId: string, permissionId: string, now = Date.now()): boolean { const replied = repliedPermissionIdsByInstance.get(instanceId) if (!replied) return false + pruneExpiredRepliedPermissions(replied, now) + if (replied.size === 0) repliedPermissionIdsByInstance.delete(instanceId) return replied.has(permissionId) } diff --git a/packages/ui/src/stores/request-locations.test.ts b/packages/ui/src/stores/request-locations.test.ts index 11e259a82..bd2363db9 100644 --- a/packages/ui/src/stores/request-locations.test.ts +++ b/packages/ui/src/stores/request-locations.test.ts @@ -4,6 +4,13 @@ import { describe, it } from "node:test" import { buildV2RequestLocations } from "./request-locations.ts" describe("buildV2RequestLocations", () => { + it("treats unavailable worktree discovery as incomplete", () => { + assert.deepEqual(buildV2RequestLocations("/repo", [], new Map()), { + locations: [{ directory: "/repo" }], + complete: false, + }) + }) + it("includes root and each workspace-backed worktree location", () => { const locations = buildV2RequestLocations( "/repo", @@ -11,7 +18,6 @@ describe("buildV2RequestLocations", () => { { slug: "root" }, { slug: "feature-a" }, { slug: "feature-b" }, - { slug: "missing-workspace" }, ], new Map([ ["feature-a", "workspace-a"], @@ -19,11 +25,31 @@ describe("buildV2RequestLocations", () => { ]), ) - assert.deepEqual(locations, [ - { directory: "/repo" }, - { directory: "/repo", workspace: "workspace-a" }, - { directory: "/repo", workspace: "workspace-b" }, - ]) + assert.deepEqual(locations, { + locations: [ + { directory: "/repo" }, + { directory: "/repo", workspace: "workspace-a" }, + { directory: "/repo", workspace: "workspace-b" }, + ], + complete: true, + }) + }) + + it("keeps root and known locations while marking unresolved worktrees incomplete", () => { + assert.deepEqual( + buildV2RequestLocations( + "/repo", + [{ slug: "known" }, { slug: "missing" }], + new Map([["known", "workspace-known"]]), + ), + { + locations: [ + { directory: "/repo" }, + { directory: "/repo", workspace: "workspace-known" }, + ], + complete: false, + }, + ) }) it("deduplicates repeated workspace locations", () => { @@ -36,9 +62,12 @@ describe("buildV2RequestLocations", () => { ]), ) - assert.deepEqual(locations, [ - { directory: "/repo" }, - { directory: "/repo", workspace: "workspace-shared" }, - ]) + assert.deepEqual(locations, { + locations: [ + { directory: "/repo" }, + { directory: "/repo", workspace: "workspace-shared" }, + ], + complete: true, + }) }) }) diff --git a/packages/ui/src/stores/request-locations.ts b/packages/ui/src/stores/request-locations.ts index 0cefe5a28..5b5c4ae02 100644 --- a/packages/ui/src/stores/request-locations.ts +++ b/packages/ui/src/stores/request-locations.ts @@ -7,19 +7,28 @@ export type V2RequestLocationWorktree = { slug?: string } +export type V2RequestLocationSnapshot = { + locations: V2Location[] + complete: boolean +} + export function buildV2RequestLocations( directory: string | undefined, worktrees: V2RequestLocationWorktree[], workspaceBySlug: Map, -): V2Location[] { +): V2RequestLocationSnapshot { const rootLocation: V2Location = directory ? { directory } : {} const locations: V2Location[] = [rootLocation] const seen = new Set([JSON.stringify(rootLocation)]) + let complete = worktrees.length > 0 for (const worktree of worktrees) { if (!worktree.slug || worktree.slug === "root") continue const workspace = workspaceBySlug.get(worktree.slug) - if (!workspace) continue + if (!workspace) { + complete = false + continue + } const location: V2Location = { ...rootLocation, workspace } const key = JSON.stringify(location) if (seen.has(key)) continue @@ -27,5 +36,5 @@ export function buildV2RequestLocations( locations.push(location) } - return locations + return { locations, complete } } diff --git a/packages/ui/src/stores/session-actions-delivery.test.ts b/packages/ui/src/stores/session-actions-delivery.test.ts new file mode 100644 index 000000000..769f60fc1 --- /dev/null +++ b/packages/ui/src/stores/session-actions-delivery.test.ts @@ -0,0 +1,235 @@ +import assert from "node:assert/strict" +import { describe, it } from "node:test" + +import { sdkManager } from "../lib/sdk-manager.ts" +import { addInstance, removeInstance, updateInstance } from "./instances.ts" +import { messageStoreBus } from "./message-v2/bus.ts" +import { executeCustomCommand, runShellCommand, sendMessage } from "./session-actions.ts" +import { setSessions } from "./session-state.ts" + +function setup(instanceId: string, sessionId: string) { + const client = { session: {} } as any + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, client) + addInstance({ id: instanceId, folder: "/work", port: 0, pid: 1, proxyPath: "", status: "ready", client }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { + id: sessionId, + instanceId, + parentId: null, + title: sessionId, + agent: "build", + model: { providerId: "", modelId: "" }, + status: "idle", + retry: null, + idleSince: null, + version: "1", + } as any]]))) + return { + client, + cleanup() { + removeInstance(instanceId, { authoritative: false }) + sdkManager.destroyClientsForInstance(instanceId) + }, + } +} + +describe("dispatched command delivery classification", () => { + it("commits the fetched message when ambiguous delivery verification finds it", async () => { + const instanceId = "prompt-delivery-verified", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let messageId = "" + ;(client.session as any).promptAsync = async (input: any) => { + messageId = input.messageID + throw new TypeError("Failed to fetch") + } + ;(client.session as any).messages = async () => ({ data: [{ + info: { id: messageId, sessionID: sessionId, role: "user", time: { created: 1 } }, + parts: [{ id: "server-part", sessionID: sessionId, messageID: messageId, type: "text", text: "server prompt" }], + }] }) + + try { + await sendMessage(instanceId, sessionId, "optimistic prompt") + const message = messageStoreBus.getOrCreate(instanceId).getMessage(messageId) + assert.equal(message?.status, "complete") + assert.equal(message?.isEphemeral, false) + assert.deepEqual(message?.partIds, ["server-part"]) + assert.equal((message?.parts["server-part"]?.data as any)?.text, "server prompt") + } finally { + cleanup() + } + }) + + it("does not seed ambiguous verification from a replaced runtime", async () => { + const instanceId = "prompt-delivery-replaced-runtime", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let messageId = "" + let resolveVerification!: (value: any) => void + ;(client.session as any).promptAsync = async (input: any) => { + messageId = input.messageID + throw new TypeError("Failed to fetch") + } + ;(client.session as any).messages = () => new Promise((resolve) => { resolveVerification = resolve }) + + try { + const request = sendMessage(instanceId, sessionId, "optimistic prompt") + while (!resolveVerification) await new Promise((resolve) => setImmediate(resolve)) + updateInstance(instanceId, { client: { session: {} } as any }) + resolveVerification({ data: [{ + info: { id: messageId, sessionID: sessionId, role: "user", time: { created: 1 } }, + parts: [{ id: "stale-part", type: "text", text: "stale snapshot" }], + }] }) + const error = await request.catch((failure) => failure) + + const message = messageStoreBus.getOrCreate(instanceId).getMessage(messageId) + assert.equal((error as any)?.suppressPromptRecovery, true) + assert.equal(message, undefined) + } finally { + cleanup() + } + }) + + for (const action of ["prompt", "command", "shell"] as const) { + it(`treats a successful ${action} response from a replaced runtime as ambiguous`, async () => { + const instanceId = `successful-${action}-replaced-runtime`, sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let resolveDispatch!: (value: any) => void + let verificationCalls = 0 + ;(client.session as any)[action === "prompt" ? "promptAsync" : action] = () => new Promise((resolve) => { + resolveDispatch = resolve + }) + ;(client.session as any).messages = async () => { + verificationCalls += 1 + return { data: [] } + } + + try { + const request = action === "prompt" + ? sendMessage(instanceId, sessionId, "optimistic prompt") + : action === "command" + ? executeCustomCommand(instanceId, sessionId, "test", "") + : runShellCommand(instanceId, sessionId, "echo test") + while (!resolveDispatch) await new Promise((resolve) => setImmediate(resolve)) + updateInstance(instanceId, { client: { session: {} } as any }) + resolveDispatch({ data: undefined }) + const error = await request.catch((failure) => failure) + + assert.equal((error as any)?.suppressPromptRecovery, true) + if (action === "prompt") { + assert.equal(verificationCalls, 0) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), []) + } + } finally { + cleanup() + } + }) + } + + it("keeps a pre-response transport failure ambiguous after empty verification", async () => { + const instanceId = "prompt-delivery-empty", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let dispatchCalls = 0 + let verificationCalls = 0 + ;(client.session as any).promptAsync = async () => { + dispatchCalls += 1 + throw new TypeError("Failed to fetch") + } + ;(client.session as any).messages = async () => { + verificationCalls += 1 + return { data: [] } + } + + try { + const error = await sendMessage(instanceId, sessionId, "keep this draft").catch((failure) => failure) + assert.equal(error?.suppressPromptRecovery, true) + assert.equal(dispatchCalls, 1) + assert.equal(verificationCalls, 3) + } finally { + cleanup() + } + }) + + it("keeps a transport failure ambiguous after stale successful verification", async () => { + const instanceId = "prompt-delivery-unknown", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let verificationCalls = 0 + ;(client.session as any).promptAsync = async () => { throw new TypeError("terminated") } + ;(client.session as any).messages = async () => { + verificationCalls += 1 + if (verificationCalls === 1) return { data: [] } + throw new TypeError("verification disconnected") + } + + try { + const error = await sendMessage(instanceId, sessionId, "run it").catch((failure) => failure) + assert.equal(error?.suppressPromptRecovery, true) + assert.equal(verificationCalls, 3) + } finally { + cleanup() + } + }) + + it("keeps an explicit prompt rejection replayable after verification", async () => { + const instanceId = "prompt-delivery-rejected", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + ;(client.session as any).promptAsync = async () => ({ + error: { message: "Bad prompt" }, + response: { status: 400 }, + }) + ;(client.session as any).messages = async () => ({ data: [] }) + + try { + const error = await sendMessage(instanceId, sessionId, "bad").catch((failure) => failure) + assert.notEqual(error?.suppressPromptRecovery, true) + } finally { + cleanup() + } + }) + + it("keeps a rate-limited prompt replayable", async () => { + const instanceId = "prompt-delivery-rate-limited", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + ;(client.session as any).promptAsync = async () => ({ + error: { message: "Rate limited" }, + response: { status: 429 }, + }) + ;(client.session as any).messages = async () => ({ data: [] }) + + try { + const error = await sendMessage(instanceId, sessionId, "retry me").catch((failure) => failure) + assert.notEqual(error?.suppressPromptRecovery, true) + } finally { + cleanup() + } + }) + + it("suppresses replay for unknown command and shell failures", async () => { + const instanceId = "delivery-unknown", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + ;(client.session as any).command = async () => { throw new SyntaxError("Unexpected end of JSON input") } + ;(client.session as any).shell = async () => { throw new TypeError("terminated") } + + try { + const commandError = await executeCustomCommand(instanceId, sessionId, "test", "").catch((error) => error) + const shellError = await runShellCommand(instanceId, sessionId, "echo test").catch((error) => error) + assert.equal(commandError?.suppressPromptRecovery, true) + assert.equal(shellError?.suppressPromptRecovery, true) + } finally { + cleanup() + } + }) + + it("keeps a definitive client rejection replayable", async () => { + const instanceId = "delivery-rejected", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + ;(client.session as any).command = async () => ({ + error: { message: "Bad command" }, + response: { status: 400 }, + }) + + try { + const error = await executeCustomCommand(instanceId, sessionId, "bad", "").catch((failure) => failure) + assert.notEqual(error?.suppressPromptRecovery, true) + } finally { + cleanup() + } + }) +}) diff --git a/packages/ui/src/stores/session-actions.ts b/packages/ui/src/stores/session-actions.ts index 6c0db47c7..73bbae19d 100644 --- a/packages/ui/src/stores/session-actions.ts +++ b/packages/ui/src/stores/session-actions.ts @@ -1,5 +1,5 @@ import { preparePromptDisplayText } from "../lib/prompt-display-metadata" -import { instances } from "./instances" +import { instances, isInstanceRuntimeCurrent } from "./instances" import { getRootClient } from "./opencode-client" import { getOpenCodeWorkspaceIdForSession } from "./opencode-workspaces" @@ -10,8 +10,9 @@ import { updateSessionInfo } from "./message-v2/session-info" import { messageStoreBus } from "./message-v2/bus" import { removeMessagePartV2, removeMessageV2 } from "./message-v2/bridge" import { getLogger } from "../lib/logger" -import { requestData } from "../lib/opencode-api" +import { isDeliveryAmbiguousError, requestData } from "../lib/opencode-api" import { clearConversationPlaybackForSession } from "./conversation-speech" +import { normalizeMessagePart } from "./message-v2/normalizers" const log = getLogger("actions") @@ -42,6 +43,17 @@ const BASE62_CHARS = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuv let lastTimestamp = 0 let localCounter = 0 +function uncertainDeliveryError(error: unknown): Error { + const result = new Error(error instanceof Error ? error.message : String(error)) + ;(result as any).cause = error + ;(result as any).suppressPromptRecovery = true + return result +} + +function classifyDeliveryError(error: unknown): unknown { + return isDeliveryAmbiguousError(error) ? uncertainDeliveryError(error) : error +} + function randomBase62(length: number): string { let result = "" const cryptoObj = (globalThis as unknown as { crypto?: Crypto }).crypto @@ -216,20 +228,76 @@ async function sendMessage( try { log.info("session.promptAsync", { instanceId, sessionId, requestBody }) const workspacePayload = await getSessionWorkspacePayload(instanceId, sessionId) + if (!isInstanceRuntimeCurrent(instanceId, instance)) { + removeMessageV2(instanceId, messageId) + throw new Error("Instance no longer active") + } const admission = beginSessionGenerationAdmission(instanceId, sessionId) try { await requestData( client.session.promptAsync({ sessionID: sessionId, + messageID: messageId, ...workspacePayload, ...(requestBody as any), }), "session.promptAsync", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) { + if (store.getMessage(messageId)?.isEphemeral) removeMessageV2(instanceId, messageId) + throw uncertainDeliveryError(new Error("Instance changed after prompt delivery")) + } admission.complete() } catch (error) { - admission.rollback() - throw error + if (!isInstanceRuntimeCurrent(instanceId, instance)) { + if (store.getMessage(messageId)?.isEphemeral) removeMessageV2(instanceId, messageId) + throw uncertainDeliveryError(error) + } + for (let attempt = 0; attempt < 3; attempt += 1) { + try { + const messages = await requestData( + client.session.messages({ sessionID: sessionId, ...workspacePayload }), + "session.messages", + ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) { + if (store.getMessage(messageId)?.isEphemeral) removeMessageV2(instanceId, messageId) + throw uncertainDeliveryError(error) + } + const verified = messages.find((message) => (message?.info ?? message)?.id === messageId) + if (verified) { + const info = verified.info ?? verified + const createdAt = info.time?.created ?? Date.now() + const completedAt = info.time?.end ?? info.time?.completed ?? createdAt + store.upsertMessage({ + id: messageId, + sessionId, + role: info.role === "assistant" ? "assistant" : "user", + status: info.error ? "error" : "complete", + parts: (verified.parts ?? []).map(normalizeMessagePart), + createdAt, + updatedAt: completedAt, + isEphemeral: false, + }) + store.setMessageInfo(messageId, info) + if (isInstanceRuntimeCurrent(instanceId, instance)) admission.complete() + return + } + } catch { + if (!isInstanceRuntimeCurrent(instanceId, instance)) { + if (store.getMessage(messageId)?.isEphemeral) removeMessageV2(instanceId, messageId) + throw uncertainDeliveryError(error) + } + // A failed verification leaves delivery ambiguous. + } + if (attempt < 2) await new Promise((resolve) => setTimeout(resolve, 100)) + } + if (store.getMessageInfo(messageId)) { + if (isInstanceRuntimeCurrent(instanceId, instance)) admission.complete() + return + } + if (store.getMessage(messageId)?.isEphemeral) removeMessageV2(instanceId, messageId) + if (isInstanceRuntimeCurrent(instanceId, instance)) admission.rollback() + throw classifyDeliveryError(error) } } catch (error) { log.error("Failed to send prompt", error) @@ -279,6 +347,7 @@ async function executeCustomCommand( } const workspacePayload = await getSessionWorkspacePayload(instanceId, sessionId) + if (!isInstanceRuntimeCurrent(instanceId, instance)) throw new Error("Instance no longer active") const admission = beginSessionGenerationAdmission(instanceId, sessionId) try { await requestData( @@ -289,10 +358,13 @@ async function executeCustomCommand( }), "session.command", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) { + throw uncertainDeliveryError(new Error("Instance changed after command delivery")) + } admission.complete() } catch (error) { - admission.rollback() - throw error + if (isInstanceRuntimeCurrent(instanceId, instance)) admission.rollback() + throw classifyDeliveryError(error) } } @@ -312,6 +384,7 @@ async function runShellCommand(instanceId: string, sessionId: string, command: s const agent = session.agent || "build" const workspacePayload = await getSessionWorkspacePayload(instanceId, sessionId) + if (!isInstanceRuntimeCurrent(instanceId, instance)) throw new Error("Instance no longer active") const admission = beginSessionGenerationAdmission(instanceId, sessionId) try { await requestData( @@ -323,10 +396,13 @@ async function runShellCommand(instanceId: string, sessionId: string, command: s }), "session.shell", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) { + throw uncertainDeliveryError(new Error("Instance changed after shell delivery")) + } admission.complete() } catch (error) { - admission.rollback() - throw error + if (isInstanceRuntimeCurrent(instanceId, instance)) admission.rollback() + throw classifyDeliveryError(error) } } @@ -342,10 +418,12 @@ async function abortSession(instanceId: string, sessionId: string): Promise { + const instance = instances().get(instanceId) + if (!instance?.client) throw new Error("Instance not ready") const instanceSessions = sessions().get(instanceId) const session = instanceSessions?.get(sessionId) if (!session) { @@ -364,6 +444,7 @@ async function updateSessionAgent(instanceId: string, sessionId: string, agent: } const nextModel = await getDefaultModel(instanceId, agent) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return const shouldApplyModel = isModelValid(instanceId, nextModel) withSession(instanceId, sessionId, (current) => { @@ -375,6 +456,7 @@ async function updateSessionAgent(instanceId: string, sessionId: string, agent: if (agent && shouldApplyModel) { await setAgentModelPreference(instanceId, agent, nextModel) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return } if (shouldApplyModel) { @@ -387,6 +469,8 @@ async function updateSessionModel( sessionId: string, model: { providerId: string; modelId: string }, ): Promise { + const instance = instances().get(instanceId) + if (!instance?.client) throw new Error("Instance not ready") const instanceSessions = sessions().get(instanceId) const session = instanceSessions?.get(sessionId) if (!session) { @@ -404,6 +488,7 @@ async function updateSessionModel( if (session.agent) { await setAgentModelPreference(instanceId, session.agent, model) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return } addRecentModelPreference(model) @@ -428,14 +513,17 @@ async function renameSession(instanceId: string, sessionId: string, nextTitle: s throw new Error("Session title is required") } + const workspace = await getSessionWorkspacePayload(instanceId, sessionId) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return await requestData( client.session.update({ sessionID: sessionId, - ...(await getSessionWorkspacePayload(instanceId, sessionId)), + ...workspace, title: trimmedTitle, }), "session.update", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return withSession(instanceId, sessionId, (current) => { current.title = trimmedTitle @@ -453,16 +541,19 @@ async function deleteMessagePart(instanceId: string, sessionId: string, messageI } const client = getRootClient(instanceId) + const workspace = await getSessionWorkspacePayload(instanceId, sessionId) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return await requestData( client.part.delete({ sessionID: sessionId, - ...(await getSessionWorkspacePayload(instanceId, sessionId)), + ...workspace, messageID: messageId, partID: partId, }), "part.delete", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return // Optimistic removal; SSE will also broadcast a part-removed event. removeMessagePartV2(instanceId, messageId, partId) @@ -477,16 +568,19 @@ async function deleteMessage(instanceId: string, sessionId: string, messageId: s } const client = getRootClient(instanceId) + const workspace = await getSessionWorkspacePayload(instanceId, sessionId) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return // The SDK generator does not currently expose a typed method for deleting a message, // but the API is available at DELETE /session/:sessionID/message/:messageID. await requestData( (client as any).client.delete({ url: `/session/${encodeURIComponent(sessionId)}/message/${encodeURIComponent(messageId)}`, - query: await getSessionWorkspacePayload(instanceId, sessionId), + query: workspace, }), "session.message.delete", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return // Optimistic removal; SSE will also broadcast a message-removed event. removeMessageV2(instanceId, messageId) diff --git a/packages/ui/src/stores/session-api.ts b/packages/ui/src/stores/session-api.ts index 28b84ce64..de146a7b0 100644 --- a/packages/ui/src/stores/session-api.ts +++ b/packages/ui/src/stores/session-api.ts @@ -6,10 +6,11 @@ import { type Session, type SessionStatus, } from "../types/session" -import type { Message } from "../types/message" +import type { Message, MessageInfo } from "../types/message" +import type { Instance } from "../types/instance" import type { Session as SDKSession, SessionListResponse } from "@opencode-ai/sdk/v2/client" -import { instances, reconcilePendingSessionIndicators } from "./instances" +import { clearSessionInterruptions, instances, isInstanceRuntimeCurrent, reconcilePendingSessionIndicators } from "./instances" import { preferences, setAgentModelPreference } from "./preferences" import { activeSessionId, @@ -22,14 +23,15 @@ import { cancelSessionGenerationAdmissions, markSessionDeletedAuthoritative, getAuthoritativelyDeletedSessionIdsForInstance, - getDescendantSessions, isBlankSession, messagesLoaded, getSessionMessagesLoadError, providers, setAgents, setMessagesLoaded, - advanceMessageLoadEpoch, + beginSessionMessageLoad, + finishSessionMessageLoad, + invalidateSessionMessageLoad, isCurrentMessageLoad, setSessionMessagesLoadError, setProviders, @@ -47,17 +49,21 @@ import { prependSessionListId, removeSessionListId, beginSessionSearch, - clearSessionSearch, + clearSessionSearch as clearSessionSearchState, isLatestSessionSearch, setSessionSearchResults, setSessionListError, setSessionExpanded, + markSessionMetadataMutation, + snapshotSessionMetadataMutationVersion, + wasSessionMetadataMutatedAfter, } from "./session-state" import { deleteSessionAttachments } from "./attachments" import { DEFAULT_MODEL_OUTPUT_LIMIT, getDefaultModel, isModelValid } from "./session-models" import { normalizeMessagePart } from "./message-v2/normalizers" import { updateSessionInfo } from "./message-v2/session-info" -import { seedSessionMessagesV2, reconcilePendingPermissionsV2, reconcilePendingQuestionsV2 } from "./message-v2/bridge" +import { seedSessionMessagesV2, reconcilePendingPermissionsV2, reconcilePendingQuestionsV2, setSessionRevertV2 } from "./message-v2/bridge" +import { clearPendingDeltasForSession, getPendingDeltasForMessage, hasPendingDeltasForMessage, requestDeltaRecovery } from "./delta-buffer" import { messageStoreBus } from "./message-v2/bus" import { clearCacheForSession } from "../lib/global-cache" import { getLogger } from "../lib/logger" @@ -88,9 +94,70 @@ import { mergeFetchedSessionRuntimeState, resolveAuthoritativeGenerationRecovery const log = getLogger("api") const sessionListRequestIds = new Map() let nextSessionListRequestId = 0 -const pendingMetadataHydrations = new Map>() +const pendingMetadataHydrations = new Map }>() +const pendingSessionSearches = new Map() const sessionWorkspaceHints = new Map>() -messageStoreBus.onInstanceDestroyed((instanceId) => sessionWorkspaceHints.delete(instanceId)) +type BufferedDeltaExpectation = { partId: string; field: string; value: string; staleSnapshotsRemaining: number } +const bufferedDeltaSnapshotFences = new Map>() +// ponytail: tolerate one settled stale snapshot; repeated omission/replacement is authoritative. +const BUFFERED_DELTA_STALE_SNAPSHOT_LIMIT = 1 + +class SessionMessageLoadTimeoutError extends Error { + constructor(timeoutMs: number) { + super(`Session message load timed out after ${timeoutMs}ms`) + this.name = "SessionMessageLoadTimeoutError" + } +} + +function clearBufferedDeltaSnapshotFence(instanceId: string, sessionId: string, messageId: string, partId: string): void { + const key = `${instanceId}:${sessionId}` + const fence = bufferedDeltaSnapshotFences.get(key) + const expectations = fence?.get(messageId) + if (!fence || !expectations) return + const remaining = expectations.filter((expectation) => expectation.partId !== partId) + if (remaining.length > 0) fence.set(messageId, remaining) + else fence.delete(messageId) + if (fence.size === 0) bufferedDeltaSnapshotFences.delete(key) +} + +messageStoreBus.onInstanceDestroyed((instanceId) => { + pendingSessionSearches.get(instanceId)?.abort() + pendingSessionSearches.delete(instanceId) + sessionWorkspaceHints.delete(instanceId) + const prefix = `${instanceId}:` + for (const key of pendingMetadataHydrations.keys()) if (key.startsWith(prefix)) pendingMetadataHydrations.delete(key) + for (const key of bufferedDeltaSnapshotFences.keys()) if (key.startsWith(prefix)) bufferedDeltaSnapshotFences.delete(key) +}) +messageStoreBus.onSessionCleared((instanceId, sessionId) => bufferedDeltaSnapshotFences.delete(`${instanceId}:${sessionId}`)) + +function adaptApiMessages( + sessionId: string, + apiMessages: any[], + sessionStatus: SessionStatus = "idle", +): { messages: Message[]; infos: Map } { + const infos = new Map() + const messages = apiMessages.map((apiMessage: any, index: number) => { + const info = (apiMessage.info || apiMessage) as MessageInfo + const messageId = info.id || String(Date.now()) + infos.set(messageId, info) + return { + id: messageId, + sessionId, + type: info.role === "user" ? "user" as const : "assistant" as const, + parts: (apiMessage.parts || []).map((part: any) => normalizeMessagePart(part)), + timestamp: info.time?.created || Date.now(), + status: (info as any).error + ? "error" as const + : info.role === "assistant" && index === apiMessages.length - 1 && + !info.time?.completed && !(info.time as { end?: number } | undefined)?.end && + (sessionStatus === "working" || sessionStatus === "compacting") + ? "streaming" as const + : "complete" as const, + version: 0, + } + }) + return { messages, infos } +} function beginSessionListRequest(instanceId: string): number { const requestId = ++nextSessionListRequestId @@ -151,7 +218,11 @@ function rememberSessionWorkspace(instanceId: string, sessionId: string, workspa sessionWorkspaceHints.set(instanceId, hints) } -async function recordSessionWorkspaceHints(instanceId: string, apiSessions: SDKSession[]): Promise { +async function recordSessionWorkspaceHints( + instanceId: string, + apiSessions: SDKSession[], + hasCommitAuthority: () => boolean, +): Promise { const hints = new Map(sessionWorkspaceHints.get(instanceId) ?? new Map()) const workspaceBySlug = new Map>() await Promise.all(apiSessions.map(async (session) => { @@ -166,6 +237,7 @@ async function recordSessionWorkspaceHints(instanceId: string, apiSessions: SDKS const workspaceId = await workspace if (workspaceId) hints.set(session.id, workspaceId) })) + if (!hasCommitAuthority()) return sessionWorkspaceHints.set(instanceId, hints) } @@ -223,10 +295,10 @@ function hasMissingParentChain(session: SDKSession, loaded: Map { +async function fetchV2Sessions(instanceId: string, options: V2SessionListOptions, signal?: AbortSignal): Promise { const client = getRootClient(instanceId) const listOptions = buildProjectSessionListOptions(options) - const data = await requestData(client.session.list(listOptions), "session.list") + const data = await requestData((client.session.list as any)(listOptions, signal ? { signal } : undefined), "session.list") const allowedDirectories = [options.directory, ...getWorktrees(instanceId).map((worktree) => worktree.directory)] return { @@ -263,8 +335,9 @@ async function hydrateMissingSessionMetadata(instanceId: string, sessionIds: str function hydrateSessionMetadata(instanceId: string, sessionId: string, client = getRootClient(instanceId)): Promise { const key = `${instanceId}:${sessionId}` + const instance = instances().get(instanceId) const current = pendingMetadataHydrations.get(key) - if (current) return current + if (current && current.runtimeToken === instance?.runtimeToken) return current.promise const hydration = (async () => { const candidates = await getSessionWorkspaceCandidates(instanceId, sessionId) let lastError: unknown @@ -272,7 +345,8 @@ function hydrateSessionMetadata(instanceId: string, sessionId: string, client = if (delayMs > 0) await new Promise((resolve) => setTimeout(resolve, delayMs)) for (const candidate of candidates) { try { - await hydrateSessionMetadataWithClient(client, instanceId, sessionId, candidate) + await hydrateSessionMetadataWithClient(client, instanceId, sessionId, candidate, () => isInstanceRuntimeCurrent(instanceId, instance)) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return rememberSessionWorkspace(instanceId, sessionId, candidate.workspace) return } catch (error) { @@ -281,8 +355,10 @@ function hydrateSessionMetadata(instanceId: string, sessionId: string, client = } } throw lastError - })().finally(() => pendingMetadataHydrations.delete(key)) - pendingMetadataHydrations.set(key, hydration) + })().finally(() => { + if (pendingMetadataHydrations.get(key)?.promise === hydration) pendingMetadataHydrations.delete(key) + }) + pendingMetadataHydrations.set(key, { runtimeToken: instance?.runtimeToken, promise: hydration }) return hydration } @@ -290,13 +366,19 @@ async function hydrateRestoredSessionChain( instanceId: string, requestedIds: Array, signal?: AbortSignal, + options?: { hasCommitAuthority?: () => boolean; hydrateKnownMetadata?: boolean }, ): Promise { + const instance = instances().get(instanceId) + if (!instance) throw new Error("Instance not ready") const client = getRootClient(instanceId) + const isCurrentInstance = () => isInstanceRuntimeCurrent(instanceId, instance) && + !signal?.aborted && (options?.hasCommitAuthority?.() ?? true) const pending = requestedIds.filter((id): id is string => Boolean(id) && id !== "info") const visited = new Set() let chainWorkspacePayload: { workspace?: string } = {} while (pending.length > 0) { signal?.throwIfAborted() + if (!isCurrentInstance()) return const sessionId = pending.shift()! if (visited.has(sessionId)) continue visited.add(sessionId) @@ -305,20 +387,27 @@ async function hydrateRestoredSessionChain( let session = sessions().get(instanceId)?.get(sessionId) if (!session) { try { + const metadataMutationFence = snapshotSessionMetadataMutationVersion() const workspaceCandidates = await getSessionWorkspaceCandidates(instanceId, sessionId, chainWorkspacePayload) signal?.throwIfAborted() + if (!isCurrentInstance()) return let apiSession: SDKSession | undefined let hydratedWorkspace: string | undefined let lastError: unknown for (const workspacePayload of workspaceCandidates) { try { apiSession = await requestData( - client.session.get({ sessionID: sessionId, ...workspacePayload }), + (client.session.get as any)( + { sessionID: sessionId, ...workspacePayload }, + signal ? { signal } : undefined, + ), "session.get", ) + if (!isCurrentInstance()) return hydratedWorkspace = workspacePayload.workspace break } catch (error) { + if (signal?.aborted) throw error lastError = error } } @@ -326,9 +415,10 @@ async function hydrateRestoredSessionChain( signal?.throwIfAborted() rememberSessionWorkspace(instanceId, sessionId, hydratedWorkspace) setSessions((prev) => { - if (getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(sessionId) || signal?.aborted) return prev + if (!isCurrentInstance() || getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(sessionId) || signal?.aborted) return prev const next = new Map(prev) const instanceSessions = new Map(next.get(instanceId) ?? new Map()) + if (instanceSessions.has(sessionId) && wasSessionMetadataMutatedAfter(instanceId, sessionId, metadataMutationFence)) return prev instanceSessions.set(sessionId, toClientSessionV2(instanceId, apiSession, instanceSessions.get(sessionId))) next.set(instanceId, instanceSessions) return next @@ -340,7 +430,7 @@ async function hydrateRestoredSessionChain( log.warn("Failed to hydrate restored session", { instanceId, sessionId, error }) continue } - } else if (shouldReplaceSessionMetadata(session.metadata)) { + } else if (options?.hydrateKnownMetadata !== false && shouldReplaceSessionMetadata(session.metadata)) { try { await hydrateSessionMetadata(instanceId, sessionId, client) } catch (error) { @@ -348,6 +438,7 @@ async function hydrateRestoredSessionChain( log.warn("Failed to hydrate restored session metadata", { instanceId, sessionId, error }) } } + session = sessions().get(instanceId)?.get(sessionId) if (session?.parentId === null) { const rootWorkspacePayload = await getSessionWorkspacePayload(instanceId, session.id) if (rootWorkspacePayload.workspace) chainWorkspacePayload = rootWorkspacePayload @@ -356,35 +447,75 @@ async function hydrateRestoredSessionChain( } } -async function ensureV2ParentChainsLoaded(instanceId: string, apiSessions: SDKSession[], directory?: string): Promise { +async function ensureV2ParentChainsLoaded( + instanceId: string, + apiSessions: SDKSession[], + hasCommitAuthority: () => boolean, + directory?: string, + signal?: AbortSignal, +): Promise { const currentSessions = sessions().get(instanceId) ?? new Map() const loaded = new Map(currentSessions) for (const session of apiSessions) loaded.set(session.id, session) - if (!apiSessions.some((session) => hasMissingParentChain(session, loaded))) return + if (!hasCommitAuthority() || !apiSessions.some((session) => hasMissingParentChain(session, loaded))) return - const page = await fetchV2Sessions(instanceId, { directory }) + const metadataMutationFence = snapshotSessionMetadataMutationVersion() + const page = await fetchV2Sessions(instanceId, { directory }, signal) + if (!hasCommitAuthority()) return const items = getV2SessionItems(page) - if (items.length === 0) return + const supplementalById = new Map(items.map((session) => [session.id, session])) + const missingAncestorIds = new Set() + for (const apiSession of apiSessions) { + let current: SDKSession | Session = apiSession + const seen = new Set() + while (getKnownParentId(current)) { + const parentId = getKnownParentId(current) + if (!parentId || seen.has(parentId)) break + seen.add(parentId) + const existingParent = loaded.get(parentId) + if (existingParent) { + current = existingParent + continue + } + const supplementalParent = supplementalById.get(parentId) + if (!supplementalParent) break + missingAncestorIds.add(parentId) + loaded.set(parentId, supplementalParent) + current = supplementalParent + } + } + if (missingAncestorIds.size > 0) { + setSessions((prev) => { + if (!hasCommitAuthority()) return prev + const next = new Map(prev) + const instanceSessions = new Map(next.get(instanceId) ?? new Map()) + const deletedSessionIds = getAuthoritativelyDeletedSessionIdsForInstance(instanceId) - setSessions((prev) => { - const next = new Map(prev) - const instanceSessions = new Map(next.get(instanceId) ?? new Map()) - const deletedSessionIds = getAuthoritativelyDeletedSessionIdsForInstance(instanceId) + for (const sessionId of missingAncestorIds) { + if (deletedSessionIds.has(sessionId)) continue + const apiSession = supplementalById.get(sessionId) + if (!apiSession) continue + const existingSession = instanceSessions.get(sessionId) + if (existingSession && wasSessionMetadataMutatedAfter(instanceId, sessionId, metadataMutationFence)) continue + instanceSessions.set(sessionId, toClientSessionV2(instanceId, apiSession, existingSession)) + } - for (const apiSession of items) { - if (deletedSessionIds.has(apiSession.id)) continue - const existingSession = instanceSessions.get(apiSession.id) - instanceSessions.set(apiSession.id, toClientSessionV2(instanceId, apiSession, existingSession)) - loaded.set(apiSession.id, apiSession) - } + next.set(instanceId, instanceSessions) + return next + }) + } - next.set(instanceId, instanceSessions) - return next + await hydrateRestoredSessionChain(instanceId, apiSessions.map((session) => session.id), signal, { + hasCommitAuthority, + hydrateKnownMetadata: false, }) } -async function fetchSessions(instanceId: string, options?: { reset?: boolean }): Promise { +async function fetchSessions( + instanceId: string, + options?: { reset?: boolean; authoritativeDeletes?: boolean; signal?: AbortSignal; hasCommitAuthority?: () => boolean }, +): Promise> { const instance = instances().get(instanceId) if (!instance || !instance.client) { throw new Error("Instance not ready") @@ -392,6 +523,11 @@ async function fetchSessions(instanceId: string, options?: { reset?: boolean }): const rootClient = getRootClient(instanceId) const requestId = beginSessionListRequest(instanceId) + const metadataMutationFence = snapshotSessionMetadataMutationVersion() + const hasCommitAuthority = () => !options?.signal?.aborted && + (options?.hasCommitAuthority?.() ?? true) && + isInstanceRuntimeCurrent(instanceId, instance) && + isLatestSessionListRequest(instanceId, requestId) setLoading((prev) => { const next = { ...prev } @@ -403,16 +539,20 @@ async function fetchSessions(instanceId: string, options?: { reset?: boolean }): try { const sessionListOptions = instance.folder ? { directory: instance.folder } : {} const existingSessions = new Map(sessions().get(instanceId) ?? new Map()) + const store = messageStoreBus.getOrCreate(instanceId) + const residentSessionIds = new Set(store.getResidentSessionIds()) + const revertReloadIds = new Set() log.info("session.list", { instanceId, limit: PROJECT_SESSION_LIST_LIMIT, directory: sessionListOptions.directory, scope: "project" }) - const response = await fetchV2Sessions(instanceId, sessionListOptions) - if (!isLatestSessionListRequest(instanceId, requestId)) return - await recordSessionWorkspaceHints(instanceId, getV2SessionItems(response)) + const response = await fetchV2Sessions(instanceId, sessionListOptions, options?.signal) + if (!hasCommitAuthority()) return new Set() + await recordSessionWorkspaceHints(instanceId, getV2SessionItems(response), hasCommitAuthority) + if (!hasCommitAuthority()) return new Set() let statusById: Record = {} let statusResponseKnown = false try { - const statusResponse = await rootClient.session.status() + const statusResponse = await (rootClient.session.status as any)(undefined, options?.signal ? { signal: options.signal } : undefined) if (statusResponse.data && typeof statusResponse.data === "object") { statusResponseKnown = true statusById = statusResponse.data as Record @@ -420,26 +560,53 @@ async function fetchSessions(instanceId: string, options?: { reset?: boolean }): } catch (error) { log.error("Failed to fetch session status:", error) } - if (!isLatestSessionListRequest(instanceId, requestId)) return + if (!hasCommitAuthority()) return new Set() const sessionMap = new Map() + const refreshedSessionIds = new Set() for (const apiSession of getV2SessionItems(response)) { const existingSession = existingSessions?.get(apiSession.id) - const existingStatus = existingSession?.status - const rawStatus = (apiSession as any)?.status ?? statusById[apiSession.id] + const latestSession = sessions().get(instanceId)?.get(apiSession.id) + const metadataMutated = wasSessionMetadataMutatedAfter(instanceId, apiSession.id, metadataMutationFence) + const statusBase = metadataMutated && latestSession ? latestSession : existingSession + const existingStatus = statusBase?.status + const rawStatus = statusResponseKnown ? statusById[apiSession.id] : (apiSession as any)?.status ?? statusById[apiSession.id] const hasType = rawStatus && typeof rawStatus === "object" && typeof rawStatus.type === "string" - const runtimeStatusKnown = Boolean(hasType || statusResponseKnown || existingSession?.runtimeStatusKnown) + const runtimeStatusKnown = Boolean(hasType || statusResponseKnown || statusBase?.runtimeStatusKnown) let status: SessionStatus - let retry = existingSession?.retry ?? null - if (existingStatus === "compacting") { + let retry = statusBase?.retry ?? null + if (existingStatus === "compacting" && !statusResponseKnown) { status = "compacting" retry = null } else { - status = hasType ? mapSdkSessionStatus(rawStatus) : existingStatus ?? "idle" + status = hasType ? mapSdkSessionStatus(rawStatus) : statusResponseKnown ? "idle" : existingStatus ?? "idle" retry = hasType ? mapSdkSessionRetry(rawStatus) : retry } + + if (metadataMutated && latestSession) { + refreshedSessionIds.add(apiSession.id) + if (hasType || statusResponseKnown) { + sessionMap.set(apiSession.id, { + ...latestSession, + status, + retry, + idleSince: getIdleSinceForStatusTransition(existingStatus, status, latestSession.idleSince), + runtimeStatusKnown, + generationRecovery: resolveAuthoritativeGenerationRecovery(latestSession.generationRecovery, status), + }) + } + continue + } + + const incomingRevert = apiSession.revert ?? null + if (residentSessionIds.has(apiSession.id) && ( + !sameSessionRevert(existingSession?.revert, incomingRevert) || + !sameSessionRevert(store.getSessionRevert(apiSession.id), incomingRevert) + )) { + revertReloadIds.add(apiSession.id) + } sessionMap.set(apiSession.id, { ...toClientSessionV2(instanceId, apiSession, existingSession), status, @@ -455,10 +622,14 @@ async function fetchSessions(instanceId: string, options?: { reset?: boolean }): const remotelyDeletedSessionIds = getAuthoritativelyMissingSessionIds( existingSessions.keys(), response.listedIds, - response.complete, + response.complete && options?.authoritativeDeletes !== false, ) - for (const sessionId of remotelyDeletedSessionIds) removeSessionRuntimeState(instanceId, sessionId) + for (const sessionId of remotelyDeletedSessionIds) { + if (wasSessionMetadataMutatedAfter(instanceId, sessionId, metadataMutationFence)) continue + removeSessionRuntimeState(instanceId, sessionId) + } + const committedSessionIds: string[] = [] setSessions((prev) => { const next = new Map(prev) const instanceSessions = new Map(next.get(instanceId) ?? new Map()) @@ -472,11 +643,25 @@ async function fetchSessions(instanceId: string, options?: { reset?: boolean }): latestSession, deletedSessionIds.has(session.id), ) - if (merged) instanceSessions.set(session.id, merged) + if (merged) { + instanceSessions.set(session.id, merged) + committedSessionIds.push(session.id) + } } next.set(instanceId, instanceSessions) return next }) + for (const sessionId of committedSessionIds) { + markSessionMetadataMutation(instanceId, sessionId) + refreshedSessionIds.add(sessionId) + } + await Promise.all([...revertReloadIds] + .filter((sessionId) => refreshedSessionIds.has(sessionId)) + .map(async (sessionId) => { + invalidateSessionMessageLoad(instanceId, sessionId) + await loadMessages(instanceId, sessionId, { force: true, signal: options?.signal }) + })) + if (!hasCommitAuthority()) return new Set() const rootIds: string[] = [] const seenRootIds = new Set() @@ -527,14 +712,16 @@ async function fetchSessions(instanceId: string, options?: { reset?: boolean }): })().catch((error) => { log.warn("Failed to finish legacy worktree map migration", { instanceId, error }) }) + const currentSessions = sessions().get(instanceId) + return new Set([...refreshedSessionIds].filter((sessionId) => currentSessions?.has(sessionId))) } catch (error) { log.error("Failed to fetch sessions:", error) - if (isLatestSessionListRequest(instanceId, requestId)) { + if (hasCommitAuthority()) { setSessionListError(instanceId, getOpencodeErrorMessage(error, tGlobal("sessionList.loadError.detail"))) } throw error } finally { - if (isLatestSessionListRequest(instanceId, requestId)) { + if (isInstanceRuntimeCurrent(instanceId, instance) && isLatestSessionListRequest(instanceId, requestId)) { setLoading((prev) => { const next = { ...prev } next.fetchingSessions.set(instanceId, false) @@ -548,6 +735,19 @@ async function loadMoreSessions(instanceId: string): Promise { return } +function clearSessionSearch(instanceId: string): void { + pendingSessionSearches.get(instanceId)?.abort() + pendingSessionSearches.delete(instanceId) + clearSessionSearchState(instanceId) +} + +type ComparableSessionRevert = { messageID?: string; partID?: string; snapshot?: string; diff?: string } + +function sameSessionRevert(left: ComparableSessionRevert | null | undefined, right: ComparableSessionRevert | null | undefined): boolean { + return left?.messageID === right?.messageID && left?.partID === right?.partID && + left?.snapshot === right?.snapshot && left?.diff === right?.diff +} + async function searchSessions(instanceId: string, query: string): Promise { const trimmedQuery = query.trim() if (!trimmedQuery) return @@ -557,15 +757,21 @@ async function searchSessions(instanceId: string, query: string): Promise throw new Error("Instance not ready") } + pendingSessionSearches.get(instanceId)?.abort() + const controller = new AbortController() + pendingSessionSearches.set(instanceId, controller) const requestId = beginSessionSearch(instanceId, trimmedQuery) + const metadataMutationFence = snapshotSessionMetadataMutationVersion() + const hasCommitAuthority = () => !controller.signal.aborted && pendingSessionSearches.get(instanceId) === controller && + isInstanceRuntimeCurrent(instanceId, instance) && isLatestSessionSearch(instanceId, trimmedQuery, requestId) try { log.info("v2.session.search", { instanceId, query: trimmedQuery, directory: instance.folder }) const response = await fetchV2Sessions(instanceId, { search: trimmedQuery, directory: instance.folder, - }) - if (!isLatestSessionSearch(instanceId, trimmedQuery, requestId)) return + }, controller.signal) + if (!hasCommitAuthority()) return const searchResults = getV2SessionItems(response) @@ -574,7 +780,11 @@ async function searchSessions(instanceId: string, query: string): Promise return } + const revertReloadIds = new Set() + const committedSessionIds: string[] = [] + const residentSessionIds = new Set(messageStoreBus.getOrCreate(instanceId).getResidentSessionIds()) setSessions((prev) => { + if (!hasCommitAuthority()) return prev const next = new Map(prev) const instanceSessions = new Map(next.get(instanceId) ?? new Map()) const deletedSessionIds = getAuthoritativelyDeletedSessionIdsForInstance(instanceId) @@ -582,43 +792,61 @@ async function searchSessions(instanceId: string, query: string): Promise for (const apiSession of searchResults) { if (deletedSessionIds.has(apiSession.id)) continue const existingSession = instanceSessions.get(apiSession.id) + if (existingSession && wasSessionMetadataMutatedAfter(instanceId, apiSession.id, metadataMutationFence)) continue + const incomingRevert = apiSession.revert ?? null + const storeRevert = messageStoreBus.getOrCreate(instanceId).getSessionRevert(apiSession.id) + if ((existingSession && !sameSessionRevert(existingSession.revert, incomingRevert)) || + (residentSessionIds.has(apiSession.id) && !sameSessionRevert(storeRevert, incomingRevert))) { + revertReloadIds.add(apiSession.id) + } instanceSessions.set(apiSession.id, toClientSessionV2(instanceId, apiSession, existingSession)) + committedSessionIds.push(apiSession.id) } next.set(instanceId, instanceSessions) return next }) - void hydrateMissingSessionMetadata(instanceId, searchResults.map((session) => session.id)) + for (const sessionId of committedSessionIds) markSessionMetadataMutation(instanceId, sessionId) - await ensureV2ParentChainsLoaded(instanceId, searchResults, instance.folder) + await ensureV2ParentChainsLoaded( + instanceId, + searchResults, + hasCommitAuthority, + instance.folder, + controller.signal, + ) - if (!isLatestSessionSearch(instanceId, trimmedQuery, requestId)) return + if (!hasCommitAuthority()) return - const hydratedSessions = sessions().get(instanceId) - const deletedSessionIds = getAuthoritativelyDeletedSessionIdsForInstance(instanceId) - const currentSearchResults = searchResults.filter((session) => !deletedSessionIds.has(session.id)) - const hasUnrenderableChildResult = currentSearchResults.some((session) => { - const parentId = session.parentID - return Boolean(parentId && !hydratedSessions?.has(parentId)) - }) + await Promise.all([...revertReloadIds].map(async (sessionId) => { + invalidateSessionMessageLoad(instanceId, sessionId) + await loadMessages(instanceId, sessionId, { force: true, signal: controller.signal }) + })) + if (!hasCommitAuthority()) return - if (hasUnrenderableChildResult) { - clearSessionSearch(instanceId) - return - } + const deletedSessionIds = getAuthoritativelyDeletedSessionIdsForInstance(instanceId) + const currentSearchResults = searchResults.filter((session) => + !deletedSessionIds.has(session.id) && Boolean(getSessionRoot(instanceId, session.id))) syncInstanceSessionIndicator(instanceId) setSessionSearchResults(instanceId, trimmedQuery, currentSearchResults.map((session) => session.id), requestId) } catch (error) { + if (controller.signal.aborted) return log.error("Failed to search sessions:", error) - if (isLatestSessionSearch(instanceId, trimmedQuery, requestId)) { - clearSessionSearch(instanceId) + if (hasCommitAuthority()) { + clearSessionSearchState(instanceId) } throw error + } finally { + if (pendingSessionSearches.get(instanceId) === controller) pendingSessionSearches.delete(instanceId) } } -function toClientSessionV2(instanceId: string, apiSession: SDKSession, existingSession?: Session): Session { +function toClientSessionV2( + instanceId: string, + apiSession: SDKSession, + existingSession?: Session, +): Session { const incomingMetadata = (apiSession as SDKSession & { metadata?: Session["metadata"] }).metadata return { id: apiSession.id, @@ -643,7 +871,7 @@ function toClientSessionV2(instanceId: string, apiSession: SDKSession, existingS ...apiSession.time, }, metadata: preferSessionMetadata(incomingMetadata, existingSession?.metadata), - revert: existingSession?.revert, + revert: apiSession.revert ? { ...apiSession.revert } : undefined, pendingPermission: existingSession?.pendingPermission, pendingQuestion: existingSession?.pendingQuestion, } @@ -666,9 +894,11 @@ async function createSession(instanceId: string, agent?: string): Promise 0 ? primaryAgents[0].name : "") const defaultModel = await getDefaultModel(instanceId, selectedAgent) + if (!isInstanceRuntimeCurrent(instanceId, instance)) throw new Error("Instance no longer active") if (selectedAgent && isModelValid(instanceId, defaultModel)) { await setAgentModelPreference(instanceId, selectedAgent, defaultModel) + if (!isInstanceRuntimeCurrent(instanceId, instance)) throw new Error("Instance no longer active") } setLoading((prev) => { @@ -680,6 +910,7 @@ async function createSession(instanceId: string, agent?: string): Promise { - const next = { ...prev } - next.creatingSession.set(instanceId, false) - return next - }) + if (isInstanceRuntimeCurrent(instanceId, instance)) { + setLoading((prev) => { + const next = { ...prev } + next.creatingSession.set(instanceId, false) + return next + }) + } } } @@ -789,12 +1023,14 @@ async function forkSession( ...(await getSessionWorkspacePayload(instanceId, sourceSessionId)), messageID: options?.messageId, } + if (!isInstanceRuntimeCurrent(instanceId, instance)) throw new Error("Instance no longer active") log.info(`[HTTP] POST /session.fork for instance ${instanceId}`, request) const info = await requestData( client.session.fork(request), "session.fork", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) throw new Error("Instance no longer active") const forkedSession = { id: info.id, instanceId, @@ -881,7 +1117,10 @@ async function deleteSession(instanceId: string, sessionId: string): Promise { - const next = { ...prev } - const deleting = next.deletingSession.get(instanceId) - if (deleting) { - deleting.delete(sessionId) - } - return next - }) + if (isInstanceRuntimeCurrent(instanceId, instance)) { + setLoading((prev) => { + const next = { ...prev } + const deleting = next.deletingSession.get(instanceId) + if (deleting) deleting.delete(sessionId) + return next + }) + } } } function removeSessionRuntimeState(instanceId: string, sessionId: string): void { + clearPendingDeltasForSession(instanceId, sessionId) sessionWorkspaceHints.get(instanceId)?.delete(sessionId) cancelSessionGenerationAdmissions(instanceId, sessionId) markSessionDeletedAuthoritative(instanceId, sessionId) deleteSessionAttachments(instanceId, sessionId) clearSessionDraftPrompt(instanceId, sessionId) setSessionExpanded(instanceId, sessionId, false) + clearSessionInterruptions(instanceId, sessionId) setSessions((prev) => { const next = new Map(prev) @@ -959,7 +1200,7 @@ function removeSessionRuntimeState(instanceId: string, sessionId: string): void } } -async function fetchAgents(instanceId: string): Promise { +async function fetchAgents(instanceId: string, signal?: AbortSignal, hasCommitAuthority: () => boolean = () => true): Promise { const instance = instances().get(instanceId) if (!instance || !instance.client) { throw new Error("Instance not ready") @@ -969,7 +1210,8 @@ async function fetchAgents(instanceId: string): Promise { try { log.info(`[HTTP] GET /app.agents for instance ${instanceId}`) - const response = await rootClient.app.agents() + const response = await rootClient.app.agents(undefined, signal ? { signal } : undefined) + if (signal?.aborted || !hasCommitAuthority() || !isInstanceRuntimeCurrent(instanceId, instance)) return const agentList = (response.data ?? []).map((agent) => ({ name: agent.name, description: agent.description || "", @@ -993,7 +1235,7 @@ async function fetchAgents(instanceId: string): Promise { } } -async function fetchProviders(instanceId: string): Promise { +async function fetchProviders(instanceId: string, signal?: AbortSignal, hasCommitAuthority: () => boolean = () => true): Promise { const instance = instances().get(instanceId) if (!instance || !instance.client) { throw new Error("Instance not ready") @@ -1003,7 +1245,8 @@ async function fetchProviders(instanceId: string): Promise { try { log.info(`[HTTP] GET /config.providers for instance ${instanceId}`) - const response = await rootClient.config.providers() + const response = await rootClient.config.providers(undefined, signal ? { signal } : undefined) + if (signal?.aborted || !hasCommitAuthority() || !isInstanceRuntimeCurrent(instanceId, instance)) return if (!response.data) return const providerList = response.data.providers.map((provider) => ({ @@ -1033,11 +1276,10 @@ async function fetchProviders(instanceId: string): Promise { async function loadMessages( instanceId: string, sessionId: string, - options?: { force?: boolean; skipChildren?: boolean }, + options?: { force?: boolean; timeoutMs?: number; applySessionRevert?: boolean; signal?: AbortSignal }, ): Promise { + options?.signal?.throwIfAborted() const force = options?.force ?? false - const skipChildren = options?.skipChildren ?? false - if (force) { setMessagesLoaded((prev) => { const next = new Map(prev) @@ -1071,15 +1313,33 @@ async function loadMessages( const client = getRootClient(instanceId) - const instanceSessions = sessions().get(instanceId) - const session = instanceSessions?.get(sessionId) + let session = sessions().get(instanceId)?.get(sessionId) if (!session) { - throw new Error("Session not found") + await hydrateRestoredSessionChain(instanceId, [sessionId], options?.signal, { hydrateKnownMetadata: false }) + options?.signal?.throwIfAborted() + session = sessions().get(instanceId)?.get(sessionId) + if (!session) throw new Error("Session not found") } - const loadEpoch = advanceMessageLoadEpoch(instanceId, sessionId) - const messageRevision = messageStoreBus.getOrCreate(instanceId).getSessionRevision(sessionId) + const { epoch: loadEpoch, signal: loadSignal, abort: abortLoad } = beginSessionMessageLoad(instanceId, sessionId) + const abortFromCaller = () => abortLoad(options?.signal?.reason) + options?.signal?.addEventListener("abort", abortFromCaller, { once: true }) + if (options?.signal?.aborted) abortFromCaller() + let loadTimeout: ReturnType | undefined + const loadTimeoutPromise = options?.timeoutMs + ? new Promise((_resolve, reject) => { + loadTimeout = setTimeout(() => { + const error = new SessionMessageLoadTimeoutError(options.timeoutMs!) + abortLoad(error) + reject(error) + }, options.timeoutMs) + }) + : undefined + const awaitLoad = (promise: Promise): Promise => loadTimeoutPromise ? Promise.race([promise, loadTimeoutPromise]) : promise + const store = messageStoreBus.getOrCreate(instanceId) + const expectedRevision = store.getSessionRevision(sessionId) let retryAfterRevisionConflict = false + const sessionForV2 = session setLoading((prev) => { const next = { ...prev } @@ -1092,12 +1352,19 @@ async function loadMessages( try { log.info(`[HTTP] GET /session.${"messages"} for instance ${instanceId}`, { sessionId }) - const apiMessages = await requestData( - client.session.messages({ sessionID: sessionId, ...(await getSessionWorkspacePayload(instanceId, sessionId)) }), - "session.messages", - ) + let apiMessages: any[] + try { + const workspacePayload = await awaitLoad(getSessionWorkspacePayload(instanceId, sessionId)) + apiMessages = await awaitLoad(requestData( + client.session.messages({ sessionID: sessionId, ...workspacePayload }, { signal: loadSignal }), + "session.messages", + )) + } catch (error) { + if (!isInstanceRuntimeCurrent(instanceId, instance) || !isCurrentMessageLoad(instanceId, sessionId, loadEpoch) || !sessions().get(instanceId)?.has(sessionId)) return + throw error + } - if (!isCurrentMessageLoad(instanceId, sessionId, loadEpoch) || !sessions().get(instanceId)?.has(sessionId)) return + if (!isInstanceRuntimeCurrent(instanceId, instance) || !isCurrentMessageLoad(instanceId, sessionId, loadEpoch) || !sessions().get(instanceId)?.has(sessionId)) return if (!Array.isArray(apiMessages)) { return @@ -1105,45 +1372,13 @@ async function loadMessages( setSessionMessagesLoadError(instanceId, sessionId, null) - if (apiMessages.length === 0) { - if (messageStoreBus.getOrCreate(instanceId).getSessionRevision(sessionId) !== messageRevision) { - retryAfterRevisionConflict = true - } else { - setMessagesLoaded((prev) => { - const next = new Map(prev) - const loadedSet = next.get(instanceId) || new Set() - loadedSet.add(sessionId) - next.set(instanceId, loadedSet) - return next - }) - } - } else { - const messagesInfo = new Map() - const messages: Message[] = apiMessages.map((apiMessage: any) => { - const info = apiMessage.info || apiMessage - const role = info.role || "assistant" - const messageId = info.id || String(Date.now()) - - messagesInfo.set(messageId, info) - - const parts: any[] = (apiMessage.parts || []).map((part: any) => normalizeMessagePart(part)) - - const message: Message = { - id: messageId, - sessionId, - type: role === "user" ? "user" : "assistant", - parts, - timestamp: info.time?.created || Date.now(), - status: "complete" as const, - version: 0, - } - - return message - }) + const latestStatus = sessions().get(instanceId)?.get(sessionId)?.status ?? sessionForV2.status + const adapted = adaptApiMessages(sessionId, apiMessages, latestStatus) - let agentName = "" - let providerID = "" - let modelID = "" + let agentName = "" + let providerID = "" + let modelID = "" + if (apiMessages.length > 0) { for (let i = apiMessages.length - 1; i >= 0; i--) { const apiMessage = apiMessages[i] @@ -1158,58 +1393,119 @@ async function loadMessages( } if (!agentName && !providerID && !modelID) { - const defaultModel = await getDefaultModel(instanceId, session.agent) - if (!isCurrentMessageLoad(instanceId, sessionId, loadEpoch) || !sessions().get(instanceId)?.has(sessionId)) return + const defaultModel = await awaitLoad(getDefaultModel(instanceId, session.agent)) + if (!isInstanceRuntimeCurrent(instanceId, instance) || !isCurrentMessageLoad(instanceId, sessionId, loadEpoch) || !sessions().get(instanceId)?.has(sessionId)) return agentName = session.agent providerID = defaultModel.providerId modelID = defaultModel.modelId } - setSessions((prev) => { - if (!isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) return prev - const next = new Map(prev) - const nextInstanceSessions = next.get(instanceId) - if (!nextInstanceSessions) return next - const existingSession = nextInstanceSessions.get(sessionId) - if (!existingSession) return next - nextInstanceSessions.set(sessionId, { - ...existingSession, - agent: agentName || existingSession.agent, - model: providerID && modelID ? { providerId: providerID, modelId: modelID } : existingSession.model, - }) - next.set(instanceId, nextInstanceSessions) - return next - }) + } - const sessionForV2 = sessions().get(instanceId)?.get(sessionId) ?? { - id: sessionId, title: session?.title, parentId: session?.parentId ?? null, revert: session?.revert, + const latestSession = sessions().get(instanceId)?.get(sessionId) ?? sessionForV2 + const applySessionRevert = options?.applySessionRevert !== false + if (!isInstanceRuntimeCurrent(instanceId, instance) || !isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) return + const snapshotFenceKey = `${instanceId}:${sessionId}` + const residentMessageIds = new Set(store.getSessionMessageIds(sessionId)) + const snapshotFence = bufferedDeltaSnapshotFences.get(snapshotFenceKey) ?? new Map() + for (const messageId of snapshotFence.keys()) if (!residentMessageIds.has(messageId)) snapshotFence.delete(messageId) + for (const messageId of residentMessageIds) { + const pendingDeltas = getPendingDeltasForMessage(instanceId, messageId) + if (pendingDeltas.length === 0) continue + const record = store.getMessage(messageId) + snapshotFence.set(messageId, pendingDeltas.map(({ partId, field, delta }) => { + const current = (record?.parts[partId]?.data as any)?.[field] + return { + partId, + field, + value: `${current ?? ""}${delta}`, + staleSnapshotsRemaining: BUFFERED_DELTA_STALE_SNAPSHOT_LIMIT, + } + })) + } + if (snapshotFence.size > 0) bufferedDeltaSnapshotFences.set(snapshotFenceKey, snapshotFence) + else bufferedDeltaSnapshotFences.delete(snapshotFenceKey) + + const incomingMessages = new Map(adapted.messages.map((message) => [message.id, message])) + let snapshotFenced = adapted.messages.some((message) => hasPendingDeltasForMessage(instanceId, message.id)) + for (const [messageId, expectations] of snapshotFence) { + if (hasPendingDeltasForMessage(instanceId, messageId)) { + snapshotFenced = true + continue } - if (!isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) return - if (!seedSessionMessagesV2(instanceId, sessionForV2, messages, messagesInfo, messageRevision)) { - retryAfterRevisionConflict = true + const incoming = incomingMessages.get(messageId) + const mismatches = expectations.filter(({ partId, field, value }) => { + if (!incoming) return true + const part = incoming.parts.find((candidate: any) => candidate.id === partId) as any + const incomingValue = part?.[field] + return typeof incomingValue !== "string" || !incomingValue.startsWith(value) + }) + if (mismatches.length === 0) { + snapshotFence.delete(messageId) + } else if (mismatches.some((expectation) => expectation.staleSnapshotsRemaining > 0)) { + for (const expectation of mismatches) expectation.staleSnapshotsRemaining = Math.max(0, expectation.staleSnapshotsRemaining - 1) + snapshotFenced = true } else { - setMessagesLoaded((prev) => { + snapshotFence.delete(messageId) + } + } + if (snapshotFence.size > 0) bufferedDeltaSnapshotFences.set(snapshotFenceKey, snapshotFence) + else bufferedDeltaSnapshotFences.delete(snapshotFenceKey) + if (snapshotFenced) { + retryAfterRevisionConflict = true + } else if (!seedSessionMessagesV2( + instanceId, + applySessionRevert ? latestSession : { id: latestSession.id, title: latestSession.title, parentId: latestSession.parentId }, + adapted.messages, + adapted.infos, + expectedRevision, + )) { + retryAfterRevisionConflict = true + } else { + bufferedDeltaSnapshotFences.delete(snapshotFenceKey) + if (apiMessages.length > 0) { + setSessions((prev) => { + if (!isInstanceRuntimeCurrent(instanceId, instance) || !isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) return prev const next = new Map(prev) - const loadedSet = next.get(instanceId) || new Set() - loadedSet.add(sessionId) - next.set(instanceId, loadedSet) + const nextInstanceSessions = next.get(instanceId) + if (!nextInstanceSessions) return next + const existingSession = nextInstanceSessions.get(sessionId) + if (!existingSession) return next + nextInstanceSessions.set(sessionId, { + ...existingSession, + agent: agentName || existingSession.agent, + model: providerID && modelID ? { providerId: providerID, modelId: modelID } : existingSession.model, + }) + next.set(instanceId, nextInstanceSessions) return next }) - reconcilePendingPermissionsV2(instanceId, sessionId) - reconcilePendingQuestionsV2(instanceId, sessionId) } + if (applySessionRevert) setSessionRevertV2(instanceId, sessionId, latestSession.revert ?? null) + setMessagesLoaded((prev) => { + const next = new Map(prev) + const loadedSet = next.get(instanceId) || new Set() + loadedSet.add(sessionId) + next.set(instanceId, loadedSet) + return next + }) + reconcilePendingPermissionsV2(instanceId, sessionId) + reconcilePendingQuestionsV2(instanceId, sessionId) } } catch (error) { + if (options?.signal?.aborted) return log.error("Failed to load messages:", error) - if (isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) { + if (isInstanceRuntimeCurrent(instanceId, instance) && isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) { setSessionMessagesLoadError(instanceId, sessionId, getOpencodeErrorMessage(error, tGlobal("messageSection.loadError.detail"))) } throw error } finally { - if (isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) { + options?.signal?.removeEventListener("abort", abortFromCaller) + if (loadTimeout) clearTimeout(loadTimeout) + finishSessionMessageLoad(instanceId, sessionId, loadEpoch) + if (isInstanceRuntimeCurrent(instanceId, instance) && isCurrentMessageLoad(instanceId, sessionId, loadEpoch)) { setLoading((prev) => { const next = { ...prev } const loadingSet = next.loadingMessages.get(instanceId) @@ -1220,28 +1516,22 @@ async function loadMessages( } if (retryAfterRevisionConflict && sessions().get(instanceId)?.has(sessionId)) { - await new Promise((resolve) => setTimeout(resolve, 50)) - return loadMessages(instanceId, sessionId, { force: true, skipChildren }) + setMessagesLoaded((prev) => { + const next = new Map(prev) + next.get(instanceId)?.delete(sessionId) + return next + }) + requestDeltaRecovery({ instanceId, sessionId, messageId: "reconcile", partId: "reconcile", field: "text" }) + return } - if (!isCurrentMessageLoad(instanceId, sessionId, loadEpoch) || !sessions().get(instanceId)?.has(sessionId)) return + if (!isInstanceRuntimeCurrent(instanceId, instance) || !isCurrentMessageLoad(instanceId, sessionId, loadEpoch) || !sessions().get(instanceId)?.has(sessionId)) return updateSessionInfo(instanceId, sessionId) - if (!skipChildren && session.parentId === null) { - for (const child of getDescendantSessions(instanceId, sessionId)) { - void loadMessages(instanceId, child.id, { skipChildren: true }).catch((error) => - log.error("Failed to load child session messages", { - instanceId, - sessionId: child.id, - parentSessionId: sessionId, - error, - }), - ) - } - } } export { + clearBufferedDeltaSnapshotFence, createSession, deleteSession, removeSessionRuntimeState, @@ -1251,8 +1541,10 @@ export { fetchSessions, hydrateRestoredSessionChain, loadMoreSessions, + clearSessionSearch, searchSessions, forkSession, loadMessages, + SessionMessageLoadTimeoutError, clearSessionListRequestState, } diff --git a/packages/ui/src/stores/session-events.ts b/packages/ui/src/stores/session-events.ts index 427579d04..efd91133f 100644 --- a/packages/ui/src/stores/session-events.ts +++ b/packages/ui/src/stores/session-events.ts @@ -14,14 +14,20 @@ import type { EventSessionStatus, } from "@opencode-ai/sdk" import type { MessageStatus } from "./message-v2/types" +import type { Instance } from "../types/instance" import { getLogger } from "../lib/logger" import type { EventSessionDeleted } from "../lib/sse-manager" import { requestData } from "../lib/opencode-api" import { enqueueDelta, + clearPendingDeltasForInstance, + clearPendingDeltasForMessage, clearPendingDeltasForPart, + clearPendingDeltasForSession, flushPendingDeltasForMessage, + requestDeltaRecovery, + setRecoveryCallback, setFlushCallback, } from "./delta-buffer" import { @@ -45,12 +51,15 @@ import { sendOsNotification } from "../lib/os-notifications" import { preferences } from "./preferences" import { instances, + isInstanceRuntimeCurrent, addPermissionToQueue, getPermissionQueue, removePermissionFromQueue, markPermissionReplied, hasRepliedPermission, addQuestionToQueue, + hasAnsweredQuestion, + markQuestionAnswered, removeQuestionFromQueue, } from "./instances" import { showAlertDialog } from "./alerts" @@ -63,13 +72,13 @@ import { type SessionRetryState, type SessionStatus, } from "../types/session" -import { ensureSessionAncestorsExpanded, getAuthoritativelyDeletedSessionIdsForInstance, prependSessionListId, sessions, setSessionStatus, setSessions, syncInstanceSessionIndicator, withSession } from "./session-state" +import { activeSessionId, ensureSessionAncestorsExpanded, getAuthoritativelyDeletedSessionIdsForInstance, invalidateSessionMessageLoad, markSessionMetadataMutation, prependSessionListId, sessions, setSessionStatus, setSessions, syncInstanceSessionIndicator, withSession } from "./session-state" import { mergeFetchedSessionRuntimeState } from "./session-generation-recovery" import { normalizeMessagePart } from "./message-v2/normalizers" import { updateSessionInfo } from "./message-v2/session-info" import { tGlobal } from "../lib/i18n" -import { loadMessages, removeSessionRuntimeState } from "./session-api" +import { clearBufferedDeltaSnapshotFence, loadMessages, removeSessionRuntimeState, SessionMessageLoadTimeoutError } from "./session-api" import { getRootClient } from "./opencode-client" import { getWorktreeSlugForDirectory, getWorktreeSlugForSession } from "./worktrees" import { getOpenCodeWorkspaceIdForWorktree } from "./opencode-workspaces" @@ -91,11 +100,44 @@ import { import { messageStoreBus } from "./message-v2/bus" import type { InstanceMessageStore } from "./message-v2/instance-store" import { handleConversationAssistantPartUpdated } from "./conversation-speech" +import { scheduleSessionMemorySweep } from "./session-memory" const log = getLogger("sse") const pendingSessionFetches = new Map>() +const pendingSessionFetchRuntimes = new Map() +const pendingSessionStatuses = new Map() +const MAX_DELTA_RECOVERY_ATTEMPTS = 3 +const DELTA_RECOVERY_BACKOFF_MS = 100 +const DELTA_RECOVERY_LOAD_TIMEOUT_MS = 10_000 +type DeltaRecovery = { + dirty: boolean + running: boolean + attempts: number + deltas: Map + requireActive: boolean + idleReconcileRequested: boolean + instance: Instance +} +const pendingDeltaRecoveries = new Map() +const pendingIdleMessageReconciliations = new Map() let activeRetryToast: ToastHandle | null = null +messageStoreBus.onInstanceDestroyed((instanceId) => { + const prefix = `${instanceId}:` + for (const key of pendingSessionFetches.keys()) if (key.startsWith(prefix)) pendingSessionFetches.delete(key) + for (const key of pendingSessionFetchRuntimes.keys()) if (key.startsWith(prefix)) pendingSessionFetchRuntimes.delete(key) + for (const key of pendingSessionStatuses.keys()) if (key.startsWith(prefix)) pendingSessionStatuses.delete(key) + for (const key of pendingDeltaRecoveries.keys()) if (key.startsWith(prefix)) pendingDeltaRecoveries.delete(key) + for (const key of pendingIdleMessageReconciliations.keys()) if (key.startsWith(prefix)) pendingIdleMessageReconciliations.delete(key) + clearPendingDeltasForInstance(instanceId) +}) + +messageStoreBus.onSessionCleared((instanceId, sessionId) => { + const key = `${instanceId}:${sessionId}` + pendingDeltaRecoveries.delete(key) + pendingIdleMessageReconciliations.delete(key) +}) + function shouldSendOsNotification(kind: "needsInput" | "idle"): boolean { if (typeof document === "undefined") return false const pref = preferences() @@ -167,22 +209,26 @@ async function fetchSessionInfo(instanceId: string, sessionId: string, directory const slug = slugFromDirectory ?? getWorktreeSlugForSession(instanceId, sessionId) const client = getRootClient(instanceId) const workspace = await getOpenCodeWorkspaceIdForWorktree(instanceId, slug) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return null try { const info = await requestData( client.session.get({ sessionID: sessionId, ...(workspace ? { workspace } : {}) }), "session.get", ) + if (!isInstanceRuntimeCurrent(instanceId, instance)) return null let rawStatus = (info as any)?.status let fetchedStatusKnown = false try { const statuses = await requestData>(client.session.status(), "session.status") + if (!isInstanceRuntimeCurrent(instanceId, instance)) return null rawStatus ??= statuses?.[sessionId] fetchedStatusKnown = true } catch (error) { log.error("Failed to fetch session status", error) } + if (!isInstanceRuntimeCurrent(instanceId, instance)) return null const hasStatus = rawStatus && typeof rawStatus === "object" && typeof rawStatus.type === "string" fetchedStatusKnown ||= Boolean(hasStatus) const fetchedStatus: SessionStatus = hasStatus ? mapSdkSessionStatus(rawStatus) : "idle" @@ -248,21 +294,36 @@ function ensureSessionStatus( ) { const existing = sessions().get(instanceId)?.get(sessionId) if (existing) { + const key = `${instanceId}:${sessionId}` + if (pendingSessionFetches.has(key)) pendingSessionStatuses.set(key, { status, retry }) setSessionStatus(instanceId, sessionId, status, { retry }) + scheduleSessionMemorySweep() return } const key = `${instanceId}:${sessionId}` - if (pendingSessionFetches.has(key)) return + pendingSessionStatuses.set(key, { status, retry }) + const runtime = instances().get(instanceId) + const pendingRuntime = pendingSessionFetchRuntimes.get(key) + if (pendingSessionFetches.has(key) && pendingRuntime && isInstanceRuntimeCurrent(instanceId, pendingRuntime)) return const pending = (async () => { const fetched = await fetchSessionInfo(instanceId, sessionId, directory) if (!fetched) return - setSessionStatus(instanceId, sessionId, status, { retry }) + const latest = pendingSessionStatuses.get(key) ?? { status, retry } + setSessionStatus(instanceId, sessionId, latest.status, { retry: latest.retry, force: true }) + scheduleSessionMemorySweep() })() pendingSessionFetches.set(key, pending) - void pending.finally(() => pendingSessionFetches.delete(key)) + if (runtime) pendingSessionFetchRuntimes.set(key, runtime) + void pending.finally(() => { + if (pendingSessionFetches.get(key) === pending) { + pendingSessionFetches.delete(key) + pendingSessionFetchRuntimes.delete(key) + pendingSessionStatuses.delete(key) + } + }) } type MessageRole = "user" | "assistant" @@ -306,6 +367,7 @@ function handleMessageUpdate(instanceId: string, event: MessageUpdateEvent | Mes const sessionId = typeof part.sessionID === "string" ? part.sessionID : fallbackSessionId const messageId = typeof part.messageID === "string" ? part.messageID : fallbackMessageId if (!sessionId || !messageId) return + if (getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(sessionId)) return if (part.type === "compaction") { ensureSessionStatus(instanceId, sessionId, "compacting", (event as any)?.directory) } @@ -345,6 +407,7 @@ function handleMessageUpdate(instanceId: string, event: MessageUpdateEvent | Mes // deltas would be stale and cause duplication if flushed later. if (part.id) { clearPendingDeltasForPart(instanceId, messageId, part.id) + clearBufferedDeltaSnapshotFence(instanceId, sessionId, messageId, part.id) } applyPartUpdateV2(instanceId, { ...part, sessionID: sessionId, messageID: messageId }) handleConversationAssistantPartUpdated(instanceId, { ...part, sessionID: sessionId, messageID: messageId }, messageInfo) @@ -363,6 +426,7 @@ function handleMessageUpdate(instanceId: string, event: MessageUpdateEvent | Mes const sessionId = typeof info.sessionID === "string" ? info.sessionID : undefined const messageId = typeof info.id === "string" ? info.id : undefined if (!sessionId || !messageId) return + if (getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(sessionId)) return // Flush any pending deltas for this message before applying the update. // Deltas are buffered for up to 50ms; if message.updated arrives before @@ -372,15 +436,17 @@ function handleMessageUpdate(instanceId: string, event: MessageUpdateEvent | Mes // message status/metadata update runs on the complete content. flushPendingDeltasForMessage(instanceId, messageId, applyPartDeltaV2) - const timeInfo = (info.time ?? {}) as { created?: number; updated?: number; end?: number } + const timeInfo = (info.time ?? {}) as { created?: number; updated?: number; end?: number; completed?: number } const nextUpdated = typeof timeInfo.end === "number" && timeInfo.end > 0 ? timeInfo.end - : typeof timeInfo.updated === "number" && timeInfo.updated > 0 - ? timeInfo.updated - : typeof timeInfo.created === "number" && timeInfo.created > 0 - ? timeInfo.created - : Date.now() + : typeof timeInfo.completed === "number" && timeInfo.completed > 0 + ? timeInfo.completed + : typeof timeInfo.updated === "number" && timeInfo.updated > 0 + ? timeInfo.updated + : typeof timeInfo.created === "number" && timeInfo.created > 0 + ? timeInfo.created + : Date.now() withSession(instanceId, sessionId, (session) => { const currentUpdated = session.time?.updated ?? 0 @@ -392,7 +458,9 @@ function handleMessageUpdate(instanceId: string, event: MessageUpdateEvent | Mes const role: MessageRole = info.role === "user" ? "user" : "assistant" const hasError = Boolean((info as any).error) - const hasEnded = typeof timeInfo.end === "number" && timeInfo.end > 0 + const hasEnded = + (typeof timeInfo.end === "number" && timeInfo.end > 0) || + (typeof timeInfo.completed === "number" && timeInfo.completed > 0) const status: MessageStatus = hasError ? "error" : hasEnded ? "complete" : "streaming" let record = store.getMessage(messageId) @@ -406,7 +474,7 @@ function handleMessageUpdate(instanceId: string, event: MessageUpdateEvent | Mes if (!record) { const createdAt = info.time?.created ?? Date.now() - const endAt = (info.time as { end?: number } | undefined)?.end + const endAt = timeInfo.end ?? timeInfo.completed store.upsertMessage({ id: messageId, sessionId, @@ -425,17 +493,137 @@ function handleMessageUpdate(instanceId: string, event: MessageUpdateEvent | Mes // Delta buffer callback setup setFlushCallback((batch) => { - for (const { instanceId, messageId, partId, field, delta } of batch) { - applyPartDeltaV2(instanceId, { messageId, partId, field, delta }) + for (const { instanceId, sessionId, messageId, partId, field, delta } of batch) { + if (!applyPartDeltaV2(instanceId, { messageId, partId, field, delta })) { + requestDeltaRecovery({ instanceId, ...(sessionId ? { sessionId } : {}), messageId, partId, field, delta }) + } + } +}) + +setRecoveryCallback(({ instanceId, sessionId, messageId, partId, field, delta }) => { + if (!sessionId) { + log.warn("Dropped orphan delta without a session", { instanceId, messageId, partId }) + return + } + const instance = instances().get(instanceId) + if (!instance) return + const key = `${instanceId}:${sessionId}` + const existing = pendingDeltaRecoveries.get(key) + if (existing && isInstanceRuntimeCurrent(instanceId, existing.instance)) { + if (delta !== undefined) { + const deltaKey = `${messageId}:${partId}:${field}` + const pending = existing.deltas.get(deltaKey) + existing.deltas.set(deltaKey, { + messageId, + partId, + field, + delta: `${pending?.delta ?? ""}${delta}`, + expectedValue: pending?.expectedValue === undefined ? undefined : `${pending.expectedValue}${delta}`, + }) + } + existing.dirty = true + runDeltaRecovery(instanceId, sessionId, existing) + return + } + if (existing) pendingDeltaRecoveries.delete(key) + const recovery: DeltaRecovery = { + dirty: true, + running: false, + attempts: 0, + deltas: new Map(delta === undefined ? [] : [[`${messageId}:${partId}:${field}`, { messageId, partId, field, delta }]]), + requireActive: activeSessionId().get(instanceId) === sessionId, + idleReconcileRequested: false, + instance, } + pendingDeltaRecoveries.set(key, recovery) + runDeltaRecovery(instanceId, sessionId, recovery) }) +function runDeltaRecovery(instanceId: string, sessionId: string, recovery: DeltaRecovery): void { + if (recovery.running) return + const key = `${instanceId}:${sessionId}` + if (!isInstanceRuntimeCurrent(instanceId, recovery.instance)) { + if (pendingDeltaRecoveries.get(key) === recovery) pendingDeltaRecoveries.delete(key) + return + } + const session = sessions().get(instanceId)?.get(sessionId) + if (session?.status === "working" || session?.status === "compacting") return + recovery.running = true + void (async () => { + try { + while (recovery.dirty && recovery.attempts < MAX_DELTA_RECOVERY_ATTEMPTS) { + if (!isInstanceRuntimeCurrent(instanceId, recovery.instance) || !sessions().get(instanceId)?.has(sessionId)) return + if (recovery.requireActive && activeSessionId().get(instanceId) !== sessionId) { + invalidateSessionMessageLoad(instanceId, sessionId) + return + } + const currentStatus = sessions().get(instanceId)?.get(sessionId)?.status + if (currentStatus === "working" || currentStatus === "compacting") return + + await new Promise((resolve) => setTimeout(resolve, (recovery.attempts + 1) * DELTA_RECOVERY_BACKOFF_MS)) + if (!isInstanceRuntimeCurrent(instanceId, recovery.instance) || !sessions().get(instanceId)?.has(sessionId)) return + if (recovery.requireActive && activeSessionId().get(instanceId) !== sessionId) { + invalidateSessionMessageLoad(instanceId, sessionId) + return + } + const statusAfterBackoff = sessions().get(instanceId)?.get(sessionId)?.status + if (statusAfterBackoff === "working" || statusAfterBackoff === "compacting") return + + recovery.dirty = false + recovery.attempts += 1 + try { + await loadMessages(instanceId, sessionId, { force: true, timeoutMs: DELTA_RECOVERY_LOAD_TIMEOUT_MS }) + if (!isInstanceRuntimeCurrent(instanceId, recovery.instance)) return + for (const [deltaKey, pending] of recovery.deltas) { + const part = messageStoreBus.getInstance(instanceId)?.getMessage(pending.messageId)?.parts[pending.partId]?.data as any + const value = part?.[pending.field] + if (pending.expectedValue !== undefined && typeof value === "string" && value.startsWith(pending.expectedValue)) { + recovery.deltas.delete(deltaKey) + } else if (recovery.attempts < MAX_DELTA_RECOVERY_ATTEMPTS) { + if (pending.expectedValue === undefined && typeof value === "string") { + pending.expectedValue = `${value}${pending.delta}` + } + applyPartDeltaV2(instanceId, pending) + recovery.dirty = true + } else { + recovery.deltas.delete(deltaKey) + } + } + } catch (error) { + recovery.dirty = true + if (recovery.attempts >= MAX_DELTA_RECOVERY_ATTEMPTS) throw error + } + } + if (recovery.dirty) { + log.warn("Revision recovery exhausted", { instanceId, sessionId, attempts: recovery.attempts }) + } + } catch (error) { + log.warn("Failed to recover orphan delta", { instanceId, sessionId, error }) + } finally { + recovery.running = false + const latestStatus = sessions().get(instanceId)?.get(sessionId)?.status + const waitingForIdle = isInstanceRuntimeCurrent(instanceId, recovery.instance) && + recovery.attempts < MAX_DELTA_RECOVERY_ATTEMPTS && recovery.dirty && + (latestStatus === "working" || latestStatus === "compacting") + const forceIdleReconcile = isInstanceRuntimeCurrent(instanceId, recovery.instance) && + (!recovery.requireActive || activeSessionId().get(instanceId) === sessionId) && + recovery.idleReconcileRequested && !waitingForIdle + if (!waitingForIdle && pendingDeltaRecoveries.get(key) === recovery) pendingDeltaRecoveries.delete(key) + if (forceIdleReconcile) reconcileIdleSessionMessages(instanceId, sessionId, true, recovery.instance) + } + })() +} + function handleMessagePartDelta(instanceId: string, event: MessagePartDeltaEvent): void { const props = event.properties if (!props) return const { messageID, partID, field, delta } = props if (!messageID || !partID || !field || typeof delta !== "string") return - enqueueDelta(instanceId, messageID, partID, field, delta) + const sessionId = props.sessionID ?? messageStoreBus.getInstance(instanceId)?.getMessage(messageID)?.sessionId + if (sessionId) { + if (getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(sessionId)) return + } + enqueueDelta(instanceId, messageID, partID, field, delta, sessionId) } function handleSessionUpdate(instanceId: string, event: EventSessionUpdated): void { @@ -443,10 +631,20 @@ function handleSessionUpdate(instanceId: string, event: EventSessionUpdated): vo if (!info) return if (getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(info.id)) return - + markSessionMetadataMutation(instanceId, info.id) const instanceSessions = sessions().get(instanceId) ?? new Map() - const existingSession = instanceSessions.get(info.id) + const incomingRevert = info.revert ?? null + const previousRevert = existingSession?.revert ?? null + const revertChanged = ( + incomingRevert?.messageID !== previousRevert?.messageID || + incomingRevert?.partID !== previousRevert?.partID || + incomingRevert?.snapshot !== previousRevert?.snapshot || + incomingRevert?.diff !== previousRevert?.diff + ) + if (revertChanged) { + invalidateSessionMessageLoad(instanceId, info.id) + } if (!existingSession) { const newSession = { @@ -492,7 +690,7 @@ function handleSessionUpdate(instanceId: string, event: EventSessionUpdated): vo }) syncInstanceSessionIndicator(instanceId, updatedInstanceSessions) - setSessionRevertV2(instanceId, info.id, info.revert ?? null) + setSessionRevertV2(instanceId, info.id, incomingRevert) if (!newSession.parentId) { prependSessionListId(instanceId, newSession.id) } @@ -518,7 +716,7 @@ function handleSessionUpdate(instanceId: string, event: EventSessionUpdated): vo snapshot: info.revert.snapshot, diff: info.revert.diff, } - : existingSession.revert, + : undefined, } let updatedInstanceSessions: Map | undefined @@ -533,7 +731,7 @@ function handleSessionUpdate(instanceId: string, event: EventSessionUpdated): vo }) syncInstanceSessionIndicator(instanceId, updatedInstanceSessions) - setSessionRevertV2(instanceId, info.id, info.revert ?? null) + setSessionRevertV2(instanceId, info.id, incomingRevert) } } @@ -541,6 +739,7 @@ function handleSessionDeleted(instanceId: string, event: EventSessionDeleted): v const properties = event.properties const sessionId = properties?.info?.id ?? properties?.sessionID ?? properties?.id if (!sessionId) return + clearPendingDeltasForSession(instanceId, sessionId) log.info(`[SSE] Session deleted: ${sessionId}`) removeSessionRuntimeState(instanceId, sessionId) @@ -558,6 +757,8 @@ function handleSessionIdle(instanceId: string, event: EventSessionIdle): void { } ensureSessionStatus(instanceId, sessionId, "idle", (event as any)?.directory) + reconcileIdleSessionMessages(instanceId, sessionId) + scheduleSessionMemorySweep() log.info(`[SSE] Session idle: ${sessionId}`) } @@ -569,6 +770,8 @@ function handleSessionStatus(instanceId: string, event: EventSessionStatus): voi const status = mapSdkSessionStatus(rawStatus) const retry = mapSdkSessionRetry(rawStatus) ensureSessionStatus(instanceId, sessionId, status, (event as any)?.directory, retry) + if (status === "idle") reconcileIdleSessionMessages(instanceId, sessionId) + scheduleSessionMemorySweep() if (retry) { const remainingSeconds = Math.max(0, Math.round((retry.next - Date.now()) / 1000)) const countdown = @@ -591,6 +794,41 @@ function handleSessionStatus(instanceId: string, event: EventSessionStatus): voi log.info(`[SSE] Session status updated: ${sessionId}`, { status }) } +function reconcileIdleSessionMessages(instanceId: string, sessionId: string, force = false, expectedInstance?: Instance): void { + const instance = expectedInstance ?? instances().get(instanceId) + if (!instance || !isInstanceRuntimeCurrent(instanceId, instance)) return + const recoveryKey = `${instanceId}:${sessionId}` + const recovery = pendingDeltaRecoveries.get(recoveryKey) + let forceReconcile = false + if (recovery) { + if ((recovery.running || recovery.attempts < MAX_DELTA_RECOVERY_ATTEMPTS) && isInstanceRuntimeCurrent(instanceId, recovery.instance)) { + recovery.idleReconcileRequested = true + runDeltaRecovery(instanceId, sessionId, recovery) + return + } + forceReconcile = recovery.dirty && isInstanceRuntimeCurrent(instanceId, recovery.instance) + pendingDeltaRecoveries.delete(recoveryKey) + } + if (!force && !forceReconcile && !messageStoreBus.getInstance(instanceId)?.hasSessionActiveWork(sessionId)) return + const key = `${instanceId}:${sessionId}` + const pendingRuntime = pendingIdleMessageReconciliations.get(key) + if (pendingRuntime && isInstanceRuntimeCurrent(instanceId, pendingRuntime)) return + pendingIdleMessageReconciliations.set(key, instance) + void loadMessages(instanceId, sessionId, { + force: true, + timeoutMs: DELTA_RECOVERY_LOAD_TIMEOUT_MS, + }) + .catch((error) => { + if (error instanceof SessionMessageLoadTimeoutError) { + messageStoreBus.getInstance(instanceId)?.interruptSessionActiveMessages(sessionId) + } + log.warn("Failed to reconcile idle session messages", { instanceId, sessionId, error }) + }) + .finally(() => { + if (pendingIdleMessageReconciliations.get(key) === instance) pendingIdleMessageReconciliations.delete(key) + }) +} + function handleSessionCompacted(instanceId: string, event: EventSessionCompacted): void { const sessionID = event.properties?.sessionID if (!sessionID) return @@ -631,6 +869,7 @@ function handleSessionError(_instanceId: string, event: EventSessionError): void message = error.message } } + if (message.length > 10_000) message = `${message.slice(0, 10_000)}...` showAlertDialog(tGlobal("sessionEvents.sessionError.message", { message }), { title: tGlobal("sessionEvents.sessionError.title"), @@ -641,8 +880,10 @@ function handleSessionError(_instanceId: string, event: EventSessionError): void function handleMessageRemoved(instanceId: string, event: MessageRemovedEvent): void { const { sessionID, messageID } = event.properties if (!sessionID || !messageID) return + if (getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(sessionID)) return log.info(`[SSE] Message removed from session ${sessionID}`, { messageID }) + clearPendingDeltasForMessage(instanceId, messageID) removeMessageV2(instanceId, messageID, sessionID) updateSessionInfo(instanceId, sessionID) } @@ -650,8 +891,10 @@ function handleMessageRemoved(instanceId: string, event: MessageRemovedEvent): v function handleMessagePartRemoved(instanceId: string, event: MessagePartRemovedEvent): void { const { sessionID, messageID, partID } = event.properties if (!sessionID || !messageID || !partID) return + if (getAuthoritativelyDeletedSessionIdsForInstance(instanceId).has(sessionID)) return log.info(`[SSE] Message part removed from session ${sessionID}`, { messageID, partID }) + clearPendingDeltasForPart(instanceId, messageID, partID) removeMessagePartV2(instanceId, messageID, partID, sessionID) updateSessionInfo(instanceId, sessionID) } @@ -715,6 +958,7 @@ function handlePermissionReplied(instanceId: string, event: EventPermissionV2Rep function handleQuestionAsked(instanceId: string, event: EventQuestionV2Asked | LegacyQuestionAskedEvent): void { const request = event?.properties as QuestionRequest | undefined if (!request) return + if (hasAnsweredQuestion(instanceId, request.id)) return const source = event.type === "question.asked" ? "legacy" : "v2" log.info(`[SSE] Question asked: ${getQuestionId(request)}`) @@ -740,6 +984,7 @@ function handleQuestionAnswered( if (!requestId) return log.info(`[SSE] Question answered: ${requestId}`) + markQuestionAnswered(instanceId, requestId) removeQuestionFromQueue(instanceId, requestId) removeQuestionV2(instanceId, requestId) } diff --git a/packages/ui/src/stores/session-generation-recovery.test.ts b/packages/ui/src/stores/session-generation-recovery.test.ts index ec5e3988e..1cc3a9cf6 100644 --- a/packages/ui/src/stores/session-generation-recovery.test.ts +++ b/packages/ui/src/stores/session-generation-recovery.test.ts @@ -51,10 +51,16 @@ describe("session generation recovery", () => { }], ["active authority clears a captured admission token", { captured: session({ runtimeStatusKnown: false, generationRecovery: "pending", generationAdmissionToken: 1 }), - fetched: session({ status: "working", runtimeStatusKnown: true, generationRecovery: null }), + fetched: session({ status: "working", runtimeStatusKnown: true, generationRecovery: null, generationAdmissionToken: 1 }), latest: session({ runtimeStatusKnown: false, generationRecovery: "pending", generationAdmissionToken: undefined }), expected: { title: "Session", status: "working", runtimeStatusKnown: true, generationRecovery: null, token: undefined, source: undefined, updated: 1 }, }], + ["active fetch overrides an unresolved admission token", { + captured: session({ runtimeStatusKnown: false, generationRecovery: "pending", generationAdmissionToken: 1 }), + fetched: session({ status: "working", runtimeStatusKnown: true, generationRecovery: null, generationAdmissionToken: 1 }), + latest: null, + expected: { title: "Session", status: "working", runtimeStatusKnown: true, generationRecovery: null, token: undefined, source: undefined, updated: 1 }, + }], ["newer local state preserves optional field deletion", { captured: session({ retry: { attempt: 1, message: "retrying", next: 10 } }), fetched: session({ retry: { attempt: 2, message: "stale", next: 20 } }), @@ -74,4 +80,14 @@ describe("session generation recovery", () => { assert.equal(mergeFetchedSessionRuntimeState(fetched, session(), undefined), null) assert.equal(mergeFetchedSessionRuntimeState(fetched, undefined, undefined, true), null) }) + it("keeps fetched revert authority while preserving an admission token", () => { + const captured = session({ generationAdmissionToken: 1, revert: { messageID: "stale" } }) + const merged = mergeFetchedSessionRuntimeState( + session({ revert: { messageID: "authoritative" } }), + captured, + captured, + ) + assert.equal(merged?.generationAdmissionToken, 1) + assert.equal(merged?.revert?.messageID, "authoritative") + }) }) diff --git a/packages/ui/src/stores/session-generation-recovery.ts b/packages/ui/src/stores/session-generation-recovery.ts index e0a767ebe..3fe3e1425 100644 --- a/packages/ui/src/stores/session-generation-recovery.ts +++ b/packages/ui/src/stores/session-generation-recovery.ts @@ -37,11 +37,14 @@ export function mergeFetchedSessionRuntimeState( ): Session | null { if (deleted) return null if (captured && !latest) return null - if (!latest) return fetched + const fetchedActive = fetched.status === "working" || fetched.status === "compacting" + const authoritativeFetched = fetchedActive ? { ...fetched, generationAdmissionToken: undefined } : fetched + if (!latest) return authoritativeFetched if (latest === captured) { - return latest.generationAdmissionToken === undefined ? fetched : { ...fetched, ...latest } + if (fetchedActive) return authoritativeFetched + return latest.generationAdmissionToken === undefined ? fetched : { ...fetched, ...latest, revert: fetched.revert } } - const merged = { ...fetched } + const merged = { ...authoritativeFetched } const keys = new Set([ ...(Object.keys(captured ?? {}) as (keyof Session)[]), ...(Object.keys(latest) as (keyof Session)[]), @@ -52,11 +55,10 @@ export function mergeFetchedSessionRuntimeState( else delete (merged as any)[key] } - const fetchedActive = fetched.status === "working" || fetched.status === "compacting" if (captured && fetchedActive && latest.generationAdmissionToken === undefined && latest.runtimeStatusKnown === false && latest.generationRecovery === "pending") { for (const key of ["status", "runtimeStatusKnown", "generationRecovery", "generationAdmissionToken", "retry", "idleSince"] as const) { - (merged as any)[key] = fetched[key] + (merged as any)[key] = authoritativeFetched[key] } } return merged diff --git a/packages/ui/src/stores/session-memory.test.ts b/packages/ui/src/stores/session-memory.test.ts new file mode 100644 index 000000000..308aed0a1 --- /dev/null +++ b/packages/ui/src/stores/session-memory.test.ts @@ -0,0 +1,400 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { createRoot, createSignal } from "solid-js" +import { + getPromptDisplayOverride, + resetPromptDisplayOverrideStateForTests, + setPromptDisplayOverride, +} from "./message-prompt-display.ts" +import { messageStoreBus } from "./message-v2/bus.ts" +import { MAX_HOT_SESSION_MESSAGE_BYTES } from "../lib/session-memory-budget.ts" +import { evictResidentSessionMessages, getVisibleSessionMemoryIds, runSessionMemorySweep, setVisibleSessionMemory } from "./session-memory.ts" +import { setSessions } from "./session-state.ts" +import { useSessionCache } from "../components/instance/shell/useSessionCache.ts" + +function addMessage(instanceId: string, sessionId: string, status: "complete" | "streaming" = "complete", text = sessionId.repeat(100)) { + messageStoreBus.getOrCreate(instanceId).upsertMessage({ + id: `${sessionId}-message`, + sessionId, + role: "assistant", + status, + parts: [{ id: `${sessionId}-part`, type: "text", text }] as any, + }) +} + +const settleMeasurements = () => new Promise((resolve) => setTimeout(resolve, 75)) + +test("mounted cached views stay protected until they unmount", async () => { + const instanceId = "memory-six-session-cache" + const sessionIds = Array.from({ length: 6 }, (_, index) => `session-${index + 1}`) + const [activeSessionId, setActiveSessionId] = createSignal(sessionIds[0]!) + const [visible, setVisible] = createSignal(true) + let dispose!: () => void + try { + sessionIds.forEach((sessionId) => addMessage(instanceId, sessionId)) + const cache = createRoot((rootDispose) => { + dispose = rootDispose + return useSessionCache({ + instanceId: () => instanceId, + instanceSessions: () => new Map(sessionIds.map((sessionId) => [sessionId, {}])), + activeSessionId, + visible, + }) + }) + for (const sessionId of sessionIds.slice(1)) setActiveSessionId(sessionId) + await Promise.resolve() + + assert.deepEqual(cache.cachedSessionIds(), sessionIds.slice(1).reverse()) + assert.deepEqual(getVisibleSessionMemoryIds(instanceId), sessionIds.slice(1)) + assert.equal(evictResidentSessionMessages(instanceId, sessionIds[1]!), false) + assert.equal(evictResidentSessionMessages(instanceId, sessionIds[0]!), true) + + setVisible(false) + await Promise.resolve() + assert.deepEqual(cache.cachedSessionIds(), []) + assert.deepEqual(getVisibleSessionMemoryIds(instanceId), []) + await settleMeasurements() + assert.deepEqual( + runSessionMemorySweep(0).map((key) => key.split("\u0000")[1]), + sessionIds.slice(1), + ) + } finally { + dispose?.() + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("resident eviction purges prompt-display overrides from memory and persistence", () => { + const instanceId = "memory-prompt-display", sessionId = "session", messageId = `${sessionId}-message` + const entries = new Map() + const originalWindow = (globalThis as any).window + ;(globalThis as any).window = { localStorage: { + getItem: (key: string) => entries.get(key) ?? null, + setItem: (key: string, value: string) => entries.set(key, value), + removeItem: (key: string) => entries.delete(key), + } } + resetPromptDisplayOverrideStateForTests() + try { + addMessage(instanceId, sessionId) + setPromptDisplayOverride(instanceId, sessionId, messageId, { segments: [{ kind: "pasted", length: 100 }] }) + + assert.equal(evictResidentSessionMessages(instanceId, sessionId), true) + assert.equal(getPromptDisplayOverride(instanceId, sessionId, messageId), undefined) + assert.deepEqual(JSON.parse(entries.get("codenomad:prompt-display:v3") ?? "{}"), {}) + + resetPromptDisplayOverrideStateForTests() + assert.equal(getPromptDisplayOverride(instanceId, sessionId, messageId), undefined) + } finally { + messageStoreBus.unregisterInstance(instanceId) + resetPromptDisplayOverrideStateForTests() + if (originalWindow === undefined) delete (globalThis as any).window + else (globalThis as any).window = originalWindow + } +}) + +test("resident message budget evicts globally across five workspaces while preserving visible and streaming sessions", async () => { + const instanceIds = Array.from({ length: 5 }, (_, index) => `memory-workspace-${index}`) + try { + for (const instanceId of instanceIds) { + addMessage(instanceId, "parent") + addMessage(instanceId, "subagent") + } + addMessage(instanceIds[0], "streaming", "streaming") + setVisibleSessionMemory(instanceIds[4], "parent", true) + await settleMeasurements() + + const evicted = runSessionMemorySweep(0) + + assert.equal(evicted.length, 9) + assert.deepEqual(messageStoreBus.getInstance(instanceIds[4])!.getSessionMessageIds("parent"), ["parent-message"]) + assert.deepEqual(messageStoreBus.getInstance(instanceIds[0])!.getSessionMessageIds("streaming"), ["streaming-message"]) + assert.deepEqual(messageStoreBus.getInstance(instanceIds[0])!.getSessionMessageIds("parent"), []) + assert.deepEqual(messageStoreBus.getInstance(instanceIds[0])!.getSessionMessageIds("subagent"), []) + } finally { + setVisibleSessionMemory(instanceIds[4], "parent", false) + for (const instanceId of instanceIds) messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("visible session leases keep a child resident until every visible owner releases it", async () => { + const instanceId = "memory-visible-leases", sessionId = "child" + try { + addMessage(instanceId, sessionId) + setVisibleSessionMemory(instanceId, sessionId, true) + setVisibleSessionMemory(instanceId, sessionId, true) + setVisibleSessionMemory(instanceId, sessionId, false) + await settleMeasurements() + + assert.deepEqual(runSessionMemorySweep(0), []) + setVisibleSessionMemory(instanceId, sessionId, false) + assert.deepEqual(runSessionMemorySweep(0), [`${instanceId}\u0000${sessionId}`]) + } finally { + setVisibleSessionMemory(instanceId, sessionId, false) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("authoritative working status protects a resident session", async () => { + const instanceId = "memory-working-status", sessionId = "session" + try { + addMessage(instanceId, sessionId) + await settleMeasurements() + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { id: sessionId, status: "working" } as any]]))) + assert.deepEqual(runSessionMemorySweep(0), []) + + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { id: sessionId, status: "idle" } as any]]))) + assert.deepEqual(runSessionMemorySweep(0), [`${instanceId}\u0000${sessionId}`]) + } finally { + setSessions((prev) => { const next = new Map(prev); next.delete(instanceId); return next }) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("memory sweeps never synchronously measure resident transcripts", async () => { + const instanceId = "memory-visible-measurement", sessionId = "session", hiddenSessionId = "hidden" + const store = messageStoreBus.getOrCreate(instanceId) + const measure = store.getSessionApproximateByteSizeIncrementally + try { + addMessage(instanceId, sessionId) + addMessage(instanceId, hiddenSessionId) + setVisibleSessionMemory(instanceId, sessionId, true) + ;(store as any).getSessionApproximateByteSizeIncrementally = () => { throw new Error("sweep started a measurement") } + assert.deepEqual(runSessionMemorySweep(0), []) + ;(store as any).getSessionApproximateByteSizeIncrementally = measure + await settleMeasurements() + assert.deepEqual(runSessionMemorySweep(0), [`${instanceId}\u0000${hiddenSessionId}`]) + } finally { + ;(store as any).getSessionApproximateByteSizeIncrementally = measure + setVisibleSessionMemory(instanceId, sessionId, false) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("protected memory still contributes to the global eviction budget", async () => { + const instanceId = "memory-protected-budget", visibleSessionId = "visible", hiddenSessionId = "hidden" + try { + addMessage(instanceId, visibleSessionId) + addMessage(instanceId, hiddenSessionId) + const store = messageStoreBus.getOrCreate(instanceId) + const byteLimit = Math.max( + await store.getSessionApproximateByteSizeIncrementally(visibleSessionId), + await store.getSessionApproximateByteSizeIncrementally(hiddenSessionId), + ) + setVisibleSessionMemory(instanceId, visibleSessionId, true) + await settleMeasurements() + + assert.deepEqual(runSessionMemorySweep(byteLimit), [`${instanceId}\u0000${hiddenSessionId}`]) + assert.deepEqual(store.getSessionMessageIds(visibleSessionId), [`${visibleSessionId}-message`]) + } finally { + setVisibleSessionMemory(instanceId, visibleSessionId, false) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("deferred measurement refreshes a protected session after it grows", async () => { + const instanceId = "memory-protected-refresh", visibleSessionId = "visible", hiddenSessionId = "hidden" + const originalRequestIdleCallback = (globalThis as any).requestIdleCallback + const originalCancelIdleCallback = (globalThis as any).cancelIdleCallback + const callbacks: Array<() => void> = [] + ;(globalThis as any).requestIdleCallback = (callback: () => void) => { callbacks.push(callback); return callbacks.length } + ;(globalThis as any).cancelIdleCallback = () => undefined + try { + addMessage(instanceId, visibleSessionId) + setVisibleSessionMemory(instanceId, visibleSessionId, true) + callbacks.shift()?.() + await new Promise((resolve) => setImmediate(resolve)) + addMessage(instanceId, hiddenSessionId) + callbacks.shift()?.() + await new Promise((resolve) => setImmediate(resolve)) + + const store = messageStoreBus.getOrCreate(instanceId) + const initialLimit = await store.getSessionApproximateByteSizeIncrementally(visibleSessionId) + + await store.getSessionApproximateByteSizeIncrementally(hiddenSessionId) + assert.deepEqual(runSessionMemorySweep(initialLimit), []) + + addMessage(instanceId, visibleSessionId, "complete", "x".repeat(100_000)) + callbacks.shift()?.() + await new Promise((resolve) => setImmediate(resolve)) + assert.deepEqual(runSessionMemorySweep(initialLimit), [`${instanceId}\u0000${hiddenSessionId}`]) + } finally { + ;(globalThis as any).requestIdleCallback = originalRequestIdleCallback + ;(globalThis as any).cancelIdleCallback = originalCancelIdleCallback + setVisibleSessionMemory(instanceId, visibleSessionId, false) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("a hidden session remains evictable from its previous measurement while remeasurement is pending", async () => { + const instanceId = "memory-pending-remeasurement", sessionId = "hidden" + const originalRequestIdleCallback = (globalThis as any).requestIdleCallback + const originalCancelIdleCallback = (globalThis as any).cancelIdleCallback + try { + addMessage(instanceId, sessionId, "complete", "x".repeat(100_000)) + await settleMeasurements() + + let pendingMeasurement: (() => void) | undefined + ;(globalThis as any).requestIdleCallback = (callback: () => void) => { pendingMeasurement = callback; return 1 } + ;(globalThis as any).cancelIdleCallback = () => { pendingMeasurement = undefined } + addMessage(instanceId, sessionId, "complete", "new content") + + assert.equal(typeof pendingMeasurement, "function") + assert.deepEqual(runSessionMemorySweep(0), [`${instanceId}\u0000${sessionId}`]) + } finally { + ;(globalThis as any).requestIdleCallback = originalRequestIdleCallback + ;(globalThis as any).cancelIdleCallback = originalCancelIdleCallback + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("hidden growth crossing the budget remains evictable while remeasurement is pending", async () => { + const instanceId = "memory-growth-crosses-budget", sessionId = "hidden" + const originalRequestIdleCallback = (globalThis as any).requestIdleCallback + const originalCancelIdleCallback = (globalThis as any).cancelIdleCallback + let pendingMeasurement: (() => void) | undefined + try { + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { id: sessionId, status: "idle" } as any]]))) + addMessage(instanceId, sessionId, "complete", "small") + await settleMeasurements() + + ;(globalThis as any).requestIdleCallback = (callback: () => void) => { pendingMeasurement = callback; return 1 } + ;(globalThis as any).cancelIdleCallback = () => { pendingMeasurement = undefined } + addMessage(instanceId, sessionId, "complete", "x".repeat(Math.ceil(MAX_HOT_SESSION_MESSAGE_BYTES / 6) + 1_000)) + + const store = messageStoreBus.getOrCreate(instanceId) + assert.ok(await store.getSessionApproximateByteSizeIncrementally(sessionId) > MAX_HOT_SESSION_MESSAGE_BYTES) + assert.equal(typeof pendingMeasurement, "function") + assert.deepEqual(runSessionMemorySweep(), [`${instanceId}\u0000${sessionId}`]) + assert.deepEqual(store.getSessionMessageIds(sessionId), []) + } finally { + ;(globalThis as any).requestIdleCallback = originalRequestIdleCallback + ;(globalThis as any).cancelIdleCallback = originalCancelIdleCallback + setSessions((prev) => { const next = new Map(prev); next.delete(instanceId); return next }) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("superseded protected measurements abort and cannot publish stale bytes", async () => { + const instanceId = "memory-measurement-generation", visibleSessionId = "visible", hiddenSessionId = "hidden" + const originalRequestIdleCallback = (globalThis as any).requestIdleCallback + const originalCancelIdleCallback = (globalThis as any).cancelIdleCallback + const callbacks = new Map void>() + let nextHandle = 0 + ;(globalThis as any).requestIdleCallback = (callback: () => void) => { + const handle = ++nextHandle + callbacks.set(handle, callback) + return handle + } + ;(globalThis as any).cancelIdleCallback = (handle: number) => callbacks.delete(handle) + const runNext = () => { + const entry = callbacks.entries().next().value as [number, () => void] | undefined + assert.ok(entry) + callbacks.delete(entry[0]) + entry[1]() + } + try { + addMessage(instanceId, visibleSessionId) + setVisibleSessionMemory(instanceId, visibleSessionId, true) + runNext() + await new Promise((resolve) => setImmediate(resolve)) + addMessage(instanceId, hiddenSessionId) + runNext() + await new Promise((resolve) => setImmediate(resolve)) + + const store = messageStoreBus.getOrCreate(instanceId) + const originalMeasure = store.getSessionApproximateByteSizeIncrementally + let firstSignal: AbortSignal | undefined + let resolveStale!: (bytes: number) => void + let calls = 0 + ;(store as any).getSessionApproximateByteSizeIncrementally = (_sessionId: string, signal?: AbortSignal) => { + calls += 1 + if (calls === 1) { + firstSignal = signal + return new Promise((resolve) => { resolveStale = resolve }) + } + return Promise.resolve(1) + } + + addMessage(instanceId, visibleSessionId) + runNext() + addMessage(instanceId, visibleSessionId) + assert.equal(firstSignal?.aborted, true) + runNext() + await new Promise((resolve) => setImmediate(resolve)) + resolveStale(1_000_000) + await new Promise((resolve) => setImmediate(resolve)) + + const limit = await originalMeasure.call(store, hiddenSessionId) + 1 + assert.deepEqual(runSessionMemorySweep(limit), []) + ;(store as any).getSessionApproximateByteSizeIncrementally = originalMeasure + } finally { + ;(globalThis as any).requestIdleCallback = originalRequestIdleCallback + ;(globalThis as any).cancelIdleCallback = originalCancelIdleCallback + setVisibleSessionMemory(instanceId, visibleSessionId, false) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("instance teardown aborts an in-flight protected measurement", async () => { + const instanceId = "memory-measurement-teardown", sessionId = "session" + const originalRequestIdleCallback = (globalThis as any).requestIdleCallback + const originalCancelIdleCallback = (globalThis as any).cancelIdleCallback + let callback: (() => void) | undefined + ;(globalThis as any).requestIdleCallback = (next: () => void) => { callback = next; return 1 } + ;(globalThis as any).cancelIdleCallback = () => undefined + try { + addMessage(instanceId, sessionId) + const store = messageStoreBus.getOrCreate(instanceId) + let signal: AbortSignal | undefined + ;(store as any).getSessionApproximateByteSizeIncrementally = (_sessionId: string, nextSignal?: AbortSignal) => { + signal = nextSignal + return new Promise(() => undefined) + } + callback?.() + messageStoreBus.unregisterInstance(instanceId) + assert.equal(signal?.aborted, true) + } finally { + ;(globalThis as any).requestIdleCallback = originalRequestIdleCallback + ;(globalThis as any).cancelIdleCallback = originalCancelIdleCallback + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("idle metadata cannot evict a normalized streaming message", async () => { + const instanceId = "memory-idle-streaming", sessionId = "session" + try { + addMessage(instanceId, sessionId, "streaming") + await settleMeasurements() + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, { id: sessionId, status: "idle" } as any]]))) + assert.deepEqual(runSessionMemorySweep(0), []) + + addMessage(instanceId, sessionId, "complete") + await settleMeasurements() + assert.deepEqual(runSessionMemorySweep(0), [`${instanceId}\u0000${sessionId}`]) + } finally { + setSessions((prev) => { const next = new Map(prev); next.delete(instanceId); return next }) + messageStoreBus.unregisterInstance(instanceId) + } +}) + +test("interruption-only sessions contribute protected bytes to the global budget", async () => { + const instanceId = "memory-pending-only", pendingSessionId = "pending", hiddenSessionId = "hidden" + try { + addMessage(instanceId, hiddenSessionId) + const store = messageStoreBus.getOrCreate(instanceId) + store.upsertQuestion({ + request: { + id: "large-question", + sessionID: pendingSessionId, + questions: [{ header: "Confirm", question: "q".repeat(100_000), options: [] }], + } as any, + enqueuedAt: 1, + }) + await new Promise((resolve) => setTimeout(resolve, 75)) + + const limit = await store.getSessionApproximateByteSizeIncrementally(pendingSessionId) + assert.ok(store.getResidentSessionIds().includes(pendingSessionId)) + assert.deepEqual(runSessionMemorySweep(limit), [`${instanceId}\u0000${hiddenSessionId}`]) + } finally { + messageStoreBus.unregisterInstance(instanceId) + } +}) diff --git a/packages/ui/src/stores/session-memory.ts b/packages/ui/src/stores/session-memory.ts new file mode 100644 index 000000000..b5b8cda30 --- /dev/null +++ b/packages/ui/src/stores/session-memory.ts @@ -0,0 +1,167 @@ +import { getLogger } from "../lib/logger" +import { MAX_HOT_SESSION_MESSAGE_BYTES, selectSessionMemoryEvictions, type SessionMemoryEntry } from "../lib/session-memory-budget" +import { messageStoreBus } from "./message-v2/bus" +import { isSessionMessagesLoading, sessions } from "./session-state" + +const log = getLogger("session") +const SWEEP_DELAY_MS = 1_000 +const touched = new Map() +const visibleLeases = new Map() +const measuredBytes = new Map() +const pendingMeasurements = new Map + controller: AbortController +}>() +let sequence = 0 +let sweepTimer: ReturnType | undefined + +function cancelSessionMeasurement(key: string): void { + const pending = pendingMeasurements.get(key) + if (!pending) return + pending.controller.abort(new Error("Session memory measurement superseded")) + if (pending.idle) (globalThis as any).cancelIdleCallback?.(pending.handle) + else clearTimeout(pending.handle) + pendingMeasurements.delete(key) +} + +function scheduleSessionMeasurement(instanceId: string, sessionId: string): void { + const key = sessionKey(instanceId, sessionId) + cancelSessionMeasurement(key) + const pending = { idle: false, handle: 0 as number | ReturnType, controller: new AbortController() } + const measure = async () => { + try { + const store = messageStoreBus.getInstance(instanceId) + if (!store?.getResidentSessionIds().includes(sessionId)) return + const bytes = await store.getSessionApproximateByteSizeIncrementally(sessionId, pending.controller.signal) + if (pendingMeasurements.get(key) !== pending || pending.controller.signal.aborted) return + measuredBytes.set(key, bytes) + scheduleSessionMemorySweep() + } catch (error) { + if (!pending.controller.signal.aborted) log.warn("Failed to measure resident session messages", { instanceId, sessionId, error }) + } finally { + if (pendingMeasurements.get(key) === pending) pendingMeasurements.delete(key) + } + } + if (typeof (globalThis as any).requestIdleCallback === "function") { + pending.idle = true + pending.handle = (globalThis as any).requestIdleCallback(() => void measure(), { timeout: 2_000 }) as number + } else { + pending.handle = setTimeout(() => void measure(), 50) + } + pendingMeasurements.set(key, pending) +} + +function sessionKey(instanceId: string, sessionId: string): string { + return `${instanceId}\u0000${sessionId}` +} + +function splitSessionKey(key: string): [string, string] { + const separator = key.indexOf("\u0000") + return [key.slice(0, separator), key.slice(separator + 1)] +} + +function hasProtectedSessionWork( + store: ReturnType, + sessionId: string, +): boolean { + return store.hasSessionActiveWork(sessionId) +} + +export function scheduleSessionMemorySweep(): void { + if (sweepTimer) return + sweepTimer = setTimeout(() => { + sweepTimer = undefined + runSessionMemorySweep() + }, SWEEP_DELAY_MS) +} + +export function setVisibleSessionMemory(instanceId: string, sessionId: string, isVisible: boolean): void { + const key = sessionKey(instanceId, sessionId) + if (isVisible) { + visibleLeases.set(key, (visibleLeases.get(key) ?? 0) + 1) + touched.set(key, ++sequence) + } else { + const leases = visibleLeases.get(key) ?? 0 + if (leases <= 1) visibleLeases.delete(key) + else visibleLeases.set(key, leases - 1) + } + scheduleSessionMemorySweep() +} + +export function getVisibleSessionMemoryIds(instanceId: string): string[] { + const prefix = `${instanceId}\u0000` + const result: string[] = [] + for (const key of visibleLeases.keys()) if (key.startsWith(prefix)) result.push(key.slice(prefix.length)) + return result +} + +export function evictResidentSessionMessages(instanceId: string, sessionId: string): boolean { + const store = messageStoreBus.getInstance(instanceId) + const status = sessions().get(instanceId)?.get(sessionId)?.status + if ( + !store || + visibleLeases.has(sessionKey(instanceId, sessionId)) || + status === "working" || + status === "compacting" || + isSessionMessagesLoading(instanceId, sessionId) || + hasProtectedSessionWork(store, sessionId) + ) return false + store.clearSession(sessionId, { preserveScroll: true }) + log.info("Evicted resident session messages", { instanceId, sessionId }) + return true +} + +export function runSessionMemorySweep(byteLimit = MAX_HOT_SESSION_MESSAGE_BYTES): string[] { + const entries: SessionMemoryEntry[] = [] + for (const [instanceId, store] of messageStoreBus.entries()) { + for (const sessionId of store.getResidentSessionIds()) { + const key = sessionKey(instanceId, sessionId) + const status = sessions().get(instanceId)?.get(sessionId)?.status + const byteSize = measuredBytes.get(key) + const awaitingMeasurement = byteSize === undefined + if (awaitingMeasurement && !pendingMeasurements.has(key)) scheduleSessionMeasurement(instanceId, sessionId) + const protectedSession = awaitingMeasurement || visibleLeases.has(key) || status === "working" || status === "compacting" || + hasProtectedSessionWork(store, sessionId) || isSessionMessagesLoading(instanceId, sessionId) + entries.push({ + key, + // ponytail: unknown sessions stay protected until their first yielding measurement completes. + byteSize: byteSize ?? byteLimit, + lastTouched: touched.get(key) ?? 0, + protected: protectedSession, + }) + } + } + + const evicted: string[] = [] + for (const key of selectSessionMemoryEvictions(entries, byteLimit)) { + const [instanceId, sessionId] = splitSessionKey(key) + if (evictResidentSessionMessages(instanceId, sessionId)) evicted.push(key) + } + return evicted +} + +messageStoreBus.onSessionChanged((instanceId, sessionId) => { + const key = sessionKey(instanceId, sessionId) + touched.set(key, ++sequence) + const measured = measuredBytes.get(key) + if (measured !== undefined) measuredBytes.set(key, Math.max(measured, MAX_HOT_SESSION_MESSAGE_BYTES + 1)) + scheduleSessionMeasurement(instanceId, sessionId) + scheduleSessionMemorySweep() +}) + +messageStoreBus.onSessionCleared((instanceId, sessionId) => { + const key = sessionKey(instanceId, sessionId) + touched.delete(key) + visibleLeases.delete(key) + measuredBytes.delete(key) + cancelSessionMeasurement(key) +}) + +messageStoreBus.onInstanceDestroyed((instanceId) => { + const prefix = `${instanceId}\u0000` + for (const key of touched.keys()) if (key.startsWith(prefix)) touched.delete(key) + for (const key of visibleLeases.keys()) if (key.startsWith(prefix)) visibleLeases.delete(key) + for (const key of measuredBytes.keys()) if (key.startsWith(prefix)) measuredBytes.delete(key) + for (const key of pendingMeasurements.keys()) if (key.startsWith(prefix)) cancelSessionMeasurement(key) +}) diff --git a/packages/ui/src/stores/session-metadata.ts b/packages/ui/src/stores/session-metadata.ts index bd4f4ab85..aeb960d6d 100644 --- a/packages/ui/src/stores/session-metadata.ts +++ b/packages/ui/src/stores/session-metadata.ts @@ -44,10 +44,12 @@ export async function hydrateSessionMetadataWithClient( instanceId: string, sessionId: string, query?: { workspace?: string }, + isCurrent: () => boolean = () => true, ): Promise { const expectedMetadata = sessions().get(instanceId)?.get(sessionId)?.metadata const latest = await requestData(client.session.get({ sessionID: sessionId, ...query }), "session.get") const metadata = normalizeMetadata(latest?.metadata) + if (!isCurrent()) return metadata withSession(instanceId, sessionId, (session) => { if (session.metadata !== expectedMetadata || !shouldReplaceSessionMetadata(session.metadata)) return false diff --git a/packages/ui/src/stores/session-request-authority.test.ts b/packages/ui/src/stores/session-request-authority.test.ts index 99bfeb78f..e5996ba5d 100644 --- a/packages/ui/src/stores/session-request-authority.test.ts +++ b/packages/ui/src/stores/session-request-authority.test.ts @@ -3,15 +3,19 @@ import { describe, it } from "node:test" import { sdkManager } from "../lib/sdk-manager.ts" import type { Session } from "../types/session.ts" -import { addInstance, removeInstance } from "./instances.ts" +import { addInstance, clearReloadableInstanceState, instances, isInstanceRuntimeCurrent, removeInstance, setActiveInstanceId, updateInstance } from "./instances.ts" import { messageStoreBus } from "./message-v2/bus.ts" -import { loadMessages, removeSessionRuntimeState, searchSessions } from "./session-api.ts" +import { clearSessionSearch, fetchSessions, loadMessages, removeSessionRuntimeState, searchSessions } from "./session-api.ts" +import { handleSessionUpdate } from "./session-events.ts" import { clearInstanceDeletedSessionAuthority, + getSessionDraftPrompt, getSessionSearchResultIds, invalidateSessionMessageLoad, loading, messagesLoaded, + setActiveSession, + setSessionDraftPrompt, sessions, setSessions, } from "./session-state.ts" @@ -34,13 +38,13 @@ function apiSession(id: string, parentID?: string) { return { id, parentID, title: id, version: "1", time: { created: 1, updated: 1 } } } -function apiMessage(id: string, sessionId: string) { +function apiMessage(id: string, sessionId: string, text?: string) { return { info: { id, sessionID: sessionId, role: "assistant", agent: "build", providerID: "provider", modelID: "model", time: { created: 1 }, }, - parts: [], + parts: text === undefined ? [] : [{ id: `${id}-part`, type: "text", text }], } } @@ -61,6 +65,20 @@ function setup(instanceId: string) { } describe("session request authority", () => { + it("keeps metadata refreshes in the same runtime but fences client replacement", () => { + const instanceId = "instance-runtime-token" + const { cleanup } = setup(instanceId) + try { + const captured = instances().get(instanceId)! + updateInstance(instanceId, { projectName: "refreshed" }) + assert.equal(isInstanceRuntimeCurrent(instanceId, captured), true) + updateInstance(instanceId, { client: { session: {} } as any }) + assert.equal(isInstanceRuntimeCurrent(instanceId, captured), false) + } finally { + cleanup() + } + }) + it("does not restore deleted search results or their parent chain", async () => { const instanceId = "late-search-delete" const { client, cleanup } = setup(instanceId) @@ -153,6 +171,878 @@ describe("session request authority", () => { } }) + it("cancels the previous session load when selection moves to another session", async () => { + const instanceId = "switched-message-load", firstSessionId = "first", secondSessionId = "second" + const { client, cleanup } = setup(instanceId) + const firstRequestStarted = deferred() + let firstSignal: AbortSignal | undefined + ;(client.session as any).messages = ({ sessionID }: { sessionID: string }, options?: { signal?: AbortSignal }) => { + if (sessionID !== firstSessionId) return Promise.resolve({ data: [apiMessage("second-message", secondSessionId)] }) + firstSignal = options?.signal + firstRequestStarted.resolve() + return new Promise((_resolve, reject) => { + options?.signal?.addEventListener("abort", () => reject(new DOMException("aborted", "AbortError")), { once: true }) + }) + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([ + [firstSessionId, session(instanceId, firstSessionId)], + [secondSessionId, session(instanceId, secondSessionId)], + ]))) + + try { + setActiveSession(instanceId, firstSessionId) + const request = loadMessages(instanceId, firstSessionId) + await firstRequestStarted.promise + setActiveSession(instanceId, secondSessionId) + assert.equal(firstSignal?.aborted, true) + await loadMessages(instanceId, secondSessionId) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(secondSessionId), ["second-message"]) + + await request + + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(firstSessionId), []) + assert.equal(messagesLoaded().get(instanceId)?.has(firstSessionId) ?? false, false) + assert.equal(loading().loadingMessages.get(instanceId)?.has(firstSessionId) ?? false, false) + } finally { + cleanup() + } + }) + + it("invalidates an in-flight message load during same-runtime rehydration", async () => { + const instanceId = "same-runtime-rehydration", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const response = deferred() + let loadSignal: AbortSignal | undefined + ;(client.session as any).messages = (_input: unknown, options?: { signal?: AbortSignal }) => { + loadSignal = options?.signal + return response.promise + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + + try { + const request = loadMessages(instanceId, sessionId) + while (!loadSignal) await new Promise((resolve) => setImmediate(resolve)) + clearReloadableInstanceState(instanceId) + assert.equal(loadSignal.aborted, true) + response.resolve({ data: [apiMessage("stale-message", sessionId)] }) + await request + + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), []) + } finally { + cleanup() + } + }) + + it("cancels a revision retry while the session is hidden", async () => { + const instanceId = "hidden-message-retry", sessionId = "session", nextSessionId = "next" + const { client, cleanup } = setup(instanceId) + const firstResponse = deferred() + let calls = 0 + ;(client.session as any).messages = () => { + calls += 1 + return calls === 1 ? firstResponse.promise : Promise.resolve({ data: [apiMessage("retried-message", sessionId)] }) + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([ + [sessionId, session(instanceId, sessionId)], + [nextSessionId, session(instanceId, nextSessionId)], + ]))) + + try { + setActiveSession(instanceId, sessionId) + const request = loadMessages(instanceId, sessionId) + messageStoreBus.getOrCreate(instanceId).upsertMessage({ + id: "live-message", + sessionId, + role: "assistant", + status: "complete", + parts: [], + }) + firstResponse.resolve({ data: [apiMessage("http-message", sessionId)] }) + await new Promise((resolve) => setImmediate(resolve)) + setActiveSession(instanceId, nextSessionId) + await request + + assert.equal(calls, 1) + assert.equal(messagesLoaded().get(instanceId)?.has(sessionId) ?? false, false) + } finally { + cleanup() + } + }) + + it("cancels child loads owned by the session being left", async () => { + const instanceId = "switched-child-load", parentSessionId = "parent", childSessionId = "child" + const { client, cleanup } = setup(instanceId) + const started = deferred() + let childSignal: AbortSignal | undefined + ;(client.session as any).messages = (_parameters: unknown, options?: { signal?: AbortSignal }) => { + childSignal = options?.signal + started.resolve() + return new Promise((_resolve, reject) => options?.signal?.addEventListener( + "abort", + () => reject(new DOMException("aborted", "AbortError")), + { once: true }, + )) + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([ + [parentSessionId, session(instanceId, parentSessionId)], + [childSessionId, session(instanceId, childSessionId, parentSessionId)], + ["next", session(instanceId, "next")], + ]))) + + try { + setActiveSession(instanceId, parentSessionId) + const request = loadMessages(instanceId, childSessionId) + await started.promise + setActiveSession(instanceId, "next") + assert.equal(childSignal?.aborted, true) + await request + assert.equal(loading().loadingMessages.get(instanceId)?.has(childSessionId) ?? false, false) + } finally { + cleanup() + } + }) + + it("cancels a selected child load through its root ownership", async () => { + const instanceId = "selected-child-load", parentSessionId = "parent", childSessionId = "child" + const { client, cleanup } = setup(instanceId) + const started = deferred() + let childSignal: AbortSignal | undefined + ;(client.session as any).messages = (_parameters: unknown, options?: { signal?: AbortSignal }) => { + childSignal = options?.signal + started.resolve() + return new Promise((_resolve, reject) => options?.signal?.addEventListener( + "abort", + () => reject(new DOMException("aborted", "AbortError")), + { once: true }, + )) + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([ + [parentSessionId, session(instanceId, parentSessionId)], + [childSessionId, session(instanceId, childSessionId, parentSessionId)], + ["next", session(instanceId, "next")], + ]))) + + try { + setActiveSession(instanceId, childSessionId) + const request = loadMessages(instanceId, childSessionId) + await started.promise + setActiveSession(instanceId, "next") + assert.equal(childSignal?.aborted, true) + await request + } finally { + cleanup() + } + }) + + it("cancels child loads when their workspace loses visibility", async () => { + const instanceId = "hidden-child-load", parentSessionId = "parent", childSessionId = "child" + const { client, cleanup } = setup(instanceId) + const started = deferred() + let childSignal: AbortSignal | undefined + ;(client.session as any).messages = (_parameters: unknown, options?: { signal?: AbortSignal }) => { + childSignal = options?.signal + started.resolve() + return new Promise((_resolve, reject) => options?.signal?.addEventListener( + "abort", + () => reject(new DOMException("aborted", "AbortError")), + { once: true }, + )) + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([ + [parentSessionId, session(instanceId, parentSessionId)], + [childSessionId, session(instanceId, childSessionId, parentSessionId)], + ]))) + + try { + setActiveInstanceId(instanceId) + setActiveSession(instanceId, parentSessionId) + const request = loadMessages(instanceId, childSessionId) + await started.promise + setActiveInstanceId("another-instance") + assert.equal(childSignal?.aborted, true) + await request + } finally { + setActiveInstanceId(null) + cleanup() + } + }) + + it("rejects message hydration after in-place client replacement", async () => { + const instanceId = "replaced-message-client", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const response = deferred() + ;(client.session as any).messages = () => response.promise + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + + try { + const request = loadMessages(instanceId, sessionId) + updateInstance(instanceId, { client: { session: {} } as any }) + response.resolve({ data: [apiMessage("stale-message", sessionId)] }) + await request + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), []) + } finally { + cleanup() + } + }) + + it("rejects session-list deletion after in-place client replacement", async () => { + const instanceId = "replaced-session-list", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const response = deferred() + ;(client.session as any).list = () => response.promise + ;(client.session as any).status = async () => ({ data: {} }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + + try { + const request = fetchSessions(instanceId) + updateInstance(instanceId, { client: { session: {} } as any }) + response.resolve({ data: [] }) + await request + assert.equal(sessions().get(instanceId)?.has(sessionId), true) + } finally { + cleanup() + } + }) + + it("uses the fetched revert anchor before pruning reconnected history", async () => { + const instanceId = "authoritative-reconnect-revert", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const staleRevert = { messageID: "stale-anchor" } + const fetchedRevert = { messageID: "fetched-anchor" } + ;(client.session as any).list = async () => ({ data: [{ ...apiSession(sessionId), revert: fetchedRevert }] }) + ;(client.session as any).status = async () => ({ data: {} }) + ;(client.session as any).messages = async () => ({ data: [ + apiMessage("before", sessionId, "before"), + apiMessage("stale-anchor", sessionId, "stale"), + apiMessage("fetched-anchor", sessionId, "fetched"), + apiMessage("after", sessionId, "after"), + ] }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[ + sessionId, + { ...session(instanceId, sessionId), revert: staleRevert }, + ]]))) + messageStoreBus.getOrCreate(instanceId).setSessionRevert(sessionId, staleRevert) + for (const id of ["before", "stale-anchor", "fetched-anchor", "after"]) { + messageStoreBus.getOrCreate(instanceId).upsertMessage({ id, sessionId, role: "assistant", status: "complete", parts: [] }) + } + messageStoreBus.getOrCreate(instanceId).setSessionRevert(sessionId, staleRevert) + + try { + await fetchSessions(instanceId) + + assert.deepEqual(sessions().get(instanceId)?.get(sessionId)?.revert, fetchedRevert) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionRevert(sessionId), fetchedRevert) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), ["before", "stale-anchor"]) + } finally { + cleanup() + } + }) + + it("repairs resident history when a session list clears its revert", async () => { + const instanceId = "session-list-cleared-revert", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const oldRevert = { messageID: "anchor" } + ;(client.session as any).list = async () => ({ data: [apiSession(sessionId)] }) + ;(client.session as any).status = async () => ({ data: {} }) + ;(client.session as any).messages = async () => ({ data: [ + apiMessage("before", sessionId, "before"), + apiMessage("anchor", sessionId, "anchor"), + apiMessage("after", sessionId, "after"), + ] }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[ + sessionId, + { ...session(instanceId, sessionId), revert: oldRevert }, + ]]))) + const store = messageStoreBus.getOrCreate(instanceId) + for (const id of ["before", "anchor", "after"]) { + store.upsertMessage({ id, sessionId, role: "assistant", status: "complete", parts: [] }) + } + store.setSessionRevert(sessionId, oldRevert) + assert.deepEqual(store.getSessionMessageIds(sessionId), ["before"]) + + try { + await fetchSessions(instanceId) + assert.deepEqual(store.getSessionMessageIds(sessionId), ["before", "anchor", "after"]) + assert.equal(store.getSessionRevert(sessionId), null) + } finally { + cleanup() + } + }) + + it("keeps newer SSE revert authority over a deferred session list", async () => { + const instanceId = "session-list-revert-fence", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const response = deferred() + const staleRevert = { messageID: "stale" } + const currentRevert = { messageID: "current" } + ;(client.session as any).list = () => response.promise + ;(client.session as any).status = async () => ({ data: {} }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[ + sessionId, + { ...session(instanceId, sessionId), revert: staleRevert }, + ]]))) + + try { + const request = fetchSessions(instanceId) + handleSessionUpdate(instanceId, { + type: "session.updated", + properties: { info: { ...apiSession(sessionId), revert: currentRevert } }, + } as any) + response.resolve({ data: [{ ...apiSession(sessionId), revert: staleRevert }] }) + await request + + assert.equal(sessions().get(instanceId)?.get(sessionId)?.revert?.messageID, currentRevert.messageID) + } finally { + cleanup() + } + }) + + it("applies authoritative status while preserving newer SSE metadata", async () => { + const instanceId = "session-list-field-authority", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const status = deferred() + let statusStarted = false + ;(client.session as any).list = async () => ({ data: [{ ...apiSession(sessionId), title: "Stale title" }] }) + ;(client.session as any).status = () => { + statusStarted = true + return status.promise + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([[ + sessionId, + { ...session(instanceId, sessionId), title: "Original title", status: "working" }, + ]]))) + + try { + const request = fetchSessions(instanceId) + while (!statusStarted) await new Promise((resolve) => setImmediate(resolve)) + handleSessionUpdate(instanceId, { + type: "session.updated", + properties: { info: { ...apiSession(sessionId), title: "Current SSE title" } }, + } as any) + status.resolve({ data: {} }) + await request + + assert.equal(sessions().get(instanceId)?.get(sessionId)?.title, "Current SSE title") + assert.equal(sessions().get(instanceId)?.get(sessionId)?.status, "idle") + } finally { + cleanup() + } + }) + + it("keeps newer SSE revert authority over a deferred session search", async () => { + const instanceId = "session-search-revert-fence", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const response = deferred() + const staleRevert = { messageID: "stale" } + const currentRevert = { messageID: "current" } + ;(client.session as any).list = () => response.promise + setSessions((prev) => new Map(prev).set(instanceId, new Map([[ + sessionId, + { ...session(instanceId, sessionId), revert: staleRevert }, + ]]))) + + try { + const request = searchSessions(instanceId, "session") + handleSessionUpdate(instanceId, { + type: "session.updated", + properties: { info: { ...apiSession(sessionId), revert: currentRevert } }, + } as any) + response.resolve({ data: [{ ...apiSession(sessionId), revert: staleRevert }] }) + await request + + assert.equal(sessions().get(instanceId)?.get(sessionId)?.revert?.messageID, currentRevert.messageID) + } finally { + cleanup() + } + }) + + it("keeps a committed search authoritative over an older reconnect list and history prune", async () => { + const instanceId = "search-over-reconnect-fence", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const reconnect = deferred() + const staleRevert = { messageID: "anchor" } + ;(client.session as any).list = ({ search }: { search?: string }) => search + ? Promise.resolve({ data: [apiSession(sessionId)] }) + : reconnect.promise + ;(client.session as any).status = async () => ({ data: {} }) + ;(client.session as any).messages = async () => ({ data: [ + apiMessage("before", sessionId, "before"), + apiMessage("anchor", sessionId, "anchor"), + apiMessage("after", sessionId, "after"), + ] }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + + try { + const reconnectRequest = fetchSessions(instanceId) + await searchSessions(instanceId, "session") + reconnect.resolve({ data: [{ ...apiSession(sessionId), revert: staleRevert }] }) + await reconnectRequest + await loadMessages(instanceId, sessionId, { force: true }) + + assert.equal(sessions().get(instanceId)?.get(sessionId)?.revert, undefined) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), ["before", "anchor", "after"]) + } finally { + cleanup() + } + }) + + it("does not delete a session mutated by a newer search after a complete list starts", async () => { + const instanceId = "search-over-omitting-list", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const status = deferred() + let statusStarted = false + ;(client.session as any).list = ({ search }: { search?: string }) => Promise.resolve({ + data: search ? [{ ...apiSession(sessionId), title: "Current search", metadata: { source: "search" } }] : [], + }) + ;(client.session as any).status = () => { + statusStarted = true + return status.promise + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + setSessionDraftPrompt(instanceId, sessionId, "unsent draft") + + try { + const reconnect = fetchSessions(instanceId, { authoritativeDeletes: true }) + while (!statusStarted) await new Promise((resolve) => setImmediate(resolve)) + await searchSessions(instanceId, "session") + status.resolve({ data: {} }) + await reconnect + + assert.equal(sessions().get(instanceId)?.get(sessionId)?.title, "Current search") + assert.equal(getSessionDraftPrompt(instanceId, sessionId), "unsent draft") + } finally { + cleanup() + } + }) + + it("keeps a newer list commit authoritative over an older search", async () => { + const instanceId = "list-over-search-fence", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const search = deferred() + const currentRevert = { messageID: "current-anchor" } + const staleRevert = { messageID: "stale-anchor" } + ;(client.session as any).list = ({ search: query }: { search?: string }) => query + ? search.promise + : Promise.resolve({ data: [{ + ...apiSession(sessionId), + title: "Current list", + metadata: { source: "list" }, + revert: currentRevert, + }] }) + ;(client.session as any).status = async () => ({ data: {} }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + + try { + const staleSearch = searchSessions(instanceId, "session") + await fetchSessions(instanceId) + search.resolve({ data: [{ + ...apiSession(sessionId), + title: "Stale search", + metadata: { source: "search" }, + revert: staleRevert, + }] }) + await staleSearch + + assert.equal(sessions().get(instanceId)?.get(sessionId)?.title, "Current list") + assert.equal(sessions().get(instanceId)?.get(sessionId)?.revert?.messageID, currentRevert.messageID) + } finally { + cleanup() + } + }) + + it("returns newer SSE metadata as reconnect authority so reload reapplies its revert", async () => { + const instanceId = "sse-revert-reconnect-authority", sessionId = "session" + const { client, cleanup } = setup(instanceId) + const status = deferred() + let statusStarted = false + ;(client.session as any).list = async () => ({ data: [{ + ...apiSession(sessionId), + metadata: { source: "stale-list" }, + }] }) + ;(client.session as any).status = () => { + statusStarted = true + return status.promise + } + ;(client.session as any).messages = async () => ({ data: [ + apiMessage("before", sessionId, "before"), + apiMessage("anchor", sessionId, "anchor"), + apiMessage("after", sessionId, "after"), + ] }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + const store = messageStoreBus.getOrCreate(instanceId) + for (const id of ["before", "anchor", "after"]) { + store.upsertMessage({ id, sessionId, role: "assistant", status: "complete", parts: [] }) + } + + try { + const reconnect = fetchSessions(instanceId) + while (!statusStarted) await new Promise((resolve) => setImmediate(resolve)) + handleSessionUpdate(instanceId, { + type: "session.updated", + properties: { info: { ...apiSession(sessionId), revert: { messageID: "anchor" } } }, + } as any) + status.resolve({ data: {} }) + const refreshed = await reconnect + await loadMessages(instanceId, sessionId, { force: true, applySessionRevert: refreshed.has(sessionId) }) + + assert.equal(refreshed.has(sessionId), true) + assert.deepEqual(store.getSessionMessageIds(sessionId), ["before"]) + } finally { + cleanup() + } + }) + + it("hydrates only missing ancestors without overwriting the owning search result", async () => { + const instanceId = "search-parent-self-overwrite", sessionId = "child" + const { client, cleanup } = setup(instanceId) + const staleRevert = { messageID: "anchor" } + let calls = 0 + ;(client.session as any).list = () => Promise.resolve({ + data: ++calls === 1 + ? [{ ...apiSession(sessionId, "parent"), title: "current child" }] + : [ + apiSession("parent"), + { ...apiSession(sessionId, "parent"), title: "stale child", revert: staleRevert }, + apiSession("unrelated"), + ], + }) + ;(client.session as any).messages = async () => ({ data: [ + apiMessage("before", sessionId, "before"), + apiMessage("anchor", sessionId, "anchor"), + apiMessage("after", sessionId, "after"), + ] }) + + try { + await searchSessions(instanceId, "child") + await loadMessages(instanceId, sessionId, { force: true }) + + assert.equal(sessions().get(instanceId)?.get(sessionId)?.title, "current child") + assert.equal(sessions().get(instanceId)?.get(sessionId)?.revert, undefined) + assert.equal(sessions().get(instanceId)?.has("unrelated"), false) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(sessionId), ["before", "anchor", "after"]) + } finally { + cleanup() + } + }) + + it("keeps a newer search revert authoritative over a deferred parent-chain response", async () => { + const instanceId = "superseded-parent-chain-revert" + const { client, cleanup } = setup(instanceId) + const parents = deferred() + const staleRevert = { messageID: "stale" } + const currentRevert = { messageID: "current" } + let calls = 0 + ;(client.session as any).list = () => { + calls += 1 + if (calls === 1) return Promise.resolve({ data: [apiSession("child", "parent")] }) + if (calls === 2) return parents.promise + return Promise.resolve({ data: [{ ...apiSession("target"), revert: currentRevert }] }) + } + + try { + const staleSearch = searchSessions(instanceId, "child") + while (calls < 2) await new Promise((resolve) => setImmediate(resolve)) + await searchSessions(instanceId, "target") + assert.equal(sessions().get(instanceId)?.get("target")?.revert?.messageID, currentRevert.messageID) + + parents.resolve({ data: [apiSession("parent"), { ...apiSession("target"), revert: staleRevert }] }) + await staleSearch + + assert.equal(sessions().get(instanceId)?.get("target")?.revert?.messageID, currentRevert.messageID) + } finally { + cleanup() + } + }) + + it("keeps SSE session metadata authoritative over deferred parent hydration", async () => { + const instanceId = "parent-chain-metadata-fence" + const { client, cleanup } = setup(instanceId) + const parents = deferred() + let calls = 0 + ;(client.session as any).list = () => ++calls === 1 + ? Promise.resolve({ data: [apiSession("child", "parent")] }) + : parents.promise + + try { + const request = searchSessions(instanceId, "child") + while (calls < 2) await new Promise((resolve) => setImmediate(resolve)) + handleSessionUpdate(instanceId, { + type: "session.updated", + properties: { info: { ...apiSession("parent"), title: "current parent" } }, + } as any) + parents.resolve({ data: [{ ...apiSession("parent"), title: "stale parent" }] }) + await request + + assert.equal(sessions().get(instanceId)?.get("parent")?.title, "current parent") + assert.deepEqual(getSessionSearchResultIds(instanceId), ["child"]) + } finally { + cleanup() + } + }) + + it("aborts superseded searches and their parent requests", async () => { + const instanceId = "abort-superseded-search" + const { client, cleanup } = setup(instanceId) + let calls = 0 + let searchSignal: AbortSignal | undefined + let parentSignal: AbortSignal | undefined + ;(client.session as any).list = ( + { search }: { search?: string }, + options?: { signal?: AbortSignal }, + ) => { + calls += 1 + if (calls === 1) return Promise.resolve({ data: [apiSession("child", "parent")] }) + if (calls === 2) { + parentSignal = options?.signal + return new Promise((_resolve, reject) => options?.signal?.addEventListener( + "abort", + () => reject(options.signal?.reason), + { once: true }, + )) + } + if (search === "stale") { + searchSignal = options?.signal + return new Promise((_resolve, reject) => options?.signal?.addEventListener( + "abort", + () => reject(options.signal?.reason), + { once: true }, + )) + } + return Promise.resolve({ data: [apiSession("current")] }) + } + + try { + const parentSearch = searchSessions(instanceId, "child") + while (!parentSignal) await new Promise((resolve) => setImmediate(resolve)) + await searchSessions(instanceId, "current") + await parentSearch + assert.equal(parentSignal.aborted, true) + + const staleSearch = searchSessions(instanceId, "stale") + while (!searchSignal) await new Promise((resolve) => setImmediate(resolve)) + clearSessionSearch(instanceId) + await staleSearch + assert.equal(searchSignal.aborted, true) + } finally { + cleanup() + } + }) + + it("hydrates a complete deep chain from a capped search page before publishing", async () => { + const instanceId = "deep-capped-search-chain" + const { client, cleanup } = setup(instanceId) + let listCalls = 0 + const getCalls: string[] = [] + ;(client.session as any).list = () => { + listCalls += 1 + if (listCalls === 1) return Promise.resolve({ data: [apiSession("leaf", "middle")] }) + return Promise.resolve({ data: [ + apiSession("middle", "root"), + ...Array.from({ length: 199 }, (_, index) => apiSession(`filler-${index}`)), + ] }) + } + ;(client.session as any).get = ({ sessionID }: { sessionID: string }) => { + getCalls.push(sessionID) + return Promise.resolve({ data: apiSession(sessionID) }) + } + + try { + await searchSessions(instanceId, "leaf") + assert.deepEqual(getCalls, ["root"]) + assert.equal(sessions().get(instanceId)?.get("leaf")?.parentId, "middle") + assert.equal(sessions().get(instanceId)?.get("middle")?.parentId, "root") + assert.deepEqual(getSessionSearchResultIds(instanceId), ["leaf"]) + } finally { + cleanup() + } + }) + + it("does not publish a deep search result whose root chain cannot be resolved", async () => { + const instanceId = "unresolved-search-chain" + const { client, cleanup } = setup(instanceId) + let calls = 0 + ;(client.session as any).list = () => Promise.resolve({ data: ++calls === 1 ? [apiSession("leaf", "missing")] : [] }) + ;(client.session as any).get = async () => { throw new Error("missing parent") } + + try { + await searchSessions(instanceId, "leaf") + assert.deepEqual(getSessionSearchResultIds(instanceId), []) + } finally { + cleanup() + } + }) + + it("reloads history pruned by an older revert when search returns newer metadata", async () => { + const instanceId = "search-revert-history-repair", sessionId = "session" + const { client, cleanup } = setup(instanceId) + ;(client.session as any).list = async () => ({ data: [apiSession(sessionId)] }) + ;(client.session as any).messages = async () => ({ data: [ + apiMessage("before", sessionId, "before"), + apiMessage("anchor", sessionId, "anchor"), + apiMessage("after", sessionId, "after"), + ] }) + const oldRevert = { messageID: "anchor" } + setSessions((prev) => new Map(prev).set(instanceId, new Map([[ + sessionId, + { ...session(instanceId, sessionId), revert: oldRevert }, + ]]))) + const store = messageStoreBus.getOrCreate(instanceId) + for (const id of ["before", "anchor", "after"]) { + store.upsertMessage({ id, sessionId, role: "assistant", status: "complete", parts: [] }) + } + store.setSessionRevert(sessionId, oldRevert) + assert.deepEqual(store.getSessionMessageIds(sessionId), ["before"]) + + try { + await searchSessions(instanceId, "session") + assert.deepEqual(store.getSessionMessageIds(sessionId), ["before", "anchor", "after"]) + assert.equal(store.getSessionRevert(sessionId), null) + } finally { + cleanup() + } + }) + + it("repairs stale message-store revert during search even when metadata already matches", async () => { + const instanceId = "search-message-store-revert-repair", sessionId = "session" + const { client, cleanup } = setup(instanceId) + ;(client.session as any).list = async () => ({ data: [apiSession(sessionId)] }) + ;(client.session as any).messages = async () => ({ data: [ + apiMessage("before", sessionId, "before"), + apiMessage("anchor", sessionId, "anchor"), + apiMessage("after", sessionId, "after"), + ] }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + const store = messageStoreBus.getOrCreate(instanceId) + for (const id of ["before", "anchor", "after"]) { + store.upsertMessage({ id, sessionId, role: "assistant", status: "complete", parts: [] }) + } + store.setSessionRevert(sessionId, { messageID: "anchor" }) + assert.equal(sessions().get(instanceId)?.get(sessionId)?.revert, undefined) + assert.deepEqual(store.getSessionMessageIds(sessionId), ["before"]) + + try { + await searchSessions(instanceId, "session") + assert.deepEqual(store.getSessionMessageIds(sessionId), ["before", "anchor", "after"]) + assert.equal(store.getSessionRevert(sessionId), null) + } finally { + cleanup() + } + }) + + it("starts metadata hydration on a replacement runtime instead of reusing the stale promise", async () => { + const instanceId = "metadata-hydration-runtime", sessionId = "session" + const { client: oldClient, cleanup } = setup(instanceId) + const oldGet = deferred() + let oldCalls = 0 + ;(oldClient.session as any).list = async () => ({ data: [apiSession(sessionId)] }) + ;(oldClient.session as any).status = async () => ({ data: {} }) + ;(oldClient.session as any).get = () => { + oldCalls += 1 + return oldGet.promise + } + + try { + await fetchSessions(instanceId) + while (oldCalls === 0) await new Promise((resolve) => setImmediate(resolve)) + + let newCalls = 0 + const newClient = { + session: { + list: async () => ({ data: [apiSession(sessionId)] }), + status: async () => ({ data: {} }), + get: async () => { + newCalls += 1 + return { data: { ...apiSession(sessionId), metadata: { owner: "new-runtime" } } } + }, + }, + } as any + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, newClient) + updateInstance(instanceId, { client: newClient }) + + await fetchSessions(instanceId) + while (newCalls === 0) await new Promise((resolve) => setImmediate(resolve)) + assert.equal(sessions().get(instanceId)?.get(sessionId)?.metadata?.owner, "new-runtime") + + oldGet.resolve({ data: { ...apiSession(sessionId), metadata: { owner: "old-runtime" } } }) + await new Promise((resolve) => setImmediate(resolve)) + assert.equal(sessions().get(instanceId)?.get(sessionId)?.metadata?.owner, "new-runtime") + } finally { + cleanup() + } + }) + + it("hydrates unknown task-child metadata before loading its transcript", async () => { + const instanceId = "unknown-task-child", childId = "child" + const { client, cleanup } = setup(instanceId) + const calls: string[] = [] + ;(client.session as any).get = ({ sessionID }: { sessionID: string }) => { + calls.push(`get:${sessionID}`) + return Promise.resolve({ data: apiSession(sessionID, sessionID === childId ? "parent" : undefined) }) + } + ;(client.session as any).messages = ({ sessionID }: { sessionID: string }) => { + calls.push(`messages:${sessionID}`) + return Promise.resolve({ data: [apiMessage("child-message", sessionID)] }) + } + + try { + await loadMessages(instanceId, childId) + assert.deepEqual(calls, ["get:child", "get:parent", "messages:child"]) + assert.equal(sessions().get(instanceId)?.get(childId)?.parentId, "parent") + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(childId), ["child-message"]) + } finally { + cleanup() + } + }) + + it("keeps SSE metadata received during unknown task-child workspace lookup", async () => { + const instanceId = "unknown-task-child-workspace-fence", childId = "child" + const { client, cleanup } = setup(instanceId) + ;(client.session as any).get = ({ sessionID }: { sessionID: string }) => Promise.resolve({ + data: { ...apiSession(sessionID, sessionID === childId ? "parent" : undefined), title: `stale ${sessionID}` }, + }) + ;(client.session as any).messages = async () => ({ data: [] }) + + try { + const request = loadMessages(instanceId, childId) + handleSessionUpdate(instanceId, { + type: "session.updated", + properties: { info: { ...apiSession(childId, "parent"), title: "current child" } }, + } as any) + await request + + assert.equal(sessions().get(instanceId)?.get(childId)?.title, "current child") + assert.equal(sessions().get(instanceId)?.get(childId)?.parentId, "parent") + } finally { + cleanup() + } + }) + + it("does not reuse search authority after an instance reopens", async () => { + const instanceId = "reopened-session-search" + const { client, cleanup } = setup(instanceId) + const oldResponse = deferred() + const newResponse = deferred() + let calls = 0 + ;(client.session as any).list = () => (++calls === 1 ? oldResponse.promise : newResponse.promise) + + try { + const oldRequest = searchSessions(instanceId, "same") + removeInstance(instanceId, { authoritative: false }) + addInstance({ id: instanceId, folder: "/work", port: 0, pid: 0, proxyPath: "", status: "ready", client }) + const newRequest = searchSessions(instanceId, "same") + + newResponse.resolve({ data: [apiSession("new-session")] }) + await newRequest + oldResponse.resolve({ data: [apiSession("old-session")] }) + await oldRequest + + assert.deepEqual(getSessionSearchResultIds(instanceId), ["new-session"]) + assert.equal(sessions().get(instanceId)?.has("old-session") ?? false, false) + } finally { + cleanup() + } + }) + it("keeps a newer load authoritative when an older request finishes last", async () => { const instanceId = "newer-message-load", sessionId = "session" const { client, cleanup } = setup(instanceId) @@ -176,4 +1066,42 @@ describe("session request authority", () => { cleanup() } }) + + it("preserves historical assistant error status during hydration", async () => { + const instanceId = "errored-message-load", sessionId = "session" + const { client, cleanup } = setup(instanceId) + ;(client.session as any).messages = async () => ({ + data: [{ ...apiMessage("errored-message", sessionId), info: { ...apiMessage("errored-message", sessionId).info, error: { name: "ProviderError" } } }], + }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + + try { + await loadMessages(instanceId, sessionId) + assert.equal(messageStoreBus.getOrCreate(instanceId).getMessage("errored-message")?.status, "error") + } finally { + cleanup() + } + }) + + it("loads subagent messages on demand instead of hydrating an entire family", async () => { + const instanceId = "lazy-subagent-messages", parentId = "parent", childId = "child" + const { client, cleanup } = setup(instanceId) + const calls: string[] = [] + ;(client.session as any).messages = async ({ sessionID }: { sessionID: string }) => { + calls.push(sessionID) + return { data: [apiMessage(`${sessionID}-message`, sessionID)] } + } + setSessions((prev) => new Map(prev).set(instanceId, new Map([ + [parentId, session(instanceId, parentId)], + [childId, session(instanceId, childId, parentId)], + ]))) + + try { + await loadMessages(instanceId, parentId) + assert.deepEqual(calls, [parentId]) + assert.deepEqual(messageStoreBus.getOrCreate(instanceId).getSessionMessageIds(childId), []) + } finally { + cleanup() + } + }) }) diff --git a/packages/ui/src/stores/session-revision-recovery.test.ts b/packages/ui/src/stores/session-revision-recovery.test.ts new file mode 100644 index 000000000..dbe0be33a --- /dev/null +++ b/packages/ui/src/stores/session-revision-recovery.test.ts @@ -0,0 +1,588 @@ +import assert from "node:assert/strict" +import { describe, it } from "node:test" + +import { sdkManager } from "../lib/sdk-manager.ts" +import type { Session } from "../types/session.ts" +import { enqueueDelta, requestDeltaRecovery } from "./delta-buffer.ts" +import { addInstance, removeInstance, updateInstance } from "./instances.ts" +import { messageStoreBus } from "./message-v2/bus.ts" +import { loadMessages } from "./session-api.ts" +import { handleMessageUpdate, handleSessionDeleted, handleSessionIdle } from "./session-events.ts" +import { clearInstanceDeletedSessionAuthority, loading, messagesLoaded, sessions, setActiveSession, setMessagesLoaded, setSessions } from "./session-state.ts" +import { evictResidentSessionMessages } from "./session-memory.ts" + +function session(instanceId: string, id: string, status: Session["status"] = "idle"): Session { + return { + id, instanceId, parentId: null, title: id, agent: "build", + model: { providerId: "provider", modelId: "model" }, status, retry: null, + idleSince: null, generationRecovery: null, runtimeStatusKnown: true, + version: "1", time: { created: 1, updated: 1 }, + } +} + +function apiMessage(id: string, sessionId: string, modelId = "model", text?: string) { + return { + info: { + id, sessionID: sessionId, role: "assistant", agent: "build", + providerID: "provider", modelID: modelId, time: { created: 1, completed: 2 }, + }, + parts: text === undefined ? [] : [{ id: `${id}-part`, type: "text", text }], + } +} + +function deferred() { + let resolve!: (value: T) => void + const promise = new Promise((done) => { resolve = done }) + return { promise, resolve } +} + +async function waitFor(check: () => boolean): Promise { + const deadline = Date.now() + 1_500 + while (!check()) { + if (Date.now() >= deadline) throw new Error("Timed out waiting for revision recovery") + await new Promise((resolve) => setTimeout(resolve, 5)) + } +} + +function setup(instanceId: string, sessionId: string, status: Session["status"] = "idle") { + const client = { session: {} } as any + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, client) + addInstance({ id: instanceId, folder: "/work", port: 0, pid: 1, proxyPath: "", status: "ready", client }) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId, status)]]))) + return { + client, + cleanup() { + messageStoreBus.unregisterInstance(instanceId) + setSessions((prev) => { const next = new Map(prev); next.delete(instanceId); return next }) + clearInstanceDeletedSessionAuthority(instanceId) + removeInstance(instanceId, { authoritative: false }) + sdkManager.destroyClientsForInstance(instanceId) + }, + } +} + +describe("revision conflict recovery", () => { + it("bounds sustained full-history conflicts with backoff", async () => { + const instanceId = "bounded-revision-recovery", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + messageStoreBus.getOrCreate(instanceId).upsertMessage({ + id: `conflict-${calls}`, sessionId, role: "assistant", status: "streaming", + createdAt: calls, updatedAt: calls, + }) + return { data: [apiMessage("authoritative", sessionId, "stale-model")] } + } + + try { + setActiveSession(instanceId, sessionId) + await loadMessages(instanceId, sessionId) + await waitFor(() => calls === 4) + await new Promise((resolve) => setTimeout(resolve, 350)) + assert.equal(calls, 4, "one initial load plus three bounded recovery attempts") + assert.equal(sessions().get(instanceId)?.get(sessionId)?.model.modelId, "model") + } finally { + cleanup() + } + }) + + it("defers recovery while streaming and resumes from the idle event", async () => { + const instanceId = "streaming-revision-recovery", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId, "working") + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + return { data: [apiMessage("authoritative", sessionId)] } + } + + try { + requestDeltaRecovery({ instanceId, sessionId, messageId: "message", partId: "part", field: "text" }) + await new Promise((resolve) => setTimeout(resolve, 30)) + assert.equal(calls, 0) + handleSessionIdle(instanceId, { type: "session.idle", properties: { sessionID: sessionId } } as any) + await waitFor(() => calls === 2) + } finally { + cleanup() + } + }) + + it("invalidates hidden recovery so reactivation reloads", async () => { + const instanceId = "hidden-revision-recovery", sessionId = "session", nextSessionId = "next" + const { client, cleanup } = setup(instanceId, sessionId) + setSessions((prev) => { + const next = new Map(prev) + next.get(instanceId)?.set(nextSessionId, session(instanceId, nextSessionId)) + return next + }) + let calls = 0 + let signal: AbortSignal | undefined + ;(client.session as any).messages = (_input: unknown, options?: { signal?: AbortSignal }) => { + calls += 1 + if (calls > 1) return Promise.resolve({ data: [apiMessage("reloaded", sessionId, "model", "complete")] }) + signal = options?.signal + return new Promise((_resolve, reject) => signal?.addEventListener("abort", () => reject(signal?.reason), { once: true })) + } + + try { + setActiveSession(instanceId, sessionId) + setMessagesLoaded((prev) => new Map(prev).set(instanceId, new Set([sessionId]))) + requestDeltaRecovery({ instanceId, sessionId, messageId: "message", partId: "part", field: "text" }) + await waitFor(() => calls === 1 && Boolean(signal)) + handleSessionIdle(instanceId, { type: "session.idle", properties: { sessionID: sessionId } } as any) + requestDeltaRecovery({ instanceId, sessionId, messageId: "new-delta", partId: "part", field: "text" }) + setActiveSession(instanceId, nextSessionId) + await waitFor(() => signal?.aborted === true) + await new Promise((resolve) => setTimeout(resolve, 350)) + assert.equal(calls, 1) + assert.equal(messagesLoaded().get(instanceId)?.has(sessionId), false) + + setActiveSession(instanceId, sessionId) + await loadMessages(instanceId, sessionId) + assert.equal(calls, 2) + } finally { + cleanup() + } + }) + + it("clears a deferred working recovery when the session is deleted", async () => { + const instanceId = "deleted-deferred-recovery", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId, "working") + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + return { data: [apiMessage("authoritative", sessionId)] } + } + + try { + setActiveSession(instanceId, sessionId) + requestDeltaRecovery({ instanceId, sessionId, messageId: "old", partId: "part", field: "text" }) + handleSessionDeleted(instanceId, { type: "session.deleted", properties: { info: { id: sessionId } } } as any) + setSessions((prev) => new Map(prev).set(instanceId, new Map([[sessionId, session(instanceId, sessionId)]]))) + requestDeltaRecovery({ instanceId, sessionId, messageId: "new", partId: "part", field: "text" }) + await waitFor(() => calls === 1) + } finally { + cleanup() + } + }) + + it("replaces a sleeping recovery when the runtime changes", async () => { + const instanceId = "replacement-runtime-recovery", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let oldCalls = 0 + let replacementCalls = 0 + ;(client.session as any).messages = async () => { oldCalls += 1; return { data: [] } } + const replacementClient = { + session: { messages: async () => { replacementCalls += 1; return { data: [] } }, status: async () => ({ data: {} }) }, + } as any + + try { + requestDeltaRecovery({ instanceId, sessionId, messageId: "message", partId: "part", field: "text" }) + updateInstance(instanceId, { client: replacementClient }) + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, replacementClient) + requestDeltaRecovery({ instanceId, sessionId, messageId: "replacement-message", partId: "part", field: "text" }) + await waitFor(() => replacementCalls === 1) + assert.equal(oldCalls, 0) + } finally { + cleanup() + } + }) + + it("does not carry an idle fallback into a replacement runtime", async () => { + const instanceId = "idle-fallback-runtime-fence", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + const started = deferred() + const release = deferred() + let replacementCalls = 0 + ;(client.session as any).messages = async () => { + started.resolve() + await release.promise + return { data: [apiMessage("old-runtime", sessionId)] } + } + const replacementClient = { + session: { messages: async () => { replacementCalls += 1; return { data: [] } }, status: async () => ({ data: {} }) }, + } as any + + try { + requestDeltaRecovery({ instanceId, sessionId, messageId: "message", partId: "part", field: "text" }) + await started.promise + handleSessionIdle(instanceId, { type: "session.idle", properties: { sessionID: sessionId } } as any) + requestDeltaRecovery({ instanceId, sessionId, messageId: "new-delta", partId: "part", field: "text" }) + updateInstance(instanceId, { client: replacementClient }) + ;(sdkManager as any).clients.set(`${instanceId}:/workspaces/${instanceId}/instance`, replacementClient) + release.resolve() + await new Promise((resolve) => setTimeout(resolve, 50)) + + assert.equal(replacementCalls, 0) + } finally { + cleanup() + } + }) + + it("aborts a message load at its requested timeout", async () => { + const instanceId = "revision-recovery-timeout", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + let signal: AbortSignal | undefined + ;(client.session as any).messages = (_input: unknown, options?: { signal?: AbortSignal }) => { + signal = options?.signal + return new Promise((_resolve, reject) => signal?.addEventListener("abort", () => reject(signal?.reason), { once: true })) + } + + try { + await assert.rejects(loadMessages(instanceId, sessionId, { force: true, timeoutMs: 10 }), /timed out/) + assert.equal(signal?.aborted, true) + } finally { + cleanup() + } + }) + + it("settles an idle reconciliation timeout when the client ignores abort", async () => { + const instanceId = "idle-reconciliation-timeout", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + const originalSetTimeout = globalThis.setTimeout + let signal: AbortSignal | undefined + ;(client.session as any).messages = (_input: unknown, options?: { signal?: AbortSignal }) => { + signal = options?.signal + return new Promise(() => undefined) + } + const store = messageStoreBus.getOrCreate(instanceId) + store.upsertMessage({ + id: "streaming", sessionId, role: "assistant", status: "streaming", + parts: [{ id: "part", type: "text", text: "partial" }] as any, + }) + + try { + globalThis.setTimeout = ((callback: TimerHandler, delay?: number, ...args: any[]) => + originalSetTimeout(callback, delay === 10_000 ? 0 : delay, ...args)) as typeof setTimeout + handleSessionIdle(instanceId, { type: "session.idle", properties: { sessionID: sessionId } } as any) + await waitFor(() => signal?.aborted === true && !(loading().loadingMessages.get(instanceId)?.has(sessionId) ?? false)) + assert.equal(store.getMessage("streaming")?.status, "error") + assert.equal(store.hasSessionActiveWork(sessionId), false) + assert.equal(evictResidentSessionMessages(instanceId, sessionId), true) + } finally { + globalThis.setTimeout = originalSetTimeout + cleanup() + } + }) + + it("falls through to idle reconciliation after recovery attempts are exhausted", async () => { + const instanceId = "exhausted-idle-recovery", sessionId = "session" + const { client, cleanup } = setup(instanceId, sessionId) + const thirdAttempt = deferred() + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + if (calls <= 3) { + messageStoreBus.getOrCreate(instanceId).upsertMessage({ + id: `conflict-${calls}`, sessionId, role: "assistant", status: "streaming", + createdAt: calls, updatedAt: calls, + }) + } + if (calls === 3) { + setSessions((prev) => { + const next = new Map(prev) + const current = next.get(instanceId)?.get(sessionId) + if (current) next.get(instanceId)?.set(sessionId, { ...current, status: "working" }) + return next + }) + await thirdAttempt.promise + } + return { data: [apiMessage("authoritative", sessionId)] } + } + + try { + requestDeltaRecovery({ instanceId, sessionId, messageId: "message", partId: "part", field: "text" }) + await waitFor(() => calls === 3) + const store = messageStoreBus.getOrCreate(instanceId) + for (let attempt = 1; attempt <= 3; attempt += 1) store.removeMessage(`conflict-${attempt}`) + assert.equal(store.hasSessionActiveWork(sessionId), false) + handleSessionIdle(instanceId, { type: "session.idle", properties: { sessionID: sessionId } } as any) + thirdAttempt.resolve() + await waitFor(() => calls === 4) + } finally { + cleanup() + } + }) + + it("runs one post-recovery reconciliation when idle arrives during a successful recovery", async () => { + const instanceId = "idle-during-recovery", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + const first = deferred() + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + if (calls === 1) { + await first.promise + return { data: [apiMessage(messageId, sessionId, "model", "pre-idle")] } + } + return { data: [apiMessage(messageId, sessionId, "model", "post-idle")] } + } + messageStoreBus.getOrCreate(instanceId).upsertMessage({ + id: messageId, + sessionId, + role: "assistant", + status: "streaming", + parts: [{ id: `${messageId}-part`, type: "text", text: "old" }] as any, + }) + + try { + setActiveSession(instanceId, sessionId) + requestDeltaRecovery({ instanceId, sessionId, messageId: "missing", partId: "part", field: "text" }) + await waitFor(() => calls === 1) + handleSessionIdle(instanceId, { type: "session.idle", properties: { sessionID: sessionId } } as any) + first.resolve() + await waitFor(() => calls === 2) + await new Promise((resolve) => setTimeout(resolve, 200)) + + const part = messageStoreBus.getOrCreate(instanceId).getMessage(messageId)?.parts[`${messageId}-part`]?.data as any + assert.equal(part?.text, "post-idle") + assert.equal(calls, 2) + } finally { + cleanup() + } + }) + + it("waits for buffered deltas before accepting a history snapshot", async () => { + const instanceId = "buffered-delta-recovery", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + const first = deferred() + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + if (calls === 1) return first.promise + return { data: [apiMessage(messageId, sessionId, "model", "base tail")] } + } + messageStoreBus.getOrCreate(instanceId).upsertMessage({ + id: messageId, + sessionId, + role: "assistant", + status: "streaming", + parts: [{ id: `${messageId}-part`, type: "text", text: "base" }] as any, + }) + + try { + const load = loadMessages(instanceId, sessionId) + await waitFor(() => calls === 1) + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", " tail", sessionId) + first.resolve({ data: [apiMessage(messageId, sessionId, "model", "base")] }) + await load + await waitFor(() => calls === 2) + const part = messageStoreBus.getOrCreate(instanceId).getMessage(messageId)?.parts[`${messageId}-part`]?.data as any + assert.equal(part?.text, "base tail") + } finally { + cleanup() + } + }) + + it("retains an orphan delta across a predating snapshot until authority includes it", async () => { + const instanceId = "orphan-delta-recovery", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + return { data: [apiMessage(messageId, sessionId, "model", calls === 1 ? "base" : "base tail")] } + } + + try { + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", " tail", sessionId) + await waitFor(() => calls === 2) + await new Promise((resolve) => setTimeout(resolve, 250)) + + const part = messageStoreBus.getOrCreate(instanceId).getMessage(messageId)?.parts[`${messageId}-part`]?.data as any + assert.equal(part?.text, "base tail") + assert.equal(calls, 2) + } finally { + cleanup() + } + }) + + it("does not lose a repeated orphan append already present at the end of its base", async () => { + const instanceId = "orphan-delta-repeated-append", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + return { data: [apiMessage(messageId, sessionId, "model", calls === 1 ? "tail" : "tailtail")] } + } + + try { + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", "tail", sessionId) + await waitFor(() => calls === 2) + await new Promise((resolve) => setTimeout(resolve, 250)) + + const part = messageStoreBus.getOrCreate(instanceId).getMessage(messageId)?.parts[`${messageId}-part`]?.data as any + assert.equal(part?.text, "tailtail") + assert.equal(calls, 2) + } finally { + cleanup() + } + }) + + it("bounds orphan-delta recovery when authority keeps omitting the message", async () => { + const instanceId = "orphan-delta-omission", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + let calls = 0 + ;(client.session as any).messages = async () => { calls += 1; return { data: [] } } + + try { + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", "orphan", sessionId) + await waitFor(() => calls === 3) + await new Promise((resolve) => setTimeout(resolve, 250)) + + assert.equal(calls, 3) + assert.equal(messageStoreBus.getOrCreate(instanceId).getMessage(messageId), undefined) + } finally { + cleanup() + } + }) + + it("accepts a newer history snapshot that supersedes a buffered append", async () => { + const instanceId = "buffered-delta-newer-snapshot", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + const first = deferred() + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + if (calls === 1) return first.promise + return { data: [apiMessage(messageId, sessionId, "model", "base tail plus")] } + } + const store = messageStoreBus.getOrCreate(instanceId) + store.upsertMessage({ + id: messageId, + sessionId, + role: "assistant", + status: "streaming", + parts: [{ id: `${messageId}-part`, type: "text", text: "base" }] as any, + }) + + try { + const load = loadMessages(instanceId, sessionId) + await waitFor(() => calls === 1) + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", " tail", sessionId) + first.resolve({ data: [apiMessage(messageId, sessionId, "model", "base")] }) + await load + await waitFor(() => calls === 2) + await new Promise((resolve) => setTimeout(resolve, 200)) + + const part = store.getMessage(messageId)?.parts[`${messageId}-part`]?.data as any + assert.equal(part?.text, "base tail plus") + assert.equal(calls, 2) + } finally { + cleanup() + } + }) + + it("accepts an authoritative part replacement after fencing a buffered append", async () => { + const instanceId = "buffered-delta-part-replacement", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + const first = deferred() + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + if (calls === 1) return first.promise + return { data: [apiMessage(messageId, sessionId, "model", "replacement")] } + } + const store = messageStoreBus.getOrCreate(instanceId) + store.upsertMessage({ + id: messageId, + sessionId, + role: "assistant", + status: "streaming", + parts: [{ id: `${messageId}-part`, type: "text", text: "base" }] as any, + }) + + try { + const load = loadMessages(instanceId, sessionId) + await waitFor(() => calls === 1) + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", " tail", sessionId) + first.resolve({ data: [apiMessage(messageId, sessionId, "model", "base")] }) + await load + handleMessageUpdate(instanceId, { + type: "message.part.updated", + properties: { + part: { id: `${messageId}-part`, sessionID: sessionId, messageID: messageId, type: "text", text: "replacement" }, + }, + } as any) + await waitFor(() => calls === 2) + await new Promise((resolve) => setTimeout(resolve, 200)) + + const part = store.getMessage(messageId)?.parts[`${messageId}-part`]?.data as any + assert.equal(part?.text, "replacement") + assert.equal(calls, 2) + } finally { + cleanup() + } + }) + + it("preserves a buffered delta transiently then accepts repeated authoritative omission", async () => { + const instanceId = "buffered-delta-empty-recovery", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + const first = deferred() + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + if (calls === 1) return first.promise + if (calls === 2) return { data: [apiMessage(messageId, sessionId, "model", "base")] } + return { data: [] } + } + const store = messageStoreBus.getOrCreate(instanceId) + store.upsertMessage({ + id: messageId, + sessionId, + role: "assistant", + status: "streaming", + parts: [{ id: `${messageId}-part`, type: "text", text: "base" }] as any, + }) + + try { + const load = loadMessages(instanceId, sessionId) + await waitFor(() => calls === 1) + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", " tail", sessionId) + first.resolve({ data: [] }) + await load + assert.deepEqual(store.getSessionMessageIds(sessionId), [messageId]) + await waitFor(() => calls === 3) + assert.deepEqual(store.getSessionMessageIds(sessionId), []) + await new Promise((resolve) => setTimeout(resolve, 250)) + assert.equal(calls, 3) + } finally { + cleanup() + } + }) + + it("preserves a buffered delta transiently then accepts repeated authoritative replacement", async () => { + const instanceId = "buffered-delta-repeated-replacement", sessionId = "session", messageId = "message" + const { client, cleanup } = setup(instanceId, sessionId) + const first = deferred() + let calls = 0 + ;(client.session as any).messages = async () => { + calls += 1 + if (calls === 1) return first.promise + return { data: [apiMessage(messageId, sessionId, "model", "replacement")] } + } + const store = messageStoreBus.getOrCreate(instanceId) + store.upsertMessage({ + id: messageId, + sessionId, + role: "assistant", + status: "streaming", + parts: [{ id: `${messageId}-part`, type: "text", text: "base" }] as any, + }) + + try { + const load = loadMessages(instanceId, sessionId) + await waitFor(() => calls === 1) + enqueueDelta(instanceId, messageId, `${messageId}-part`, "text", " tail", sessionId) + first.resolve({ data: [apiMessage(messageId, sessionId, "model", "base")] }) + await load + await waitFor(() => calls === 3) + + const part = store.getMessage(messageId)?.parts[`${messageId}-part`]?.data as any + assert.equal(part?.text, "replacement") + await new Promise((resolve) => setTimeout(resolve, 250)) + assert.equal(calls, 3) + } finally { + cleanup() + } + }) +}) diff --git a/packages/ui/src/stores/session-state-purge.test.ts b/packages/ui/src/stores/session-state-purge.test.ts new file mode 100644 index 000000000..f13ded013 --- /dev/null +++ b/packages/ui/src/stores/session-state-purge.test.ts @@ -0,0 +1,36 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { + agents, + loading, + messagesLoaded, + providers, + purgeInstanceSessionState, + sessionInfoByInstance, + sessions, + setAgents, + setLoading, + setMessagesLoaded, + setProviders, + setSessionInfoByInstance, + setSessions, +} from "./session-state.ts" + +test("purging an instance removes its session metadata without touching other workspaces", () => { + const removed = "purged-instance" + const retained = "retained-instance" + setSessions(new Map([[removed, new Map()], [retained, new Map()]])) + setAgents(new Map([[removed, []], [retained, []]])) + setProviders(new Map([[removed, []], [retained, []]])) + setMessagesLoaded(new Map([[removed, new Set(["session"])], [retained, new Set()]])) + setSessionInfoByInstance(new Map([[removed, new Map()], [retained, new Map()]])) + setLoading((current) => ({ ...current, loadingMessages: new Map([[removed, new Set(["session"])], [retained, new Set()]]) })) + + purgeInstanceSessionState(removed) + + for (const state of [sessions(), agents(), providers(), messagesLoaded(), sessionInfoByInstance(), loading().loadingMessages]) { + assert.equal(state.has(removed), false) + assert.equal(state.has(retained), true) + } + purgeInstanceSessionState(retained) +}) diff --git a/packages/ui/src/stores/session-state.ts b/packages/ui/src/stores/session-state.ts index e9cf30bd0..49a4e98dd 100644 --- a/packages/ui/src/stores/session-state.ts +++ b/packages/ui/src/stores/session-state.ts @@ -39,11 +39,49 @@ interface GenerationAdmission { baseline: Pick } const generationAdmissions = new Map() +let sessionMetadataMutationSequence = 0 +const sessionMetadataMutationVersions = new Map() + +function snapshotSessionMetadataMutationVersion(): number { + return sessionMetadataMutationSequence +} + +function markSessionMetadataMutation(instanceId: string, sessionId: string): void { + sessionMetadataMutationVersions.set(`${instanceId}:${sessionId}`, ++sessionMetadataMutationSequence) +} + +function wasSessionMetadataMutatedAfter(instanceId: string, sessionId: string, version: number): boolean { + return (sessionMetadataMutationVersions.get(`${instanceId}:${sessionId}`) ?? 0) > version +} function cancelSessionGenerationAdmissions(instanceId: string, sessionId: string): void { generationAdmissions.delete(`${instanceId}:${sessionId}`) } +function resetInstanceSessionRequestState(instanceId: string): void { + const prefix = `${instanceId}:` + const activeMessageSessionIds = [...messageLoadControllers.keys()] + .filter((key) => key.startsWith(prefix)) + .map((key) => key.slice(prefix.length)) + for (const sessionId of activeMessageSessionIds) advanceMessageLoadEpoch(instanceId, sessionId) + for (const [key, admission] of generationAdmissions) { + if (!key.startsWith(prefix)) continue + const sessionId = key.slice(prefix.length) + withSession(instanceId, sessionId, (session) => { + if (session.generationAdmissionToken !== admission.token) return false + session.generationAdmissionToken = undefined + Object.assign(session, admission.baseline) + }) + generationAdmissions.delete(key) + } + setLoading((prev) => ({ + fetchingSessions: removeInstanceMapEntry(prev.fetchingSessions, instanceId), + creatingSession: removeInstanceMapEntry(prev.creatingSession, instanceId), + deletingSession: removeInstanceMapEntry(prev.deletingSession, instanceId), + loadingMessages: removeInstanceMapEntry(prev.loadingMessages, instanceId), + })) +} + export interface SessionInfo { cost: number contextWindow: number @@ -80,6 +118,7 @@ const [messagesLoaded, setMessagesLoaded] = createSignal const [messageLoadErrors, setMessageLoadErrors] = createSignal>>(new Map()) const [sessionListErrors, setSessionListErrors] = createSignal>(new Map()) const messageLoadEpochs = new Map() +const messageLoadControllers = new Map() let nextMessageLoadEpoch = 0 const [sessionInfoByInstance, setSessionInfoByInstance] = createSignal>>(new Map()) const [threadTotalsByInstance, setThreadTotalsByInstance] = createSignal>>(new Map()) @@ -337,11 +376,38 @@ function clearLoadedFlag(instanceId: string, sessionId: string) { function advanceMessageLoadEpoch(instanceId: string, sessionId: string): number { const key = getDraftKey(instanceId, sessionId) + messageLoadControllers.get(key)?.controller.abort() + messageLoadControllers.delete(key) const epoch = ++nextMessageLoadEpoch messageLoadEpochs.set(key, epoch) return epoch } +function beginSessionMessageLoad(instanceId: string, sessionId: string): { + epoch: number + signal: AbortSignal + abort: (reason?: unknown) => void +} { + const epoch = advanceMessageLoadEpoch(instanceId, sessionId) + const controller = new AbortController() + const ownerId = getSessionRoot(instanceId, sessionId)?.id ?? sessionId + messageLoadControllers.set(getDraftKey(instanceId, sessionId), { epoch, ownerSessionId: ownerId, controller }) + return { epoch, signal: controller.signal, abort: (reason) => controller.abort(reason) } +} + +function finishSessionMessageLoad(instanceId: string, sessionId: string, epoch: number): void { + const key = getDraftKey(instanceId, sessionId) + if (messageLoadControllers.get(key)?.epoch === epoch) messageLoadControllers.delete(key) +} + +function invalidateOwnedSessionMessageLoads(instanceId: string, ownerSessionId: string): void { + const prefix = `${instanceId}:` + const sessionIds = [...messageLoadControllers] + .filter(([key, load]) => key.startsWith(prefix) && load.ownerSessionId === ownerSessionId) + .map(([key]) => key.slice(prefix.length)) + for (const sessionId of sessionIds) invalidateSessionMessageLoad(instanceId, sessionId) +} + function isCurrentMessageLoad(instanceId: string, sessionId: string, epoch: number): boolean { return messageLoadEpochs.get(getDraftKey(instanceId, sessionId)) === epoch } @@ -724,6 +790,10 @@ function writeSessionSelection( } function writeActiveSession(instanceId: string, sessionId: string | null): void { + const previousSessionId = activeSessionId().get(instanceId) + if (previousSessionId && previousSessionId !== sessionId) { + invalidateOwnedSessionMessageLoads(instanceId, getSessionRoot(instanceId, previousSessionId)?.id ?? previousSessionId) + } writeSessionSelection(setActiveSessionId, instanceId, sessionId) if (sessionId) { // Backfill authoritative Yolo state for the now-active session so the badge @@ -1207,6 +1277,58 @@ async function cleanupBlankSessions(instanceId: string, excludeSessionId?: strin } } +function removeInstanceMapEntry(map: Map, instanceId: string): Map { + if (!map.has(instanceId)) return map + const next = new Map(map) + next.delete(instanceId) + return next +} + +function purgeInstanceSessionState(instanceId: string): void { + if (!instanceId) return + const prefix = `${instanceId}:` + batch(() => { + setSessions((prev) => removeInstanceMapEntry(prev, instanceId)) + setActiveSessionId((prev) => removeInstanceMapEntry(prev, instanceId)) + setActiveParentSessionId((prev) => removeInstanceMapEntry(prev, instanceId)) + setAgents((prev) => removeInstanceMapEntry(prev, instanceId)) + setProviders((prev) => removeInstanceMapEntry(prev, instanceId)) + setMessagesLoaded((prev) => removeInstanceMapEntry(prev, instanceId)) + setMessageLoadErrors((prev) => removeInstanceMapEntry(prev, instanceId)) + setSessionListErrors((prev) => removeInstanceMapEntry(prev, instanceId)) + setSessionInfoByInstance((prev) => removeInstanceMapEntry(prev, instanceId)) + setThreadTotalsByInstance((prev) => removeInstanceMapEntry(prev, instanceId)) + setExpandedSessions((prev) => removeInstanceMapEntry(prev, instanceId)) + setSessionPagination((prev) => removeInstanceMapEntry(prev, instanceId)) + setSessionSearch((prev) => removeInstanceMapEntry(prev, instanceId)) + setInstanceIndicatorCounts((prev) => removeInstanceMapEntry(prev, instanceId)) + setLoading((prev) => ({ + fetchingSessions: removeInstanceMapEntry(prev.fetchingSessions, instanceId), + creatingSession: removeInstanceMapEntry(prev.creatingSession, instanceId), + deletingSession: removeInstanceMapEntry(prev.deletingSession, instanceId), + loadingMessages: removeInstanceMapEntry(prev.loadingMessages, instanceId), + })) + setAuthoritativeSessionSelectionInstanceIds((prev) => { + if (!prev.has(instanceId)) return prev + const next = new Set(prev) + next.delete(instanceId) + return next + }) + setSessionDraftPrompts((prev) => new Map([...prev].filter(([key]) => !key.startsWith(prefix)))) + setAuthoritativeDraftKeys((prev) => new Set([...prev].filter((key) => !key.startsWith(prefix)))) + setAuthoritativelyDeletedSessionKeys((prev) => new Set([...prev].filter((key) => !key.startsWith(prefix)))) + setAuthoritativeSessionExpansionKeys((prev) => new Set([...prev].filter((key) => !key.startsWith(prefix)))) + }) + for (const key of generationAdmissions.keys()) if (key.startsWith(prefix)) generationAdmissions.delete(key) + for (const [key, load] of messageLoadControllers) { + if (!key.startsWith(prefix)) continue + load.controller.abort() + messageLoadControllers.delete(key) + } + for (const key of messageLoadEpochs.keys()) if (key.startsWith(prefix)) messageLoadEpochs.delete(key) + for (const key of sessionMetadataMutationVersions.keys()) if (key.startsWith(prefix)) sessionMetadataMutationVersions.delete(key) +} + export { sessions, setSessions, @@ -1223,6 +1345,9 @@ export { getSessionListError, setSessionListError, advanceMessageLoadEpoch, + beginSessionMessageLoad, + finishSessionMessageLoad, + invalidateOwnedSessionMessageLoads, isCurrentMessageLoad, invalidateSessionMessageLoad, setSessionMessagesLoadError, @@ -1257,6 +1382,10 @@ export { hydrateSessionGenerationRecovery, beginSessionGenerationAdmission, cancelSessionGenerationAdmissions, + resetInstanceSessionRequestState, + snapshotSessionMetadataMutationVersion, + markSessionMetadataMutation, + wasSessionMetadataMutatedAfter, setSessionStatus, setActiveSession, setActiveParentSession, @@ -1291,6 +1420,7 @@ export { getSessionInfo, isBlankSession, cleanupBlankSessions, + purgeInstanceSessionState, SESSION_PAGE_SIZE, sessionPagination, sessionSearch, diff --git a/packages/ui/src/stores/sessions.ts b/packages/ui/src/stores/sessions.ts index 9fd27961d..fef50f29b 100644 --- a/packages/ui/src/stores/sessions.ts +++ b/packages/ui/src/stores/sessions.ts @@ -59,7 +59,6 @@ import { setSessionExpanded, setSessionStatus, toggleSessionExpanded, - clearSessionSearch, getSessionHasMore, hydrateSessionDraftPrompt, isSessionSearchLoading, @@ -79,6 +78,7 @@ import { fetchSessions, hydrateRestoredSessionChain, loadMoreSessions, + clearSessionSearch, searchSessions, forkSession, loadMessages, diff --git a/packages/ui/src/stores/worktree-ready.test.ts b/packages/ui/src/stores/worktree-ready.test.ts index e78456e00..4e9249b42 100644 --- a/packages/ui/src/stores/worktree-ready.test.ts +++ b/packages/ui/src/stores/worktree-ready.test.ts @@ -127,6 +127,7 @@ describe("handleWorktreeReady", () => { await Promise.resolve() assert.equal(requestCount, 2) + assert.deepEqual(getWorktrees(instanceId).map((worktree) => worktree.slug), ["root"]) resolveReload({ isGitRepo: true, @@ -142,4 +143,38 @@ describe("handleWorktreeReady", () => { serverApi.fetchWorktrees = originalFetchWorktrees } }) + + it("does not let an aborted waiter expose an older request to overwrite a newer reload", async () => { + const instanceId = "instance-aborted-waiter" + const originalFetchWorktrees = serverApi.fetchWorktrees + const requests: Array<(value: Awaited>) => void> = [] + serverApi.fetchWorktrees = () => new Promise((resolve) => requests.push(resolve)) + + try { + const first = reloadWorktrees(instanceId) + await Promise.resolve() + const controller = new AbortController() + const aborted = reloadWorktrees(instanceId, controller.signal) + controller.abort(new Error("reload timed out")) + await assert.rejects(aborted, /reload timed out/) + + const latest = reloadWorktrees(instanceId) + await Promise.resolve() + await Promise.resolve() + requests[1]({ + isGitRepo: true, + worktrees: [{ slug: "new", directory: "/new", kind: "worktree" }], + }) + await latest + requests[0]({ + isGitRepo: true, + worktrees: [{ slug: "old", directory: "/old", kind: "worktree" }], + }) + await first + + assert.deepEqual(getWorktrees(instanceId).map((worktree) => worktree.slug), ["new"]) + } finally { + serverApi.fetchWorktrees = originalFetchWorktrees + } + }) }) diff --git a/packages/ui/src/stores/worktrees.ts b/packages/ui/src/stores/worktrees.ts index be7c3860c..7e8fe084d 100644 --- a/packages/ui/src/stores/worktrees.ts +++ b/packages/ui/src/stores/worktrees.ts @@ -12,12 +12,13 @@ const [worktreesByInstance, setWorktreesByInstance] = createSignal>(new Map()) const [gitRepoStatusByInstance, setGitRepoStatusByInstance] = createSignal>(new Map()) -const worktreeRequests = new Map>() +const worktreeRequests = new Map>() +const worktreeRequestGenerations = new Map() const worktreeReadyRefreshes = new Map>() const mapLoads = new Map>() const mapMigrations = new Map>() -type WorktreeReadyRefresh = (instanceId: string) => Promise +type WorktreeReadyRefresh = (instanceId: string) => Promise function normalizeMap(input?: WorktreeMap | null): WorktreeMap { if (!input || typeof input !== "object") { @@ -30,11 +31,29 @@ function normalizeMap(input?: WorktreeMap | null): WorktreeMap { } } -async function queueWorktreeRequest(instanceId: string, initial: boolean): Promise { +async function queueWorktreeRequest(instanceId: string, initial: boolean, signal?: AbortSignal): Promise { const previous = worktreeRequests.get(instanceId) - const task = (previous?.catch(() => undefined) ?? Promise.resolve()).then(async () => { + let generation = 0 + const isCurrent = () => generation > 0 && worktreeRequestGenerations.get(instanceId) === generation + const task = (async () => { + if (previous) { + if (signal) { + await Promise.race([ + previous.catch(() => undefined), + new Promise((_, reject) => { + if (signal.aborted) reject(signal.reason) + else signal.addEventListener("abort", () => reject(signal.reason), { once: true }) + }), + ]) + } else { + await previous.catch(() => undefined) + } + } + generation = (worktreeRequestGenerations.get(instanceId) ?? 0) + 1 + worktreeRequestGenerations.set(instanceId, generation) try { - const response = await serverApi.fetchWorktrees(instanceId) + const response = await serverApi.fetchWorktrees(instanceId, signal) + if (!isCurrent()) return false setWorktreesByInstance((prev) => { const next = new Map(prev) next.set(instanceId, response.worktrees ?? []) @@ -51,9 +70,11 @@ async function queueWorktreeRequest(instanceId: string, initial: boolean): Promi if (worktreeMapByInstance().has(instanceId)) { void pruneWorktreeMap(instanceId).catch(() => undefined) } + return true } catch (error) { + if (!isCurrent() || signal?.aborted) return false log.warn(initial ? "Failed to load worktrees" : "Failed to reload worktrees", { instanceId, error }) - if (!initial) return + if (!initial) return false setWorktreesByInstance((prev) => { const next = new Map(prev) @@ -68,33 +89,40 @@ async function queueWorktreeRequest(instanceId: string, initial: boolean): Promi next.set(instanceId, null) return next }) + return false } - }) + })() worktreeRequests.set(instanceId, task) - await task.finally(() => { - if (worktreeRequests.get(instanceId) === task) { + return task.finally(() => { + if (isCurrent() && worktreeRequests.get(instanceId) === task) { worktreeRequests.delete(instanceId) } }) } -async function ensureWorktreesLoaded(instanceId: string): Promise { +async function ensureWorktreesLoaded(instanceId: string, signal?: AbortSignal): Promise { if (!instanceId) return if (worktreesByInstance().has(instanceId) && gitRepoStatusByInstance().has(instanceId)) return const existing = worktreeRequests.get(instanceId) if (existing) { - await existing + await Promise.race([ + existing, + new Promise((_, reject) => { + if (signal?.aborted) reject(signal.reason) + else signal?.addEventListener("abort", () => reject(signal.reason), { once: true }) + }), + ]) if (worktreesByInstance().has(instanceId) && gitRepoStatusByInstance().has(instanceId)) return } - await queueWorktreeRequest(instanceId, true) + await queueWorktreeRequest(instanceId, true, signal) } -async function reloadWorktrees(instanceId: string): Promise { - if (!instanceId) return - await queueWorktreeRequest(instanceId, false) +async function reloadWorktrees(instanceId: string, signal?: AbortSignal): Promise { + if (!instanceId) return false + return queueWorktreeRequest(instanceId, false, signal) } async function handleWorktreeReady( diff --git a/packages/ui/src/styles/messaging/message-section.css b/packages/ui/src/styles/messaging/message-section.css index 361bf7029..5a635ec8e 100644 --- a/packages/ui/src/styles/messaging/message-section.css +++ b/packages/ui/src/styles/messaging/message-section.css @@ -381,6 +381,15 @@ border-top: 1px solid var(--border-base); } +.message-search-preview { + padding: 0.625rem 0.75rem 0; + border-top: 1px solid var(--border-base); + margin-top: 0.625rem; + color: var(--text-secondary); + font-size: var(--font-size-xs); + overflow-wrap: anywhere; +} + .message-stream-block[data-search-result="true"] { border-radius: 0.75rem; } diff --git a/packages/ui/src/types/instance.ts b/packages/ui/src/types/instance.ts index 4ba2be593..57c8a1991 100644 --- a/packages/ui/src/types/instance.ts +++ b/packages/ui/src/types/instance.ts @@ -45,4 +45,6 @@ export interface Instance { binaryLabel?: string binaryVersion?: string environmentVariables?: Record + runtimeToken?: symbol + attachedPid?: number }