#!/usr/bin/env python
# Contact: f.smai@brgm.fr
# version: 1.3.1

""" conversion des fichiers VTK depuis les sorties TOUGH2

Les fichiers VTK sont lisibles par Paraview.
Paraview est un logiciel libre et multi-plateforme. Use it !

TOUGH2 fournit des informations tres pauvres sur le maillage, ce script
fonctionne bien avec les maillages generes par TOUGH2 lui-meme.
"""

__all__ = []


import sys
import os
import itertools
import collections
import re
import xml.etree.cElementTree as ET
import argparse
try:
    import tkFileDialog as filedialog
except ImportError:
    import tkinter.filedialog as filedialog  # NOQA

try:
    zip = itertools.izip
except AttributeError:
    pass
try:
    range = xrange
except NameError:
    pass


DEFAULT_OUTPUT = 'OUTPUT'
DEFAULT_MESH = 'MESH'

BLOCK_SEP = '@' * 131
ELEM_INDEX = r'ELEM\.\s+INDEX'


class EndOfOutFile(Exception):
    pass


class NotAGridError(Exception):
    def __init__(self, mesh_filename):
        super(Exception, self).__init__(
            "'%s' file does not describe a valid grid" % mesh_filename)


def _float(s):
    try:
        return float(s)
    except ValueError:
        return 0.

###############################################################################
###############################################################################
##  vtk_file management
###############################################################################

def build_vtkfile(mesh_filename, out_filename, vtk_basename):
    coordinates = build_coordinates(mesh_filename)
    convert_idx = build_convert_idx(coordinates)
    coordinates_data = build_coordinates_DataArrays(coordinates)
    node_shape = tuple(len(c) for c in coordinates)
    create_data_directory(vtk_basename)
    with open(out_filename) as outfile:
        pvd_data = []
        for it, (time, fields, data) in enumerate(iter_time_data(outfile)):
            vtk_filename = write_vtk_RectilinearGrid(
                vtk_basename, it, time, fields, data,
                node_shape, coordinates_data, convert_idx
            )
            pvd_data.append((it, time, vtk_filename))
        write_pvd_file(pvd_data, vtk_basename)


def create_data_directory(newpath):
    if not os.path.exists(newpath):
        os.makedirs(newpath)


def write_pvd_file(pvd_data, vtk_basename):
    vtkfile = ET.Element("VTKFile")
    vtkfile.set('type', 'Collection')
    vtkfile.set('version', '0.1')
    vtkfile.set('byte_order', 'LittleEndian')

    collection = ET.SubElement(vtkfile, 'Collection')
    path = os.path.dirname(vtk_basename)
    path += '/' if path else ''
    for it, time, fname in pvd_data:
        fname = re.sub(r'^%s' % path, '', fname)
        _dar = ET.SubElement(collection, 'DataSet')
        _dar.set('timestep', str(time))
        _dar.set('file', fname)
    pvd_filename = '%s.pvd' % vtk_basename
    write_xml(vtkfile, pvd_filename, True)


def write_vtk_RectilinearGrid(vtk_basename, it, time, fields, data,
                              node_shape, coordinates_data, convert_idx):
    cell_shape = tuple(max(1, s - 1) for s in node_shape)
    extent = '0 %i 0 %i 0 %i' % cell_shape

    vtkfile = ET.Element("VTKFile")
    vtkfile.set('type', 'RectilinearGrid')
    vtkfile.set('version', '0.1')
    vtkfile.set('byte_order', 'LittleEndian')

    grid = ET.SubElement(vtkfile, 'RectilinearGrid')
    grid.set('WholeExtent', extent)

    piece = ET.SubElement(grid, 'Piece')
    piece.set('Extent', extent)

    coordinates = ET.SubElement(piece, 'Coordinates')
    xyz = []
    for axis, name in [(0, 'x'), (1, 'y'), (2, 'z')]:
        _dar = ET.SubElement(coordinates, 'DataArray')
        xyz.append(_dar)
        _dar.set('type', 'Float32')
        _dar.set('format', 'ascii')
        _dar.set('Name', name)
        _dar.text = coordinates_data[axis]

    celldata = ET.SubElement(piece, 'CellData')
    celldata.set('Scalars', fields[0])
    for name, values in zip(fields, data):
        _dar = ET.SubElement(celldata, 'DataArray')
        _dar.set('type', 'Float32')
        _dar.set('format', 'ascii')
        _dar.set('Name', name)
        _dar.text = ' '.join(str(values[idx]) for idx in convert_idx)

    vtk_filename = '%s/%i.vtr' % (vtk_basename, it)
    write_xml(vtkfile, vtk_filename, True)
    return vtk_filename


