#!/usr/bin/env python
# -*- coding: utf-8 -*-

import matplotlib.pyplot as plt
import numpy as np
import sys

# # Paramètres physiques et géométriques de base
# L0 = 0.1  # Longueur de la face du cube en mètres
# K = 1.2e-4  # Diffusivité thermique en m²/s
#
# # Calcul du temps caractéristique de diffusion thermique
# tcar = (L0**2) / K
# print(f"L'anomalie de température atteint la face opposée du cube en cuivre en t={tcar:.2f}s")

# Coefficients de diffusivité thermique pour différentes conditions
Kc = 8.33e-7  # Diffusivité thermique dans la croûte (m²/s)
def kappa(z, T):
    """
    Calcul de la diffusivité thermique en fonction de la profondeur et de la température.
    """
    if z > 11e3:  # Si la profondeur dépasse 11 km
        K = Kc
    else:
        if T < 725:  # Pour les températures inférieures à 725°C
            K = 9e-7
        else:  # Pour les températures supérieures ou égales à 725°C
            K = 4.5e-7
    A = K * dt / (dz**2)  # Coefficient du schéma explicite en fonction du pas de temps et du pas spatial
    return A

# Discrétisation spatiale et initialisation des conditions
zmax = 80e3  # Profondeur maximale en mètres
T0 = 825     # Température initiale dans la région magmatique en °C
Tzmax = 100  # Température de de la croute en °C

dz = 500  # Pas en profondeur (m)
Nz = int(zmax / dz)  # Nombre de points en profondeur
z = np.arange(0, zmax+dz, dz)  # Discrétisation spatiale
indice_limite = int(11e3 / dz)   # Indice limite pour la transition magma/croûte

# Calcul du pas de temps maximum pour stabilité numérique
y2s = 60 * 60 * 24 * 365  # Conversion d'années en secondes
dtmax = (dz**2) / (2 * Kc)
print(f'Valeur maximale du pas de temps permise : {dtmax:.2f} sec = {dtmax/y2s:.2f} ans')
tcar = zmax**2 / Kc
print(f'Temps de diffusion sur toute la profondeur: {tcar:.2e} sec = {tcar/(y2s*1e6):.2f} Ma')

# Discrétisation temporelle
dt = 4000 * y2s  # Pas de temps (4000 ans)
Nt = 80000  # Nombre de pas de temps
tmax = Nt * dt
t = np.arange(0, tmax + dt, dt)  # Discrétisation temporelle
print(f'Durée totale du calcul : {tmax:.2e} sec = {tmax/(y2s*1e6):.2f} Ma')

# Vérification de la condition de stabilité
if dt > dtmax:
    print("Pas de temps trop grand. Stop")
    sys.exit()

# Initialisation de la grille de température
T = np.ones((Nt + 1, Nz + 1)) * Tzmax
T[0, :indice_limite] = T0  # Conditions initiales : température constante T0 jusqu'à l'interface magma/croûte

# Schéma explicite pour la diffusion thermique
for i in range(1, Nt + 1):
    for j in range(1, Nz):
        A = kappa(j * dz, T[i - 1, j])
        T[i, j] = A * T[i - 1, j + 1] + (1 - 2 * A) * T[i - 1, j] + A * T[i - 1, j - 1]
    T[i, 0] = T[i, 1]  # solution 2: flux nul en surface
    # T[i, 0] = T0 # solution 1: T constante

# Tracé de la distribution de température en fonction de la profondeur pour différents instants
fig1, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 6))

# Sélection de 6 instants de temps pour affichage
tp = np.linspace(0, Nt, 6, dtype=int)
for i in tp:
    ax1.plot(T[i, :], z / 1e3, label=f't = {i * dt / (y2s * 1e6):.2f} Ma')
ax1.set_ylabel('Profondeur z (km)')
ax1.set_xlabel('Température T (°C)')
ax1.invert_yaxis()
ax1.legend(loc='best')

# Tracé de la température en fonction du temps à différentes profondeurs
profondeurs_a_etudier = [1, 2, 5, 20, 40, 50, 70]  # Profondeurs en km
for profondeur in profondeurs_a_etudier:
    tt = t / (y2s * 1e6)  # Conversion du temps en Ma
    indice = int(profondeur * 1e3 / dz) + 1  # Conversion de la profondeur en indice
    ax2.plot(tt, T[:, indice], label=f'z = {profondeur} km')
ax2.set_xlabel('Temps (Ma)')
ax2.set_ylabel('Température T (°C)')
ax2.legend(loc='best')

plt.tight_layout()
plt.show()
