import sys
sys.path.append('./sfepy')
import numpy as np
from sfepy.discrete.fem import FEDomain
from sfepy.mesh.mesh_generators import gen_block_mesh
from sfepy.discrete import(
    FieldVariable, Material, Integral, Function, Equation, Equations, Problem
)
from sfepy.discrete.fem import FEDomain, Field, Mesh
from sfepy.discrete.conditions import Conditions, EssentialBC
from sfepy.base.base import IndexedStruct, Struct
from sfepy.solvers.ls import ScipyDirect, ScipyUmfpack, PETScKrylovSolver
from sfepy.solvers.nls import Newton
from sfepy.solvers.ts_solvers import SimpleTimeSteppingSolver
from sfepy.terms import Term
from sfepy.base.base import dict_to_struct
from sfepy.discrete.evaluate import Evaluator

def gen_mesh(length, height, size, write = True):
    from sfepy.mesh.mesh_generators import gen_block_mesh
    dim = [length, height, height]
    size_mesh = int(height / size + 1)
    shape = [int(length / size + 1), size_mesh, size_mesh]
    center = [length / 2, (height / 2), (height / 2)]
    mesh = gen_block_mesh(dim, shape, center)
    if write:
        mesh.write(
            "./week1_verification/mesh/mesh.vtk", file_format = "vtk-ascii"
        )
    return mesh

def get_displacement(ts, coors, bc=None, problem=None):
    """
    Define the time-dependent displacement.
    """
    out = -.2 * ts.time * coors[:, 0]
    return out


def deform(mesh, curve):
    [coors, vertex_groups, conns, mat_ids, descs] = mesh._get_io_data()
    # coors[:,0] = 1

    return Mesh.from_data('mesh', coors, vertex_groups, conns, mat_ids, descs)







def main():

    order = 2

    ### Mesh and regions ###
    mesh = gen_mesh(30, 1, 0.5, write = False)
    mesh = deform(mesh, 1)
    domain = FEDomain('domain', mesh)
    omega = domain.create_region('Omega', 'all')
    left = domain.create_region(
        'left', 'vertices in x < 0.01', 'facet'
    )
    right = domain.create_region(
        'right', 'vertices in x > 29.99', 'facet'
    )
    bottom = domain.create_region(
        'bottom', 'vertices in y < 0.01', 'facet'
    )
    front = domain.create_region(
        'front', 'vertices in z < 0.01', 'facet'
    )


    ### Fields ###

    scalar_field = Field.from_args(
        'fu', np.float64, 'scalar', omega, approx_order=order-1)
    vector_field = Field.from_args(
        'fv', np.float64, 'vector', omega, approx_order=order)

    u = FieldVariable('u', 'unknown', vector_field, history=1)
    v = FieldVariable('v', 'test', vector_field, primary_var_name='u')
    p = FieldVariable('p', 'unknown', scalar_field, history=1)
    q = FieldVariable('q', 'test', scalar_field, primary_var_name='p')

    ts_vals = ['0.0', '1.0e-1', '161']
    ts = {
        't0' : float(ts_vals[0]), 't1' : float(ts_vals[1]),
        'n_step' : int(ts_vals[2])}

    ### Material ###
    material_parameters = [1.0, 0.5]
    c10, c01 = material_parameters
    m = Material(
        'm', iK=1.0 / 1e3,                                                      # bulk modulus
        mu=20e0,                                                                # shear modulus of neoHookean term
        kappa=10e0                                                              # shear modulus of Mooney-Rivlin term
    )

    ### Boundary conditions ###
    x_sym = EssentialBC('fix', left, {'u.all' : 0.0})
    disp_fun = Function('disp_fun', get_displacement)
    displacement = EssentialBC(
        'displacement', right, {'u.0' : disp_fun, 'u.[1,2]' : 0.0})
    ebcs = Conditions([x_sym, displacement])

    ### Terms and equations ###
    integral = Integral('i', order=2*order)

    term_neohook = Term.new(
        'dw_ul_he_neohook(m.mu, v, u)',
        integral, omega, m=m, v=v, u=u)
    term_mooney = Term.new(
        'dw_ul_he_mooney_rivlin(m.kappa, v, u)',
        integral, omega, m=m, v=v, u=u)
    term_pressure = Term.new(
        'dw_ul_bulk_pressure(v, u, p)',
        integral, omega, v=v, u=u, p=p)

    term_volume_change = Term.new(
        'dw_ul_volume(q, u)',
        integral, omega, q=q, u=u, term_mode='volume')
    term_volume = Term.new(
        'dw_ul_compressible(m.iK, q, p, u)',
        integral, omega, m=m, q=q, p=p, u=u)

    eq_balance = Equation('balance', term_neohook+term_mooney+term_pressure)
    eq_volume = Equation('volume', term_volume_change-term_volume)
    equations = Equations([eq_balance, eq_volume])

    ### Solvers ###
    ls = ScipyUmfpack({})
    nls_status = IndexedStruct()
    nls = Newton(
        {'i_max' : 10},
        lin_solver=ls, status=nls_status
    )

    ### Problem ###
    pb = Problem('hyper', equations=equations)
    dict = {'ulf' : True, 'mesh_update_variables' : 'u'}
    pb.conf.options = dict_to_struct(dict)
    pb.set_bcs(ebcs=ebcs)
    pb.set_ics(ics=Conditions([]))
    tss = SimpleTimeSteppingSolver(ts, nls=nls, context=pb)
    pb.set_solver(tss)
    ev = pb.get_evaluator()
    pb.nls_iter_hook = ev.new_ulf_iteration
    # nls.iter_hook = pb.nls_iter_hook
    vec = pb.get_initial_state().vec
    nls.iter_hook = Evaluator.new_ulf_iteration(pb, nls, vec, None, None, None)
    pb.solve(save_results=True)







if __name__ == "__main__":
    main()


# python3 ./sfepy/postproc.py *.vtk --wireframe -b -d'u,plot_displacements,rel_scaling=1' --step=-1