#!/usr/bin/env python3

"""Convert HDF5 particle track to VTK poly data.

"""

import argparse

import openmc
import vtk


def _parse_args():
    # Create argument parser.
    parser = argparse.ArgumentParser(
        description='Convert particle track file(s) 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()

    # 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'

    # 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.
        track_file = openmc.Tracks(fname)
        for track in track_file:
            for particle in track:
                for state in particle.states:
                    points.InsertNextPoint(state['r'])

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

                # Add line to cell array
                cells.InsertNextCell(line)

    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()
