#!/usr/bin/env python2
"""Convert binary particle track to VTK poly data.

Usage information can be obtained by running 'track.py --help':

    usage: track.py [-h] [-o OUT] IN [IN ...]

    Convert particle track file to a .pvtp file.

    positional arguments:
      IN                 Input particle track data filename(s).

    optional arguments:
      -h, --help         show this help message and exit
      -o OUT, --out OUT  Output VTK poly data filename.

"""

import os
import argparse
import struct
import vtk


def _parse_args():
    # Create argument parser.
    parser = argparse.ArgumentParser(
        description='Convert particle track file to a .pvtp file.')
    parser.add_argument('input', metavar='IN', type=str, nargs='+',
                        help='Input particle track data filename(s).')
    parser.add_argument('-o', '--out', metavar='OUT', type=str, dest='out',
                        help='Output VTK poly data filename.')

    # Parse and return commandline arguments.
    return parser.parse_args()


def main():
    # Parse commandline arguments.
    args = _parse_args()

    # Check input file extensions.
    for fname in args.input:
        if not (fname.endswith('.h5') or fname.endswith('.binary')):
            raise ValueError("Input file names must either end with '.h5' or"
                             "'.binary'.")

    # Make sure that the output filename ends with '.pvtp'.
    if not args.out:
        args.out = 'tracks.pvtp'
    elif not args.out.endswith('.pvtp'):
        args.out += '.pvtp'

    # Import HDF library if HDF files are present
    for fname in args.input:
        if fname.endswith('.h5'):
            import h5py
            break

    # Initialize data arrays and offset.
    points = vtk.vtkPoints()
    cells = vtk.vtkCellArray()
    point_offset = 0
    for fname in args.input:
        # Write coordinate values to points array.
        if fname.endswith('.binary'):
            track = open(fname, 'rb').read()
            coords = [struct.unpack("ddd", track[24*i : 24*(i+1)])
                      for i in range(len(track)/24)]
            n_points = len(coords)
            for triplet in coords:
                points.InsertNextPoint(triplet)
        else:
            coords = h5py.File(fname).get('coordinates')
            n_points = coords.shape[0]
            for i in range(n_points):
                points.InsertNextPoint(coords[i, :])

        # Create VTK line and assign points to line.
        line = vtk.vtkPolyLine()
        line.GetPointIds().SetNumberOfIds(n_points)
        for i in range(n_points):
            line.GetPointIds().SetId(i, point_offset+i)

        cells.InsertNextCell(line)
        point_offset += n_points
    data = vtk.vtkPolyData()
    data.SetPoints(points)
    data.SetLines(cells)

    writer = vtk.vtkXMLPPolyDataWriter()
    if vtk.vtkVersion.GetVTKMajorVersion() > 5:
        writer.SetInputData(data)
    else:
        writer.SetInput(data)
    writer.SetFileName(args.out)
    writer.Write()


if __name__ == '__main__':
    main()
