from __future__ import print_function, absolute_import
from scipy.interpolate import interp1d
from scipy.interpolate import CloughTocher2DInterpolator
import underworld as uw
from underworld.scaling import non_dimensionalise as nd

class FreeSurfaceProcessor_ALEIB(object):
    """FreeSurfaceProcessor"""

    def __init__(self, model):
        """Create a Freesurface processor

        Parameters
        ----------

        model : UWGeodynamics Model

        """
        self.model = model

        minCoord = tuple([nd(val) for val in self.model.minCoord])
        maxCoord = tuple([nd(val) for val in self.model.maxCoord])

        # Initialize model mesh
        self._init_mesh = uw.mesh.FeMesh_Cartesian(elementType=self.model.elementType,
                                                   elementRes=self.model.elementRes,
                                                   minCoord=minCoord,
                                                   maxCoord=maxCoord,
                                                   periodic=self.model.periodic)
                                     
        # Create the tools
        self.TField = self._init_mesh.add_variable(nodeDofCount=1)
        self.TField.data[:, 0] = self._init_mesh.data[:, -1].copy()

        self.top = self.model.top_wall
        self.bottom = self.model.bottom_wall
        self.internal = self.model.inter_wall

        # Create boundary condition
        self._conditions = uw.conditions.DirichletCondition(
            variable=self.TField,
            indexSetsPerDof=(self.top + self.bottom + self.internal,))

        # Create Eq System
        self._system = uw.systems.SteadyStateHeat(
            temperatureField=self.TField,
            fn_diffusivity=1.0,
            conditions=self._conditions)

        self._solver = uw.systems.Solver(self._system)

    def _solve_sle(self):
        self._solver.solve()

    def _advect_surface(self, dt):

        if self.internal:
            if self.model.mesh.dim == 2:
                # Extract internalsurface
                x = self.model.mesh.data[self.internal.data][:, 0]
                y = self.model.mesh.data[self.internal.data][:, 1]
    
                # Extract velocities from top
                vx = self.model.velocityField.data[self.internal.data][:, 0]
                vy = self.model.velocityField.data[self.internal.data][:, 1]
    
                # Advect internal surface
                x2 = x + vx * nd(dt)
                y2 = y + vy * nd(dt)
    
                # Spline internal surface
                f = interp1d(x2, y2, kind='cubic', fill_value='extrapolate')
    
                self.TField.data[self.internal.data, 0] = f(x)
            else:
                # Extract internal surface
                x = self.model.mesh.data[self.internal.data][:, 0]
                y = self.model.mesh.data[self.internal.data][:, 1]
                z = self.model.mesh.data[self.internal.data][:, -1]
    
                # Extract velocities from top
                vx = self.model.velocityField.data[self.internal.data][:, 0]
                vy = self.model.velocityField.data[self.internal.data][:, 1]
                vz = self.model.velocityField.data[self.internal.data][:, -1]
    
                # Advect top surface
                x2 = x + vx * nd(dt)
                y2 = y + vy * nd(dt)
                z2 = z + vz * nd(dt)
    
                # Spline top surface
                f = CloughTocher2DInterpolator((x2, y2), z2)
                self.TField.data[self.internal.data, 0] = f((x,y))
        uw.mpi.barrier()
        self.TField.syncronise()

    def _update_mesh(self):

        with self.model.mesh.deform_mesh():
            # Last dimension is the vertical dimension
            self.model.mesh.data[:, -1] = self.TField.data[:, 0].copy()

    def solve(self, dtime):
        """ Advect free surface through dt and update the mesh """

        # First we advect the surface
        self._advect_surface(dtime)
        # Then we solve the system of linear equation
        self._solve_sle()
        # Finally we update the mesh
        self._update_mesh()
