import numpy as np
from numpy import matlib
from mpl_toolkits import mplot3d 
import matplotlib.pyplot as plt
import matplotlib.tri as mtri
from mpl_toolkits import mplot3d
from mpl_toolkits.mplot3d import Axes3D
from mpl_toolkits.mplot3d import proj3d
import math
import pylab
from red_refine import red_refine
import scipy.sparse
import scipy.sparse.linalg
from scipy.sparse import csr_matrix
from scipy.sparse import hstack
from scipy.sparse import vstack

def FEM(coord,triangles,dirichlet,f,uD):
    nelems=np.size(triangles,0)
    nnodes=np.size(coord,0)

    A=stiffness_matrix(coord,triangles)
    b=RHS_vector(coord,triangles,f)
    
    x=np.zeros(3*nnodes+2*nelems+1)
    
    dbnodes=np.unique(dirichlet)
    for j in dbnodes:
        coord_loc=(coord[j,:])
        tmp=uD(coord_loc[0],coord_loc[1])
        x[j]=tmp[0]
        x[nnodes+j]=tmp[1]
  
    b=b-A.dot(x)
    
    inodes=np.setdiff1d(range(0,nnodes),dbnodes)
    dof=np.concatenate((inodes,nnodes+inodes,\
                        2*nnodes+np.array(range(0,2*nelems+nnodes+1))\
                        ),axis=0)
    A_inner=A[np.ix_(dof,dof)]
    b_inner=b[dof]   
    x[dof]=scipy.sparse.linalg.spsolve(A_inner,b_inner)
    
    u1=np.array([x[0:nnodes]]).T
    u2=np.array([x[nnodes:2*nnodes]]).T
    u=np.concatenate((u1,u2),axis=1)
    return x,u
    
def get_geom():
    coord = np.asarray([[-1,-1],[1,-1],[1,1],[-1,1]])
    triangles = np.asarray([[2,0,1],[0,2,3]])
    dirichlet= np.array([[0,1],[1,2],[2,3],[3,0]])
    neumann= np.zeros([0, 2])
    return coord, triangles, dirichlet, neumann


