"""
CTF SOLVER v2 - Localisation par analyse solaire + droites d'azimut
Village cible : Verdun-sur-le-Doubs
 
DÉPENDANCES :
    pip install opencv-python numpy matplotlib pysolar pillow

USAGE :
    1. Mets tous tes .mp4 + village.png dans le même dossier
    2. Lance : python ctf_solver_v2.py
    3. Le script détecte automatiquement le bon décor et extrait la lumière
"""

import cv2
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.patches as patches
from datetime import datetime, timedelta, timezone
import json
import os
import re
import sys
import glob

try:
    from pysolar.solar import get_altitude, get_azimuth
    PYSOLAR_OK = True
except ImportError:
    print("[WARN] pysolar non installé. Lance : pip install pysolar")
    PYSOLAR_OK = False


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

# Coordonnées GPS de Verdun-sur-le-Doubs
LAT = 46.8945173
LON = 5.0232693

# ROI fenêtre (y_start, x_start, hauteur, largeur)
ROI = (306, 0, 50, 50)

# Segments de référence pour construire le modèle de décor
# (dans la première vidéo disponible)
REFERENCE_SEGMENTS = [
    ("00:41:12", "00:51:12"),
    ("01:28:16", "01:34:41"),
    ("05:02:17", "05:14:02"),
    ("07:02:07", "07:15:00"),
]

# Seuil de similarité décor (0.0 à 1.0 — plus haut = plus strict)
SIMILARITY_THRESHOLD = 0.75

# Pas d'échantillonnage en secondes pour l'extraction
STEP = 5

# Nombre de frames de référence à extraire pour modéliser le décor
NB_REF_FRAMES = 30

# Image satellite
SATELLITE_IMAGE = "village.png"
SAT_GPS_TL = (46.914627, 4.993286)
SAT_GPS_BR = (46.875213, 5.050964)

# Fichiers de cache
OUTPUT_DATA    = "brightness_data_v2.json"
OUTPUT_HEATMAP = "heatmap_v2.png"


# ============================================================
#  PARSE DATE DEPUIS NOM DE FICHIER
# ============================================================

def parse_date_from_filename(filepath: str) -> datetime:
    """
    Parse la date/heure depuis un nom de fichier au format :
    'ghostofatale - 2026-02-22_11-19-28.mp4'
    Retourne un datetime UTC.
    """
    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:
        raise ValueError(f"Impossible de parser la date depuis : {basename}")
    
    date_str = match.group(1)          # 2026-02-22
    time_str = match.group(2).replace('-', ':')  # 11:19:28
    
    dt = datetime.strptime(f"{date_str} {time_str}", "%Y-%m-%d %H:%M:%S")
    dt = dt.replace(tzinfo=timezone.utc)
    
    print(f"  [PARSE] {basename} → {dt.strftime('%Y-%m-%d %H:%M:%S UTC')}")
    return dt


# ============================================================
#  ÉTAPE 1 — CONSTRUCTION DU MODÈLE DE DÉCOR (référence)
# ============================================================

def to_sec(t_str: str) -> int:
    h, m, s = map(int, t_str.split(':'))
    return h*3600 + m*60 + s


