Skip to content
2 changes: 1 addition & 1 deletion plugins/codex-security/scripts/workbench_scan_usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -456,7 +456,7 @@ def _read_rollout_usage(
continue
if event.get("type") != "event_msg" or payload.get("type") != "token_count":
continue
if payload.get("info") is None:
if "info" in payload and payload["info"] is None:
continue
timestamp = _timestamp(event.get("timestamp"))
snapshot = _token_snapshot(payload)
Expand Down
28 changes: 28 additions & 0 deletions plugins/codex-security/tests/test_workbench_scan_usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -323,6 +323,34 @@ def test_completion_counts_only_scan_owned_parent_and_descendants(
assert json.loads(stored[0]) == {"usage": expected}


@pytest.mark.parametrize("measurement", ["unknown", "malformed", "missing"])
def test_completion_distinguishes_unknown_usage_from_invalid_measurements(
tmp_path: Path, measurement: str
) -> None:
fixture = _start_scan(tmp_path)
counted = fixture.started_at + timedelta(microseconds=1)
payload: dict[str, Any] = {"type": "token_count"}
if measurement == "unknown":
payload["info"] = None
payload["rate_limits"] = {"primary": {"used_percent": 5}}
elif measurement == "malformed":
payload["info"] = {"total_token_usage": {"input_tokens": "invalid"}}
parent = _rollout(
tmp_path,
"scan-parent",
[
_token_event(counted, 10, 2),
_event(counted, "event_msg", payload),
_token_event(counted, 20, 4),
],
)
_state_graph(fixture.environment, {"scan-parent": parent}, [])
usage = _complete_scan(fixture)["scan"]["usage"]
assert usage["totalTokens"] == 24
assert usage["coverage"] == ("complete" if measurement == "unknown" else "partial")
assert ("token_record_invalid" in usage.get("warnings", [])) is (measurement != "unknown")


@pytest.mark.parametrize("cache_write_field", ["cache_write_input_tokens", "cache_write_tokens"])
def test_scan_usage_helper_preserves_cached_token_totals(
cache_write_field: str,
Expand Down
106 changes: 67 additions & 39 deletions sdk/typescript/src/cost.ts
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import { createHash } from "node:crypto";
import { open, readdir, stat } from "node:fs/promises";
import { join } from "node:path";
import { isRecord } from "./record.js";
Expand Down Expand Up @@ -38,6 +39,7 @@ interface SessionReasoning {
}

interface SessionUsage {
tracked: boolean;
offset: number;
pendingLine: Buffer[];
pendingLineBytes: number;
Expand All @@ -51,7 +53,7 @@ interface SessionUsage {
usage: ScanTokenUsage | null;
calls: Map<string, ScanActivity>;
activities: ScanActivity[];
progress: ScanProgress[];
progress?: ScanProgress[];
filesCompleted: number;
filesTotal: number | null;
prose: Set<string>;
Expand Down Expand Up @@ -85,6 +87,7 @@ const SESSION_READ_SIZE = 64 * 1_024;

function createSessionUsage(): SessionUsage {
return {
tracked: false,
offset: 0,
pendingLine: [],
pendingLineBytes: 0,
Expand All @@ -98,7 +101,6 @@ function createSessionUsage(): SessionUsage {
usage: null,
calls: new Map(),
activities: [],
progress: [],
filesCompleted: 0,
filesTotal: null,
prose: new Set(),
Expand Down Expand Up @@ -202,7 +204,6 @@ export class ScanCostTracker {

async #readSessions(): Promise<void> {
if (this.#threadId === null) return;
const unreadable: Array<{ session: SessionUsage; error: unknown }> = [];
for await (const path of sessionFiles(
join(this.#options.codexHome, "sessions"),
)) {
Expand All @@ -211,11 +212,10 @@ export class ScanCostTracker {
session = createSessionUsage();
this.#sessions.set(path, session);
}
try {
await readSessionUsage(path, session, this.#options.repository);
} catch (error) {
if (session.threadId === null) throw error;
unreadable.push({ session, error });
if (session.threadId === null) {
// Index ownership before reading a transcript. Unrelated sessions need
// only their metadata, including parents discovered by a later poll.
await readSessionUsage(path, session, undefined, true);
}
}

Expand Down Expand Up @@ -259,30 +259,34 @@ export class ScanCostTracker {
}
}
} while (included.size !== previousSize);
for (const { session, error } of unreadable) {
if (included.has(session.threadId!)) throw error;
}

const usages = new Map(this.#receipts);
for (const [path, tracked] of this.#sessions) {
const threadId = tracked.threadId;
if (threadId === null || !included.has(threadId)) continue;
let session = tracked;
if (
this.#options.onSessionEvent !== undefined &&
session.events === undefined
) {
// Replay only newly associated sessions, including their early events.
session = createSessionUsage();
session.events = [];
await readSessionUsage(path, session, this.#options.repository);
this.#sessions.set(path, session);
}
let worker: number | undefined;
if (threadId !== this.#threadId) {
worker = this.#workers.get(threadId) ?? this.#workers.size + 1;
this.#workers.set(threadId, worker);
}
if (!session.tracked) {
// Replay newly associated sessions from the start for every observer,
// not just raw session events. Their early usage and activity matter.
session = createSessionUsage();
session.tracked = true;
if (this.#options.onSessionEvent !== undefined) session.events = [];
if (worker !== undefined && this.#options.onProgress !== undefined) {
session.progress = [];
}
}
await readSessionUsage(
path,
session,
worker !== undefined && this.#options.onActivity !== undefined
? this.#options.repository
: undefined,
);
this.#sessions.set(path, session);
for (const event of session.events?.splice(0) ?? []) {
this.#options.onSessionEvent?.({
threadId,
Expand Down Expand Up @@ -331,7 +335,7 @@ export class ScanCostTracker {
if (this.#options.onProgress === undefined || session.threadId === null) {
return;
}
for (const progress of session.progress.splice(0)) {
for (const progress of session.progress?.splice(0) ?? []) {
const expectedFilesTotal = this.#expectedFilesTotal;
if (
(expectedFilesTotal !== undefined &&
Expand Down Expand Up @@ -405,6 +409,7 @@ async function readSessionUsage(
path: string,
session: SessionUsage,
repository?: string,
metadataOnly = false,
): Promise<void> {
if (session.unreadable) return;
let file;
Expand All @@ -427,13 +432,19 @@ async function readSessionUsage(
if (bytesRead === 0) return;
session.offset += bytesRead;
try {
readSessionChunk(buffer.subarray(0, bytesRead), session, repository);
readSessionChunk(
buffer.subarray(0, bytesRead),
session,
repository,
metadataOnly,
);
} catch (error) {
session.unreadable = true;
session.pendingLine = [];
session.pendingLineBytes = 0;
throw error;
}
if (metadataOnly && session.threadId !== null) return;
}
} finally {
await file.close();
Expand All @@ -444,6 +455,7 @@ function readSessionChunk(
contents: Buffer,
session: SessionUsage,
repository?: string,
metadataOnly = false,
): void {
let lineStart = 0;
while (lineStart < contents.length) {
Expand All @@ -461,17 +473,24 @@ function readSessionChunk(
}

if (session.pendingLineBytes === 0) {
readSessionEvent(fragment.toString("utf8"), session, repository);
readSessionEvent(
fragment.toString("utf8"),
session,
repository,
metadataOnly,
);
} else {
if (fragment.length > 0) session.pendingLine.push(Buffer.from(fragment));
readSessionEvent(
Buffer.concat(session.pendingLine, lineBytes).toString("utf8"),
session,
repository,
metadataOnly,
);
session.pendingLine = [];
session.pendingLineBytes = 0;
}
if (metadataOnly && session.threadId !== null) return;
lineStart = newline + 1;
}
}
Expand All @@ -480,6 +499,7 @@ function readSessionEvent(
line: string,
session: SessionUsage,
repository?: string,
metadataOnly = false,
): void {
if (line.length === 0) return;
let event: unknown;
Expand Down Expand Up @@ -507,6 +527,7 @@ function readSessionEvent(
session.events?.push(event);
return;
}
if (metadataOnly) return;
if (session.replaying) {
if (event["type"] !== "event_msg") return;
if (payload["type"] === "token_count" && isRecord(payload["info"])) {
Expand All @@ -532,7 +553,7 @@ function readSessionEvent(
}
session.events?.push(event);
if (event["type"] === "response_item") {
session.progress.push(...sessionProgressUpdates(payload));
session.progress?.push(...sessionProgressUpdates(payload));
if (repository === undefined) return;
if (
payload["type"] === "reasoning" &&
Expand All @@ -553,10 +574,7 @@ function readSessionEvent(
},
repository,
);
if (
activity === null ||
session.prose.has(`${activity.kind}:${activity.description}`)
) {
if (activity === null || session.prose.has(proseKey(activity))) {
continue;
}
session.reasoning = {
Expand Down Expand Up @@ -591,12 +609,12 @@ function readSessionEvent(
session.reasoning = null;
if (
activity.kind === "message" &&
session.prose.has(`${activity.kind}:${activity.description}`)
session.prose.has(proseKey(activity))
) {
return;
}
if (activity.kind === "message") {
session.prose.add(`${activity.kind}:${activity.description}`);
session.prose.add(proseKey(activity));
}
if (activity.status === "running") {
session.calls.set(activity.id, activity);
Expand Down Expand Up @@ -632,7 +650,9 @@ function readSessionEvent(
payload["type"] === "agent_message" &&
typeof payload["message"] === "string"
) {
session.progress.push(...scanProgressUpdatesFromText(payload["message"]));
session.progress?.push(
...scanProgressUpdatesFromText(payload["message"]),
);
}
if (repository === undefined) return;
if (payload["type"] !== "agent_message") {
Expand All @@ -641,11 +661,8 @@ function readSessionEvent(
}
session.reasoning = null;
const activity = scanActivityFromSessionEvent(event, repository);
if (
activity !== null &&
!session.prose.has(`${activity.kind}:${activity.description}`)
) {
session.prose.add(`${activity.kind}:${activity.description}`);
if (activity !== null && !session.prose.has(proseKey(activity))) {
session.prose.add(proseKey(activity));
session.activities.push(activity);
}
return;
Expand Down Expand Up @@ -742,10 +759,21 @@ function recordReasoningActivity(
return;
}
reasoning.activity = activity;
session.prose.add(`${activity.kind}:${activity.description}`);
session.prose.add(proseKey(activity));
session.activities.push(activity);
}

function proseKey(activity: ScanActivity): string {
// Deduplication needs an identity, not a retained copy of every transcript
// message and every expanding reasoning prefix.
// Hash UTF-16 code units to keep distinct lone surrogates distinct too.
return createHash("sha256")
.update(activity.kind)
.update(":")
.update(activity.description, "utf16le")
.digest("hex");
}

function sessionProgressUpdates(
payload: Readonly<Record<string, unknown>>,
): ScanProgress[] {
Expand Down
Loading
Loading