"""
RECONSTRUCTION 3D CHAMBRE
Pipeline complet : filtre IRL → depth MiDaS → SfM → export .ply

DÉPENDANCES :
    pip install opencv-python numpy scikit-learn joblib torch torchvision timm matplotlib pyntcloud pandas scipy

USAGE :
    python reconstruct_3d.py

NOTES :
    - Lance train_classifier.py d'abord pour générer le classifieur
    - Le modèle MiDaS (~400MB) se télécharge automatiquement au premier lancement
    - Résultat final : chambre_3d.ply → ouvre dans MeshLab
"""

import cv2
import numpy as np
import os
import sys
import glob
import re
import joblib
import torch
import pandas as pd
from pyntcloud import PyntCloud
from scipy.spatial import cKDTree
from datetime import datetime, timedelta, timezone
from multiprocessing import Pool, cpu_count, freeze_support
import multiprocessing as mp
import json

# ============================================================
#  CONFIGURATION
# ============================================================

# Classifieur IRL vs gameplay
CLASSIFIER_FILE = "classifier_irl_vs_gameplay.joblib"
SCALER_FILE     = "scaler_irl_vs_gameplay.joblib"

# Seuil de confiance IRL (0.0 à 1.0)
IRL_CONFIDENCE = 0.75

# Pas d'échantillonnage en secondes
SAMPLE_STEP = 5

# Nombre max de frames à traiter (None = toutes)
MAX_FRAMES = 2000

# Nb de workers parallèles (on laisse 2 cores libres)
NB_WORKERS = min(26, cpu_count() - 2)

# Taille de traitement MiDaS
MIDAS_SIZE = (384, 384)

# Fichiers de sortie
OUTPUT_PLY       = "chambre_3d.ply"
OUTPUT_FRAMES    = "frames_irl/"
DEPTH_CACHE      = "depth_cache.npz"
FEATURES_CACHE   = "sfm_features.json"

# ============================================================
#  CHARGEMENT DU CLASSIFIEUR
# ============================================================

def load_classifier():
    if not os.path.exists(CLASSIFIER_FILE):
        print(f"[ERR] Classifieur introuvable : {CLASSIFIER_FILE}")
        print("      Lance d'abord : python train_classifier.py")
        sys.exit(1)
    clf    = joblib.load(CLASSIFIER_FILE)
    scaler = joblib.load(SCALER_FILE)
    print(f"[OK] Classifieur chargé")
    return clf, scaler


# ============================================================
#  CHARGEMENT MIDAS
# ============================================================

def load_midas():
    print("\n[MIDAS] Chargement du modèle de profondeur...")
    print("        (téléchargement ~400MB au premier lancement)")

    model_type = "DPT_Large"
    midas = torch.hub.load("intel-isl/MiDaS", model_type)
    midas.eval()

    # CPU uniquement
    device = torch.device("cpu")
    midas.to(device)

    transforms = torch.hub.load("intel-isl/MiDaS", "transforms")
    transform  = transforms.dpt_transform

    print(f"[OK] MiDaS chargé sur CPU")
    return midas, transform, device


# ============================================================
#  EXTRACTION FEATURES CLASSIFIEUR (même que train_classifier.py)
# ============================================================

