"""
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 glob
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 segment gameplay
IRL_VIDEO = "ghostofatale - 2026-03-07_11-53-32.mp4"
IRL_EXCLUDE_START = 59 * 60 + 9    # 59:09 en secondes
IRL_EXCLUDE_END   = 61 * 60 + 1    # 1:01:01 en secondes

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

# Nombre de frames à extraire par classe
NB_FRAMES_PER_CLASS = 200

# Pas d'échantillonnage en secondes
SAMPLE_STEP = 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)
    """
    features = []

    # Redimensionner pour uniformiser (et accélérer)
    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 (mesure de netteté/bruit) ---
    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) ---
    # Un jeu vidéo H264 a un spectre HF différent d'une caméra téléphone
    fft = np.fft.fft2(gray)
    fft_shift = np.fft.fftshift(fft)
    magnitude = np.abs(fft_shift)
    h, w = magnitude.shape
    # Zone haute fréquence = bords de la magnitude FFT
    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 (texture) ---
    # Calcul sur des 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 moyenne et écart-type ---
    sat = hsv[:, :, 1].astype(float)
    features.append(float(np.mean(sat)))
    features.append(float(np.std(sat)))

    # --- Ratio bords fins / bords épais ---
    # Bords fins = typique jeu vidéo (aliasing net)
    # Bords épais = typique caméra (flou naturel)
    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)


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

def extract_frames(video_file: str, label: int,
                   exclude_start: int = -1,
                   exclude_end: int = -1,
                   max_frames: int = NB_FRAMES_PER_CLASS) -> tuple:
    """
    Extrait max_frames frames depuis video_file.
    Exclut le segment [exclude_start, exclude_end] si défini.
    Retourne (X, y) où X = features, y = labels.
    """
    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"
    print(f"\n[EXTRACT] {video_file} → {label_name} ({duration_sec//60}min)")

    curr = 0
    extracted = 0

    while curr <= duration_sec and extracted < max_frames:
        # Exclure le segment indésirable
        if exclude_start >= 0 and exclude_start <= curr <= exclude_end:
            curr += SAMPLE_STEP
            continue

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

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

        if extracted % 20 == 0:
            print(f"  {extracted}/{max_frames} frames...", end="\r")

        curr += SAMPLE_STEP

    cap.release()
    print(f"\n  [OK] {extracted} frames extraites")
    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] {len(X)} frames ({np.sum(y==1)} IRL, {np.sum(y==0)} gameplay)")

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

    # Random Forest — optimisé pour CPU multi-core
    clf = RandomForestClassifier(
        n_estimators=200,
        max_depth=15,
        min_samples_split=5,
        n_jobs=-1,          # utilise tous les cores (28 threads sur ton Xeon)
        random_state=42,
        class_weight='balanced',
    )

    # Cross-validation 5-fold
    print("[TRAIN] Cross-validation 5-fold...")
    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% — essaie d'ajouter plus de frames de référence")
    else:
        print("[OK] Classifieur suffisamment précis ✓")

    # Entraînement final sur tout le dataset
    clf.fit(X_scaled, y)

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

    # Importance des features
    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 SUR QUELQUES FRAMES
# ============================================================

def quick_test(clf, scaler):
    """
    Teste le classifieur sur quelques frames des deux vidéos
    et affiche les résultats pour vérification visuelle.
    """
    print("\n[TEST] Vérification rapide...")

    test_cases = [
        (IRL_VIDEO,      1, 10,  "IRL"),
        (IRL_VIDEO,      1, 120, "IRL"),
        (GAMEPLAY_VIDEO, 0, 30,  "GAMEPLAY"),
        (GAMEPLAY_VIDEO, 0, 200, "GAMEPLAY"),
    ]

    for video, true_label, t_sec, label_name in test_cases:
        if not os.path.exists(video):
            continue
        cap = cv2.VideoCapture(video)
        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

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


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

if __name__ == "__main__":
    print("=" * 60)
    print("  ENTRAÎNEMENT CLASSIFIEUR IRL vs GAMEPLAY")
    print("=" * 60)

    # Vérifications
    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 frames IRL (en excluant le segment gameplay)
    X_irl, y_irl = extract_frames(
        IRL_VIDEO, label=1,
        exclude_start=IRL_EXCLUDE_START,
        exclude_end=IRL_EXCLUDE_END,
        max_frames=NB_FRAMES_PER_CLASS,
    )

    # Extraction frames gameplay
    X_gp, y_gp = extract_frames(
        GAMEPLAY_VIDEO, label=0,
        max_frames=NB_FRAMES_PER_CLASS,
    )

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

    # Entraînement
    clf, scaler = train_classifier(X, y)

    # Test rapide
    quick_test(clf, scaler)

    print("\n[OK] Classifieur prêt !")
    print("     Intègre-le dans ctf_solver_v2.py avec :")
    print("     clf    = joblib.load('classifier_irl_vs_gameplay.joblib')")
    print("     scaler = joblib.load('scaler_irl_vs_gameplay.joblib')")
