"""
CLASSIFIEUR IRL vs GAMEPLAY
Extrait des frames de référence et entraîne un classifieur léger.

DÉPENDANCES :
    pip install opencv-python numpy scikit-learn joblib

USAGE :
    python train_classifier.py
"""

import cv2
import numpy as np
import os
import sys
import joblib
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import cross_val_score
from sklearn.preprocessing import StandardScaler

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

# Vidéo IRL — tout sauf le(s) segment(s) gameplay
# Pour exclure plusieurs segments : liste de tuples (start_sec, end_sec)
IRL_VIDEO = "ghostofatale - 2026-03-07_11-53-32.mp4"
IRL_EXCLUDE_SEGMENTS = [
    (59 * 60 + 9, 61 * 60 + 1),   # 59:09 → 1:01:01
    # Ajoute d'autres segments à exclure ici si besoin :
    # (120 * 60, 125 * 60),
]

# Vidéo gameplay pur
GAMEPLAY_VIDEO = "ghostofatale - 2026-04-01_14-12-10.mp4"

# ---------------------------------------------------------------
# NOMBRE DE FRAMES — MASSIVEMENT AUGMENTÉ
# La vidéo entière sera parcourue, ce nombre est un plafond de sécurité.
# Mets une valeur très grande pour tout prendre.
# ---------------------------------------------------------------
NB_FRAMES_IRL      = 99999   # prend TOUT ce qui est disponible en IRL
NB_FRAMES_GAMEPLAY = 99999   # pareil pour le gameplay

# Pas d'échantillonnage en secondes — réduit pour plus de densité
SAMPLE_STEP = 3   # 1 frame toutes les 3 secondes (était 15)

# Fichiers de sortie
MODEL_FILE   = "classifier_irl_vs_gameplay.joblib"
SCALER_FILE  = "scaler_irl_vs_gameplay.joblib"

# ============================================================
#  EXTRACTION DE FEATURES
# ============================================================

def extract_features(frame: np.ndarray) -> np.ndarray:
    """
    Extrait un vecteur de features discriminantes IRL vs gameplay.

    Features utilisées :
    - Histogramme HSV réduit (couleur globale)
    - Variance du laplacien (netteté / bruit)
    - Énergie haute fréquence via FFT (grain capteur vs rendu jeu)
    - Variance locale (texture naturelle vs texture jeu)
    - Stats sur les gradients (bords IRL flous vs bords jeu nets)
    - Saturation moyenne et écart-type
    - Ratio bords fins / bords épais (Canny)
    """
    features = []

    frame_small = cv2.resize(frame, (128, 72))

    # --- HSV histogram (48 bins) ---
    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())

    # --- Variance du Laplacien ---
    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))))

    # --- Énergie haute fréquence (FFT) ---
    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)))

    # --- Variance locale (blocs 8x8) ---
    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)))

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

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

    # --- Ratio bords fins / bords épais ---
    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)


# ============================================================
#  VÉRIFICATION EXCLUSION
# ============================================================

def is_excluded(t_sec: int, segments: list) -> bool:
    """Retourne True si t_sec est dans un segment à exclure."""
    for (start, end) in segments:
        if start <= t_sec <= end:
            return True
    return False


# ============================================================
#  EXTRACTION DES FRAMES
# ============================================================

