using Pkg

# -----------------------------------------------------------------------------------------------
# Installation des packages : mettre à false une fois les packages installés une première fois.
# Pkg.add() vérifie le registre à chaque exécution, ce qui coûte plusieurs secondes/minutes
# à chaque lancement du script pour rien si les packages sont déjà présents.
# -----------------------------------------------------------------------------------------------
#const FIRST_RUN = false
const FIRST_RUN = false

if FIRST_RUN
    for pkg in ["Statistics", "FITSIO", "Images", "NativeFileDialog"]
        Pkg.add(pkg)
    end
end

using Statistics
using Printf

# fits file  : https://juliaastro.org/FITSIO.jl/stable/
using FITSIO

# package pour filtrage des images (imfilter, Kernel.gaussian)
using Images

# package pour sélection interactive du répertoire et du fichier (fenêtres de dialogue)
using NativeFileDialog

##################################################################################################################
# Fonctions du programme :
#   Portage Julia de l'algorithme MGN (Multi-scale Gaussian Normalization), inspiré de
#   sunkit_image.enhance.mgn (package affilié SunPy, licence BSD-3-Clause) :
#   https://github.com/sunpy/sunkit-image
#
#   Référence :
#   Morgan, Huw, and Miloslav Druckmuller. "Multi-scale Gaussian normalization for solar image
#   processing." Sol Phys (2014) 289: 2945. doi:10.1007/s11207-014-0523-9
#
# Principe :
#   pour chaque échelle sigma :
#     - calcul de la moyenne locale (convolution gaussienne)
#     - calcul de l'écart-type local (convolution gaussienne du carré de l'écart à la moyenne)
#     - normalisation du pixel par cette moyenne/écart-type local (contraste local)
#     - transformation arctan (équivalent à une transformation gamma "douce")
#   les images obtenues à chaque échelle sont ensuite moyennées (pondération possible), puis
#   combinées avec une transformation gamma globale de l'image d'origine (paramètre h = poids
#   relatif de la partie globale vs la partie multi-échelle locale)
#
# Entrées/paramètres :
#      directory et nom du fichier fit à traiter (HDR, ou tout autre fit 3 couches)
#      sigma  : liste des échelles gaussiennes (en pixels)
#      k      : sévérité de la transformation arctan
#      gamma  : exposant de la transformation gamma globale
#      h      : poids de la partie globale par rapport à la partie locale multi-échelle
###################################################################################################################

function mgn(data::AbstractMatrix{T};
             sigma::Vector{<:Real} = [1.25, 2.5, 5, 10, 20, 40],
             k::Real = 0.7,
             gamma::Real = 3.2,
             h::Real = 0.7,
             weights::Union{Nothing,Vector{<:Real}} = nothing,
             truncate::Real = 3) where {T<:AbstractFloat}

    # Fonction :
    #   applique la Multi-scale Gaussian Normalization sur une image 2D (un seul canal)
    #
    # Entrée :
    #   data     : image 2D (Float32 ou Float64)
    #   sigma    : liste des écarts-types des noyaux gaussiens (échelles spatiales, en pixels)
    #   k        : facteur multiplicatif avant la transformation arctan (contraste local)
    #   gamma    : exposant de la transformation gamma globale (entre 2.5 et 4 selon le papier)
    #   h        : poids de la partie globale (gamma) par rapport à la partie locale (multi-échelle)
    #   weights  : poids relatif de chaque échelle sigma (par défaut : poids égaux)
    #   truncate : nombre d'écarts-types utilisés pour tronquer le noyau gaussien
    #
    # Sortie :
    #   image transformée (mêmes dimensions que data), valeurs typiquement dans [0,1]

    if weights === nothing
        weights = ones(length(sigma))
    end
    @assert length(weights) == length(sigma) "weights et sigma doivent avoir la même longueur"

    # copie de travail : on ne modifie pas le tableau d'entrée
    data = copy(data)

    # 1. remplacement des pixels négatifs ou nuls par une valeur infime positive
    #    (nécessaire car la transformation gamma globale utilise data.^(1/gamma))
    data[data .<= 0] .= T(1e-15)

    image = zeros(T, size(data))

    for (s, w) in zip(sigma, weights)

        # noyau gaussien : par défaut, ImageFiltering utilise une taille de noyau de
        # 2*ceil(3*sigma)+1, ce qui correspond exactement au troncage à 3 sigma (truncate=3)
        # utilisé par défaut dans l'algorithme original - donc pas de calcul de taille à faire ici
        kernel = Kernel.gaussian(s)

        # 2 & 3. moyenne locale par convolution gaussienne (équation 1 du papier)
        local_mean = imfilter(data, kernel)

        # 4. écart à la moyenne locale, puis écart-type local (équation 2 du papier)
        diff = data .- local_mean
        local_var = imfilter(diff.^2, kernel)
        local_std = sqrt.(local_var)
        local_std[local_std .== 0] .= T(1.0)   # évite une division par 0

        # 5. normalisation par la moyenne/écart-type local => Ci
        Ci = diff ./ local_std

        # 6. transformation arctan => C'i (équation 3 du papier), pondérée par le poids de l'échelle
        Cprime_i = w .* atan.(k .* Ci)

        image .+= Cprime_i
    end

    # 8. moyenne (pondérée) des images normalisées sur toutes les échelles
    image ./= length(sigma)

    # 9. transformation gamma globale de l'image d'origine (équation 4 du papier)
    data_min = minimum(data)
    data_max = maximum(data)
    Cprime_g = data .- data_min
    if (data_max - data_min) != 0
        Cprime_g = Cprime_g ./ T(data_max - data_min)
    end
    Cprime_g = Cprime_g .^ T(1/gamma)
    Cprime_g = T(h) .* Cprime_g

    # 10. combinaison de la partie locale multi-échelle et de la partie globale
    image = T(1-h) .* image .+ Cprime_g

    return image