def build_reference_model(ref_video: str) -> dict:
    """
    Extrait NB_REF_FRAMES frames depuis les segments de référence.
    Construit :
    - mean_frame  : frame moyenne du décor
    - mean_hist   : histogramme moyen (comparaison rapide)
    - roi_stats   : stats de luminosité de la ROI sur ces frames
    """
    cache_file = "reference_model.npz"
    if os.path.exists(cache_file):
        print(f"[INFO] Modèle de référence chargé depuis cache")
        data = np.load(cache_file, allow_pickle=True)
        return {
            "mean_frame": data["mean_frame"],
            "mean_hist":  data["mean_hist"],
        }

    print(f"\n[RÉFÉRENCE] Construction du modèle depuis : {ref_video}")
    cap = cv2.VideoCapture(ref_video)
    if not cap.isOpened():
        print(f"[ERR] Impossible d'ouvrir {ref_video}")
        sys.exit(1)

    frames = []
    step_ref = max(1, sum(to_sec(e) - to_sec(s) 
                          for s, e in REFERENCE_SEGMENTS) // NB_REF_FRAMES)

    for seg_start, seg_end in REFERENCE_SEGMENTS:
        curr = to_sec(seg_start)
        limit = to_sec(seg_end)
        cap.set(cv2.CAP_PROP_POS_MSEC, curr * 1000)

        while curr <= limit and len(frames) < NB_REF_FRAMES:
            ret, frame = cap.read()
            if not ret:
                break
            frames.append(frame.astype(np.float32))
            curr += step_ref
            cap.set(cv2.CAP_PROP_POS_MSEC, curr * 1000)

    cap.release()

    if not frames:
        print("[ERR] Aucune frame de référence extraite")
        sys.exit(1)

    mean_frame = np.mean(frames, axis=0).astype(np.float32)

    # Histogramme moyen (sur frame entière en niveaux de gris)
    mean_gray = cv2.cvtColor(mean_frame.astype(np.uint8), cv2.COLOR_BGR2GRAY)
    mean_hist = cv2.calcHist([mean_gray], [0], None, [64], [0, 256])
    cv2.normalize(mean_hist, mean_hist)

    np.savez(cache_file, mean_frame=mean_frame, mean_hist=mean_hist)
    print(f"[OK] {len(frames)} frames de référence extraites et mises en cache")

    return {"mean_frame": mean_frame, "mean_hist": mean_hist}


# ============================================================
#  ÉTAPE 2 — DÉTECTION DU BON DÉCOR
# ============================================================

def frame_similarity(frame: np.ndarray, model: dict) -> float:
    """
    Calcule la similarité entre une frame et le modèle de référence.
    Combine :
    - Corrélation d'histogramme (rapide, globale)
    - SSIM simplifié sur version réduite (structurel)
    Retourne un score entre 0 et 1.
    """
    gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)

    # -- Histogramme --
    hist = cv2.calcHist([gray], [0], None, [64], [0, 256])
    cv2.normalize(hist, hist)
    hist_score = cv2.compareHist(model["mean_hist"], hist, 
                                  cv2.HISTCMP_CORREL)
    hist_score = max(0.0, float(hist_score))

    # -- Différence structurelle sur image réduite --
    ref_small = cv2.resize(model["mean_frame"].astype(np.uint8), (64, 64))
    ref_gray  = cv2.cvtColor(ref_small, cv2.COLOR_BGR2GRAY).astype(float)
    frm_small = cv2.resize(frame, (64, 64))
    frm_gray  = cv2.cvtColor(frm_small, cv2.COLOR_BGR2GRAY).astype(float)

    diff = np.abs(ref_gray - frm_gray)
    struct_score = 1.0 - float(np.mean(diff) / 255.0)

    # Score combiné (histogramme pèse plus car robuste aux variations de lumière)
    score = 0.6 * hist_score + 0.4 * struct_score
    return score


def is_valid_frame(frame: np.ndarray, model: dict, 
                   prev_gray: np.ndarray = None) -> tuple:
    """
    Retourne (valide: bool, similarite: float, variation: float)
    - Invalide si similarité trop faible (mauvais décor / plan différent)
    - Invalide si variation inter-frame trop forte (passage devant la cam)
    """
    sim = frame_similarity(frame, model)

    gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
    variation = 0.0
    if prev_gray is not None:
        variation = float(np.mean(cv2.absdiff(prev_gray, gray)))

    # Passage devant la cam = variation brutale > 30
    if variation > 30.0:
        return False, sim, variation

    if sim < SIMILARITY_THRESHOLD:
        return False, sim, variation

    return True, sim, variation


# ============================================================
#  ÉTAPE 3 — EXTRACTION LUMIÈRE SUR TOUTES LES VIDÉOS
# ============================================================

