Recover resumed context usage from rollout logs

This commit is contained in:
murashit 2026-06-06 13:52:13 +09:00
parent c9c1622b2d
commit 6a4c36d50d
9 changed files with 486 additions and 4 deletions

View file

@ -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<ClientRequestMethod, keyof ClientResponseByMethod>;
@ -114,6 +117,7 @@ export class AppServerClient {
private lifecycle: AppServerClientLifecycleState = { kind: "disconnected" };
private nextId = 1;
private pending = new Map<RequestId, PendingRequest>();
private suppressedOrphanResponses = new Set<RequestId>();
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<FsReadFileResponse> {
return this.request("fs/readFile", { path }, options);
}
unarchiveThread(threadId: string): Promise<ThreadUnarchiveResponse> {
return this.request("thread/unarchive", { threadId });
}
@ -413,13 +421,18 @@ export class AppServerClient {
this.send({ id: requestId, error: { code, message } });
}
private request<M extends TypedClientRequestMethod>(method: M, params: ClientRequestParams<M>): Promise<ClientResponseByMethod[M]> {
private request<M extends TypedClientRequestMethod>(
method: M,
params: ClientRequestParams<M>,
options: { timeoutMs?: number } = {},
): Promise<ClientResponseByMethod[M]> {
const id = this.nextId++;
const promise = new Promise<ClientResponseByMethod[M]>((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 {

View file

@ -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<string>;
export async function recoverRolloutTokenUsage(
path: string | null,
readFileBase64: RolloutReadFileBase64,
): Promise<ThreadTokenUsage | null> {
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<string, unknown> | null {
return value !== null && typeof value === "object" && !Array.isArray(value) ? (value as Record<string, unknown>) : 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;
}

View file

@ -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,

View file

@ -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;
},
};
}

View file

@ -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<ThreadTokenUsage | null>;
}
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();
}

View file

@ -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<InitializeResponse> });
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<InitializeResponse> });
await connecting;
const timedOutRequests: { id: number; rejection: Promise<void> }[] = [];
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();

View file

@ -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,
},
},
});
}

View file

@ -82,6 +82,8 @@ function restoredThreadState(overrides: Partial<ThreadLifecycleStatePort> = {}):
restorePlaceholder: vi.fn(),
displayItemsEmpty: () => true,
applyResumedThread: vi.fn(),
applyTokenUsage: vi.fn(),
applyRecoveredTokenUsage: vi.fn(),
...overrides,
};
}

View file

@ -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<ConstructorParameters<typeof ThreadResumeController>[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<ThreadTokenUsage | null>();
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<ThreadTokenUsage | null>();
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<ThreadTokenUsage | null>();
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<T>(): { promise: Promise<T>; resolveAndFlush: (value: T) => Promise<void> } {
let resolve!: (value: T) => void;
const promise = new Promise<T>((innerResolve) => {
resolve = innerResolve;
});
return {
promise,
async resolveAndFlush(value) {
resolve(value);
await Promise.resolve();
},
};
}