def build_coordinates_DataArrays(coordinates):
    return [' '.join(str(s) for s in spacing) for spacing in coordinates]


###############################################################################
###############################################################################
##  mesh_file management
###############################################################################

def build_convert_idx(coordinates):
    "convert_idx[vtk_idx] -> t2_idx"
    Nx, Ny, Nz = [max(1, len(c) - 1) for c in coordinates]
    convert_idx = [
        k + j * Nz + i * Nz * Ny for k, j, i in
        itertools.product(range(Nz), range(Ny), range(Nx))
    ]
    return convert_idx


def build_coordinates(mesh_filename):
    dim, t2mesh = read_t2mesh(mesh_filename)
    spacings, start, ends = build_spacings(dim, t2mesh, mesh_filename)
    coordinates = [list(iter_cumsum(spacing)) for spacing in spacings]
    coordinates += ([[0.]] * 3)[len(coordinates):]
    return coordinates


def read_t2mesh(mesh_filename):
    global_idx = name_to_index(mesh_filename)
    ncell = len(global_idx)

    t2mesh = [T2Cell() for i in range(ncell)]

    with open(mesh_filename) as meshfile:
        move_after(meshfile, 'CONNE')
        for line in meshfile:
            if not line.strip(' \n+'):
                break

            name1 = line[:5] + line[15:20]
            name2 = line[5:10] + line[20:25]
            idx1, idx2 = global_idx[name1], global_idx[name2]
            d1, d2 = _float(line[30:40]), _float(line[40:50])

            t2mesh[idx1].neib_idx.append(idx2)
            t2mesh[idx1].neib_dist.append(d1)

            t2mesh[idx2].neib_idx.append(idx1)
            t2mesh[idx2].neib_dist.append(d2)

    neib_tot = collections.Counter(len(c.neib_idx) for c in t2mesh)
    dim = min(neib_tot)
    if not (
        set(range(dim, 2 * dim + 1)) == set(neib_tot.keys())
        and
        neib_tot[dim] == 2 ** dim
        and
        not sum(neib_tot[i] % 2 ** (2 * dim - i) for i in range(dim, 2 * dim))
    ):
        raise NotAGridError(mesh_filename)
    return dim, t2mesh


T2Cell_type = collections.namedtuple('T2Cell', ['neib_idx', 'neib_dist'])


def T2Cell():
    return T2Cell_type([], [])


def name_to_index(mesh_filename):
    with open(mesh_filename) as meshfile:
        move_after(meshfile, 'ELEME')
        gidx = {}
        for n, line in enumerate(meshfile):
            if not line.strip(' \n+'):
                break
            cell_name = line[:5] + line[10:15]
            gidx[cell_name] = n
    return gidx


def build_spacings(dim, t2mesh, mesh_filename):
    spacings = []
    ends = []
    start = next(i for i, c in enumerate(t2mesh) if len(c.neib_idx) == dim)
    for axis in range(dim):
        spacing = [2 * t2mesh[start].neib_dist[axis]]
        old, new = start, t2mesh[start].neib_idx[axis]
        level = dim + 1
        while True:
            nextc = next_edge_node(t2mesh, new, old, level, mesh_filename)
            if nextc is None:
                # deal with corner and break
                level -= 1
                nextc = next_edge_node(t2mesh, new, old, level, mesh_filename)
                old, (new, dist) = new, nextc
                spacing.append(2 * dist)
                # add corner spacing
                endc = t2mesh[new]
                dist = next(
                    d for i, d in zip(endc.neib_idx, endc.neib_dist)
                    if i == old
                )
                spacing.append(2 * dist)
                break
            else:
                old, (new, dist) = new, nextc
                spacing.append(2 * dist)
        spacings.append(spacing)
        ends.append(new)
    return spacings, start, ends


def next_edge_node(t2mesh, new, old, level, mesh_filename):
    nexts = [(i, dist) for i, dist in
             zip(t2mesh[new].neib_idx, t2mesh[new].neib_dist)
             if i != old and len(t2mesh[i].neib_idx) == level]
    if not nexts:
        return None
    elif len(nexts) > 1:
        raise NotAGridError(mesh_filename)
    else:
        return nexts[0]


def iter_cumsum(it, start=0):
    yield start
    for x in it:
        start += x
        yield start


###############################################################################
###############################################################################
##  out_file management
###############################################################################

def iter_time_data(outfile):
    while True:
        try:
            line = move_after(outfile, 'OUTPUT DATA AFTER')
        except EndOfOutFile:
            break
        else:
            time = get_time(line)
            fields, data = get_data_block(outfile)
            yield time, fields, data