def extract_all_videos(model: dict) -> list:
    """
    Trouve tous les .mp4 du dossier, parse leur date depuis le nom,
    et extrait la luminosité ROI uniquement sur les frames valides.
    """
    mp4_files = sorted(glob.glob("*.mp4"))
    if not mp4_files:
        print("[ERR] Aucun fichier .mp4 trouvé dans le dossier courant")
        sys.exit(1)

    print(f"\n[INFO] {len(mp4_files)} vidéo(s) trouvée(s) : {mp4_files}")

    all_data = []
    ry, rx, rh, rw = ROI

    for video_file in mp4_files:
        try:
            real_start = parse_date_from_filename(video_file)
        except ValueError as e:
            print(f"  [SKIP] {e}")
            continue

        cap = cv2.VideoCapture(video_file)
        if not cap.isOpened():
            print(f"  [SKIP] Impossible d'ouvrir {video_file}")
            continue

        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)

        print(f"\n[VIDEO] {video_file} ({duration_sec//3600}h"
              f"{(duration_sec%3600)//60}m)")

        prev_gray   = None
        valid_count = 0
        total_count = 0
        curr = 0

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

            total_count += 1
            valid, sim, variation = is_valid_frame(frame, model, prev_gray)

            gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
            prev_gray = gray.copy()

            if valid:
                roi = frame[ry:ry+rh, rx:rx+rw]
                roi_gray = cv2.cvtColor(roi, cv2.COLOR_BGR2GRAY)
                # Flou gaussien pour réduire le bruit capteur
                roi_gray = cv2.GaussianBlur(roi_gray, (5, 5), 0)

                mean_val = float(np.mean(roi_gray))
                std_val  = float(np.std(roi_gray))

                real_dt = real_start + timedelta(seconds=curr)

                all_data.append({
                    "dt":        real_dt.isoformat(),
                    "mean":      mean_val,
                    "std":       std_val,
                    "variation": variation,
                    "sim":       sim,
                    "video":     video_file,
                })
                valid_count += 1

                if valid_count % 20 == 0:
                    print(f"  {real_dt.strftime('%Y-%m-%d %H:%M:%S')} | "
                          f"mean={mean_val:6.1f} sim={sim:.2f}", end="\r")

            curr += STEP

        cap.release()
        print(f"\n  [OK] {valid_count}/{total_count} frames valides")

    print(f"\n[OK] Total : {len(all_data)} samples extraits")
    return all_data


# ============================================================
#  ÉTAPE 4 — NORMALISATION
# ============================================================

def normalize_data(data: list) -> list:
    vals = np.array([d["mean"] for d in data], dtype=float)
    vmin, vmax = vals.min(), vals.max()
    denom = vmax - vmin if vmax != vmin else 1.0
    for i, d in enumerate(data):
        d["norm"] = float((vals[i] - vmin) / denom)
    return data


# ============================================================
#  ÉTAPE 5 — POSITION SOLAIRE
# ============================================================

def get_sun(dt_iso: str) -> tuple:
    if not PYSOLAR_OK:
        return 45.0, 180.0
    dt = datetime.fromisoformat(dt_iso)
    if dt.tzinfo is None:
        dt = dt.replace(tzinfo=timezone.utc)
    alt = get_altitude(LAT, LON, dt)
    az  = get_azimuth(LAT, LON, dt)
    return float(alt), float(az)


# ============================================================
#  ÉTAPE 6 — CALCUL DES DROITES D'AZIMUT
# ============================================================