def extract_frames(video_file: str, label: int,
                   exclude_segments: list = None,
                   max_frames: int = 99999) -> tuple:
    """
    Extrait jusqu'à max_frames frames depuis video_file.
    Exclut les segments listés dans exclude_segments.
    Parcourt la vidéo ENTIÈRE de début à fin avec SAMPLE_STEP secondes d'intervalle.
    Retourne (X, y).
    """
    if exclude_segments is None:
        exclude_segments = []

    X, y = [], []

    if not os.path.exists(video_file):
        print(f"[ERR] Vidéo introuvable : {video_file}")
        return np.array(X), np.array(y)

    cap = cv2.VideoCapture(video_file)
    if not cap.isOpened():
        print(f"[ERR] Impossible d'ouvrir : {video_file}")
        return np.array(X), np.array(y)

    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)

    label_name = "IRL" if label == 1 else "GAMEPLAY"
    total_possible = duration_sec // SAMPLE_STEP
    print(f"\n[EXTRACT] {os.path.basename(video_file)}")
    print(f"          Classe     : {label_name}")
    print(f"          Durée      : {duration_sec//3600:02d}h{(duration_sec%3600)//60:02d}m{duration_sec%60:02d}s")
    print(f"          Pas        : {SAMPLE_STEP}s → ~{total_possible} frames possibles")

    # Calcul du nombre de secondes exclues
    excluded_secs = sum(end - start for start, end in exclude_segments)
    if excluded_secs > 0:
        print(f"          Exclusions : {excluded_secs//60}min ({len(exclude_segments)} segment(s))")

    curr      = 0
    extracted = 0
    skipped   = 0

    while curr <= duration_sec and extracted < max_frames:

        if is_excluded(curr, exclude_segments):
            skipped += 1
            curr += SAMPLE_STEP
            continue

        cap.set(cv2.CAP_PROP_POS_MSEC, curr * 1000)
        ret, frame = cap.read()
        if not ret:
            curr += SAMPLE_STEP
            continue

        feats = extract_features(frame)
        X.append(feats)
        y.append(label)
        extracted += 1

        if extracted % 100 == 0:
            pct = (curr / duration_sec * 100) if duration_sec > 0 else 0
            print(f"  {extracted} frames extraites | {curr//60}:{curr%60:02d} / {duration_sec//60}:{duration_sec%60:02d} ({pct:.0f}%)   ", end="\r")

        curr += SAMPLE_STEP

    cap.release()
    print(f"\n  [OK] {extracted} frames extraites ({skipped} secondes exclues)")
    return np.array(X), np.array(y)


# ============================================================
#  ENTRAÎNEMENT
# ============================================================

def train_classifier(X: np.ndarray, y: np.ndarray):
    """
    Entraîne un Random Forest et évalue par cross-validation.
    Sauvegarde le modèle + scaler.
    """
    print(f"\n[TRAIN] Dataset total : {len(X)} frames")
    print(f"        IRL      : {np.sum(y==1)}")
    print(f"        Gameplay : {np.sum(y==0)}")

    scaler   = StandardScaler()
    X_scaled = scaler.fit_transform(X)

    # Random Forest avec plus d'arbres vu le volume de données
    clf = RandomForestClassifier(
        n_estimators=300,
        max_depth=20,
        min_samples_split=5,
        n_jobs=-1,
        random_state=42,
        class_weight='balanced',
    )

    print("\n[TRAIN] Cross-validation 5-fold (peut prendre quelques minutes)...")
    scores = cross_val_score(clf, X_scaled, y, cv=5, scoring='accuracy', n_jobs=-1)
    print(f"  Accuracy : {scores.mean()*100:.1f}% ± {scores.std()*100:.1f}%")

    if scores.mean() < 0.85:
        print("[WARN] Précision < 85% — les deux classes sont peut-être visuellement proches")
    elif scores.mean() < 0.92:
        print("[OK] Classifieur correct (>85%) — acceptable")
    else:
        print("[OK] Classifieur excellent (>92%) ✓")

    print("\n[TRAIN] Entraînement final sur tout le dataset...")
    clf.fit(X_scaled, y)

    joblib.dump(clf,    MODEL_FILE)
    joblib.dump(scaler, SCALER_FILE)
    print(f"[OK] Modèle sauvegardé  → {MODEL_FILE}")
    print(f"[OK] Scaler sauvegardé  → {SCALER_FILE}")

    feature_names = (
        [f"hsv_h_{i}" for i in range(16)] +
        [f"hsv_s_{i}" for i in range(16)] +
        [f"hsv_v_{i}" for i in range(16)] +
        ["lap_var", "lap_mean",
         "fft_hf_mean", "fft_hf_std",
         "block_var_mean", "block_var_std", "block_var_p90",
         "grad_mean", "grad_std", "grad_p95",
         "sat_mean", "sat_std",
         "edge_ratio", "edge_thin", "edge_thick"]
    )
    importances = clf.feature_importances_
    top_idx = np.argsort(importances)[::-1][:10]
    print("\n  Top-10 features les plus discriminantes :")
    for i in top_idx:
        name = feature_names[i] if i < len(feature_names) else f"feat_{i}"
        print(f"    {name:<25} {importances[i]*100:.2f}%")

    return clf, scaler