def get_data_block(outfile):
    move_after(outfile, BLOCK_SEP)
    move_after(outfile, BLOCK_SEP)
    line = move_after(outfile, ELEM_INDEX)
    fields = line.split()[2:]
    Nfields = len(fields)
    data = tuple([] for i in range(Nfields))
    for line in outfile:
        if BLOCK_SEP in line:
            break
        elif (len(line) > 6 and line[:6] != ' ' * 6 and
                not re.search(ELEM_INDEX, line)):
            values = [_float(s) for s in line.split()[-Nfields:]]
            for field, val in zip(data, values):
                field.append(val)
    return fields, data


def move_after(file, target, verbose=False):
    ctarget = re.compile(target)
    for line in file:
        if ctarget.search(line):
            return line
        if verbose:
            print line
    raise EndOfOutFile


def get_time(line):
    match = re.search(r'THE TIME IS (.*) DAYS', line)
    return _float(match.group(1))


###############################################################################
###############################################################################
##  write_xml
###############################################################################

def indent_element(elem, level=0, indent="  "):
    i = "\n" + level * indent
    if len(elem):
        if not elem.text or not elem.text.strip():
            elem.text = i + "  "
        if not elem.tail or not elem.tail.strip():
            elem.tail = i
        for elem in elem:
            indent_element(elem, level + 1)
        if not elem.tail or not elem.tail.strip():
            elem.tail = i
    else:
        if level and (not elem.tail or not elem.tail.strip()):
            elem.tail = i


def write_xml(elem, filename, indent=False):
    if indent:
        indent_element(elem)
    tree = ET.ElementTree(elem)
    tree.write(filename)


###############################################################################
###############################################################################
##  command line parser
###############################################################################

def parse_cl():
    parser = argparse.ArgumentParser(
        description="convert Tough2 result files to VTK RectilinearGrid files",
        epilog="WARNING: only works with meshes generated by Tough2 !",
    )
    parser.add_argument(
        "--gui",
        help="""
        launch GUI prompt instead of reading command line arguments.
        GUI prompt is invoked if no arguments provided.
        """,
        action="store_true",
    )
    parser.add_argument(
        "-m", "--t2mesh",
        help="Tough2 mesh file ('MESH') [default:%s]" % DEFAULT_MESH,
        type=argparse.FileType(),
    )
    parser.add_argument(
        "-o", "--t2out",
        help="""
            Tough2 output file ('flow.out'/'OUTPUT') [default:%s]
        """ % DEFAULT_OUTPUT,
        type=argparse.FileType(),
    )
    parser.add_argument(
        "-d", "--t2dir",
        help="prefix directory for Tough2 files",
    )
    parser.add_argument(
        "basename",
        help="""
            base name for the created file and directory
            ('BASENAME.pvd', 'BASENAME/')
        """,
        nargs='?',
    )

    args = parser.parse_args()

    if args.gui:
        return gui_prompt()

    if args.basename is None:
        parser.error("basename not found")

    mesh, output, basename = args.t2mesh, args.t2out, args.basename
    mesh = mesh.name if mesh else DEFAULT_MESH
    output = output.name if output else DEFAULT_OUTPUT
    if args.t2dir is not None:
        mesh = '%s/%s' % (args.t2dir, mesh)
        output = '%s/%s' % (args.t2dir, output)
    basename = clean_basename(basename)
    return mesh, output, basename


###############################################################################
###############################################################################
##  GUI prompt
###############################################################################

def gui_prompt():
    path = '~' if ':\WINDOWS' in os.getcwd() else None
    mesh = filedialog.askopenfile(title="Tough2 mesh file", initialdir=path)
    if mesh is None:
        exit()
    else:
        mesh = mesh.name

    path = os.path.dirname(mesh)
    output = filedialog.askopenfile(title="Tough2 out file", initialdir=path)
    if output is None:
        exit()
    else:
        output = output.name

    path = os.path.dirname(output)
    basename = filedialog.asksaveasfilename(title="VTK .pvd file / basename",
                                            initialdir=path)
    basename = clean_basename(basename)

    return mesh, output, basename


###############################################################################
###############################################################################
##  main
###############################################################################

def clean_basename(basename):
    return basename[:-4] if basename[-4:] == '.pvd' else basename


def get_inputs():
    return parse_cl() if len(sys.argv) > 1 else gui_prompt()


if __name__ == '__main__':
    inputs = get_inputs()
    build_vtkfile(*inputs)
