#!/usr/bin/env python
import argparse
import sys
import chemlab as cl
import os
from chemlab.libs.termcolor import cprint
import time

def infoprint(msg, color=None):
    cprint('INFO: ', color='yellow', end='')    
    cprint(msg, color=color, end='')
    sys.stdout.flush()
    
def view(args):
    from chemlab.io import datafile
    from chemlab.graphics import display_system, display_trajectory
    from chemlab.core import subsystem_from_molecules
    import numpy as np

    base, ext = os.path.splitext(args.filename)
    
    infoprint('Reading file %s ... ' % args.filename)
    t0 = time.time()
    
    df = datafile(args.filename)
    sys = df.read('system')
  
    cprint('DONE (%f s)'%(time.time()-t0), color='green')	
  
    if args.nowater:
        infoprint('Removing water molecules ... ')
        t0 = time.time()
        
        mols = sys.get_derived_molecule_array('formula')
        noh2o = np.logical_not(mols == 'H2O')
        noh2o_atom = sys.mol_to_atom_indices(noh2o.nonzero()[0])        
        
        sys = subsystem_from_molecules(sys, noh2o)
      
        cprint('DONE (%f s)'%(time.time() - t0), color='green')
   
    if not args.traj:
        display_system(sys)
    else:
        infoprint('Reading trajectory frames ... ')
        t0 = time.time()
        
        df = datafile(args.traj)
        times, frames = df.read('trajectory', skip=args.skip or None)
        
        cprint('DONE (%f s)'%(time.time()-t0), color='green')
        
        if args.nowater:
            infoprint('Removing water from frames ... ')
            t0 = time.time()
            
            for i, f in enumerate(frames):
                frames[i] = f[noh2o_atom]
                
            cprint('DONE (%f s)'%(time.time()-t0), color='green')
            
        if args.smooth:
            infoprint('Smoothing frames ... ')
            t0 = time.time()
            for i, f in enumerate(frames):
                nf = np.zeros_like(f)
                toaverage = range(max(0, i - 2), min(i+4, len(frames)))
                navg = 0
                for j in toaverage:
                    navg += 1
                    nf += frames[j]

                nf /= navg
                frames[i] = nf
            cprint('DONE (%f s)'%(time.time()-t0), color='green')
            
        display_trajectory(sys, times,  frames)

def plot(args):
    base, ext = os.path.splitext(args.filename)

parser = argparse.ArgumentParser()

subparsers = parser.add_subparsers()

# Contrib plugins
from chemlab.contrib import gromacs
gromacs.setup_commands(subparsers)

view_parser = subparsers.add_parser('view')
view_parser.add_argument('filename', type=str)
view_parser.add_argument('--traj', action='store', type=str)
view_parser.add_argument('--nowater', action='store_true',
                         help='Do not display water molecules')
view_parser.add_argument('--smooth', action='store_true',
                         help='Apply a very crude smoothing filter')
view_parser.add_argument('--skip', action='store', type=int,
                         help='Take only the ith frames')

view_parser.set_defaults(func=view)

args = parser.parse_args(sys.argv[1:])
args.func(args)