"""Voxel-downsample a point cloud while preserving its geometry, and show the comparison.

Usage
-----
    python voxel_downsample.py                              # defaults: bigpointcloud_001.ply
    python voxel_downsample.py --input my.ply --voxel 0.15  # force a voxel size
    python voxel_downsample.py --target-points 3000         # pick voxel size by point budget

What it does
------------
1. Sweeps a range of voxel sizes and, for each, measures how far the surface moved:
   one-sided distances original -> downsampled (mean / RMS / p95 / max = Hausdorff),
   normal deviation, and bounding-box shrinkage.
2. Auto-selects the *largest* voxel (= fewest points) whose p95 surface error still
   stays under one nearest-neighbour spacing of the original cloud, i.e. the geometry
   is perturbed no more than the original sampling noise already does.
3. Compares against random subsampling at the *same* point count, which is the usual
   way geometry gets destroyed (uneven density, holes, outliers kept).
4. Writes the downsampled .ply, a sweep .csv, side-by-side PNG renders and an
   interactive .html (fig.show() is useless on a headless node).
"""

import argparse
import csv
import ctypes
import glob
import os
import sys

import numpy as np


def _import_open3d():
    """Import open3d, preloading a libGL if the system has none (headless cluster)."""
    try:
        import open3d as o3d
        return o3d
    except OSError as err:
        candidates = []
        for prefix in (sys.prefix, os.path.dirname(os.path.dirname(sys.prefix))):
            candidates += glob.glob(os.path.join(prefix, "lib", "libGL.so.1"))
            candidates += glob.glob(os.path.join(prefix, "envs", "*", "lib", "libGL.so.1"))
        for lib in candidates:
            try:
                ctypes.CDLL(lib, mode=ctypes.RTLD_GLOBAL)
                import open3d as o3d
                return o3d
            except OSError:
                continue
        raise err


o3d = _import_open3d()

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt


# ----------------------------------------------------------------------------- metrics
def mean_nn_spacing(pcd, sample=20000, seed=0):
    """Average distance from a point to its nearest neighbour: the cloud's native resolution."""
    pts = np.asarray(pcd.points)
    rng = np.random.default_rng(seed)
    idx = rng.choice(len(pts), size=min(sample, len(pts)), replace=False)
    tree = o3d.geometry.KDTreeFlann(pcd)
    d = []
    for i in idx:
        _, _, sq = tree.search_knn_vector_3d(pts[i], 2)  # [0] is the point itself
        d.append(np.sqrt(sq[1]))
    return float(np.mean(d))


def geometry_error(original, reduced, return_per_point=False):
    """How far the surface moved, plus how well normals and extent survived."""
    d_fwd = np.asarray(original.compute_point_cloud_distance(reduced))  # orig -> reduced
    d_bwd = np.asarray(reduced.compute_point_cloud_distance(original))  # reduced -> orig

    out = {
        "n_points": len(reduced.points),
        "mean_err": float(d_fwd.mean()),
        "rms_err": float(np.sqrt((d_fwd ** 2).mean())),
        "p95_err": float(np.percentile(d_fwd, 95)),
        "hausdorff": float(max(d_fwd.max(), d_bwd.max())),
    }

    # normal deviation: angle between each original normal and the normal it collapsed into
    if original.has_normals() and reduced.has_normals():
        n_o = np.asarray(original.normals)
        n_r = np.asarray(reduced.normals)
        n_r = n_r / np.clip(np.linalg.norm(n_r, axis=1, keepdims=True), 1e-12, None)
        tree = o3d.geometry.KDTreeFlann(reduced)
        ang = []
        for p, n in zip(np.asarray(original.points), n_o):
            _, j, _ = tree.search_knn_vector_3d(p, 1)
            ang.append(np.degrees(np.arccos(np.clip(abs(float(n @ n_r[j[0]])), -1.0, 1.0))))
        out["mean_normal_deg"] = float(np.mean(ang))
    else:
        out["mean_normal_deg"] = float("nan")

    ext_o = original.get_axis_aligned_bounding_box().get_extent()
    ext_r = reduced.get_axis_aligned_bounding_box().get_extent()
    out["bbox_shrink_pct"] = float(100.0 * np.max(1.0 - ext_r / ext_o))
    return (out, d_fwd) if return_per_point else out