def extract_features(frame: np.ndarray) -> np.ndarray:
    features = []
    frame_small = cv2.resize(frame, (128, 72))

    hsv = cv2.cvtColor(frame_small, cv2.COLOR_BGR2HSV)
    for i in range(3):
        hist = cv2.calcHist([hsv], [i], None, [16], [0, 256])
        hist = cv2.normalize(hist, hist).flatten()
        features.extend(hist.tolist())

    gray = cv2.cvtColor(frame_small, cv2.COLOR_BGR2GRAY).astype(float)
    lap  = cv2.Laplacian(gray, cv2.CV_64F)
    features.append(float(lap.var()))
    features.append(float(np.mean(np.abs(lap))))

    fft       = np.fft.fft2(gray)
    fft_shift = np.fft.fftshift(fft)
    magnitude = np.abs(fft_shift)
    h, w      = magnitude.shape
    hf_ring   = magnitude.copy()
    hf_ring[h//4:3*h//4, w//4:3*w//4] = 0
    features.append(float(np.mean(hf_ring)))
    features.append(float(np.std(hf_ring)))

    block_vars = []
    bsize = 8
    for y in range(0, gray.shape[0] - bsize, bsize):
        for x in range(0, gray.shape[1] - bsize, bsize):
            block = gray[y:y+bsize, x:x+bsize]
            block_vars.append(float(block.var()))
    features.append(float(np.mean(block_vars)))
    features.append(float(np.std(block_vars)))
    features.append(float(np.percentile(block_vars, 90)))

    sobelx  = cv2.Sobel(gray, cv2.CV_64F, 1, 0, ksize=3)
    sobely  = cv2.Sobel(gray, cv2.CV_64F, 0, 1, ksize=3)
    grad    = np.sqrt(sobelx**2 + sobely**2)
    features.append(float(np.mean(grad)))
    features.append(float(np.std(grad)))
    features.append(float(np.percentile(grad, 95)))

    sat = hsv[:, :, 1].astype(float)
    features.append(float(np.mean(sat)))
    features.append(float(np.std(sat)))

    edges_thin  = cv2.Canny(frame_small, 100, 200)
    edges_thick = cv2.Canny(frame_small, 30, 80)
    thin_count  = float(np.sum(edges_thin > 0))
    thick_count = float(np.sum(edges_thick > 0))
    ratio       = thin_count / (thick_count + 1e-6)
    features.append(ratio)
    features.append(thin_count)
    features.append(thick_count)

    return np.array(features, dtype=float)


# ============================================================
#  ÉTAPE 1 — EXTRACTION DES FRAMES IRL
# ============================================================

def is_irl_frame(frame, clf, scaler) -> tuple:
    """Retourne (is_irl: bool, confidence: float)"""
    feats  = extract_features(frame).reshape(1, -1)
    scaled = scaler.transform(feats)
    proba  = clf.predict_proba(scaled)[0]
    # classe 1 = IRL
    classes = list(clf.classes_)
    irl_idx = classes.index(1) if 1 in classes else 1
    conf    = float(proba[irl_idx])
    return conf >= IRL_CONFIDENCE, conf


def parse_date_from_filename(filepath: str) -> datetime:
    basename = os.path.basename(filepath)
    pattern  = r'(\d{4}-\d{2}-\d{2})_(\d{2}-\d{2}-\d{2})'
    match    = re.search(pattern, basename)
    if not match:
        return None
    date_str = match.group(1)
    time_str = match.group(2).replace('-', ':')
    dt = datetime.strptime(f"{date_str} {time_str}", "%Y-%m-%d %H:%M:%S")
    return dt.replace(tzinfo=timezone.utc)


def extract_irl_frames_from_video(args) -> list:
    """
    Worker multiprocessing : extrait les frames IRL d'une vidéo.
    Retourne liste de chemins de frames sauvegardées.
    """
    video_file, clf, scaler, out_dir, max_per_video = args

    if not os.path.exists(video_file):
        return []

    cap = cv2.VideoCapture(video_file)
    if not cap.isOpened():
        return []

    fps          = cap.get(cv2.CAP_PROP_FPS) or 30
    total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
    duration_sec = int(total_frames / fps)

    saved_frames = []
    curr         = 0
    prev_gray    = None
    video_base   = os.path.splitext(os.path.basename(video_file))[0]

    while curr <= duration_sec and len(saved_frames) < max_per_video:
        cap.set(cv2.CAP_PROP_POS_MSEC, curr * 1000)
        ret, frame = cap.read()
        if not ret:
            break

        # Filtre classifieur
        is_irl, conf = is_irl_frame(frame, clf, scaler)

        if is_irl:
            # Filtre variation brutale (passage devant cam)
            gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
            if prev_gray is not None:
                variation = float(np.mean(cv2.absdiff(prev_gray, gray)))
                if variation > 35.0:
                    prev_gray = gray.copy()
                    curr += SAMPLE_STEP
                    continue
            prev_gray = gray.copy() if prev_gray is None else prev_gray

            # Sauvegarde frame
            frame_name = f"{video_base}_{curr:06d}.jpg"
            frame_path = os.path.join(out_dir, frame_name)
            cv2.imwrite(frame_path, frame, [cv2.IMWRITE_JPEG_QUALITY, 85])
            saved_frames.append({
                "path": frame_path,
                "video": video_file,
                "t_sec": curr,
                "conf": conf,
            })

        curr += SAMPLE_STEP

    cap.release()
    return saved_frames


def extract_all_irl_frames(clf, scaler) -> list:
    """Lance l'extraction en parallèle sur tous les .mp4"""
    os.makedirs(OUTPUT_FRAMES, exist_ok=True)

    mp4_files = sorted(glob.glob("*.mp4"))
    if not mp4_files:
        print("[ERR] Aucun .mp4 trouvé")
        sys.exit(1)

    print(f"\n[EXTRACT] {len(mp4_files)} vidéos trouvées")
    print(f"[EXTRACT] {NB_WORKERS} workers parallèles")

    max_per_video = max(10, (MAX_FRAMES or 99999) // len(mp4_files))

    args_list = [
        (f, clf, scaler, OUTPUT_FRAMES, max_per_video)
        for f in mp4_files
    ]

    # Multiprocessing sur les workers — context "spawn" pour éviter
    # les conflits avec PyTorch (pas d'AVX, pas de fork de threads BLAS)
    ctx = mp.get_context("spawn")
    with ctx.Pool(processes=NB_WORKERS) as pool:
        results = pool.map(extract_irl_frames_from_video, args_list)

    all_frames = []
    for r in results:
        all_frames.extend(r)

    print(f"\n[OK] {len(all_frames)} frames IRL extraites")

    if MAX_FRAMES and len(all_frames) > MAX_FRAMES:
        # Garde les frames les plus confiantes
        all_frames = sorted(all_frames, key=lambda x: -x["conf"])[:MAX_FRAMES]
        print(f"[OK] Réduit à {MAX_FRAMES} frames (plus confiantes)")

    return all_frames


# ============================================================
#  ÉTAPE 2 — ESTIMATION DE PROFONDEUR (MiDaS)
# ============================================================

def estimate_depth(frame_bgr: np.ndarray, midas, transform, device) -> np.ndarray:
    """Retourne une carte de profondeur normalisée [0,1]."""
    img_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)
    input_batch = transform(img_rgb).to(device)

    with torch.no_grad():
        prediction = midas(input_batch)
        prediction = torch.nn.functional.interpolate(
            prediction.unsqueeze(1),
            size=frame_bgr.shape[:2],
            mode="bicubic",
            align_corners=False,
        ).squeeze()

    depth = prediction.cpu().numpy()

    # Normalisation [0, 1]
    dmin, dmax = depth.min(), depth.max()
    if dmax - dmin > 1e-6:
        depth = (depth - dmin) / (dmax - dmin)

    return depth.astype(np.float32)


def compute_all_depths(frames: list, midas, transform, device) -> dict:
    """
    Calcule la carte de profondeur pour chaque frame.
    Retourne un dict {frame_path: depth_array}.
    Cache les résultats dans depth_cache.npz.
    """
    if os.path.exists(DEPTH_CACHE):
        print(f"\n[DEPTH] Cache trouvé : {DEPTH_CACHE}")
        cache = np.load(DEPTH_CACHE, allow_pickle=True)
        return dict(cache["depths"].item())

    print(f"\n[DEPTH] Estimation profondeur sur {len(frames)} frames (CPU)...")
    print("        (peut prendre du temps sur CPU — ~1-3s par frame)")

    depths = {}
    for i, frame_info in enumerate(frames):
        path  = frame_info["path"]
        frame = cv2.imread(path)
        if frame is None:
            continue

        depth = estimate_depth(frame, midas, transform, device)
        depths[path] = depth

        if (i + 1) % 10 == 0:
            pct = 100 * (i + 1) / len(frames)
            print(f"  {i+1}/{len(frames)} ({pct:.0f}%)...", end="\r")

    np.savez(DEPTH_CACHE, depths=depths)
    print(f"\n[OK] Profondeurs calculées et mises en cache")
    return depths


# ============================================================
#  ÉTAPE 3 — STRUCTURE FROM MOTION (SfM léger)
# ============================================================

def compute_camera_poses(frames: list) -> list:
    """
    Estime les poses caméra entre frames consécutives via ORB + homographie.
    Retourne une liste de matrices de transformation 4x4.
    """
    print(f"\n[SFM] Calcul des poses caméra...")

    orb      = cv2.ORB_create(nfeatures=1000)
    matcher  = cv2.BFMatcher(cv2.NORM_HAMMING, crossCheck=True)

    poses    = [np.eye(4)]  # première pose = identité
    curr_pose = np.eye(4)

    # Paramètres caméra approximatifs pour 480p
    # Focale approximée pour un téléphone standard
    h, w  = 480, 854
    fx    = w * 0.8   # approximation focale
    fy    = fx
    cx, cy = w / 2, h / 2
    K     = np.array([[fx, 0, cx],
                      [0, fy, cy],
                      [0,  0,  1]], dtype=float)

    prev_frame = None
    prev_kp    = None
    prev_des   = None

    for i, frame_info in enumerate(frames[:-1]):
        frame = cv2.imread(frame_info["path"])
        if frame is None:
            poses.append(curr_pose.copy())
            continue

        gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
        kp, des = orb.detectAndCompute(gray, None)

        if prev_frame is not None and prev_des is not None and des is not None:
            matches = matcher.match(prev_des, des)
            matches = sorted(matches, key=lambda x: x.distance)

            if len(matches) >= 8:
                pts1 = np.float32([prev_kp[m.queryIdx].pt for m in matches])
                pts2 = np.float32([kp[m.trainIdx].pt for m in matches])

                # Matrice essentielle
                E, mask = cv2.findEssentialMat(
                    pts1, pts2, K,
                    method=cv2.RANSAC,
                    prob=0.999,
                    threshold=1.0
                )

                if E is not None:
                    _, R, t, _ = cv2.recoverPose(E, pts1, pts2, K, mask=mask)

                    # Construire transformation relative
                    T_rel      = np.eye(4)
                    T_rel[:3, :3] = R
                    T_rel[:3, 3]  = t.flatten()

                    curr_pose = curr_pose @ T_rel

        poses.append(curr_pose.copy())
        prev_frame = gray
        prev_kp    = kp
        prev_des   = des

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

    print(f"\n[OK] {len(poses)} poses calculées")
    return poses, K


# ============================================================
#  ÉTAPE 4 — FUSION NUAGE DE POINTS
# ============================================================

def build_point_cloud(frames: list, depths: dict, poses: list, K: np.ndarray) -> dict:
    """
    Fusionne les cartes de profondeur avec les poses caméra
    pour construire un nuage de points 3D global.
    Retourne un dict {"points": np.ndarray (N,3), "colors": np.ndarray (N,3)}
    """
    print(f"\n[CLOUD] Construction du nuage de points...")

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

    all_points = []
    all_colors = []

    # Grille de pixels (sous-échantillonné pour performance)
    skip = 4  # un pixel sur 4

    for i, frame_info in enumerate(frames):
        path = frame_info["path"]

        if path not in depths:
            continue

        frame = cv2.imread(path)
        if frame is None:
            continue

        depth = depths[path]
        pose  = poses[i] if i < len(poses) else np.eye(4)

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

        # Pixels valides (profondeur non nulle)
        ys, xs = np.mgrid[0:h_f:skip, 0:w_f:skip]
        ys     = ys.flatten()
        xs     = xs.flatten()
        ds     = depth[ys, xs]

        # Filtre profondeur nulle
        valid  = ds > 0.01
        ys, xs, ds = ys[valid], xs[valid], ds[valid]

        # Rétroprojection en 3D (coordonnées caméra)
        # On scale la profondeur MiDaS (relative) par une constante
        depth_scale = 3.0  # mètres approximatifs (chambre ~3m)
        Z = ds * depth_scale
        X = (xs - cx) / fx * Z
        Y = (ys - cy) / fy * Z

        pts_cam = np.stack([X, Y, Z, np.ones_like(Z)], axis=1)

        # Transformation en coordonnées monde
        pts_world = (pose @ pts_cam.T).T[:, :3]

        # Couleurs
        colors = frame_rgb[ys, xs]

        all_points.append(pts_world)
        all_colors.append(colors)

        if (i + 1) % 20 == 0:
            total_pts = sum(len(p) for p in all_points)
            print(f"  {i+1}/{len(frames)} frames | {total_pts:,} points...", end="\r")

    if not all_points:
        print("[ERR] Aucun point généré")
        sys.exit(1)

    all_points = np.vstack(all_points)
    all_colors = np.vstack(all_colors)

    print(f"\n[OK] {len(all_points):,} points générés")

    return {"points": all_points, "colors": all_colors}


# ============================================================
#  ÉTAPE 5 — NETTOYAGE ET EXPORT
# ============================================================

def clean_and_export(pcd: dict):
    """
    Nettoie le nuage (supprime outliers, voxel downsampling) et exporte en .ply
    Utilise pyntcloud + numpy/scipy (pas besoin d'AVX/Open3D).
    """
    print(f"\n[CLEAN] Nettoyage du nuage de points...")

    points = pcd["points"]
    colors = pcd["colors"]

    print(f"  Avant nettoyage : {len(points):,} pts")

    # --- Suppression des outliers statistiques (équivalent remove_statistical_outlier) ---
    # Pour chaque point, calcule la distance moyenne à ses k voisins
    k = 20
    tree = cKDTree(points)
    dists, _ = tree.query(points, k=k + 1)  # k+1 car le point lui-même est inclus
    mean_dists = dists[:, 1:].mean(axis=1)  # exclure distance à soi-même (0)
    threshold = mean_dists.mean() + 2.0 * mean_dists.std()
    mask = mean_dists < threshold
    points = points[mask]
    colors = colors[mask]
    print(f"  Après outliers  : {len(points):,} pts")

    # --- Voxel downsampling (équivalent voxel_down_sample) ---
    voxel_size = 0.02
    voxel_indices = np.floor(points / voxel_size).astype(np.int32)
    # Utilise un dict pour garder un seul point par voxel
    voxel_dict = {}
    for i, vi in enumerate(map(tuple, voxel_indices)):
        if vi not in voxel_dict:
            voxel_dict[vi] = i
    keep = np.array(list(voxel_dict.values()))
    points = points[keep]
    colors = colors[keep]
    print(f"  Après voxel     : {len(points):,} pts")

    # --- Estimation des normales (PCA locale, pour MeshLab) ---
    print("  Estimation des normales...")
    normals = np.zeros_like(points)
    radius  = 0.1
    max_nn  = 30
    tree2   = cKDTree(points)
    for i in range(len(points)):
        idxs = tree2.query_ball_point(points[i], r=radius)
        if len(idxs) > max_nn:
            # Garder les max_nn plus proches
            d = np.linalg.norm(points[idxs] - points[i], axis=1)
            idxs = [idxs[j] for j in np.argsort(d)[:max_nn]]
        if len(idxs) < 3:
            continue
        neighbors = points[idxs]
        cov = np.cov((neighbors - neighbors.mean(axis=0)).T)
        eigvals, eigvecs = np.linalg.eigh(cov)
        normals[i] = eigvecs[:, 0]  # vecteur propre minimal = normale

    # --- Export .ply via pyntcloud ---
    colors_uint8 = np.clip(colors * 255, 0, 255).astype(np.uint8)
    normals_f32  = normals.astype(np.float32)

    df = pd.DataFrame({
        "x": points[:, 0].astype(np.float32),
        "y": points[:, 1].astype(np.float32),
        "z": points[:, 2].astype(np.float32),
        "red":   colors_uint8[:, 0],
        "green": colors_uint8[:, 1],
        "blue":  colors_uint8[:, 2],
        "nx": normals_f32[:, 0],
        "ny": normals_f32[:, 1],
        "nz": normals_f32[:, 2],
    })

    cloud = PyntCloud(df)
    cloud.to_file(OUTPUT_PLY)

    print(f"\n[OK] Nuage exporté → {OUTPUT_PLY}")
    print(f"     Ouvre dans MeshLab : File → Import Mesh → {OUTPUT_PLY}")

    # Stats finales
    bbox = points.max(axis=0) - points.min(axis=0)
    print(f"\n  Dimensions reconstituées :")
    print(f"    Largeur   : {bbox[0]:.2f}m")
    print(f"    Profondeur: {bbox[1]:.2f}m")
    print(f"    Hauteur   : {bbox[2]:.2f}m")

    return {"points": points, "colors": colors}


# ============================================================
#  MAIN
# ============================================================

if __name__ == "__main__":
    # Nécessaire sur Windows et conseillé sur Linux avec PyTorch
    freeze_support()
    # "fork" est le défaut sur Linux mais cause des conflits avec PyTorch
    # On force "spawn" pour isoler proprement les workers du process PyTorch
    mp.set_start_method("spawn", force=True)

    print("=" * 60)
    print("  RECONSTRUCTION 3D CHAMBRE")
    print(f"  Workers : {NB_WORKERS} / {cpu_count()} cores disponibles")
    print("=" * 60)

    # 1. Classifieur
    clf, scaler = load_classifier()

    # 2. Extraction frames IRL (multiprocessing)
    print("\n[ÉTAPE 1/5] Extraction des frames IRL...")
    all_frames = extract_all_irl_frames(clf, scaler)

    if not all_frames:
        print("[ERR] Aucune frame IRL trouvée — vérifie le classifieur")
        sys.exit(1)

    # 3. MiDaS
    print("\n[ÉTAPE 2/5] Chargement MiDaS...")
    midas, transform, device = load_midas()

    # 4. Profondeur
    print("\n[ÉTAPE 3/5] Estimation de profondeur...")
    depths = compute_all_depths(all_frames, midas, transform, device)

    # 5. SfM
    print("\n[ÉTAPE 4/5] Structure from Motion...")
    poses, K = compute_camera_poses(all_frames)

    # 6. Nuage de points
    print("\n[ÉTAPE 5/5] Construction et export nuage 3D...")
    pcd = build_point_cloud(all_frames, depths, poses, K)
    pcd_final = clean_and_export(pcd)

    print("\n" + "=" * 60)
    print("  TERMINÉ !")
    print(f"  → Ouvre {OUTPUT_PLY} dans MeshLab")
    print("=" * 60)
