#!/usr/bin/env python

import os
import glob
import argparse
import torch
import numpy as np
import surfa as sf

# import voxelmorph with pytorch backend
os.environ['VXM_BACKEND'] = 'pytorch'
os.environ['NEURITE_BACKEND'] = 'pytorch'
import voxelmorph as vxm 


def proc_image(basepath, timepoint, contrast, affine):
    filepath = f'{basepath}_{timepoint:02d}_????_{contrast}.nii.gz'
    img = sf.load_volume(glob.glob(filepath)[0])
    if timepoint == 0:
        img = img.transform(affine=affine.inv())
    img -= img.min()
    img = img.astype(np.float32)
    img /= img.percentile(99)
    img = img.clip(0, 1)
    return img


def proc_timepoint(basepath, timepoint, affine):
    contrasts = ['t1ce', 't1', 't2', 'flair']
    sv = [proc_image(basepath, timepoint, c, affine) for c in contrasts]
    mv = min(sv, key=lambda v : np.prod(v.shape))
    sv = [v.resample_like(mv, method='nearest') for v in sv]
    return sf.stack(sv)


def point_sample(coords, features, image_size):
    half_size = (image_size - 1) / 2
    coords = (coords - half_size) / half_size
    coords = torch.reshape(coords, (1, coords.shape[-2], 1, 1, coords.shape[-1]))
    point_features = torch.nn.functional.grid_sample(features.unsqueeze(0).swapaxes(-1, -3), coords, align_corners=True, mode='bilinear')
    point_features = point_features.squeeze(0).squeeze(-1).squeeze(-1).swapaxes(-1, -2)
    return point_features


def generate_output(args):
    '''
    Generates landmarks, detJ, deformation fields (optional), and followup_registered_to_baseline images (optional) for challenge submission
    '''
    input_path = os.path.abspath(args["input"])
    output_path = os.path.abspath(args["output"])

    # Output will be written to= /output
    print("Output will be written to:", output_path)

    os.makedirs(output_path, exist_ok=True)

    # input shape: (240, 240, 155)
    # temp_shape = (160, 192, 160)
    temp_shape = (256, 256, 160)
    image_size = torch.tensor(temp_shape, dtype=torch.float32).to(device)

    # Get script path
    cpath = os.path.dirname(os.path.realpath(__file__))

    # Affine path
    apath = os.path.join(output_path, 'affine')
    os.makedirs(apath, exist_ok=True)

    # Affine align
    err = sf.system.run(f'python3 {cpath}/align {input_path} {apath}')
    if err != 0:
        print('could not run affine alignment')
        exit(1)

    # Read in model
    replace = dict(inshape=temp_shape)
    model = vxm.networks.VxmDense.load(f'{cpath}/ncc.pt', device, replace)
    model.bidir = True
    model.to(device)
    model.train(False)

    # Now we iterate through each subject folder under input_path
    for subj_path in glob.glob(os.path.join(input_path, "BraTSReg*")):

        subj = os.path.basename(subj_path)
        print(f"Performing deformation {subj}")

        # Load affine
        geom = sf.load_volume(glob.glob(f'{subj_path}/{subj}_00_????_t1ce.nii.gz')[0]).geom
        affine = sf.Affine(np.loadtxt(f'{apath}/{subj}_affine.txt'), space='world', source=geom, target=geom)

        # Read in data
        orig_source = proc_timepoint(f'{subj_path}/{subj}', 0, affine=affine)
        orig_target = proc_timepoint(f'{subj_path}/{subj}', 1, affine=affine)

        source = orig_source.fit_to_shape(temp_shape)
        target = orig_target.fit_to_shape(temp_shape)

        with torch.no_grad():
            source_tensor = torch.from_numpy(np.moveaxis(source.data.astype(np.float32, copy=False), -1, 0)[np.newaxis]).to(device)
            target_tensor = torch.from_numpy(np.moveaxis(target.data.astype(np.float32, copy=False), -1, 0)[np.newaxis]).to(device)
            registered_b2f, registered_f2b, warp_f2b, warp_b2f = model(source_tensor, target_tensor, registration=True)

        # load landmarks
        landmark_file = glob.glob(f'{subj_path}/{subj}_01_????_landmarks.csv')[0]

        # convert to voxel coordinates
        input_landmark = np.loadtxt(landmark_file, delimiter=',', skiprows=1)[:, 1:] @ np.diag([-1, -1, 1])
        input_landmark = source.geom.world2vox(input_landmark)

        # warp landmarks
        input_landmark = torch.from_numpy(input_landmark.astype(np.float32, copy=False)).to(device)
        sampled = point_sample(input_landmark, warp_f2b.squeeze(0), image_size)
        output_landmark = input_landmark + sampled
        output_landmark = output_landmark.cpu().numpy().squeeze()
        output_landmark = affine.convert(space='vox', source=source, target=source).transform(output_landmark)
        output_landmark = source.geom.vox2world(output_landmark)
        output_landmark = output_landmark @ np.diag([-1, -1, 1])

        # write landmarks
        with open(os.path.join(output_path, f"{subj}.csv"), 'w') as file:
            file.write('Landmark,X,Y,Z\n')
            for i, coord in enumerate(output_landmark):
                row = ','.join([f'{c:.4f}' for c in coord])
                file.write(f'{i+1},{row}\n')

        warp_f2b = source.new(np.moveaxis(warp_f2b.cpu().numpy().squeeze(), 0, -1)).resample_like(orig_source, method='nearest')
        warp_b2f = source.new(np.moveaxis(warp_b2f.cpu().numpy().squeeze(), 0, -1)).resample_like(orig_source, method='nearest')

        # calculate the determinant of jacobian of the deformation field
        detj = orig_source.new(vxm.py.utils.jacobian_determinant(warp_b2f.data))
        detj.save(os.path.join(output_path, f"{subj}_detj.nii.gz"))

        # 
        if args["def"] or args["reg"]:
            aff = affine.convert(space='vox')
            shape = orig_source.baseshape
            mesh = np.mgrid[tuple([slice(0, s) for s in shape])]
            # mesh = mesh - ((np.array(shape) - 1) / 2)[:, None, None, None]
            m = [f.flatten() for f in mesh]
            m.append(np.ones(m[0].shape))
            m = np.stack(m, axis=1).T
            s = np.moveaxis(mesh, 0, -1)

            A = (aff.inv().matrix @ m)[:3].T.reshape((*shape, 3)) - s
            B = (aff.matrix @ m)[:3].T.reshape((*shape, 3)) - s
            A = orig_source.new(np.ascontiguousarray(A)).astype('float32')
            B = orig_source.new(np.ascontiguousarray(B)).astype('float32')

            warp_f2b = warp_f2b + B.transform(disp=warp_f2b)
            warp_b2f = A + warp_b2f.transform(disp=A)

        if args["def"]:
            # write both the forward and backward deformation fields to the output/ folder
            warp_f2b.save(os.path.join(output_path, f"{subj}_df_f2b.nii.gz"))
            warp_b2f.save(os.path.join(output_path, f"{subj}_df_b2f.nii.gz"))

        if args["reg"]:
            # write the follow-ups-registered-to-baseline resampled images
            warp_b2f.data = np.asfortranarray(warp_b2f.data)
            warp = lambda c : sf.load_volume(glob.glob(f'{subj_path}/{subj}_01_????_{c}.nii.gz')[0]).transform(disp=warp_b2f)

            warp('t1ce').save(os.path.join(output_path, f"{subj}_t1ce_f2b.nii.gz"))
            warp('t1').save(os.path.join(output_path, f"{subj}_t1_f2b.nii.gz"))
            warp('t2').save(os.path.join(output_path, f"{subj}_t2_f2b.nii.gz"))
            warp('flair').save(os.path.join(output_path, f"{subj}_flair_f2b.nii.gz"))