def random_subsample(pcd, n, seed=0):
    rng = np.random.default_rng(seed)
    idx = rng.choice(len(pcd.points), size=min(n, len(pcd.points)), replace=False)
    return pcd.select_by_index(np.sort(idx).tolist())


# ----------------------------------------------------------------------------- rendering
VIEWS = [("iso", 24, -62), ("top-down", 88, -90)]


def scatter(ax, pts, colors, title, size, extent, cmap=None, vmax=None):
    art = ax.scatter(pts[:, 0], pts[:, 1], pts[:, 2], c=colors, s=size, linewidths=0,
                     cmap=cmap, vmin=0 if cmap else None, vmax=vmax)
    ax.set_box_aspect(extent, zoom=1.45)
    ax.set_title(title, fontsize=10)
    ax.set_axis_off()
    return art


def cloud_colors(pcd):
    if pcd.has_colors():
        return np.clip(np.asarray(pcd.colors), 0, 1)
    return np.tile([0.20, 0.40, 0.75], (len(pcd.points), 1))


def render_comparison(clouds, path, point_size=1.0):
    """clouds: list of (label, pcd). One row per viewpoint."""
    extent = clouds[0][1].get_axis_aligned_bounding_box().get_extent()
    fig = plt.figure(figsize=(4.6 * len(clouds), 4.0 * len(VIEWS)))
    for r, (vname, elev, azim) in enumerate(VIEWS):
        for c, (label, pcd) in enumerate(clouds):
            ax = fig.add_subplot(len(VIEWS), len(clouds), r * len(clouds) + c + 1, projection="3d")
            scatter(ax, np.asarray(pcd.points), cloud_colors(pcd), f"{label}\n[{vname}]",
                    point_size, extent)
            ax.view_init(elev=elev, azim=azim)
    fig.subplots_adjust(wspace=0.02, hspace=0.02)
    fig.savefig(path, dpi=150, bbox_inches="tight")
    plt.close(fig)


def render_error_maps(original, entries, path, point_size=1.0):
    """Original points coloured by distance to each reduced cloud (shared scale)."""
    pts = np.asarray(original.points)
    extent = original.get_axis_aligned_bounding_box().get_extent()
    vmax = max(float(np.percentile(d, 99)) for _, d in entries)
    fig = plt.figure(figsize=(4.6 * len(entries), 4.0 * len(VIEWS)))
    art = None
    for r, (vname, elev, azim) in enumerate(VIEWS):
        for c, (label, d) in enumerate(entries):
            ax = fig.add_subplot(len(VIEWS), len(entries), r * len(entries) + c + 1,
                                 projection="3d")
            art = scatter(ax, pts, d, f"{label}\n[{vname}]", point_size, extent,
                          cmap="inferno", vmax=vmax)
            ax.view_init(elev=elev, azim=azim)
    fig.subplots_adjust(wspace=0.02, hspace=0.02, right=0.90)
    cax = fig.add_axes([0.92, 0.25, 0.012, 0.5])
    fig.colorbar(art, cax=cax).set_label("distance from original point to kept surface", fontsize=9)
    fig.savefig(path, dpi=150, bbox_inches="tight")
    plt.close(fig)


def render_sweep(rows, chosen, nn, path):
    v = [r["voxel"] for r in rows]
    fig, axes = plt.subplots(1, 3, figsize=(15, 4.2))

    axes[0].plot(v, [r["n_points"] for r in rows], "o-", color="#2c7fb8")
    axes[0].set_yscale("log")
    axes[0].set_ylabel("points kept")

    axes[1].plot(v, [r["mean_err"] for r in rows], "o-", label="mean", color="#2c7fb8")
    axes[1].plot(v, [r["p95_err"] for r in rows], "s-", label="p95", color="#e6844a")
    axes[1].plot(v, [r["hausdorff"] for r in rows], "^-", label="Hausdorff", color="#c0392b")
    axes[1].axhline(nn, ls="--", c="gray", lw=1, label=f"orig. NN spacing = {nn:.3f}")
    axes[1].set_ylabel("surface error (units)")
    axes[1].legend(fontsize=8)

    axes[2].plot(v, [r["mean_normal_deg"] for r in rows], "o-", color="#2c7fb8")
    axes[2].set_ylabel("mean normal deviation (deg)")

    for ax in axes:
        ax.set_xlabel("voxel size")
        ax.axvline(chosen, ls=":", c="green", lw=1.5)
        ax.grid(alpha=0.3)
    fig.suptitle(f"voxel-size sweep (green = selected {chosen:g})", fontsize=11)
    fig.tight_layout()
    fig.savefig(path, dpi=150, bbox_inches="tight")
    plt.close(fig)


