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

from ngsolve import *
from netgen.occ import *
import os
import matplotlib.pyplot as plt


### Create new result folder if it does not exist
foldername = "res"
myfolderpath = "./"+foldername
os.makedirs(myfolderpath, exist_ok = True)


### GENERAL FUNCTIONS
### colors
colordict = {"black": (1, 1, 1), "blue": (0.2, 0.4, 1), \
           "green": (0, 0.5, 0), "red": (1, 0, 0), \
            "orange" : (1, 0.6, 0), "cyan": (0, 1, 1)}

### Transient solver
def ImplicitEuler(invmdta, t0, tend, nbs):
    
    data_list, time_list = list(), list()
    sample = int(floor(tend / dt / nbs)+1)

    gfut = GridFunction(gfu.space, multidim = 0) # copy the solution on a multidimensional object to store the time
    gfut.AddMultiDimComponent(gfu.vec) # adding the time dependent solution

    cnt = 0; tt = t0
    data_list.append(gfu(mesh(0.,0., lcb/2)))
    time_list.append(tt)
    while tt <= tend:
        res = dt * L.vec - dt * A.mat * gfu.vec
        gfu.vec.data += invmdta * res
        print("iter: {0:4d}, time: {1:.3g}\r".format(cnt,tt))
        if (cnt % sample) == 0:
            gfut.AddMultiDimComponent(gfu.vec)
        cnt += 1; tt = cnt * dt
        vtk.Do(time = tt)
        data_list.append(gfu(mesh(0.,0., lcb/2)))
        time_list.append(tt)
        
    return gfut, data_list, time_list

### PARAMETERS
### geometry
Rcb = 0.4125e-3 # radius of the wire
lcb = Rcb/10 # length of wire

### Thermophysical
T_inf = 77 # Temperature of LN2
T_0 = 300 # initial temperature
k = 100.0 # thermal conductivity
rho_m = 8966.0 # mass density
cp = 0.1 # specific heat capacity
Cc = rho_m*cp # heat capacity
dif = k / Cc # diffusivity
tau = Rcb**2 / dif # time characterisitics of diffusivity

### mesh
lc_cb = Rcb/10 # Mesh density

### time
dt = 1e-7
nbstp = 100
t_ini = 0
t_final = t_ini+nbstp*dt

### sources and sinks
source = 0.0 / (pi*Rcb**2*lcb)  # dissipation in W/m^3
h = 131.0 # heat exchange coefficient
biot = h*(Rcb/2) / k # Biot number


### GEOMETRY
### Coordinates
x0 = 0.; y0 = 0.; z0 = 0.

### Points or vertices
p0 = Pnt(x0, y0, z0)

### Geometry
cable = Cylinder(p0, Z, r = Rcb, h = lcb)
cable.solids.name = "cable"
cable.solids.col = colordict["orange"]
cable.faces[0].name = "cableside"
cable.faces.Max(Z).name = "cableoutlet"
cable.faces.Min(Z).name = "cableinlet"
cable.faces.col = colordict["blue"]
cable.faces.maxh = lc_cb

geo = OCCGeometry(cable)

# ~ geo.shape.WriteStep("cable.step")

### MESHING
ngmesh = geo.GenerateMesh(maxh = lc_cb)
# ~ ngmesh.Export("cable.msh", "Gmsh Format")
mesh = Mesh(ngmesh).Curve(3)


### FUNCTION SPACE
fes = H1(mesh, order = 1, dirichlet = "cableside")
u = fes.TrialFunction()
v = fes.TestFunction()

gfu = GridFunction(fes)
#gfu.Set(T_inf, BND) # Dirichlet condition on boundary
gfu.Set(T_0) # initial condition on domain

q = CoefficientFunction(source)


### M and A are same sparsity pattern to be sum up as flattened vectors: 
### "symmetric = False" to get the same dimensions
with TaskManager():
    ### FORMS
    M = BilinearForm(fes, symmetric = False) # Mass matrix
    A = BilinearForm(fes, symmetric = False) # Rigidity matrix
    L = LinearForm(fes) # Forces or loads and some boundary conditions
    
    ### WEAK FORMULATION
    M += (Cc * u) * v * dx # Mass
    
    A += (k * grad(u)) * grad(v) * dx # Conduction
    A += (h * u) * v * ds(definedon = mesh.Boundaries("cableside")) # convection
    
    L += q * v * dx # load
    L += (h * T_inf) * v * ds(definedon = mesh.Boundaries("cableside")) # convection
    
    ### ASSEMBLY
    M.Assemble()
    A.Assemble()
    L.Assemble()

MdtA = M.mat.CreateMatrix() # New temporal matrix: M+dt*A
MdtA.AsVector().data = M.mat.AsVector() + dt*A.mat.AsVector()
InvMdtA = MdtA.Inverse(freedofs = fes.FreeDofs())


### VISUALIZATION
vtk = VTKOutput(mesh, coefs = [gfu], names = ["Temperature"], \
    filename = foldername+"/solution", subdivision = 2)
vtk.Do(time = t_ini)


### SOLVING
gfut, data_list, time_list = ImplicitEuler(InvMdtA, t_ini, t_final, nbstp)


### PRINTINGS
fig, axs = plt.subplots(1, figsize=(6, 4))
fig.tight_layout()
axs.grid()
axs.set_xlabel(r'$t$ (s)', fontsize = 18)
axs.set_ylabel(r'$T$ (K)', fontsize = 18)
axs.plot(time_list, data_list, color = 'red', linestyle='-', linewidth = 2)

# ~ print("\nSave figure")
# ~ plt.savefig(myfolderpath+'/figure.png', format = 'png', dpi = 75, bbox_inches = 'tight')
plt.show()

print("\nMesh boundaries ", mesh.GetBoundaries())
print(f"Check matrix sizes: M.mat.nze = {M.mat.nze}, A.mat.nze={A.mat.nze}, Mstar.nze={MdtA.nze}")
print("Time constant {0:.1g} s".format(tau))
print("Biot number {0:.1g}\n".format(biot))