def stiffness_matrix(coord,triangles):
    nelems=np.size(triangles,0)
    nnodes=np.size(coord,0)
    Alocal=np.zeros((nelems,3,3))
    Ab_loc=np.zeros((nelems,1)) #bubble part
    Blocal=np.zeros((nelems,3,8))
    I1=np.zeros((nelems,3,3))
    I2=np.zeros((nelems,3,3))

    J1=np.zeros((nelems,3,8))
    J2=np.zeros((nelems,3,8))
    #compute local matrices
    for j in range(0,nelems):
        nodes_loc=triangles[j,:]
        coord_loc=coord[nodes_loc,:]
        T=np.array([coord_loc[1,:]-coord_loc[0,:] ,
               coord_loc[2,:]-coord_loc[0,:] ])
        area = 0.5 * ( T[0,0]*T[1,1] - T[0,1]*T[1,0] )
        tmp1= np.concatenate((np.array([[1,1,1]]), coord_loc.T),axis=0)
        tmp2= np.array([[0,0],[1,0],[0,1]])
        grads = np.linalg.solve(tmp1,tmp2)
        #P1 part
        Alocal[j,:,:]=area* np.matmul(grads,grads.T)
        I1[j,:,:] = np.concatenate((np.array([nodes_loc]),np.array([nodes_loc]),np.array([nodes_loc])),axis=0)
        I2[j,:,:] = np.concatenate((np.array([nodes_loc]).T,np.array([nodes_loc]).T,np.array([nodes_loc]).T),axis=1)
        #bubble part
        Ab_loc[j]=area/180*np.sum((np.matmul(grads,grads.T)).diagonal())
        #local matrix B
        tmp3 = np.array([np.concatenate((grads[:,0],grads[:,1]),axis=0).T])
        Blocal[j,:,:]=area*(1/3)*\
                     np.concatenate(\
                         (np.concatenate((tmp3,tmp3,tmp3),axis=0),\
                         (-1/20)*grads),axis=1)
        J1[j,:,:] = np.matlib.repmat(np.array([nodes_loc]).T, 1, 8)
        J2[j,:,:] = np.concatenate(\
                    (np.array([nodes_loc]),np.array([nnodes+nodes_loc]),\
                        np.array([[2*nnodes+j]]),\
                        np.array([[2*nnodes+nelems+j]])),axis=1)
                                           
    
    #assemble
    Alocal=np.reshape(Alocal,(9*nelems,1)).T
    I1=np.reshape(I1,(9*nelems,1)).T
    I2=np.reshape(I2,(9*nelems,1)).T
    A=csr_matrix((Alocal[0,:],(I1[0,:],I2[0,:])),shape = (nnodes,nnodes))
    
    #assemble
    Blocal=np.reshape(Blocal,(24*nelems,1)).T
    J1=np.reshape(J1,(24*nelems,1)).T
    J2=np.reshape(J2,(24*nelems,1)).T
    B=csr_matrix((Blocal[0,:],(J1[0,:],J2[0,:])),shape = (nnodes,2*(nnodes+nelems)))
    
    #assemble
    Ab_local = np.reshape(np.concatenate((Ab_loc,Ab_loc),axis=1).T,(2*nelems,1)).T
    K=np.reshape(range(0,2*nelems),(2*nelems,1)).T
    Ab=csr_matrix((Ab_local[0,:],(K[0,:],K[0,:])),\
        shape = (2*nelems,2*nelems))
    
    #compute nodal areas
    nodalareas=np.zeros((nnodes,1))
    for j in range(0,nelems-1):
        nodes_loc=triangles[j,:]
        coord_loc=coord[nodes_loc,:]
        T=np.array([coord_loc[1,:]-coord_loc[0,:] ,
               coord_loc[2,:]-coord_loc[0,:] ])
        area = 0.5 * ( T[0,0]*T[1,1] - T[0,1]*T[1,0] )
        nodalareas[nodes_loc] = nodalareas[nodes_loc]+area;    
    
    N=nnodes
    E=nelems
    A = vstack( (hstack((A,csr_matrix((N,N+2*E)))), \
                 hstack((csr_matrix((N,N)),A,csr_matrix((N,2*E)))), \
                 hstack((csr_matrix((2*E,2*N)),Ab))  ))
    
    
    A = vstack((hstack((A,B.T,csr_matrix((2*N+2*E,1)))),\
                hstack((B,csr_matrix((N,N)),1/3*nodalareas)),\
                hstack((csr_matrix((1,2*N+2*E)),1/3*nodalareas.T,csr_matrix((1,1))))  ))
    A = A.tocsr()
    return A


def RHS_vector(coord,triangles,f):
    nelems=np.size(triangles,0)
    nnodes=np.size(coord,0)
    b=np.zeros(3*nnodes+2*nelems+1)


    for j in range(0,nelems):
        nodes_loc=triangles[j,:]
        coord_loc=coord[nodes_loc,:]
        tmp=np.array([coord_loc[1,:]-coord_loc[0,:] ,
               coord_loc[2,:]-coord_loc[0,:] ])
        area = 0.5 * ( tmp[0,0]*tmp[1,1] - tmp[0,1]*tmp[1,0] )
        T1= np.array([[0,0],[1,0],[0,1]])
        mid=1/3*(coord_loc[0,:]+coord_loc[1,:]+coord_loc[2,:])
        f_mid = f(mid[0],mid[1])
        b[nodes_loc]=b[nodes_loc]+area/3*f_mid[0]
        b[nnodes+nodes_loc]=b[nnodes+nodes_loc]+area/3*f_mid[1]
        b[2*nnodes+j]=9*area*f_mid[0]
        b[2*nnodes+nelems+j]=9*area*f_mid[1]
    
    return b