def render_interactive(orig, down, path):
    try:
        import plotly.graph_objects as go
        from plotly.subplots import make_subplots
    except ImportError:
        return None

    def trace(pcd):
        p = np.asarray(pcd.points)
        c = np.asarray(pcd.colors) if pcd.has_colors() else None
        marker = dict(size=1.4)
        if c is not None:
            marker["color"] = [f"rgb({int(r*255)},{int(g*255)},{int(b*255)})" for r, g, b in c]
        return go.Scatter3d(x=p[:, 0], y=p[:, 1], z=p[:, 2], mode="markers", marker=marker)

    fig = make_subplots(rows=1, cols=2, specs=[[{"type": "scene"}, {"type": "scene"}]],
                        subplot_titles=(f"original ({len(orig.points)} pts)",
                                        f"voxel downsampled ({len(down.points)} pts)"))
    fig.add_trace(trace(orig), row=1, col=1)
    fig.add_trace(trace(down), row=1, col=2)
    hidden = dict(xaxis=dict(visible=False), yaxis=dict(visible=False), zaxis=dict(visible=False),
                  aspectmode="data")
    fig.update_layout(scene=hidden, scene2=hidden, showlegend=False,
                      margin=dict(l=0, r=0, t=30, b=0))
    fig.write_html(path)
    return path


