"""An artistic tomographic image filter

Execute this script using

  python tomographic_filter.py --help

to get a summary of the built-in command line interface.


Author: Paul Müller (https://pycache.de)
License: CC BY-SA 4.0 (https://creativecommons.org/licenses/by-sa/4.0/)
Version: 1.1 (2018)

Changes:
1.1
 - raise ValueError when input image does not have square shape
 - include condition of square shape in command line docs
 - during filtering, print note that larger images require longer
   computation times
 - print output file name after filter operation is completed
1.0
 - initial release
"""
import argparse
import pathlib

import imageio
import numpy as np
import radontea

# default parameters
dft_num_angles = 37  # number of projections to use (int)
dft_ival_coverage = (0, 180)  # angular coverage (tuple of floats)
dft_randomness = 0  # distribute projections randomly (False or int)
dft_weight_angles = False  # angular weighting during reconstruction (bool)
dft_angle_offset = 0  # offset added to the angular interval (float)
dft_circular_mask = False  # only use image data within a circle (bool)
dft_normalize = True  # normalize input image data (bool)


def adjust_colors(orig, recon):
    """Adjust brightness of input image to that of output

    Normalization uses data from within a circular area (excluding
    the image corners).

    Parameters
    ----------
    orig: ndarray of shape (M,M)  for grayscale or (M,M,3) for color
        The original image
    recon: ndarray, same shape as orig
        The reconstructed image

    Returns
    -------
    corr: ndarray, same shape as orig
        Color-corrected version `recon` (using mean and standard
        deviation)
    """
    if len(orig.shape) == 3:
        # treat each color separately
        update = np.zeros_like(orig)
        for ii in range(3):
            update[..., ii] = adjust_colors(orig[..., ii], recon[..., ii])
    # circular aperture of input image
    valid = get_circular_mask(orig.shape[0])

    iavg = np.average(orig[valid])
    istd = np.std(orig[valid])
    ravg = np.average(recon[valid])
    rstd = np.std(recon[valid])

    upd = (recon - ravg) / rstd * istd + iavg
    return upd


def get_angles(num_angles=dft_num_angles, ival_coverage=dft_ival_coverage,
               angle_offset=dft_angle_offset, randomness=dft_randomness):
    """Generate a set of acquisition angles

    Parameters
    ----------
    num_angles: int
        Total number of angles
    ival_coverage: tuple of floats
        The angular coverage in degrees. A full angular coverage is
        given by (0, 180).
    angle_offset: float
        An angle in degrees that is added ontop of `ival_coverage`
    randomness: int
        If zero, angles are distributed equidistant in `ival_coverage`.
        If non-zero, specifies a random seed for angle determination.

    Returns
    -------
    angles: 1d ndarray
        Angles in radians
    """
    ivrad = np.deg2rad(ival_coverage) + np.deg2rad(angle_offset)
    if randomness:
        rdnst = np.random.RandomState(randomness)
        np.random.set_state(rdnst.get_state())
        fullang = np.linspace(ivrad[0], ivrad[1],
                              num_angles*20, endpoint=False)
        angles = np.random.choice(fullang, num_angles, replace=False)
    else:
        angles = np.linspace(ivrad[0], ivrad[1], num_angles, endpoint=False)
    return angles


def get_circular_mask(size):
    """Return a circular mask for a square image

    Parameters
    ----------
    size: int
        Side length of the square image in pixels

    Returns
    -------
    mask: ndarray of shape (size,size)
        Binary mask of a circle with radius size/2.
    """
    x = np.linspace(-size/2, size/2, size, endpoint=True)
    xx, yy = np.meshgrid(x, x)
    mask = (xx**2 + yy**2 <= (size/2)**2)
    return mask


