#!/usr/bin/env python3
"""Python memory profiling helper.

This script provides a compact, operationally useful workflow for investigating
suspected Python memory growth using tracemalloc and objgraph.

Features:
- Run a representative command under tracemalloc
- Capture and compare two snapshots
- Report the largest allocation deltas by line or traceback
- Optionally inspect object growth with objgraph
- Optionally show backreferences for a suspicious object type

Usage examples:

  # Compare snapshots around a workload command
  python memory_profile_helper.py run -- python your_app.py --arg value

  # Compare two existing snapshots
  python memory_profile_helper.py compare before.snap after.snap --limit 20

  # Show object growth and most common types
  python memory_profile_helper.py objgraph --limit 15

Notes:
- objgraph is optional; install it if you want object graph inspection.
- tracemalloc only tracks allocations managed by Python's memory allocator.
- This script is intended for repeatable investigations, not one-off tuning.
"""

from __future__ import annotations

import argparse
import os
import subprocess
import sys
import textwrap
import tracemalloc
from pathlib import Path
from typing import Iterable, List, Optional


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 <= 0:
        raise argparse.ArgumentTypeError("value must be greater than zero")
    return parsed


def existing_file(path: str) -> Path:
    p = Path(path)
    if not p.exists():
        raise argparse.ArgumentTypeError(f"file does not exist: {path}")
    if not p.is_file():
        raise argparse.ArgumentTypeError(f"not a file: {path}")
    return p


def parse_args(argv: Optional[List[str]] = None) -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Investigate Python memory growth with tracemalloc and objgraph.",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog=textwrap.dedent(
            """
            Common workflows:
              1) Start with a representative workload.
              2) Compare snapshots before and after the workload.
              3) Use objgraph to inspect retained object types.
              4) Repeat after a fix to verify the change.
            """
        ),
    )

    subparsers = parser.add_subparsers(dest="command", required=True)

    run_parser = subparsers.add_parser(
        "run",
        help="Run a command under tracemalloc and compare two snapshots.",
    )
    run_parser.add_argument(
        "--frames",
        type=positive_int,
        default=25,
        help="tracemalloc traceback depth (default: 25)",
    )
    run_parser.add_argument(
        "--limit",
        type=positive_int,
        default=10,
        help="number of top allocation deltas to display (default: 10)",
    )
    run_parser.add_argument(
        "--stat",
        choices=("lineno", "filename", "traceback"),
        default="lineno",
        help="grouping key for snapshot comparison (default: lineno)",
    )
    run_parser.add_argument(
        "--snapshot-prefix",
        default="memsnap",
        help="prefix for saved snapshot files (default: memsnap)",
    )
    run_parser.add_argument(
        "--keep-snapshots",
        action="store_true",
        help="keep snapshot files after reporting",
    )
    run_parser.add_argument(
        "--",
        dest="double_dash",
        action="store_true",
        help=argparse.SUPPRESS,
    )
    run_parser.add_argument(
        "cmd",
        nargs=argparse.REMAINDER,
        help="command to execute after --",
    )

    compare_parser = subparsers.add_parser(
        "compare",
        help="Compare two existing tracemalloc snapshot files.",
    )
    compare_parser.add_argument("before", type=existing_file, help="baseline snapshot file")
    compare_parser.add_argument("after", type=existing_file, help="later snapshot file")
    compare_parser.add_argument(
        "--limit",
        type=positive_int,
        default=10,
        help="number of top deltas to display (default: 10)",
    )
    compare_parser.add_argument(
        "--stat",
        choices=("lineno", "filename", "traceback"),
        default="lineno",
        help="grouping key for snapshot comparison (default: lineno)",
    )

    objgraph_parser = subparsers.add_parser(
        "objgraph",
        help="Show growth and common object types using objgraph.",
    )
    objgraph_parser.add_argument(
        "--limit",
        type=positive_int,
        default=10,
        help="number of object types to display (default: 10)",
    )
    objgraph_parser.add_argument(
        "--backrefs",
        metavar="TYPE_NAME",
        help="show a backreference graph for objects of this type name",
    )
    objgraph_parser.add_argument(
        "--max-depth",
        type=positive_int,
        default=5,
        help="maximum backreference depth when using --backrefs (default: 5)",
    )

    return parser.parse_args(argv)


def save_snapshot(snapshot: tracemalloc.Snapshot, path: Path) -> None:
    snapshot.dump(str(path))


def load_snapshot(path: Path) -> tracemalloc.Snapshot:
    return tracemalloc.Snapshot.load(str(path))


def print_top_stats(stats: Iterable[tracemalloc.StatisticDiff], limit: int) -> None:
    print(f"Top {limit} allocation deltas:")
    for idx, stat in enumerate(stats):
        if idx >= limit:
            break
        print(stat)


def run_command_under_tracemalloc(args: argparse.Namespace) -> int:
    if not args.cmd:
        raise SystemExit("error: provide a command after --")

    command = args.cmd
    if command and command[0] == "--":
        command = command[1:]
    if not command:
        raise SystemExit("error: no command provided")

    tracemalloc.start(args.frames)
    before = tracemalloc.take_snapshot()

    print("Running command:")
    print(" ".join(command))
    completed = subprocess.run(command)

    after = tracemalloc.take_snapshot()
    stats = after.compare_to(before, args.stat)
    print_top_stats(stats, args.limit)

    prefix = args.snapshot_prefix
    before_path = Path(f"{prefix}-before.snap")
    after_path = Path(f"{prefix}-after.snap")
    save_snapshot(before, before_path)
    save_snapshot(after, after_path)
    print(f"Saved snapshots: {before_path} {after_path}")

    if not args.keep_snapshots:
        # Keep snapshots only when requested; otherwise clean them up after reporting.
        for path in (before_path, after_path):
            try:
                path.unlink()
            except FileNotFoundError:
                pass
        print("Removed snapshot files (use --keep-snapshots to retain them).")

    return completed.returncode


def compare_snapshots(args: argparse.Namespace) -> int:
    tracemalloc.start()
    before = load_snapshot(args.before)
    after = load_snapshot(args.after)
    stats = after.compare_to(before, args.stat)
    print_top_stats(stats, args.limit)
    return 0


def run_objgraph(args: argparse.Namespace) -> int:
    try:
        import objgraph  # type: ignore
    except ImportError:
        print(
            "objgraph is not installed. Install it with 'pip install objgraph' and try again.",
            file=sys.stderr,
        )
        return 2

    print("Most common object types:")
    objgraph.show_most_common_types(limit=args.limit)
    print("\nObject growth:")
    objgraph.show_growth(limit=args.limit)

    if args.backrefs:
        import gc

        gc.collect()
        candidates = [obj for obj in gc.get_objects() if type(obj).__name__ == args.backrefs]
        if not candidates:
            print(f"\nNo live objects found with type name: {args.backrefs}")
            return 0

        target = candidates[0]
        print(f"\nShowing backreferences for one {args.backrefs} instance...")
        # This may open a graph visualization if supported by the environment.
        objgraph.show_backrefs([target], max_depth=args.max_depth)

    return 0


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

    if args.command == "run":
        return run_command_under_tracemalloc(args)
    if args.command == "compare":
        return compare_snapshots(args)
    if args.command == "objgraph":
        return run_objgraph(args)

    raise SystemExit(f"unknown command: {args.command}")


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