# ----------------------------------------------------------------------------- main
def main():
    here = os.path.dirname(os.path.abspath(__file__))
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--input", default=os.path.join(here, "bigpointcloud_001.ply"))
    ap.add_argument("--outdir", default=os.path.join(here, "voxel_downsample_out"))
    ap.add_argument("--voxel", type=float, default=None,
                    help="force this voxel size instead of auto-selecting")
    ap.add_argument("--target-points", type=int, default=None,
                    help="pick the sweep voxel whose point count is closest to this")
    ap.add_argument("--err-tol", type=float, default=1.0,
                    help="safe pick: accept a voxel while p95 surface error <= err-tol * NN spacing")
    ap.add_argument("--aggressive-tol", type=float, default=1.0,
                    help="aggressive pick: accept while MEAN surface error <= tol * NN spacing")
    args = ap.parse_args()
    os.makedirs(args.outdir, exist_ok=True)

    pcd = o3d.io.read_point_cloud(args.input)
    if len(pcd.points) == 0:
        sys.exit(f"no points read from {args.input}")
    ext = pcd.get_axis_aligned_bounding_box().get_extent()
    diag = float(np.linalg.norm(ext))
    nn = mean_nn_spacing(pcd)

    print(f"input            : {args.input}")
    print(f"points           : {len(pcd.points)}  normals={pcd.has_normals()} colors={pcd.has_colors()}")
    print(f"bbox extent      : {ext} (diagonal {diag:.3f})")
    print(f"mean NN spacing  : {nn:.4f}  <- native resolution of the cloud\n")

    # sweep from "no-op" (below the native spacing) up to clearly-too-coarse
    voxels = [round(m * nn, 6) for m in (0.5, 0.75, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, 8.0, 12.0)]
    rows = []
    for v in voxels:
        d = pcd.voxel_down_sample(voxel_size=v)
        if len(d.points) < 10:
            continue
        m = geometry_error(pcd, d)
        m["voxel"] = v
        m["keep_pct"] = 100.0 * m["n_points"] / len(pcd.points)
        rows.append(m)

    hdr = f"{'voxel':>8} {'points':>8} {'keep%':>7} {'mean':>8} {'rms':>8} {'p95':>8} {'hausdorff':>10} {'normal°':>8} {'bbox↓%':>7}"
    print("voxel-size sweep (errors in cloud units; 'mean' = avg distance from an original point to the kept surface)")
    print(hdr)
    print("-" * len(hdr))
    for r in rows:
        print(f"{r['voxel']:8.3f} {r['n_points']:8d} {r['keep_pct']:7.1f} {r['mean_err']:8.4f} "
              f"{r['rms_err']:8.4f} {r['p95_err']:8.4f} {r['hausdorff']:10.4f} "
              f"{r['mean_normal_deg']:8.2f} {r['bbox_shrink_pct']:7.2f}")

    # ---- choose the voxel size(s)
    def largest_where(pred):
        ok = [r for r in rows if pred(r)]
        return (max(ok, key=lambda r: r["voxel"]) if ok else rows[0])["voxel"]

    if args.voxel is not None:
        voxel, why, tag = args.voxel, "user-specified", "chosen"
    elif args.target_points is not None:
        voxel = min(rows, key=lambda r: abs(r["n_points"] - args.target_points))["voxel"]
        why, tag = f"closest to target of {args.target_points} points", "chosen"
    else:
        voxel = largest_where(lambda r: r["p95_err"] <= args.err_tol * nn)
        why = f"largest voxel with p95 error <= {args.err_tol:g} x NN spacing"
        tag = "safe"
    # a second, coarser operating point: surface moves on average <= one point spacing
    voxel_aggr = largest_where(lambda r: r["mean_err"] <= args.aggressive_tol * nn)

    print(f"\n{tag} voxel size      : {voxel:.4f}  ({why})")
    print(f"aggressive voxel size: {voxel_aggr:.4f}  (largest voxel with mean error "
          f"<= {args.aggressive_tol:g} x NN spacing)")

    down = pcd.voxel_down_sample(voxel_size=voxel)
    down_m, d_down = geometry_error(pcd, down, return_per_point=True)
    aggr = pcd.voxel_down_sample(voxel_size=voxel_aggr)
    aggr_m, d_aggr = geometry_error(pcd, aggr, return_per_point=True)
    rnd = random_subsample(pcd, len(aggr.points))
    rnd_m, d_rnd = geometry_error(pcd, rnd, return_per_point=True)

    print(f"\nvoxel vs random subsampling at the same budget ({len(aggr.points)} points):")
    print(f"{'':>20}{'mean':>9}{'rms':>9}{'p95':>9}{'hausdorff':>11}{'normal°':>9}")
    for name, m in ((f"voxel {voxel_aggr:.3f}", aggr_m), ("random (baseline)", rnd_m)):
        print(f"{name:>20}{m['mean_err']:9.4f}{m['rms_err']:9.4f}{m['p95_err']:9.4f}"
              f"{m['hausdorff']:11.4f}{m['mean_normal_deg']:9.2f}")

    # ---- write outputs
    stem = os.path.splitext(os.path.basename(args.input))[0]
    outs = []
    for t, v, c in ((tag, voxel, down), ("aggressive", voxel_aggr, aggr)):
        p = os.path.join(args.outdir, f"{stem}_voxel{v:.3f}_{t}.ply")
        o3d.io.write_point_cloud(p, c)
        outs.append(p)

    csv_out = os.path.join(args.outdir, "sweep.csv")
    with open(csv_out, "w", newline="") as fh:
        w = csv.DictWriter(fh, fieldnames=["voxel", "n_points", "keep_pct", "mean_err", "rms_err",
                                           "p95_err", "hausdorff", "mean_normal_deg",
                                           "bbox_shrink_pct"])
        w.writeheader()
        w.writerows(rows)
    outs.append(csv_out)

    def pct(c):
        return f"{100 * len(c.points) / len(pcd.points):.0f}%"

    cmp_png = os.path.join(args.outdir, "comparison.png")
    render_comparison([(f"original — {len(pcd.points)} pts", pcd),
                       (f"voxel {voxel:.3f} ({tag}) — {len(down.points)} pts, {pct(down)}", down),
                       (f"voxel {voxel_aggr:.3f} (aggressive) — {len(aggr.points)} pts, "
                        f"{pct(aggr)}", aggr),
                       (f"random {len(rnd.points)} pts (baseline)", rnd)], cmp_png)
    outs.append(cmp_png)

    err_png = os.path.join(args.outdir, "error_map.png")
    render_error_maps(pcd, [(f"voxel {voxel:.3f} ({tag})", d_down),
                            (f"voxel {voxel_aggr:.3f} (aggressive)", d_aggr),
                            (f"random {len(rnd.points)} pts (baseline)", d_rnd)], err_png)
    outs.append(err_png)

    sweep_png = os.path.join(args.outdir, "sweep.png")
    render_sweep(rows, voxel, nn, sweep_png)
    outs.append(sweep_png)

    outs.append(render_interactive(pcd, aggr, os.path.join(args.outdir, "comparison.html")))

    print("\nwrote:")
    for p in outs:
        if p:
            print(f"  {p}")


if __name__ == "__main__":
    main()