def get_sinogram(orig, angles, circular_mask=dft_circular_mask,
                 normalize=dft_normalize):
    """Compute the sinogram of an image

    Parameters
    ----------
    orig: ndarray of shape (M,M)  for grayscale or (M,M,3) for color
        The original image
    angles: 1d ndarray of shape (A,)
        Angles in radians corresponding to the sinogram slices.
    circular_mask: bool
        Only use image data from within a circle to reduce artifacts
        from the image corners.

    Returns
    -------
    sinogram: list of ndarrays of shape (A,M)
        Sinogram data. The length of the list is 1 (gray scale)
        or 3 (color).
    """
    data = np.array(orig, dtype=float)
    if data.shape[0] != data.shape[1]:
        raise ValueError("Input image must have square shape. Please crop "
                         + "or extend it with your favorite image "
                         + "manipulation program!")
    valid = get_circular_mask(data.shape[0])
    if normalize:
        data -= np.mean(data[~valid])
    if circular_mask:
        data[~valid] = 0
    # treat RGB data as individual images
    if len(data.shape) == 3:
        trdata = [data[..., 0], data[..., 1], data[..., 2]]
    else:
        trdata = [data]

    sino = []
    for td in trdata:
        sino.append(radontea.radon_parallel(td, angles=angles))

    return sino


def reconstruct(sinogram, angles, weight_angles=dft_weight_angles):
    """Perform tomographic reconstruction from a sinogram

    Parameters
    ----------
    sinogram: list of ndarrays of shape (A,M)
        Sinogram data. The length of the list is 1 (gray scale)
        or 3 (color).
    angles: 1d ndarray of shape (A,)
        Angles in radians corresponding to the sinogram slices.
    weight_angles: bool
        Whether to perform angular weighting or not.

    Returns
    -------
    recon: list of ndarrays of shape (M,M)
        Reconstructions corresponding to the sinogram list items.
    """
    recons = []
    for sn in sinogram:
        rec = radontea.backproject(sn, angles,
                                   filtering="ramp",
                                   weight_angles=weight_angles,
                                   padding=True,
                                   padval=0)
        recons.append(rec)
    if len(recons) == 3:
        # convert to RGB array
        recim = np.dstack(recons)
    else:
        recim = recons[0]
    return recim


def tomographic_filter(path_in, path_out=None, num_angles=dft_num_angles,
                       ival_coverage=dft_ival_coverage,
                       angle_offset=dft_angle_offset,
                       weight_angles=dft_weight_angles,
                       circular_mask=dft_circular_mask,
                       normalize=True,
                       randomness=dft_randomness):
    """Apply a artistic tomographic image filter

    Parameters
    ----------
    path_in: str
        Input file name
    path_out: str
        Output file name
    num_angles: int
        Total number of angles
    ival_coverage: tuple of floats
        The angular coverage in degrees. A full angular coverage is
        given by (0, 180).
    angle_offset: float
        An angle in degrees that is added ontop of `ival_coverage`
    weight_angles: bool
        Whether to perform angular weighting or not.
    circular_mask: bool
        Only use image data from within a circle to reduce artifacts
        from the image corners.
    randomness: int
        If zero, angles are distributed equidistant in `ival_coverage`.
        If non-zero, specifies a random seed for angle determination.
    """
    # load data from file
    path_in = pathlib.Path(path_in).resolve()
    data = np.array(imageio.imread(str(path_in)))
    # compute angles
    angles = get_angles(num_angles=num_angles, ival_coverage=ival_coverage,
                        randomness=randomness, angle_offset=angle_offset)
    # compute sinogram
    sino = get_sinogram(orig=data, angles=angles, circular_mask=circular_mask,
                        normalize=normalize)
    # perform reconstruction
    rec = reconstruct(sino, angles=angles, weight_angles=weight_angles)
    # adjust colors to match those if input
    rec = adjust_colors(orig=data, recon=rec)
    rec[rec < 0] = 0
    rec[rec > 254] = 254
    rec = np.asarray(rec, dtype=np.uint8)
    # save image
    if path_out is None:
        namelist = [path_in.stem, "_tf"]
        if num_angles != dft_num_angles:
            namelist.append("_nang-{:d}".format(num_angles))
        if ival_coverage != dft_ival_coverage:
            namelist.append("_ival-{}-{}".format(*ival_coverage))
        if angle_offset != dft_angle_offset:
            namelist.append("_off-{}".format(angle_offset))
        if weight_angles != dft_weight_angles:
            namelist.append("_weight-{}".format(weight_angles))
        if circular_mask != dft_circular_mask:
            namelist.append("_mask-{}".format(circular_mask))
        if randomness != dft_randomness:
            namelist.append("_rand-{}".format(randomness))
        if normalize != dft_normalize:
            namelist.append("_norm-{}".format(normalize))
        namelist.append(path_in.suffix)
        path_out = path_in.with_name("".join(namelist))
    imageio.imsave(str(path_out), rec)
    return path_out


