"""
EXPORT PLY UNIQUEMENT — Repart du cache sans tout recalculer.

Lance ce script si reconstruct_3d_fixed.py a planté à l'étape [CLEAN] / export.
Il relit les depth maps en cache + frames IRL déjà extraites et réexporte le .ply.

Usage :
    python export_ply_only.py
"""

import os
import sys
import glob
import numpy as np
import cv2
from scipy.spatial import cKDTree

# ============================================================
#  CONFIG — doit correspondre à reconstruct_3d_fixed.py
# ============================================================
OUTPUT_FRAMES   = "frames_irl/"
DEPTH_CACHE_DIR = "depth_cache_frames/"
OUTPUT_PLY      = "chambre_3d.ply"

SKIP        = 4      # sous-échantillonnage pixels (même valeur que dans le script principal)
VOXEL_SIZE  = 0.02
K_OUTLIER   = 10
K_NORMAL    = 15
BATCH_SIZE  = 2000
DEPTH_SCALE = 3.0

# Paramètres caméra approx (mêmes que dans le script principal)
H, W   = 480, 854
FX     = W * 0.8
FY     = FX
CX, CY = W / 2, H / 2
K_CAM  = np.array([[FX, 0, CX],
                   [0, FY, CY],
                   [0,  0,  1]], dtype=float)

# ============================================================
#  CHARGEMENT DES FRAMES IRL DÉJÀ EXTRAITES
# ============================================================
def load_cached_frames():
    frames = sorted(glob.glob(os.path.join(OUTPUT_FRAMES, "*.jpg")))
    if not frames:
        print(f"[ERR] Aucune frame trouvée dans {OUTPUT_FRAMES}")
        sys.exit(1)
    print(f"[OK] {len(frames)} frames IRL trouvées dans le cache")
    return [{"path": f} for f in frames]


# ============================================================
#  CONSTRUCTION NUAGE DE POINTS (lazy depth loading)
# ============================================================
def build_point_cloud(frames):
    print(f"\n[CLOUD] Reconstruction du nuage depuis le cache...")

    fx, fy = K_CAM[0, 0], K_CAM[1, 1]
    cx, cy = K_CAM[0, 2], K_CAM[1, 2]

    CHUNK = 100
    chunk_pts, chunk_col = [], []
    merged_pts, merged_col = [], []
    total = 0

    def _flush():
        if chunk_pts:
            merged_pts.append(np.vstack(chunk_pts))
            merged_col.append(np.vstack(chunk_col))
            chunk_pts.clear()
            chunk_col.clear()

    for i, fi in enumerate(frames):
        path = fi["path"]
        basename = os.path.splitext(os.path.basename(path))[0]
        cache_file = os.path.join(DEPTH_CACHE_DIR, basename + ".npy")

        if not os.path.exists(cache_file):
            continue

        depth = np.load(cache_file)
        frame = cv2.imread(path)
        if frame is None:
            continue

        h_f, w_f = depth.shape
        frame_rgb = cv2.cvtColor(
            cv2.resize(frame, (w_f, h_f)), cv2.COLOR_BGR2RGB
        ).astype(np.float32) / 255.0
        del frame

        ys, xs = np.mgrid[0:h_f:SKIP, 0:w_f:SKIP]
        ys, xs = ys.flatten(), xs.flatten()
        ds = depth[ys, xs]
        del depth

        valid = ds > 0.01
        ys, xs, ds = ys[valid], xs[valid], ds[valid]

        Z = ds * DEPTH_SCALE
        X = (xs - cx) / fx * Z
        Y = (ys - cy) / fy * Z

        pts = np.stack([X, Y, Z], axis=1).astype(np.float32)
        col = frame_rgb[ys, xs].astype(np.float32)

        chunk_pts.append(pts)
        chunk_col.append(col)
        total += len(pts)

        if (i + 1) % CHUNK == 0:
            _flush()

        if (i + 1) % 50 == 0:
            print(f"  {i+1}/{len(frames)} frames | {total:,} pts...", end="\r")

    _flush()

    if not merged_pts:
        print("[ERR] Aucun point généré — vérifie que depth_cache_frames/ existe et n'est pas vide")
        sys.exit(1)

    points = np.vstack(merged_pts)
    colors = np.vstack(merged_col)
    del merged_pts, merged_col
    print(f"\n[OK] {len(points):,} points générés")
    return points, colors


