diff --git a/src/app-server/client.ts b/src/app-server/client.ts index 21e4362c..f6a523c6 100644 --- a/src/app-server/client.ts +++ b/src/app-server/client.ts @@ -5,6 +5,7 @@ import type { ReasoningEffort } from "../generated/app-server/ReasoningEffort"; import type { ApprovalsReviewer } from "../generated/app-server/v2/ApprovalsReviewer"; import type { ConfigReadResponse } from "../generated/app-server/v2/ConfigReadResponse"; import type { ConfigWriteResponse } from "../generated/app-server/v2/ConfigWriteResponse"; +import type { FsReadFileResponse } from "../generated/app-server/v2/FsReadFileResponse"; import type { GetAccountRateLimitsResponse } from "../generated/app-server/v2/GetAccountRateLimitsResponse"; import type { HookMetadata } from "../generated/app-server/v2/HookMetadata"; import type { HooksListResponse } from "../generated/app-server/v2/HooksListResponse"; @@ -42,6 +43,7 @@ import type { JsonValue } from "../generated/app-server/serde_json/JsonValue"; import type { ServiceTierRequest } from "./service-tier"; const DEFAULT_REQUEST_TIMEOUT_MS = 120_000; +const MAX_SUPPRESSED_ORPHAN_RESPONSES = 256; export interface AppServerClientHandlers { onNotification: (notification: ServerNotification) => void; @@ -96,6 +98,7 @@ interface ClientResponseByMethod { "turn/start": TurnStartResponse; "turn/steer": TurnSteerResponse; "turn/interrupt": TurnInterruptResponse; + "fs/readFile": FsReadFileResponse; } type TypedClientRequestMethod = Extract; @@ -114,6 +117,7 @@ export class AppServerClient { private lifecycle: AppServerClientLifecycleState = { kind: "disconnected" }; private nextId = 1; private pending = new Map(); + private suppressedOrphanResponses = new Set(); constructor( private readonly codexPath: string, @@ -254,6 +258,10 @@ export class AppServerClient { return this.request("thread/read", { threadId, includeTurns }); } + readFile(path: string, options: { timeoutMs?: number } = {}): Promise { + return this.request("fs/readFile", { path }, options); + } + unarchiveThread(threadId: string): Promise { return this.request("thread/unarchive", { threadId }); } @@ -413,13 +421,18 @@ export class AppServerClient { this.send({ id: requestId, error: { code, message } }); } - private request(method: M, params: ClientRequestParams): Promise { + private request( + method: M, + params: ClientRequestParams, + options: { timeoutMs?: number } = {}, + ): Promise { const id = this.nextId++; const promise = new Promise((resolve, reject) => { const timeout = window.setTimeout(() => { this.pending.delete(id); + this.suppressOrphanResponse(id); reject(new Error(`Codex app-server request timed out: ${method}`)); - }, this.requestTimeoutMs); + }, options.timeoutMs ?? this.requestTimeoutMs); this.pending.set(id, { method, resolve: resolve as (value: unknown) => void, @@ -473,6 +486,7 @@ export class AppServerClient { if ("id" in message) { const pending = this.pending.get(message.id); if (!pending) { + if (this.suppressedOrphanResponses.delete(message.id)) return; this.handlers.onLog(`Orphan app-server response: ${JSON.stringify(message)}`); return; } @@ -501,6 +515,16 @@ export class AppServerClient { pending.reject(error); } this.pending.clear(); + this.suppressedOrphanResponses.clear(); + } + + private suppressOrphanResponse(id: RequestId): void { + this.suppressedOrphanResponses.add(id); + while (this.suppressedOrphanResponses.size > MAX_SUPPRESSED_ORPHAN_RESPONSES) { + const oldest = this.suppressedOrphanResponses.values().next().value; + if (oldest === undefined) break; + this.suppressedOrphanResponses.delete(oldest); + } } private activeTransport(): AppServerTransport | null { diff --git a/src/app-server/rollout-token-usage.ts b/src/app-server/rollout-token-usage.ts new file mode 100644 index 00000000..ca7602e2 --- /dev/null +++ b/src/app-server/rollout-token-usage.ts @@ -0,0 +1,110 @@ +import type { ThreadTokenUsage } from "../generated/app-server/v2/ThreadTokenUsage"; +import type { TokenUsageBreakdown } from "../generated/app-server/v2/TokenUsageBreakdown"; + +export const ROLLOUT_TOKEN_USAGE_READ_TIMEOUT_MS = 2_000; +export const ROLLOUT_TOKEN_USAGE_MAX_BASE64_BYTES = 12 * 1024 * 1024; + +export type RolloutReadFileBase64 = (path: string, options: { timeoutMs: number }) => Promise; + +export async function recoverRolloutTokenUsage( + path: string | null, + readFileBase64: RolloutReadFileBase64, +): Promise { + if (!path || !isAbsolutePath(path)) return null; + + let dataBase64: string; + try { + dataBase64 = await readFileBase64(path, { timeoutMs: ROLLOUT_TOKEN_USAGE_READ_TIMEOUT_MS }); + } catch { + return null; + } + if (dataBase64.length > ROLLOUT_TOKEN_USAGE_MAX_BASE64_BYTES) return null; + + const text = decodeBase64Text(dataBase64); + return text ? parseRolloutTokenUsageJsonl(text) : null; +} + +export function parseRolloutTokenUsageJsonl(text: string): ThreadTokenUsage | null { + const lines = text.split(/\r?\n/); + for (let index = lines.length - 1; index >= 0; index -= 1) { + const line = lines[index]?.trim(); + if (!line) continue; + + let value: unknown; + try { + value = JSON.parse(line); + } catch { + continue; + } + + const usage = tokenUsageFromRolloutRecord(value); + if (usage) return usage; + } + return null; +} + +function tokenUsageFromRolloutRecord(value: unknown): ThreadTokenUsage | null { + const record = objectRecord(value); + if (record?.["type"] !== "event_msg") return null; + const payload = objectRecord(record["payload"]); + if (payload?.["type"] !== "token_count") return null; + const info = objectRecord(payload["info"]); + if (!info) return null; + + const last = tokenUsageBreakdownFromRecord(info["last_token_usage"]); + const total = tokenUsageBreakdownFromRecord(info["total_token_usage"]); + const modelContextWindow = nullableNonNegativeNumber(info["model_context_window"]); + if (!last || !total || modelContextWindow === undefined) return null; + + return { last, total, modelContextWindow }; +} + +function tokenUsageBreakdownFromRecord(value: unknown): TokenUsageBreakdown | null { + const record = objectRecord(value); + if (!record) return null; + const totalTokens = nonNegativeNumber(record["total_tokens"]); + const inputTokens = nonNegativeNumber(record["input_tokens"]); + const cachedInputTokens = nonNegativeNumber(record["cached_input_tokens"]); + const outputTokens = nonNegativeNumber(record["output_tokens"]); + const reasoningOutputTokens = nonNegativeNumber(record["reasoning_output_tokens"]); + if ( + totalTokens === null || + inputTokens === null || + cachedInputTokens === null || + outputTokens === null || + reasoningOutputTokens === null + ) { + return null; + } + return { totalTokens, inputTokens, cachedInputTokens, outputTokens, reasoningOutputTokens }; +} + +function decodeBase64Text(dataBase64: string): string | null { + try { + const binary = atob(dataBase64); + const bytes = new Uint8Array(binary.length); + for (let index = 0; index < binary.length; index += 1) { + bytes[index] = binary.charCodeAt(index); + } + return new TextDecoder().decode(bytes); + } catch { + return null; + } +} + +function isAbsolutePath(path: string): boolean { + return path.startsWith("/") || /^[A-Za-z]:[\\/]/.test(path); +} + +function objectRecord(value: unknown): Record | null { + return value !== null && typeof value === "object" && !Array.isArray(value) ? (value as Record) : null; +} + +function nonNegativeNumber(value: unknown): number | null { + return typeof value === "number" && Number.isFinite(value) && value >= 0 ? value : null; +} + +function nullableNonNegativeNumber(value: unknown): number | null | undefined { + if (value === null) return null; + return nonNegativeNumber(value) ?? undefined; +} diff --git a/src/features/chat/chat-view-controller-assembly.ts b/src/features/chat/chat-view-controller-assembly.ts index 8ca66b24..d81b6da6 100644 --- a/src/features/chat/chat-view-controller-assembly.ts +++ b/src/features/chat/chat-view-controller-assembly.ts @@ -5,6 +5,7 @@ import { ConnectionManager } from "../../app-server/connection-manager"; import type { ArchiveExportAdapter } from "../../domain/threads/export"; import type { RuntimeSnapshot } from "../../runtime/state"; import { currentModel } from "../../runtime/state"; +import { recoverRolloutTokenUsage } from "../../app-server/rollout-token-usage"; import type { ChatState, ChatStateStore } from "./chat-state"; import type { DisplayDetailSection } from "./display/types"; import type { ComposerMetaViewModel } from "./view-model"; @@ -526,6 +527,11 @@ export function createChatViewControllerAssembly(host: ChatViewControllerAssembl forceMessagesToBottom: host.effects.scroll.forceBottom, render: host.effects.render.now, refreshLiveState: host.effects.liveState.refresh, + recoverTokenUsageFromRollout: (path) => + recoverRolloutTokenUsage(path, async (filePath, options) => { + const response = await host.getClient()?.readFile(filePath, options); + return response?.dataBase64 ?? ""; + }), }); threadIdentity = new ThreadIdentityController({ state: threadState, diff --git a/src/features/chat/controllers/state-ports.ts b/src/features/chat/controllers/state-ports.ts index 3bc943f8..2c5d323e 100644 --- a/src/features/chat/controllers/state-ports.ts +++ b/src/features/chat/controllers/state-ports.ts @@ -1,5 +1,6 @@ import type { InitializeResponse } from "../../../generated/app-server/InitializeResponse"; import type { Thread } from "../../../generated/app-server/v2/Thread"; +import type { ThreadTokenUsage } from "../../../generated/app-server/v2/ThreadTokenUsage"; import { activeTurnId, chatTurnBusy, pendingTurnStart, type ChatStateStore, type PendingTurnStart } from "../chat-state"; import type { PendingApproval } from "../approvals/model"; import type { PendingUserInput } from "../user-input/model"; @@ -42,6 +43,8 @@ export interface ThreadLifecycleStatePort { restorePlaceholder(threadId: string, item: DisplayItem): void; displayItemsEmpty(): boolean; applyResumedThread(response: ThreadActivationResponse, displayItems: readonly DisplayItem[]): void; + applyTokenUsage(threadId: string, tokenUsage: ThreadTokenUsage): boolean; + applyRecoveredTokenUsage(threadId: string, tokenUsage: ThreadTokenUsage): boolean; } export interface SubmissionStateSnapshot { @@ -149,6 +152,17 @@ export function createThreadLifecycleStatePort(stateStore: ChatStateStore): Thre }), ); }, + applyTokenUsage(threadId, tokenUsage) { + if (stateStore.getState().activeThreadId !== threadId) return false; + stateStore.dispatch({ type: "thread/token-usage-set", tokenUsage }); + return true; + }, + applyRecoveredTokenUsage(threadId, tokenUsage) { + const state = stateStore.getState(); + if (state.activeThreadId !== threadId || state.tokenUsage !== null) return false; + stateStore.dispatch({ type: "thread/token-usage-set", tokenUsage }); + return true; + }, }; } diff --git a/src/features/chat/controllers/thread/thread-resume-controller.ts b/src/features/chat/controllers/thread/thread-resume-controller.ts index 661e9b79..e504bfcb 100644 --- a/src/features/chat/controllers/thread/thread-resume-controller.ts +++ b/src/features/chat/controllers/thread/thread-resume-controller.ts @@ -1,4 +1,5 @@ import type { AppServerClient } from "../../../../app-server/client"; +import type { ThreadTokenUsage } from "../../../../generated/app-server/v2/ThreadTokenUsage"; import type { DisplayItem } from "../../display/types"; import type { RestoredThreadController } from "./restored-thread-controller"; import type { ThreadActivationResponse } from "../../thread-resume"; @@ -23,6 +24,7 @@ export interface ThreadResumeControllerHost { forceMessagesToBottom: () => void; render: () => void; refreshLiveState: () => void; + recoverTokenUsageFromRollout?: (path: string) => Promise; } export class ThreadResumeController { @@ -42,6 +44,7 @@ export class ThreadResumeController { const response = await client.resumeThread(threadId, this.host.vaultPath); if (this.isStale(resume)) return; this.applyResumedThread(response); + this.recoverResumedThreadTokenUsage(response.thread.id, response.thread.path, resume); if (response.initialTurnsPage) { this.host.history.applyLatestPage(response.thread.id, response.initialTurnsPage); } else { @@ -71,6 +74,19 @@ export class ThreadResumeController { this.host.refreshLiveState(); } + private recoverResumedThreadTokenUsage(threadId: string, path: string | null, resume: ActiveChatResume): void { + if (!path || !this.host.recoverTokenUsageFromRollout) return; + void this.host + .recoverTokenUsageFromRollout(path) + .then((tokenUsage) => { + if (!tokenUsage || this.isStale(resume)) return; + if (!this.host.state.applyRecoveredTokenUsage(threadId, tokenUsage)) return; + this.host.refreshLiveState(); + this.host.render(); + }) + .catch(() => undefined); + } + private isStale(resume: ActiveChatResume): boolean { return this.host.resumeWork.isStale(resume) || this.host.closing(); } diff --git a/tests/app-server/app-server-client.test.ts b/tests/app-server/app-server-client.test.ts index 9073b028..14ea44bd 100644 --- a/tests/app-server/app-server-client.test.ts +++ b/tests/app-server/app-server-client.test.ts @@ -1,4 +1,4 @@ -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AppServerClient } from "../../src/app-server/client"; import type { AppServerRpcError } from "../../src/app-server/client"; @@ -90,6 +90,11 @@ describe("AppServerClient", () => { }); }); + afterEach(() => { + vi.useRealTimers(); + vi.unstubAllGlobals(); + }); + it("routes responses, notifications, and server requests", async () => { let transport: FakeTransport; const getTransport = () => transport; @@ -214,6 +219,110 @@ describe("AppServerClient", () => { await expect(steering).resolves.toEqual({ turnId: "turn-1" }); }); + it("reads files through app-server fs requests", async () => { + const { client, transport } = await connectedClient(); + + const reading = client.readFile("/tmp/rollout.jsonl"); + + await expectRequest( + transport, + reading, + { method: "fs/readFile", params: { path: "/tmp/rollout.jsonl" } }, + { dataBase64: btoa("hello") }, + ); + await expect(reading).resolves.toEqual({ dataBase64: btoa("hello") }); + }); + + it("suppresses late responses after per-request timeouts", async () => { + vi.useFakeTimers(); + vi.stubGlobal("window", { + clearTimeout, + setTimeout, + }); + const logs: string[] = []; + let transport!: FakeTransport; + const client = new AppServerClient( + "/bin/codex", + "/vault", + { + onNotification: () => undefined, + onServerRequest: () => undefined, + onLog: (message) => logs.push(message), + onExit: () => undefined, + }, + 500, + (handlers) => { + transport = new FakeTransport(handlers); + return transport; + }, + ); + const connecting = client.connect(); + transport.emitLine({ id: 1, result: { codexHome: "/tmp/codex" } satisfies Partial }); + await connecting; + + const reading = client.readFile("/tmp/slow.jsonl", { timeoutMs: 10 }); + const rejection = expect(reading).rejects.toThrow("Codex app-server request timed out: fs/readFile"); + const sent = latestSent(transport); + if (!("id" in sent) || typeof sent.id !== "number") throw new Error("Expected an app-server request id."); + + await vi.advanceTimersByTimeAsync(10); + await rejection; + + transport.emitLine({ id: sent.id, result: { dataBase64: btoa("late") } }); + expect(logs).toEqual([]); + }); + + it("bounds timed-out response suppression when responses never arrive", async () => { + vi.useFakeTimers(); + vi.stubGlobal("window", { + clearTimeout, + setTimeout, + }); + const logs: string[] = []; + let transport!: FakeTransport; + const client = new AppServerClient( + "/bin/codex", + "/vault", + { + onNotification: () => undefined, + onServerRequest: () => undefined, + onLog: (message) => logs.push(message), + onExit: () => undefined, + }, + 500, + (handlers) => { + transport = new FakeTransport(handlers); + return transport; + }, + ); + const connecting = client.connect(); + transport.emitLine({ id: 1, result: { codexHome: "/tmp/codex" } satisfies Partial }); + await connecting; + + const timedOutRequests: { id: number; rejection: Promise }[] = []; + for (let index = 0; index < 257; index += 1) { + const promise = client.readFile(`/tmp/slow-${String(index)}.jsonl`, { timeoutMs: 10 }); + const rejection = expect(promise).rejects.toThrow("Codex app-server request timed out"); + const sent = latestSent(transport); + if (!("id" in sent) || typeof sent.id !== "number") throw new Error("Expected an app-server request id."); + timedOutRequests.push({ id: sent.id, rejection }); + } + + await vi.advanceTimersByTimeAsync(10); + await Promise.all(timedOutRequests.map(({ rejection }) => rejection)); + + const firstTimedOutRequest = timedOutRequests[0]; + const lastTimedOutRequest = timedOutRequests.at(-1); + if (!firstTimedOutRequest || !lastTimedOutRequest) throw new Error("Expected timed-out requests."); + + transport.emitLine({ id: firstTimedOutRequest.id, result: { dataBase64: btoa("evicted") } }); + expect(logs).toHaveLength(1); + expect(logs[0]).toContain("Orphan app-server response"); + + transport.emitLine({ id: lastTimedOutRequest.id, result: { dataBase64: btoa("suppressed") } }); + expect(logs).toHaveLength(1); + }); + it("sends golden thread and turn request payloads", async () => { const { client, transport } = await connectedClient(); diff --git a/tests/app-server/rollout-token-usage.test.ts b/tests/app-server/rollout-token-usage.test.ts new file mode 100644 index 00000000..13030bf9 --- /dev/null +++ b/tests/app-server/rollout-token-usage.test.ts @@ -0,0 +1,101 @@ +import { describe, expect, it, vi } from "vitest"; + +import { + parseRolloutTokenUsageJsonl, + recoverRolloutTokenUsage, + ROLLOUT_TOKEN_USAGE_MAX_BASE64_BYTES, + ROLLOUT_TOKEN_USAGE_READ_TIMEOUT_MS, +} from "../../src/app-server/rollout-token-usage"; + +describe("rollout token usage recovery", () => { + it("parses the last valid token_count event", () => { + const first = tokenCountLine({ input: 100, total: 120, context: 1000 }); + const second = tokenCountLine({ input: 250, total: 300, context: 2000 }); + + expect(parseRolloutTokenUsageJsonl(["not json", first, '{"type":"response_item","payload":{}}', second, ""].join("\n"))).toEqual({ + last: { + inputTokens: 250, + cachedInputTokens: 25, + outputTokens: 10, + reasoningOutputTokens: 5, + totalTokens: 300, + }, + total: { + inputTokens: 500, + cachedInputTokens: 50, + outputTokens: 20, + reasoningOutputTokens: 10, + totalTokens: 600, + }, + modelContextWindow: 2000, + }); + }); + + it("returns null for missing or invalid token usage shapes", () => { + expect(parseRolloutTokenUsageJsonl("")).toBeNull(); + expect(parseRolloutTokenUsageJsonl('{"type":"event_msg","payload":{"type":"agent_message"}}')).toBeNull(); + expect( + parseRolloutTokenUsageJsonl( + JSON.stringify({ + type: "event_msg", + payload: { + type: "token_count", + info: { + last_token_usage: { input_tokens: -1 }, + total_token_usage: {}, + model_context_window: 1000, + }, + }, + }), + ), + ).toBeNull(); + }); + + it("recovers usage from an absolute rollout path through app-server file reads", async () => { + const readFileBase64 = vi.fn().mockResolvedValue(btoa(tokenCountLine({ input: 42, total: 50, context: 1000 }))); + + await expect(recoverRolloutTokenUsage("/tmp/rollout.jsonl", readFileBase64)).resolves.toMatchObject({ + last: { inputTokens: 42, totalTokens: 50 }, + modelContextWindow: 1000, + }); + expect(readFileBase64).toHaveBeenCalledWith("/tmp/rollout.jsonl", { timeoutMs: ROLLOUT_TOKEN_USAGE_READ_TIMEOUT_MS }); + }); + + it("skips relative paths, read failures, invalid base64, and oversized payloads", async () => { + const readFileBase64 = vi.fn().mockResolvedValue(btoa(tokenCountLine({ input: 42, total: 50, context: 1000 }))); + await expect(recoverRolloutTokenUsage("relative.jsonl", readFileBase64)).resolves.toBeNull(); + expect(readFileBase64).not.toHaveBeenCalled(); + + await expect(recoverRolloutTokenUsage("/tmp/rollout.jsonl", vi.fn().mockRejectedValue(new Error("missing")))).resolves.toBeNull(); + await expect(recoverRolloutTokenUsage("/tmp/rollout.jsonl", vi.fn().mockResolvedValue("%%%"))).resolves.toBeNull(); + await expect( + recoverRolloutTokenUsage("/tmp/rollout.jsonl", vi.fn().mockResolvedValue("a".repeat(ROLLOUT_TOKEN_USAGE_MAX_BASE64_BYTES + 1))), + ).resolves.toBeNull(); + }); +}); + +function tokenCountLine(options: { input: number; total: number; context: number }): string { + return JSON.stringify({ + type: "event_msg", + payload: { + type: "token_count", + info: { + last_token_usage: { + input_tokens: options.input, + cached_input_tokens: 25, + output_tokens: 10, + reasoning_output_tokens: 5, + total_tokens: options.total, + }, + total_token_usage: { + input_tokens: options.input * 2, + cached_input_tokens: 50, + output_tokens: 20, + reasoning_output_tokens: 10, + total_tokens: options.total * 2, + }, + model_context_window: options.context, + }, + }, + }); +} diff --git a/tests/features/chat/controllers/thread/restored-thread-controller.test.ts b/tests/features/chat/controllers/thread/restored-thread-controller.test.ts index 64c3a51f..9820d65a 100644 --- a/tests/features/chat/controllers/thread/restored-thread-controller.test.ts +++ b/tests/features/chat/controllers/thread/restored-thread-controller.test.ts @@ -82,6 +82,8 @@ function restoredThreadState(overrides: Partial = {}): restorePlaceholder: vi.fn(), displayItemsEmpty: () => true, applyResumedThread: vi.fn(), + applyTokenUsage: vi.fn(), + applyRecoveredTokenUsage: vi.fn(), ...overrides, }; } diff --git a/tests/features/chat/controllers/thread/thread-resume-controller.test.ts b/tests/features/chat/controllers/thread/thread-resume-controller.test.ts index 69f1b972..4a0d4377 100644 --- a/tests/features/chat/controllers/thread/thread-resume-controller.test.ts +++ b/tests/features/chat/controllers/thread/thread-resume-controller.test.ts @@ -10,6 +10,7 @@ import type { ThreadHistoryLoader } from "../../../../../src/features/chat/threa import { ChatResumeWorkTracker } from "../../../../../src/features/chat/view-lifecycle"; import type { ThreadItem } from "../../../../../src/generated/app-server/v2/ThreadItem"; import type { Thread } from "../../../../../src/generated/app-server/v2/Thread"; +import type { ThreadTokenUsage } from "../../../../../src/generated/app-server/v2/ThreadTokenUsage"; import type { Turn } from "../../../../../src/generated/app-server/v2/Turn"; function thread(id: string): Thread { @@ -50,7 +51,10 @@ function activation(threadId: string): ThreadActivationResponse { }; } -function createController(response: ThreadActivationResponse = activation("thread")) { +function createController( + response: ThreadActivationResponse = activation("thread"), + overrides: Partial[0]> = {}, +) { const stateStore = createChatStateStore(createChatState()); const resumeThread = vi.fn().mockResolvedValue(response); const client = { resumeThread } as unknown as AppServerClient; @@ -74,6 +78,7 @@ function createController(response: ThreadActivationResponse = activation("threa forceMessagesToBottom: vi.fn(), render: vi.fn(), refreshLiveState: vi.fn(), + ...overrides, }; return { controller: new ThreadResumeController(host), host, applyLatestPage, loadLatest, restoredClear, resumeThread, stateStore }; } @@ -129,6 +134,79 @@ describe("ThreadResumeController", () => { expect(resumeThread).not.toHaveBeenCalled(); expect(host.addSystemMessage).toHaveBeenCalledWith("Finish or interrupt the current turn before switching threads."); }); + + it("recovers rollout token usage without blocking latest history loading", async () => { + const response = activation("thread"); + response.thread.path = "/tmp/rollout.jsonl"; + const recovery = deferred(); + const recoverTokenUsageFromRollout = vi.fn().mockReturnValue(recovery.promise); + const { controller, loadLatest, stateStore } = createController(response, { recoverTokenUsageFromRollout }); + + await controller.resumeThread("thread"); + + expect(recoverTokenUsageFromRollout).toHaveBeenCalledWith("/tmp/rollout.jsonl"); + expect(loadLatest).toHaveBeenCalledWith("thread"); + expect(stateStore.getState().tokenUsage).toBeNull(); + + await recovery.resolveAndFlush(tokenUsageFixture(42)); + + expect(stateStore.getState().tokenUsage).toMatchObject({ last: { inputTokens: 42 } }); + }); + + it("ignores stale rollout token usage recovery", async () => { + const first = activation("thread"); + first.thread.path = "/tmp/thread.jsonl"; + const second = activation("other"); + const recovery = deferred(); + const recoverTokenUsageFromRollout = vi.fn().mockReturnValue(recovery.promise); + const { controller, stateStore } = createController(first, { recoverTokenUsageFromRollout }); + + await controller.resumeThread("thread"); + stateStore.dispatch({ + type: "thread/resumed", + thread: second.thread, + cwd: "/vault", + model: null, + reasoningEffort: null, + serviceTier: null, + approvalPolicy: null, + approvalsReviewer: null, + activePermissionProfile: null, + }); + + await recovery.resolveAndFlush(tokenUsageFixture(42)); + + expect(stateStore.getState().activeThreadId).toBe("other"); + expect(stateStore.getState().tokenUsage).toBeNull(); + }); + + it("does not let late rollout token usage recovery overwrite live token usage", async () => { + const response = activation("thread"); + response.thread.path = "/tmp/rollout.jsonl"; + const recovery = deferred(); + const recoverTokenUsageFromRollout = vi.fn().mockReturnValue(recovery.promise); + const { controller, stateStore } = createController(response, { recoverTokenUsageFromRollout }); + + await controller.resumeThread("thread"); + stateStore.dispatch({ type: "thread/token-usage-set", tokenUsage: tokenUsageFixture(99) }); + + await recovery.resolveAndFlush(tokenUsageFixture(42)); + + expect(stateStore.getState().tokenUsage).toMatchObject({ last: { inputTokens: 99 } }); + }); + + it("ignores rollout token usage recovery failures", async () => { + const response = activation("thread"); + response.thread.path = "/tmp/rollout.jsonl"; + const recoverTokenUsageFromRollout = vi.fn().mockRejectedValue(new Error("read failed")); + const { controller, host, stateStore } = createController(response, { recoverTokenUsageFromRollout }); + + await controller.resumeThread("thread"); + await Promise.resolve(); + + expect(stateStore.getState().tokenUsage).toBeNull(); + expect(host.addSystemMessage).not.toHaveBeenCalledWith("read failed"); + }); }); function turnFixture(items: ThreadItem[]): Turn { @@ -147,3 +225,25 @@ function turnFixture(items: ThreadItem[]): Turn { function userMessage(id: string, text: string): ThreadItem { return { type: "userMessage", id, clientId: null, content: [{ type: "text", text, text_elements: [] }] }; } + +function tokenUsageFixture(inputTokens: number): ThreadTokenUsage { + return { + last: { inputTokens, cachedInputTokens: 0, outputTokens: 2, reasoningOutputTokens: 0, totalTokens: inputTokens + 2 }, + total: { inputTokens, cachedInputTokens: 0, outputTokens: 2, reasoningOutputTokens: 0, totalTokens: inputTokens + 2 }, + modelContextWindow: 1000, + }; +} + +function deferred(): { promise: Promise; resolveAndFlush: (value: T) => Promise } { + let resolve!: (value: T) => void; + const promise = new Promise((innerResolve) => { + resolve = innerResolve; + }); + return { + promise, + async resolveAndFlush(value) { + resolve(value); + await Promise.resolve(); + }, + }; +}