def compute_azimuth_lines(data: list, sat_img: np.ndarray) -> list:
    """
    Pour chaque pic de luminosité (entrée de lumière directe) :
    - On calcule l'azimut du soleil à ce moment
    - On trace une droite perpendiculaire à cet azimut sur l'image satellite
    - Ces droites correspondent aux orientations de façades possibles
    
    L'intersection de plusieurs droites (jours différents) = position fenêtre.
    """
    h_img, w_img = sat_img.shape[:2]

    # Détection des pics (moments où la lumière entre)
    # = variation forte ET luminosité croissante
    norms = np.array([d["norm"] for d in data])

    # Gradient du signal normalisé
    grad = np.gradient(norms)

    # Seuil : top 10% des gradients positifs = entrées de lumière
    threshold = np.percentile(grad[grad > 0], 90) if (grad > 0).any() else 0.1
    peak_indices = np.where(grad > threshold)[0]

    print(f"\n[AZIMUT] {len(peak_indices)} pics de lumière détectés")

    lines = []
    for idx in peak_indices:
        d = data[idx]
        alt, az = get_sun(d["dt"])

        if alt < 5:
            continue

        # Azimut perpendiculaire à la façade (façade face au soleil)
        facade_az = (az + 180) % 360

        # Conversion azimut → vecteur direction dans l'image
        # Nord = haut de l'image satellite
        az_rad = math.radians(facade_az)
        dx = math.sin(az_rad)
        dy = -math.cos(az_rad)  # Y inversé en image

        # Droite passant par le centre de l'image
        cx, cy = w_img / 2, h_img / 2

        lines.append({
            "az":     facade_az,
            "sun_az": az,
            "alt":    alt,
            "cx": cx, "cy": cy,
            "dx": dx, "dy": dy,
            "dt": d["dt"],
            "norm": d["norm"],
        })

    return lines


import math


def gps_to_pixel(lat, lon, h_img, w_img):
    lat_tl, lon_tl = SAT_GPS_TL
    lat_br, lon_br = SAT_GPS_BR
    px = int((lon - lon_tl) / (lon_br - lon_tl) * w_img)
    py = int((lat - lat_tl) / (lat_br - lat_tl) * h_img)
    return px, py


def pixel_to_gps(px, py, h_img, w_img):
    lat_tl, lon_tl = SAT_GPS_TL
    lat_br, lon_br = SAT_GPS_BR
    lat = lat_tl + (py / h_img) * (lat_br - lat_tl)
    lon = lon_tl + (px / w_img) * (lon_br - lon_tl)
    return lat, lon


# ============================================================
#  ÉTAPE 7 — HEATMAP PAR ACCUMULATION DE DROITES
# ============================================================

def build_heatmap_lines(lines: list, sat_img: np.ndarray) -> np.ndarray:
    """
    Pour chaque droite d'azimut, on "peint" une bande sur la heatmap.
    Les zones où beaucoup de bandes se croisent = position probable.
    """
    h_img, w_img = sat_img.shape[:2]
    heatmap = np.zeros((h_img, w_img), dtype=float)

    # Largeur de la bande (pixels) — plus fin = plus précis
    BAND_WIDTH = 40

    print(f"\n[HEATMAP] Tracé de {len(lines)} droites d'azimut...")

    for line in lines:
        dx, dy = line["dx"], line["dy"]
        norm_val = line["norm"]  # pondère par la force du pic

        # Normale à la droite
        nx, ny = -dy, dx

        # Pour chaque pixel, distance à la droite
        ys, xs = np.mgrid[0:h_img, 0:w_img]
        dist = np.abs((xs - line["cx"]) * nx + (ys - line["cy"]) * ny)

        # Contribution gaussienne centrée sur la droite
        contribution = norm_val * np.exp(-(dist**2) / (2 * (BAND_WIDTH/2)**2))
        heatmap += contribution

    # Normalisation
    if heatmap.max() > 1e-9:
        heatmap /= heatmap.max()

    return heatmap


# ============================================================
#  ÉTAPE 8 — MASQUE BÂTIMENTS
# ============================================================

def apply_building_mask(heatmap: np.ndarray, sat_img: np.ndarray) -> np.ndarray:
    """
    Multiplie la heatmap par un masque qui favorise les zones construites
    (toits, bâtiments) détectées par Canny sur l'image satellite.
    """
    gray = cv2.cvtColor(sat_img, cv2.COLOR_RGB2GRAY)
    edges = cv2.Canny(gray, 40, 120)
    kernel = np.ones((15, 15), np.uint8)
    building_mask = cv2.dilate(edges, kernel, iterations=3)
    mask_norm = building_mask.astype(float) / 255.0

    # On ne supprime pas complètement les zones sans bâtiment
    # on atténue juste (0.2 = plancher minimum)
    combined = heatmap * (0.2 + 0.8 * mask_norm)

    if combined.max() > 1e-9:
        combined /= combined.max()

    return combined