def apply_deformation(args):
    '''
    Applies a deformation field on an input image and saves/returns the output
    '''
    print("apply_deformation called")
        
    # Read the field
    disp = sf.load_volume(args['field'])

    # Read the input image
    image = sf.load_volume(args['image'])

    # apply field on image and get output
    method = 'nearest' if args['interpolation'] == 'nearest_neighbour' else 'linear'
    transformed = image.transform(disp=disp, method=method)

    # If a save_path is provided then write the output there, otherwise return the output
    save_path = args.get('path_to_output_nifti')
    if save_path:
        transformed.save(save_path)
    else:
        return transformed.data


if __name__ == "__main__":
    # You can first check what devices are available to the singularity
    # setting device on GPU if available, else CPU
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print('Using device:', device)

    #Additional Info when using cuda
    if device.type == 'cuda':
        print('GPU:', torch.cuda.get_device_name(0))

    # Parse the input arguments

    parser = argparse.ArgumentParser(description='Argument parser for BraTS_Reg challenge')
    
    subparsers = parser.add_subparsers()

    command1_parser = subparsers.add_parser('generate_output')
    command1_parser.set_defaults(func=generate_output)
    command1_parser.add_argument('-i', '--input', type=str, default="/input", help='Provide full path to directory that contains input data')
    command1_parser.add_argument('-o', '--output', type=str, default="/output", help='Provide full path to directory where output will be written')
    command1_parser.add_argument('-d', '--def', action='store_true', help='Output forward and backward deformation fields')
    command1_parser.add_argument('-r', '--reg', action='store_true', help='Output followup scans registered to baseline')

    command2_parser = subparsers.add_parser('apply_deformation')
    command2_parser.set_defaults(func=apply_deformation)
    command2_parser.add_argument('-f', '--field', type=str, required=True, help='Provide full path to deformation field')
    command2_parser.add_argument('-i', '--image', type=str, required=True, help='Provide full path to image on which field will be applied')
    command2_parser.add_argument('-t', '--interpolation', type=str, required=True, help='Should be nearest_neighbour (for segmentation mask type images) or trilinear etc. (for normal scans). To be handled inside apply_deformation() function')
    command2_parser.add_argument('-p', '--path_to_output_nifti', type=str, default = None, help='Format: /path/to/output_image_after_applying_deformation_field.nii.gz')
    

    args = vars(parser.parse_args())

    print("Received the following arguments =", args) 

    args["func"](args)
