#!/usr/bin/env python3
"""Heap snapshot comparison helper for Node.js leak investigations.

This script is intentionally lightweight and vendor-neutral. It does not parse
native V8 heap snapshot formats directly. Instead, it compares two CSV or JSON
summary exports from your heap analysis tool and produces a readable report.

Expected input formats:

1) CSV with headers (recommended):
   constructor,count,shallow_size,retained_size
   Map,120,3840,812288
   Object,4500,144000,9221120

2) JSON as either:
   - a list of objects with the same keys as above, or
   - an object containing a top-level "entries" list

The script compares baseline and follow-up snapshots, identifies the largest
increases, and optionally filters to constructors that persist or grow.

Usage examples:

  python heap_snapshot_compare.py --baseline baseline.csv --followup after.csv
  python heap_snapshot_compare.py --baseline baseline.json --followup after.json --top 20
  python heap_snapshot_compare.py --baseline baseline.csv --followup after.csv --output report.txt

Exit codes:
  0 success
  1 invalid input or runtime error
"""

from __future__ import annotations

import argparse
import csv
import json
import os
import sys
from dataclasses import dataclass
from typing import Dict, Iterable, List, Optional, Tuple


@dataclass(frozen=True)
class SnapshotEntry:
    constructor: str
    count: int
    shallow_size: int
    retained_size: int


@dataclass(frozen=True)
class DeltaEntry:
    constructor: str
    baseline: SnapshotEntry
    followup: SnapshotEntry
    count_delta: int
    shallow_delta: int
    retained_delta: int


def positive_int(value: str) -> int:
    try:
        parsed = int(value)
    except ValueError as exc:
        raise argparse.ArgumentTypeError(f"invalid integer: {value!r}") from exc
    if parsed < 1:
        raise argparse.ArgumentTypeError("value must be >= 1")
    return parsed


def read_text_file(path: str) -> str:
    if not os.path.isfile(path):
        raise FileNotFoundError(f"file not found: {path}")
    with open(path, "r", encoding="utf-8") as f:
        return f.read()


def normalize_name(name: str) -> str:
    return " ".join(name.strip().split())


def parse_entry(obj: dict) -> SnapshotEntry:
    required = ["constructor", "count", "shallow_size", "retained_size"]
    for key in required:
        if key not in obj:
            raise ValueError(f"missing required field: {key}")

    constructor = normalize_name(str(obj["constructor"]))
    if not constructor:
        raise ValueError("constructor name cannot be empty")

    try:
        count = int(obj["count"])
        shallow_size = int(obj["shallow_size"])
        retained_size = int(obj["retained_size"])
    except (TypeError, ValueError) as exc:
        raise ValueError(f"invalid numeric field for constructor {constructor!r}") from exc

    if count < 0 or shallow_size < 0 or retained_size < 0:
        raise ValueError(f"negative values are not allowed for constructor {constructor!r}")

    return SnapshotEntry(constructor, count, shallow_size, retained_size)


def load_from_csv(text: str) -> List[SnapshotEntry]:
    rows = list(csv.DictReader(text.splitlines()))
    if not rows:
        raise ValueError("CSV contains no data rows")
    return [parse_entry(row) for row in rows]


def load_from_json(text: str) -> List[SnapshotEntry]:
    data = json.loads(text)
    if isinstance(data, dict):
        if "entries" not in data or not isinstance(data["entries"], list):
            raise ValueError('JSON object must contain an "entries" array')
        items = data["entries"]
    elif isinstance(data, list):
        items = data
    else:
        raise ValueError("JSON must be a list or an object with an entries array")

    if not items:
        raise ValueError("JSON contains no entries")
    return [parse_entry(item) for item in items]


def load_snapshot(path: str) -> List[SnapshotEntry]:
    text = read_text_file(path)
    ext = os.path.splitext(path)[1].lower()

    if ext == ".csv":
        entries = load_from_csv(text)
    elif ext == ".json":
        entries = load_from_json(text)
    else:
        # Best-effort detection for convenience.
        stripped = text.lstrip()
        if stripped.startswith("[") or stripped.startswith("{"):
            entries = load_from_json(text)
        else:
            entries = load_from_csv(text)

    dedup: Dict[str, SnapshotEntry] = {}
    for entry in entries:
        if entry.constructor in dedup:
            prev = dedup[entry.constructor]
            dedup[entry.constructor] = SnapshotEntry(
                constructor=entry.constructor,
                count=prev.count + entry.count,
                shallow_size=prev.shallow_size + entry.shallow_size,
                retained_size=prev.retained_size + entry.retained_size,
            )
        else:
            dedup[entry.constructor] = entry
    return list(dedup.values())