# ============================================================
#  ÉTAPE 9 — VISUALISATION
# ============================================================

def visualize(sat_img: np.ndarray, heatmap: np.ndarray,
              lines: list, data: list):

    fig, axes = plt.subplots(1, 3, figsize=(20, 7))
    fig.patch.set_facecolor('#0d1117')

    h_img, w_img = sat_img.shape[:2]

    # --- (1) Heatmap sur satellite ---
    ax = axes[0]
    ax.imshow(sat_img)
    im = ax.imshow(heatmap, alpha=0.6, cmap='inferno', vmin=0.2, vmax=1.0)
    plt.colorbar(im, ax=ax, label='Probabilité')
    ax.set_title('Heatmap — Intersection des droites d\'azimut\n'
                 '(blanc/jaune = zone très probable)',
                 color='white', fontsize=10)

    # Trace les droites d'azimut sur l'image
    # On n'affiche que les 5 droites les plus fortes pour la lisibilité
    top_lines = sorted(lines, key=lambda l: -l["norm"])[:5]
    colors_lines = plt.cm.cool(np.linspace(0, 1, len(top_lines)))

    for line, col in zip(top_lines, colors_lines):
        dx, dy = line["dx"], line["dy"]
        scale = max(w_img, h_img)
        x0 = line["cx"] - dx * scale
        y0 = line["cy"] - dy * scale
        x1 = line["cx"] + dx * scale
        y1 = line["cy"] + dy * scale
        ax.plot([x0, x1], [y0, y1], color=col, alpha=0.5, lw=1,
                label=f"az={line['az']:.0f}°")

    ax.legend(fontsize=7, loc='lower right')

    # Top-3 positions
    flat = np.argsort(heatmap.ravel())[::-1][:3]
    marker_colors = ['cyan', 'lime', 'yellow']
    for rank, idx in enumerate(flat):
        py, px = divmod(idx, w_img)
        lat, lon = pixel_to_gps(px, py, h_img, w_img)
        ax.plot(px, py, marker='*', markersize=15 - rank*3,
                color=marker_colors[rank],
                label=f"Top-{rank+1} ({lat:.5f},{lon:.5f})")
    ax.legend(fontsize=7, loc='upper right')
    ax.axis('off')

    # --- (2) Signal lumière par jour ---
    ax2 = axes[1]
    ax2.set_facecolor('#0d1117')

    # Regroupe par date
    by_date = {}
    for d in data:
        date_key = d["dt"][:10]
        by_date.setdefault(date_key, []).append(d)

    colors_days = plt.cm.tab10(np.linspace(0, 1, len(by_date)))
    for (date_key, day_data), col in zip(by_date.items(), colors_days):
        norms = [d["norm"] for d in day_data]
        ax2.plot(range(len(norms)), norms, color=col, lw=1.2,
                 label=date_key, alpha=0.8)

    ax2.set_title('Signal luminosité par jour\n(frames valides uniquement)',
                  color='white')
    ax2.set_xlabel('Sample', color='grey')
    ax2.set_ylabel('Luminosité normalisée', color='grey')
    ax2.legend(fontsize=8)
    ax2.tick_params(colors='grey')
    for sp in ax2.spines.values():
        sp.set_edgecolor('#333')

    # --- (3) Rose des azimuts détectés ---
    ax3 = axes[2]
    ax3.remove()
    ax3 = fig.add_subplot(1, 3, 3, projection='polar')
    ax3.set_facecolor('#0d1117')

    if lines:
        azs    = [math.radians(l["az"]) for l in lines]
        norms_ = [l["norm"] for l in lines]
        ax3.scatter(azs, norms_, c=norms_, cmap='inferno', alpha=0.7, s=20)
        ax3.set_title('Orientations de façade détectées\n(force des pics)',
                      color='white', pad=15)
    ax3.tick_params(colors='grey')
    ax3.set_theta_zero_location('N')
    ax3.set_theta_direction(-1)

    fig.suptitle('CTF Solver v2 — Verdun-sur-le-Doubs — Droites d\'azimut solaire',
                 color='white', fontsize=13, fontweight='bold')
    plt.tight_layout()
    plt.savefig(OUTPUT_HEATMAP, dpi=150, bbox_inches='tight',
                facecolor='#0d1117')
    print(f"\n[OK] Heatmap sauvegardée → {OUTPUT_HEATMAP}")
    plt.show()