# ============================================================
#  NETTOYAGE
# ============================================================
def clean(points, colors):
    print(f"\n[CLEAN] Nettoyage...")
    print(f"  Avant : {len(points):,} pts")

    # Outliers par batch
    tree = cKDTree(points)
    mean_dists = np.empty(len(points), dtype=np.float32)
    for s in range(0, len(points), 10000):
        e = min(s + 10000, len(points))
        d, _ = tree.query(points[s:e], k=K_OUTLIER + 1)
        mean_dists[s:e] = d[:, 1:].mean(axis=1)
    thr = mean_dists.mean() + 2.0 * mean_dists.std()
    mask = mean_dists < thr
    del mean_dists
    points, colors = points[mask], colors[mask]
    print(f"  Après outliers : {len(points):,} pts")

    # Voxel downsampling vectorisé
    vi = np.floor(points / VOXEL_SIZE).astype(np.int32)
    vi_min = vi.min(axis=0)
    vi = vi - vi_min
    vi_max = vi.max(axis=0) + 1
    flat = vi[:, 0] * vi_max[1] * vi_max[2] + vi[:, 1] * vi_max[2] + vi[:, 2]
    _, first = np.unique(flat, return_index=True)
    points, colors = points[first], colors[first]
    print(f"  Après voxel    : {len(points):,} pts")

    # Normales par batch vectorisé
    print("  Estimation des normales...")
    normals = np.zeros_like(points)
    tree2 = cKDTree(points)
    for s in range(0, len(points), BATCH_SIZE):
        e = min(s + BATCH_SIZE, len(points))
        _, idx_b = tree2.query(points[s:e], k=K_NORMAL + 1)
        for j, idxs in enumerate(idx_b):
            nb = points[idxs[1:]]
            if len(nb) < 3:
                continue
            cov = np.cov((nb - nb.mean(axis=0)).T)
            _, eigvecs = np.linalg.eigh(cov)
            normals[s + j] = eigvecs[:, 0]
        if (s // BATCH_SIZE) % 5 == 0:
            print(f"    {100*e//len(points)}%...", end="\r")
    print()

    return points, colors, normals


# ============================================================
#  EXPORT PLY (écriture manuelle — contourne le bug pyntcloud/pandas)
# ============================================================
def export_ply(points, colors, normals):
    """Écrit un fichier PLY binaire little-endian sans pyntcloud."""
    import struct

    print(f"\n[EXPORT] Création du fichier {OUTPUT_PLY}...")

    n = len(points)
    pts   = points.astype(np.float32)
    col   = np.clip(colors * 255, 0, 255).astype(np.uint8)
    norms = normals.astype(np.float32)

    header = (
        "ply\n"
        "format binary_little_endian 1.0\n"
        f"element vertex {n}\n"
        "property float x\n"
        "property float y\n"
        "property float z\n"
        "property uchar red\n"
        "property uchar green\n"
        "property uchar blue\n"
        "property float nx\n"
        "property float ny\n"
        "property float nz\n"
        "end_header\n"
    )

    with open(OUTPUT_PLY, "wb") as f:
        f.write(header.encode("ascii"))
        # Construit un array structuré pour un seul write() rapide
        dt = np.dtype([
            ("x",  np.float32), ("y",  np.float32), ("z",  np.float32),
            ("r",  np.uint8),   ("g",  np.uint8),   ("b",  np.uint8),
            ("nx", np.float32), ("ny", np.float32), ("nz", np.float32),
        ])
        buf = np.empty(n, dtype=dt)
        buf["x"],  buf["y"],  buf["z"]  = pts[:, 0],   pts[:, 1],   pts[:, 2]
        buf["r"],  buf["g"],  buf["b"]  = col[:, 0],   col[:, 1],   col[:, 2]
        buf["nx"], buf["ny"], buf["nz"] = norms[:, 0], norms[:, 1], norms[:, 2]
        f.write(buf.tobytes())

    size_mb = os.path.getsize(OUTPUT_PLY) / 1024 / 1024
    print(f"[OK] Exporté → {OUTPUT_PLY}  ({size_mb:.1f} MB, {n:,} pts)")
    bbox = points.max(axis=0) - points.min(axis=0)
    print(f"\n  Dimensions :")
    print(f"    Largeur    : {bbox[0]:.2f}m")
    print(f"    Profondeur : {bbox[1]:.2f}m")
    print(f"    Hauteur    : {bbox[2]:.2f}m")
    print(f"\n  Ouvre dans MeshLab : File → Import Mesh → {OUTPUT_PLY}")


# ============================================================
#  MAIN
# ============================================================
if __name__ == "__main__":
    print("=" * 55)
    print("  EXPORT PLY — repart du cache (pas de recalcul)")
    print("=" * 55)

    frames = load_cached_frames()
    points, colors = build_point_cloud(frames)
    points, colors, normals = clean(points, colors)
    export_ply(points, colors, normals)

    print("\n" + "=" * 55)
    print("  TERMINÉ !")
    print("=" * 55)