end

###############################################################################################
#
## début du programme principal
#
###############################################################################################

function main()

# chronométrage du temps total de calcul
t_start = time()

println("Nombre de threads Julia disponibles : ", Threads.nthreads())
if Threads.nthreads() == 1
    println("ATTENTION : un seul thread disponible - le calcul des 3 canaux R,G,B restera")
    println("            séquentiel. Relancez avec 'julia -t 3 ...' (ou -t auto) pour paralléliser.")
end

# sélection interactive du répertoire de travail
work_dir = pick_folder(pwd())
isempty(work_dir) && error("Aucun répertoire sélectionné - programme arrêté")
cd(work_dir)

println(pwd())

# sélection interactive du fichier fit à traiter (par ex. le résultat HDR/_color d'un script précédent)
selected_file = pick_file(work_dir; filterlist="fit,fits")
isempty(selected_file) && error("Aucun fichier sélectionné - programme arrêté")
file_name_in = basename(selected_file)

# lecture de l'image d'entrée (fit, 3 couches RGB)
f_in = FITS(file_name_in)
header = read_header(f_in[1])
image_in = read(f_in[1])
close(f_in)

dim_x = header["NAXIS1"]
dim_y = header["NAXIS2"]
println("dim_x = ", dim_x)
println("dim_y = ", dim_y)

# conversion en Float32 (au cas où l'image d'entrée serait dans un autre type)
image_in = Float32.(image_in)

# -----------------------------------------------------------------------------------------------
# paramètres de la MGN : à ajuster selon le rendu souhaité (voir docstring de mgn() ci-dessus)
# valeurs par défaut proches de celles utilisées pour les images AIA/EUV dans sunkit-image
# -----------------------------------------------------------------------------------------------
sigma_mgn = [7, 7, 7, 10, 20, 40]
#k_mgn     = 0.7
#gamma_mgn = 3.2
#h_mgn     = 0.7
#k_mgn     = 0.7
k_mgn     = 0.35
gamma_mgn = 4.2
h_mgn     = 0.92


image_mgn = similar(image_in)

# -----------------------------------------------------------------------------------------------
# traitement indépendant de chaque canal couleur (R,G,B) : les 3 canaux ne communiquent jamais
# entre eux dans mgn(), donc on peut les calculer en parallèle sur des threads Julia distincts.
# Chaque itération écrit dans une tranche disjointe de image_mgn (pas de risque de conflit mémoire).
#
# IMPORTANT : pour que ceci accélère réellement le calcul, Julia doit être démarré avec plusieurs
# threads, par ex. :
#     julia -t 3 TSE2026-4-eclipse-MGN-v2.jl
# ou en définissant la variable d'environnement avant de lancer Julia :
#     export JULIA_NUM_THREADS=3      (Linux/macOS)
#     set JULIA_NUM_THREADS=3         (Windows)
# "auto" fonctionne aussi (JULIA_NUM_THREADS=auto) et laisse Julia choisir selon les coeurs dispo.
# Avec 1 seul thread disponible, Threads.@threads ne fait aucun mal mais n'accélère rien non plus.
# -----------------------------------------------------------------------------------------------
Threads.@threads for i_ch in [1,2,3]
    println("calcul MGN, canal ", i_ch, " (thread ", Threads.threadid(), ")")
    image_mgn[:,:,i_ch] = mgn(image_in[:,:,i_ch]; sigma=sigma_mgn, k=k_mgn, gamma=gamma_mgn, h=h_mgn)
end

# nom du fichier de sortie : ajout de "_MGN" juste avant l'extension ".fit"
# chaîne "k09_g32_h09_" : valeur de chaque paramètre * 10, arrondie, sur 2 chiffres (point supprimé)
# chaîne "s5_k09_g32_h09_" : sigma(1) tel quel (ex: 5 -> "s5_"), puis k,gamma,h *10 arrondis sur 2 chiffres
params_str = @sprintf("s%02d_k%03d_g%02d_h%03d_", round(Int, sigma_mgn[1]), round(Int, k_mgn*100), round(Int, gamma_mgn*10), round(Int, h_mgn*100))
file_name_out = chop(file_name_in, tail=4)*"_"*params_str*"MGN.fit"

f_out = FITS(file_name_out, "w")
write(f_out, image_mgn[:,:,:]; header=header)
close(f_out)

println("fichier écrit : ", file_name_out)

# affichage du temps de calcul total
elapsed = time() - t_start
println(@sprintf("Temps de calcul total : %.1f s (%.1f min)", elapsed, elapsed/60))

end # fin de la fonction main()

main()