def compute_H1_error(coord,triangles,gruex,u):
    nelems=np.size(triangles,0)
    nnodes=np.size(coord,0)
    b=np.zeros(3*nnodes+2*nelems+1)

    H1errSq = 0
    for j in range(0,nelems):
        nodes_loc=triangles[j,:]
        coord_loc=coord[nodes_loc,:]
        T=np.array([coord_loc[1,:]-coord_loc[0,:] ,
               coord_loc[2,:]-coord_loc[0,:] ])
        area = 0.5 * ( T[0,0]*T[1,1] - T[0,1]*T[1,0] )
        tmp1= np.concatenate((np.array([[1,1,1]]), coord_loc.T),axis=0)
        tmp2= np.array([[0,0],[1,0],[0,1]])
        grads = np.linalg.solve(tmp1,tmp2)
        mid=1/3*(coord_loc[0,:]+coord_loc[1,:]+coord_loc[2,:])
        gru_ex_mid = gruex(mid[0],mid[1])
        gruh = u[nodes_loc[0]]*grads[0,:]\
               +u[nodes_loc[1]]*grads[1,:]\
               +u[nodes_loc[2]]*grads[2,:]
        e=gru_ex_mid-gruh
        H1errSq=H1errSq+area*(e[0]**2+e[1]**2)
            
    return H1errSq

#######################################################################
#######################################################################
#######################################################################



fun = lambda x, y:  (x-x**2)+(y- y**2)
fun = lambda x, y:  20*x*y**4-4*x**5
u1_exact=np.vectorize(fun)
fun = lambda x, y:  20*x**4*y-4*y**5
u2_exact=np.vectorize(fun)
f = lambda x, y:  0*2* np.array([(x-x**2)+(y- y**2) ,0])
uD = lambda x, y:  np.array([20*x*y**4-4*x**5,20*x**4*y-4*y**5])

gruex1 = lambda x, y:  np.array([20*y**4-20*x**4,80*x*y**3])
gruex2 = lambda x, y:  np.array([80*x**3*y,20*x**4-20*y**4])
coord, triangles, dirichlet, neumann = get_geom()
nref=5
max_err=np.zeros(nref)
H1_err=np.zeros(nref)
for j in range(0,nref):
    coord, triangles, dirichlet,_,_,_ = \
           red_refine(coord, triangles, dirichlet, neumann)
    x,u=FEM(coord, triangles, dirichlet,f,uD)
    u1_at_nodes=u1_exact(coord[:,0],coord[:,1])
    u2_at_nodes=u2_exact(coord[:,0],coord[:,1])
    max_err[j]=np.max(np.abs(u1_at_nodes-u[:,0]))
    H1_err[j]=np.sqrt(compute_H1_error(coord,triangles,gruex1,u[:,0])\
                +compute_H1_error(coord,triangles,gruex2,u[:,1]))
print(max_err)
print(H1_err)
    
## plot the solution graph (FEM and exact) 
fig = plt.figure(figsize =(14, 9))
ax = plt.axes(projection ='3d')
trisurf = ax.plot_trisurf(coord[:,0],coord[:,1],u[:,0],
                          triangles = triangles, 
                          cmap =plt.get_cmap('summer'),
                          edgecolor='Gray');
ax.set_title('finite element solution')
plt.show(block=False)
    
uex = lambda x, y:  (x-x**2)*(y- y**2)
func2=np.vectorize(uex)
#fig = plt.figure(figsize =(14, 9))
#ax = plt.axes(projection ='3d')
#trisurf = ax.plot_trisurf(coord[:,0],coord[:,1],u1_at_nodes,
#                          triangles = triangles, 
#                          cmap =plt.get_cmap('summer'),
#                          edgecolor='Gray');
fig, ax = plt.subplots()
ax.quiver(coord[:,0],coord[:,1], u[:,0], u[:,1])
ax.set_title('exact solution')
plt.show()