def index_by_constructor(entries: Iterable[SnapshotEntry]) -> Dict[str, SnapshotEntry]:
    return {entry.constructor: entry for entry in entries}


def compute_deltas(baseline: List[SnapshotEntry], followup: List[SnapshotEntry]) -> List[DeltaEntry]:
    base_map = index_by_constructor(baseline)
    follow_map = index_by_constructor(followup)
    all_names = sorted(set(base_map) | set(follow_map))

    deltas: List[DeltaEntry] = []
    for name in all_names:
        base = base_map.get(name, SnapshotEntry(name, 0, 0, 0))
        after = follow_map.get(name, SnapshotEntry(name, 0, 0, 0))
        deltas.append(
            DeltaEntry(
                constructor=name,
                baseline=base,
                followup=after,
                count_delta=after.count - base.count,
                shallow_delta=after.shallow_size - base.shallow_size,
                retained_delta=after.retained_size - base.retained_size,
            )
        )
    return deltas


def format_bytes(n: int) -> str:
    sign = "-" if n < 0 else ""
    value = abs(n)
    units = ["B", "KB", "MB", "GB", "TB"]
    size = float(value)
    unit = 0
    while size >= 1024 and unit < len(units) - 1:
        size /= 1024.0
        unit += 1
    return f"{sign}{size:.2f} {units[unit]}"


def sort_deltas(deltas: List[DeltaEntry]) -> List[DeltaEntry]:
    return sorted(
        deltas,
        key=lambda d: (d.retained_delta, d.count_delta, d.shallow_delta),
        reverse=True,
    )


def render_report(baseline_path: str, followup_path: str, deltas: List[DeltaEntry], top_n: int) -> str:
    lines: List[str] = []
    lines.append("Node.js Heap Snapshot Comparison Report")
    lines.append("=" * 40)
    lines.append(f"Baseline: {baseline_path}")
    lines.append(f"Follow-up: {followup_path}")
    lines.append("")
    lines.append("Top retained-size increases")
    lines.append("-" * 30)

    selected = deltas[:top_n]
    if not selected:
        lines.append("No entries found.")
        return "\n".join(lines)

    header = f"{'Constructor':40} {'Count Δ':>10} {'Shallow Δ':>14} {'Retained Δ':>14}"
    lines.append(header)
    lines.append("-" * len(header))
    for d in selected:
        lines.append(
            f"{d.constructor[:40]:40} {d.count_delta:10d} {format_bytes(d.shallow_delta):>14} {format_bytes(d.retained_delta):>14}"
        )

    total_retained = sum(d.retained_delta for d in deltas)
    total_count = sum(d.count_delta for d in deltas)
    lines.append("")
    lines.append("Summary")
    lines.append("-" * 7)
    lines.append(f"Total count delta: {total_count}")
    lines.append(f"Total retained delta: {format_bytes(total_retained)}")
    lines.append("")
    lines.append("Interpretation hints")
    lines.append("- Persistent increases in retained size across equivalent workloads deserve investigation.")
    lines.append("- Trace the retaining path in your heap analysis tool for the top growing constructors.")
    lines.append("- Validate a fix by repeating the same workload and confirming the same constructors stabilize.")
    lines.append("- A legitimate cache or queue may need bounds or TTLs rather than removal.")
    return "\n".join(lines)


def build_arg_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(
        description="Compare two Node.js heap snapshot summary exports and report growth patterns.",
    )
    parser.add_argument("--baseline", required=True, help="Path to the baseline CSV or JSON summary export")
    parser.add_argument("--followup", required=True, help="Path to the follow-up CSV or JSON summary export")
    parser.add_argument("--top", type=positive_int, default=15, help="Number of growing constructors to show")
    parser.add_argument("--output", help="Optional output file path; if omitted, print to stdout")
    return parser


def main(argv: Optional[List[str]] = None) -> int:
    parser = build_arg_parser()
    args = parser.parse_args(argv)

    try:
        baseline = load_snapshot(args.baseline)
        followup = load_snapshot(args.followup)
        deltas = sort_deltas(compute_deltas(baseline, followup))
        report = render_report(args.baseline, args.followup, deltas, args.top)

        if args.output:
            with open(args.output, "w", encoding="utf-8") as f:
                f.write(report + "\n")
        else:
            print(report)
        return 0
    except (OSError, ValueError, json.JSONDecodeError) as exc:
        print(f"error: {exc}", file=sys.stderr)
        return 1


if __name__ == "__main__":
    raise SystemExit(main())