if __name__ == "__main__":
    descr = "An artistic tomographic image filter."
    parser = argparse.ArgumentParser(
        description=descr,
        formatter_class=argparse.ArgumentDefaultsHelpFormatter)
    parser.add_argument('path',
                        type=pathlib.Path,
                        help='Input image path. The image must have square'
                             + ' shape.')
    parser.add_argument('--out',
                        type=str,
                        default="none",
                        help='Output image path.')
    parser.add_argument('--num-angles',
                        type=np.uint,
                        default=dft_num_angles,
                        help='Number of angles for tomographic reconstruction.'
                             + ' Affects the number of streak-artifacts.')
    parser.add_argument('--cov-min',
                        type=float,
                        default=dft_ival_coverage[0],
                        help='Interval start for angular coverage. If not '
                             + 'zero, leads to so-called missing-angle '
                             + 'artifacts. Should be a number between 0 '
                             + 'and 180.')
    parser.add_argument('--cov-max',
                        type=float,
                        default=dft_ival_coverage[1],
                        help='Interval start for angular coverage. If not '
                             + '180, leads to so-called missing-angle '
                             + 'artifacts. Should be a number between 0 '
                             + 'and 180.')
    parser.add_argument('--offset',
                        type=float,
                        default=dft_angle_offset,
                        help='Offset added to the coverage interval in deg.')
    parser.add_argument('--randomness',
                        type=np.uint,
                        default=dft_randomness,
                        help='If zero, the streak-artifacts are distributed '
                             + 'equidistant. If an integer, a random '
                             + 'distribution of angles is used.')
    parser.add_argument('--weight-angles',
                        dest="weight_angles",
                        action='store_true',
                        help='If set, enables angular weighting, '
                             + 'affecting the reconstruction quality when '
                             + '"randomness" or "cov-*" are set.')
    parser.set_defaults(weight_angles=dft_weight_angles)
    parser.add_argument('--circular-mask',
                        dest="circular_mask",
                        action='store_true',
                        help='Only use input data at coordinates within a '
                             + 'circle. This removes artifacts due to data '
                             + 'at the image edges.')
    parser.set_defaults(circular_mask=dft_circular_mask)
    parser.add_argument('--no-normalize',
                        dest="normalize",
                        action='store_false',
                        help='Do not subtract the average color of the corners'
                             + ' from the image. This option causes ring-'
                             + 'like artifacts for bright images.')
    parser.set_defaults(normalize=dft_normalize)
    args = parser.parse_args()
    msg = "Please wait, this may take a while (especially for large images)..."
    print(msg, end="\r", flush=True)
    pout = tomographic_filter(
        path_in=args.path,
        path_out=args.out if args.out is not "none" else None,
        num_angles=args.num_angles,
        ival_coverage=(args.cov_min, args.cov_max),
        angle_offset=args.offset,
        weight_angles=args.weight_angles,
        circular_mask=args.circular_mask,
        normalize=args.normalize,
        randomness=args.randomness)
    print(" "*len(msg), end="\r", flush=True)
    print("Filter operation completed. Output written to {}.".format(pout))
