#!/usr/bin/env python3
"""Collect de-identified Claude Code usage for the shared-quota reconciliation.

Reads local Claude Code transcripts (~/.claude/projects/**/*.jsonl) since the
current weekly cycle start, aggregates per (day, model) with exact token
components, and emits a JSON file that contains NO prompts, responses,
project paths, or account identifiers.

Usage:
    python3 collect_usage.py --cycle-start-epoch 1791269999 --out usage.json
    python3 collect_usage.py --cycle-start "2026-10-06T06:59:59Z" --out usage.json

Output schema: claude-usage.collect/v1
"""
import argparse
import datetime
import glob
import json
import os
import re
import sys
import time
from collections import defaultdict

UTC = datetime.timezone.utc


def parse_args():
    p = argparse.ArgumentParser(description=__doc__)
    g = p.add_mutually_exclusive_group(required=True)
    g.add_argument("--cycle-start-epoch", type=int,
                   help="Weekly cycle start as Unix epoch (UTC seconds).")
    g.add_argument("--cycle-start", type=str,
                   help="Weekly cycle start ISO time, e.g. 2026-10-06T06:59:59Z")
    p.add_argument("--projects-dir", default="~/.claude/projects",
                   help="Claude Code projects directory (default: ~/.claude/projects).")
    p.add_argument("--out", default="usage.json", help="Output JSON path.")
    p.add_argument("--max-file-mb", type=float, default=200.0,
                   help="Skip transcript files larger than this (defensive).")
    return p.parse_args()


def main():
    args = parse_args()
    if args.cycle_start_epoch is not None:
        cycle_start = float(args.cycle_start_epoch)
    else:
        t = args.cycle_start.replace("Z", "+00:00")
        cycle_start = datetime.datetime.fromisoformat(t).astimezone(UTC).timestamp()

    projects_dir = os.path.expanduser(args.projects_dir)
    if not os.path.isdir(projects_dir):
        sys.exit(f"error: {projects_dir} not found")

    # aggregate: (date, model) -> token components
    agg = defaultdict(lambda: {
        "input_tokens": 0,
        "cache_creation_input_tokens": 0,
        "cache_creation_5m": 0,
        "cache_creation_1h": 0,
        "cache_read_input_tokens": 0,
        "output_tokens": 0,
        "requests": 0,
    })
    files_scanned = 0
    files_with_usage = 0
    usage_records = 0
    skipped_large = 0
    models_seen = set()
    first_ts = None
    last_ts = None

    for path in glob.glob(projects_dir + "/**/*.jsonl", recursive=True):
        try:
            st = os.stat(path)
            if st.st_mtime < cycle_start - 86400:
                continue  # clearly old file, skip fast
            if st.st_size > args.max_file_mb * 1e6:
                skipped_large += 1
                continue
        except OSError:
            continue
        files_scanned += 1
        try:
            with open(path, encoding="utf-8") as f:
                for line in f:
                    try:
                        obj = json.loads(line)
                    except json.JSONDecodeError:
                        continue
                    ts = obj.get("timestamp")
                    try:
                        dt = datetime.datetime.fromisoformat(
                            (ts or "").replace("Z", "+00:00"))
                    except ValueError:
                        continue
                    ep = dt.timestamp()
                    if ep < cycle_start:
                        continue
                    msg = obj.get("message") or {}
                    usage = msg.get("usage")
                    if not isinstance(usage, dict):
                        continue
                    model = msg.get("model") or "unknown"
                    if model == "<synthetic>":
                        continue
                    inp = usage.get("input_tokens", 0) or 0
                    cc = usage.get("cache_creation") or {}
                    w5 = cc.get("ephemeral_5m_input_tokens", 0) or 0
                    w60 = cc.get("ephemeral_1h_input_tokens", 0) or 0
                    cw = usage.get("cache_creation_input_tokens", 0) or 0
                    cr = usage.get("cache_read_input_tokens", 0) or 0
                    out = usage.get("output_tokens", 0) or 0
                    if not (inp or cw or cr or out):
                        continue
                    day = dt.astimezone(UTC).strftime("%Y-%m-%d")
                    row = agg[(day, model)]
                    row["input_tokens"] += inp
                    row["cache_creation_input_tokens"] += cw
                    row["cache_creation_5m"] += w5
                    row["cache_creation_1h"] += w60
                    row["cache_read_input_tokens"] += cr
                    row["output_tokens"] += out
                    row["requests"] += 1
                    models_seen.add(model)
                    usage_records += 1
                    if first_ts is None or ep < first_ts:
                        first_ts = ep
                    if last_ts is None or ep > last_ts:
                        last_ts = ep
        except OSError:
            continue
        else:
            if usage_records:
                files_with_usage += 1

    records = []
    for (day, model), row in sorted(agg.items()):
        records.append({"date": day, "model": model, **row})

    payload = {
        "schema": "claude-usage.collect/v1",
        "collected_at_utc": datetime.datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"),
        "cycle_start_epoch": int(cycle_start),
        "cycle_start_utc": datetime.datetime.fromtimestamp(cycle_start, UTC).strftime(
            "%Y-%m-%dT%H:%M:%SZ"),
        "projects_dir_name": os.path.basename(projects_dir.rstrip("/")),
        "files_scanned": files_scanned,
        "files_with_usage": files_with_usage,
        "transcripts_skipped_large": skipped_large,
        "usage_records": usage_records,
        "models": sorted(models_seen),
        "first_usage_utc": (datetime.datetime.fromtimestamp(first_ts, UTC).strftime(
            "%Y-%m-%dT%H:%M:%SZ") if first_ts else None),
        "last_usage_utc": (datetime.datetime.fromtimestamp(last_ts, UTC).strftime(
            "%Y-%m-%dT%H:%M:%SZ") if last_ts else None),
        "records": records,
    }
    with open(args.out, "w", encoding="utf-8") as f:
        json.dump(payload, f, indent=2, ensure_ascii=False)
        f.write("\n")

    tot = defaultdict(int)
    for r in records:
        for k in ("input_tokens", "cache_creation_input_tokens",
                  "cache_read_input_tokens", "output_tokens", "requests"):
            tot[k] += r[k]
    print(f"collected {usage_records} usage records from {files_with_usage} transcript files")
    print(f"models: {sorted(models_seen)}")
    print(f"input={tot['input_tokens']:,} cache_write={tot['cache_creation_input_tokens']:,} "
          f"cache_read={tot['cache_read_input_tokens']:,} output={tot['output_tokens']:,} "
          f"requests={tot['requests']:,}")
    print(f"wrote {args.out}")


if __name__ == "__main__":
    main()