# ============================================================
#  RAPPORT CONSOLE
# ============================================================

def print_report(heatmap: np.ndarray, lines: list, data: list):
    h_img, w_img = heatmap.shape

    print("\n" + "="*60)
    print("  RAPPORT CTF SOLVER v2 — VERDUN-SUR-LE-DOUBS")
    print("="*60)

    by_date = {}
    for d in data:
        by_date.setdefault(d["dt"][:10], []).append(d)
    print(f"\n  Jours analysés       : {len(by_date)} ({', '.join(by_date.keys())})")
    print(f"  Samples valides      : {len(data)}")
    print(f"  Droites d'azimut     : {len(lines)}")

    if lines:
        azs = [l["az"] for l in lines]
        print(f"  Orientation façade   : {np.mean(azs):.1f}° ± {np.std(azs):.1f}°")

    flat = np.argsort(heatmap.ravel())[::-1][:5]
    print("\n  Top-5 positions candidates :")
    for rank, idx in enumerate(flat):
        py, px = divmod(idx, w_img)
        lat, lon = pixel_to_gps(px, py, h_img, w_img)
        sc = heatmap[py, px]
        print(f"    [{rank+1}] score={sc:.4f}  GPS=({lat:.6f}, {lon:.6f})")
        print(f"         → https://maps.google.com/?q={lat},{lon}")

    print("\n" + "="*60)


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

if __name__ == "__main__":
    print("=" * 60)
    print("  CTF SOLVER v2 — LOCALISATION PAR DROITES D'AZIMUT SOLAIRE")
    print(f"  Village : Verdun-sur-le-Doubs ({LAT}, {LON})")
    print("=" * 60)

    # 1. Trouver la première vidéo pour construire le modèle de référence
    mp4_files = sorted(glob.glob("*.mp4"))
    if not mp4_files:
        print("[ERR] Aucun .mp4 trouvé")
        sys.exit(1)

    ref_video = mp4_files[0]
    print(f"\n[INFO] Vidéo de référence : {ref_video}")

    # 2. Construire le modèle de décor
    model = build_reference_model(ref_video)

    # 3. Extraction données (avec cache)
    if os.path.exists(OUTPUT_DATA):
        print(f"\n[INFO] Chargement cache : {OUTPUT_DATA}")
        with open(OUTPUT_DATA) as f:
            data = json.load(f)
    else:
        data = extract_all_videos(model)
        with open(OUTPUT_DATA, "w") as f:
            json.dump(data, f, indent=2)

    if not data:
        print("[ERR] Aucune donnée extraite")
        sys.exit(1)

    data = normalize_data(data)
    print(f"\n[OK] {len(data)} samples chargés")

    # 4. Chargement image satellite
    if not os.path.exists(SATELLITE_IMAGE):
        print(f"[ERR] Image satellite introuvable : {SATELLITE_IMAGE}")
        print("      Lance d'abord : python download_satellite.py")
        sys.exit(1)

    sat_img = cv2.cvtColor(cv2.imread(SATELLITE_IMAGE), cv2.COLOR_BGR2RGB)
    print(f"[OK] Image satellite : {sat_img.shape[1]}x{sat_img.shape[0]} px")

    # 5. Calcul des droites d'azimut
    lines = compute_azimuth_lines(data, sat_img)

    if not lines:
        print("[WARN] Aucune droite détectée — baisse SIMILARITY_THRESHOLD")
    else:
        # 6. Heatmap par accumulation de droites
        heatmap = build_heatmap_lines(lines, sat_img)

        # 7. Masque bâtiments
        heatmap = apply_building_mask(heatmap, sat_img)

        # 8. Rapport
        print_report(heatmap, lines, data)

        # 9. Visualisation
        visualize(sat_img, heatmap, lines, data)
