coding-agent-cost 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,368 @@
1
+ """Codex CLI reader: ``state_*.sqlite`` threads + rollout JSONL -> facts.
2
+
3
+ The state DB records one row per thread (roughly: one Codex session); each
4
+ thread points at a rollout JSONL file holding a running ``token_count``
5
+ snapshot after every turn. A fact is the *increment* between two
6
+ consecutive snapshots, computed independently per token channel so that an
7
+ anomaly in one channel (e.g. a session reset that makes the cumulative
8
+ total go backwards) doesn't corrupt the others.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ import os
15
+ import re
16
+ import sqlite3
17
+ import tempfile
18
+ from datetime import datetime, timedelta, timezone
19
+ from pathlib import Path
20
+ from typing import Optional, Tuple
21
+
22
+ from . import ReadResult
23
+ from ..facts import Fact, normalize_model_key
24
+
25
+ _FAST_MODEL_RE = re.compile(r"-fast\b", re.IGNORECASE)
26
+ _NON_FAST_MODES = {"default", "standard", "normal"}
27
+
28
+
29
+ def _detect_mode(model: Optional[str], collaboration_mode: Optional[str]) -> str:
30
+ """Best-effort Fast Mode classification.
31
+
32
+ An unrecognized ``collaboration_mode`` value is left as ``"unknown"``
33
+ rather than assumed non-fast, so a future mode name doesn't silently
34
+ get costed at standard rates.
35
+ """
36
+ if model and _FAST_MODEL_RE.search(str(model)):
37
+ return "fast"
38
+ if collaboration_mode:
39
+ cm = str(collaboration_mode).strip().lower()
40
+ if "fast" in cm:
41
+ return "fast"
42
+ if cm in _NON_FAST_MODES:
43
+ return "normal"
44
+ return "unknown"
45
+
46
+
47
+ def snapshot_db(src: Path) -> Path:
48
+ """Copy the Codex state DB to a temp file using sqlite3's backup API.
49
+
50
+ The backup API produces a consistent snapshot even while a writer
51
+ holds the source DB open in WAL mode, without shelling out to the
52
+ ``sqlite3`` binary.
53
+ """
54
+ if not src.exists():
55
+ raise FileNotFoundError(f"Codex state DB not found: {src}")
56
+ fd, dst_str = tempfile.mkstemp(prefix="agent-cost-codex-snapshot-", suffix=".sqlite")
57
+ os.close(fd)
58
+ dst = Path(dst_str)
59
+ dst.unlink(missing_ok=True)
60
+ src_conn = sqlite3.connect(str(src))
61
+ try:
62
+ dst_conn = sqlite3.connect(str(dst))
63
+ try:
64
+ src_conn.backup(dst_conn)
65
+ finally:
66
+ dst_conn.close()
67
+ finally:
68
+ src_conn.close()
69
+ return dst
70
+
71
+
72
+ def _table_columns(conn: sqlite3.Connection, table: str) -> set:
73
+ return {row[1] for row in conn.execute(f"PRAGMA table_info({table})")}
74
+
75
+
76
+ def fetch_threads(
77
+ snapshot: Path,
78
+ *,
79
+ include_archived: bool = True,
80
+ since_ms: Optional[int] = None,
81
+ until_ms: Optional[int] = None,
82
+ ) -> list:
83
+ """Read rows of the ``threads`` table from a snapshot DB.
84
+
85
+ Column presence is feature-detected via ``PRAGMA table_info`` so an
86
+ older schema lacking ``archived`` or ``created_at_ms`` still loads
87
+ (those filters are simply skipped when the column is absent).
88
+
89
+ ``since_ms``/``until_ms`` are an optional, opt-in convenience for
90
+ callers that only care about thread *creation* time (e.g. ad-hoc
91
+ inspection). ``read_codex_facts`` deliberately does not use them for
92
+ its window filtering: a thread created before a report window can
93
+ still emit usage inside it, so the only correct place to apply
94
+ since/until is per-fact, against each fact's own timestamp.
95
+ """
96
+ conn = sqlite3.connect(str(snapshot))
97
+ conn.row_factory = sqlite3.Row
98
+ try:
99
+ columns = _table_columns(conn, "threads")
100
+ where = []
101
+ params: list = []
102
+ if not include_archived and "archived" in columns:
103
+ where.append("archived = 0")
104
+ if "created_at_ms" in columns:
105
+ if since_ms is not None:
106
+ where.append("created_at_ms >= ?")
107
+ params.append(since_ms)
108
+ if until_ms is not None:
109
+ where.append("created_at_ms < ?")
110
+ params.append(until_ms)
111
+ sql = "SELECT * FROM threads"
112
+ if where:
113
+ sql += " WHERE " + " AND ".join(where)
114
+ if "created_at_ms" in columns:
115
+ sql += " ORDER BY created_at_ms DESC"
116
+ return [dict(row) for row in conn.execute(sql, params)]
117
+ finally:
118
+ conn.close()
119
+
120
+
121
+ def _to_utc(ts: Optional[str]) -> Optional[datetime]:
122
+ if not ts:
123
+ return None
124
+ try:
125
+ text = ts[:-1] + "+00:00" if ts.endswith("Z") else ts
126
+ dt = datetime.fromisoformat(text)
127
+ except ValueError:
128
+ return None
129
+ if dt.tzinfo is None:
130
+ dt = dt.replace(tzinfo=timezone.utc)
131
+ return dt.astimezone(timezone.utc)
132
+
133
+
134
+ def parse_rollout_facts(
135
+ rollout_path: Path,
136
+ *,
137
+ model_raw: Optional[str],
138
+ session_id: Optional[str],
139
+ ) -> Tuple[list, int, int]:
140
+ """Convert a rollout JSONL's cumulative ``token_count`` events into
141
+ per-event delta facts.
142
+
143
+ Returns ``(facts, malformed_events, negative_deltas)``. The first
144
+ ``token_count`` event (``total: null``) is an initialization marker
145
+ and contributes no delta. When a channel's increment relative to the
146
+ previous snapshot is negative, that channel's fact for this event is
147
+ dropped and counted in ``negative_deltas`` instead of being clamped to
148
+ zero; the running baseline still advances so the anomaly isn't
149
+ repeated on every subsequent event.
150
+
151
+ ``output`` is ``output_tokens`` alone. Cross-checking real rollout
152
+ files confirms ``total_tokens == input_tokens + output_tokens`` for
153
+ every sample observed; ``reasoning_output_tokens`` is a breakdown of
154
+ (already-counted) output tokens, not an additional charge, so adding
155
+ it in would double-count reasoning tokens.
156
+
157
+ The very first successfully-computed delta in a rollout is measured
158
+ against an assumed-zero baseline (there is no earlier snapshot to
159
+ diff against), which is correct for a thread that really did start
160
+ from nothing but would overstate usage for one resumed from a prior
161
+ context the rollout doesn't capture. Facts from that first delta
162
+ carry ``source_quality="first_event_delta"`` so that caveat is
163
+ visible downstream instead of looking identical to an ordinary turn.
164
+ """
165
+ facts: list = []
166
+ malformed = 0
167
+ negative_deltas = 0
168
+ if not rollout_path.exists():
169
+ return facts, malformed, negative_deltas
170
+
171
+ model_key = normalize_model_key(model_raw)
172
+ collaboration_mode: Optional[str] = None
173
+ last_input_total = 0
174
+ last_cached_total = 0
175
+ last_output_total = 0
176
+ seen_first = False
177
+
178
+ with rollout_path.open() as fh:
179
+ for line in fh:
180
+ line = line.strip()
181
+ if not line:
182
+ continue
183
+ try:
184
+ event = json.loads(line)
185
+ except json.JSONDecodeError:
186
+ malformed += 1
187
+ continue
188
+ if event.get("type") != "event_msg":
189
+ continue
190
+ payload = event.get("payload") or {}
191
+ ptype = payload.get("type")
192
+ if ptype == "task_started":
193
+ cm = payload.get("collaboration_mode_kind")
194
+ if cm is not None and collaboration_mode is None:
195
+ collaboration_mode = str(cm)
196
+ continue
197
+ if ptype != "token_count":
198
+ continue
199
+
200
+ info = payload.get("info") or {}
201
+ total = info.get("total_token_usage")
202
+ if total is None:
203
+ continue # initialization event, no cumulative data yet
204
+
205
+ occurred_at = _to_utc(event.get("timestamp"))
206
+ if occurred_at is None:
207
+ malformed += 1
208
+ continue
209
+
210
+ input_total = int(total.get("input_tokens") or 0)
211
+ cached_total = int(total.get("cached_input_tokens") or 0)
212
+ output_total = int(total.get("output_tokens") or 0)
213
+
214
+ is_first_delta = not seen_first
215
+ if is_first_delta:
216
+ diff_input, diff_cached, diff_output = input_total, cached_total, output_total
217
+ seen_first = True
218
+ else:
219
+ diff_input = input_total - last_input_total
220
+ diff_cached = cached_total - last_cached_total
221
+ diff_output = output_total - last_output_total
222
+
223
+ mode = _detect_mode(model_raw, collaboration_mode)
224
+ quality = "first_event_delta" if is_first_delta else "ok"
225
+
226
+ if diff_input < 0 or diff_cached < 0:
227
+ negative_deltas += 1
228
+ else:
229
+ nocache = max(diff_input - diff_cached, 0)
230
+ if nocache:
231
+ facts.append(
232
+ Fact(
233
+ occurred_at_utc=occurred_at,
234
+ agent="codex",
235
+ session_id=session_id,
236
+ model_raw=model_raw,
237
+ model_key=model_key,
238
+ token_kind="input_nocache",
239
+ tokens=nocache,
240
+ mode=mode,
241
+ source_quality=quality,
242
+ )
243
+ )
244
+ if diff_cached:
245
+ facts.append(
246
+ Fact(
247
+ occurred_at_utc=occurred_at,
248
+ agent="codex",
249
+ session_id=session_id,
250
+ model_raw=model_raw,
251
+ model_key=model_key,
252
+ token_kind="cache_read",
253
+ tokens=diff_cached,
254
+ mode=mode,
255
+ source_quality=quality,
256
+ )
257
+ )
258
+
259
+ if diff_output < 0:
260
+ negative_deltas += 1
261
+ elif diff_output:
262
+ facts.append(
263
+ Fact(
264
+ occurred_at_utc=occurred_at,
265
+ agent="codex",
266
+ session_id=session_id,
267
+ model_raw=model_raw,
268
+ model_key=model_key,
269
+ token_kind="output",
270
+ tokens=diff_output,
271
+ mode=mode,
272
+ source_quality=quality,
273
+ )
274
+ )
275
+
276
+ last_input_total = input_total
277
+ last_cached_total = cached_total
278
+ last_output_total = output_total
279
+
280
+ return facts, malformed, negative_deltas
281
+
282
+
283
+ def read_codex_facts(
284
+ codex_db_path: Path,
285
+ codex_home: Path,
286
+ *,
287
+ since_utc: Optional[datetime] = None,
288
+ until_utc: Optional[datetime] = None,
289
+ include_archived: bool = True,
290
+ ) -> ReadResult:
291
+ """Read every thread's rollout into facts.
292
+
293
+ Threads are never filtered by ``created_at_ms``: a thread created
294
+ before the window can still emit usage events inside it. A rollout
295
+ file's mtime is used only as a coarse "we can skip reading this
296
+ file" optimization on the *since* side (a file untouched since well
297
+ before the window cannot contain events inside it); there is no
298
+ corresponding *until*-side skip, since a file modified after the
299
+ window can still contain earlier in-window events. The exact
300
+ boundary is always re-checked per fact below.
301
+ """
302
+ if not codex_db_path.exists():
303
+ raise FileNotFoundError(f"Codex state DB not found: {codex_db_path}")
304
+
305
+ snapshot = snapshot_db(codex_db_path)
306
+ try:
307
+ threads = fetch_threads(snapshot, include_archived=include_archived)
308
+ finally:
309
+ snapshot.unlink(missing_ok=True)
310
+
311
+ all_facts: list = []
312
+ malformed_total = 0
313
+ negative_total = 0
314
+ skipped_files = 0
315
+ tokens_used_diffs = 0
316
+
317
+ for thread in threads:
318
+ rollout_path_raw = thread.get("rollout_path")
319
+ if not rollout_path_raw:
320
+ continue
321
+ rollout_path = Path(rollout_path_raw)
322
+ if not rollout_path.is_absolute():
323
+ rollout_path = codex_home / rollout_path
324
+ if not rollout_path.exists():
325
+ skipped_files += 1
326
+ continue
327
+
328
+ if since_utc is not None:
329
+ try:
330
+ mtime = datetime.fromtimestamp(rollout_path.stat().st_mtime, tz=timezone.utc)
331
+ except OSError:
332
+ skipped_files += 1
333
+ continue
334
+ if mtime < since_utc - timedelta(days=1):
335
+ continue
336
+
337
+ facts, malformed, negative = parse_rollout_facts(
338
+ rollout_path,
339
+ model_raw=thread.get("model"),
340
+ session_id=thread.get("id"),
341
+ )
342
+ malformed_total += malformed
343
+ negative_total += negative
344
+
345
+ # Diagnostic only: how far does this thread's derived-fact total
346
+ # sit from the state DB's own `tokens_used` column. Not a source
347
+ # of truth for the fact stream itself.
348
+ tokens_used = thread.get("tokens_used") or 0
349
+ derived_total = sum(f.tokens for f in facts)
350
+ if tokens_used and derived_total:
351
+ diff = abs(tokens_used - derived_total)
352
+ if diff > max(100, tokens_used * 0.01):
353
+ tokens_used_diffs += 1
354
+
355
+ for fact in facts:
356
+ if since_utc is not None and fact.occurred_at_utc < since_utc:
357
+ continue
358
+ if until_utc is not None and fact.occurred_at_utc >= until_utc:
359
+ continue
360
+ all_facts.append(fact)
361
+
362
+ return ReadResult(
363
+ facts=all_facts,
364
+ malformed_events=malformed_total,
365
+ skipped_files=skipped_files,
366
+ negative_deltas=negative_total,
367
+ tokens_used_diffs=tokens_used_diffs,
368
+ )
@@ -0,0 +1,90 @@
1
+ """Render a report payload as a table, CSV, or JSON string."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import csv
6
+ import io
7
+ import json
8
+
9
+ _COLUMNS = (
10
+ "month",
11
+ "agent",
12
+ "model",
13
+ "token_kind",
14
+ "tokens",
15
+ "priced_tokens",
16
+ "unpriced_tokens",
17
+ "estimated_cost_usd",
18
+ "credits",
19
+ "pricing_status",
20
+ )
21
+ _HEADERS = (
22
+ "Month",
23
+ "Agent",
24
+ "Model",
25
+ "Token Kind",
26
+ "Tokens",
27
+ "Priced",
28
+ "Unpriced",
29
+ "Est. Cost (USD)",
30
+ "Credits",
31
+ "Status",
32
+ )
33
+
34
+
35
+ def _cell(row: dict, column: str) -> str:
36
+ value = row.get(column)
37
+ if value is None:
38
+ return "-"
39
+ if column == "estimated_cost_usd":
40
+ return f"{value:.4f}"
41
+ if column == "credits":
42
+ return f"{value:.2f}" if value else "-"
43
+ return str(value)
44
+
45
+
46
+ def render_table(payload: dict) -> str:
47
+ rows = payload["rows"]
48
+ formatted = [[_cell(row, c) for c in _COLUMNS] for row in rows]
49
+ widths = [len(h) for h in _HEADERS]
50
+ for frow in formatted:
51
+ for i, cell in enumerate(frow):
52
+ widths[i] = max(widths[i], len(cell))
53
+
54
+ lines = [" ".join(h.ljust(widths[i]) for i, h in enumerate(_HEADERS))]
55
+ lines.append(" ".join("-" * w for w in widths))
56
+ for frow in formatted:
57
+ lines.append(" ".join(cell.ljust(widths[i]) for i, cell in enumerate(frow)))
58
+
59
+ total_tokens = sum((r.get("tokens") or 0) for r in rows)
60
+ total_cost = sum((r.get("estimated_cost_usd") or 0) for r in rows)
61
+ lines.append("")
62
+ lines.append(f"Total tokens: {total_tokens:,} Total estimated cost: ${total_cost:.4f}")
63
+
64
+ rates = payload.get("rates") or {}
65
+ sha = rates.get("sha256") or ""
66
+ lines.append(f"Rates catalog: {rates.get('catalog_version')} (sha256={sha[:12]}...)")
67
+
68
+ dq = payload.get("data_quality") or {}
69
+ if any(dq.values()):
70
+ lines.append(
71
+ "Data quality: "
72
+ f"malformed_events={dq.get('malformed_events', 0)} "
73
+ f"skipped_files={dq.get('skipped_files', 0)} "
74
+ f"negative_deltas={dq.get('negative_deltas', 0)} "
75
+ f"unpriced_tokens={dq.get('unpriced_tokens', 0)}"
76
+ )
77
+ return "\n".join(lines)
78
+
79
+
80
+ def render_csv(payload: dict) -> str:
81
+ buf = io.StringIO()
82
+ writer = csv.DictWriter(buf, fieldnames=_COLUMNS)
83
+ writer.writeheader()
84
+ for row in payload["rows"]:
85
+ writer.writerow({c: row.get(c) for c in _COLUMNS})
86
+ return buf.getvalue()
87
+
88
+
89
+ def render_json(payload: dict) -> str:
90
+ return json.dumps(payload, ensure_ascii=False, indent=2)