# ============================================================
#  TEST RAPIDE
# ============================================================

def quick_test(clf, scaler):
    print("\n[TEST] Vérification rapide sur quelques timestamps...")

    test_cases = [
        (IRL_VIDEO,      1,  10,  "IRL"),
        (IRL_VIDEO,      1,  300, "IRL"),
        (IRL_VIDEO,      1,  900, "IRL"),
        (IRL_VIDEO,      1,  1800,"IRL"),
        (GAMEPLAY_VIDEO, 0,  30,  "GAMEPLAY"),
        (GAMEPLAY_VIDEO, 0,  300, "GAMEPLAY"),
        (GAMEPLAY_VIDEO, 0,  900, "GAMEPLAY"),
        (GAMEPLAY_VIDEO, 0,  1800,"GAMEPLAY"),
    ]

    ok = 0
    total = 0
    for video, true_label, t_sec, label_name in test_cases:
        if not os.path.exists(video):
            continue
        cap = cv2.VideoCapture(video)
        fps   = cap.get(cv2.CAP_PROP_FPS) or 30
        total_f = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
        dur   = int(total_f / fps)
        if t_sec > dur:
            cap.release()
            continue
        cap.set(cv2.CAP_PROP_POS_MSEC, t_sec * 1000)
        ret, frame = cap.read()
        cap.release()
        if not ret:
            continue

        feats  = extract_features(frame).reshape(1, -1)
        scaled = scaler.transform(feats)
        pred   = clf.predict(scaled)[0]
        proba  = clf.predict_proba(scaled)[0]

        pred_name = "IRL" if pred == 1 else "GAMEPLAY"
        correct   = "✓" if pred == true_label else "✗"
        conf      = max(proba) * 100
        total += 1
        if pred == true_label:
            ok += 1

        print(f"  {correct} t={t_sec:5d}s | réel={label_name:<8} prédit={pred_name:<8} confiance={conf:.1f}%")

    if total > 0:
        print(f"\n  Score test rapide : {ok}/{total} ({ok/total*100:.0f}%)")


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

if __name__ == "__main__":
    print("=" * 60)
    print("  ENTRAÎNEMENT CLASSIFIEUR IRL vs GAMEPLAY")
    print(f"  Pas d'échantillonnage : {SAMPLE_STEP}s")
    print("=" * 60)

    missing = []
    if not os.path.exists(IRL_VIDEO):
        missing.append(IRL_VIDEO)
    if not os.path.exists(GAMEPLAY_VIDEO):
        missing.append(GAMEPLAY_VIDEO)
    if missing:
        print(f"[ERR] Fichiers manquants : {missing}")
        sys.exit(1)

    # --- Extraction IRL (vidéo entière sauf segments exclus) ---
    print("\n[ÉTAPE 1/3] Extraction frames IRL...")
    X_irl, y_irl = extract_frames(
        IRL_VIDEO,
        label=1,
        exclude_segments=IRL_EXCLUDE_SEGMENTS,
        max_frames=NB_FRAMES_IRL,
    )

    # --- Extraction Gameplay (vidéo entière) ---
    print("\n[ÉTAPE 2/3] Extraction frames Gameplay...")
    X_gp, y_gp = extract_frames(
        GAMEPLAY_VIDEO,
        label=0,
        exclude_segments=[],
        max_frames=NB_FRAMES_GAMEPLAY,
    )

    if len(X_irl) == 0 or len(X_gp) == 0:
        print("[ERR] Une des deux classes est vide — vérifie les vidéos")
        sys.exit(1)

    # --- Fusion ---
    X = np.vstack([X_irl, X_gp])
    y = np.concatenate([y_irl, y_gp])

    # --- Entraînement ---
    print("\n[ÉTAPE 3/3] Entraînement...")
    clf, scaler = train_classifier(X, y)

    # --- Test ---
    quick_test(clf, scaler)

    print("\n" + "=" * 60)
    print("  CLASSIFIEUR PRÊT !")
    print(f"  → {MODEL_FILE}")
    print(f"  → {SCALER_FILE}")
    print("=" * 60)
