#!/usr/bin/env python3
# -*- coding: UTF-8 -*-

"""
This script reads an image file, for each pixel reads the color and counts
all pixels with the same color and creates a ranking. The list is then used
to create a stacked bar chart with the used colors. Each bar gets the width
in percentage of the amount used in the original file.

The script makes sense for schematic images with a small colors palette only.
Real photograhic images are not suited.

Usage: used-colors.py [options] <infile> [<outfile>]

Arguments:
<infile>        The image file that is analyzed.
<outfile>       The optional file where the SVG output is written to. When
                not set the SVG is written to stdout.
--alpha <int>   Pixels with an alpha value at or below this threshold (0-255)
                are treated as transparent and ignored. Default is 0, which
                drops only fully transparent pixels. Raise it to also discard
                faint anti-aliased edge pixels.
--border <str>  The color and width of the border around the SVG. Default is
                none. A border can be defined as <color>:<width> e.g. #000000:2
                for a black border with width 2. The color can be provided as
                a hex number e.g. #ffffff or as 255,255,255 for white.
--dim <WxH>|<R> The dimension of the resulting SVG. Default is 200x300. When
                the width is larger as the height, vertical bars are drawn,
                otherwise horizintal bars are drawn.
                If only one value is provided, a pie chart is drawn with the
                given radius.
--dump          Output the list of colors and occurrence, no SVG is written.
--help          Print this help screen.
--ignore <str>  Ignore a color (e.g. the background color from anaylzing).
                The color can be provided as a hex number e.g. #ffffff or
                (255, 255, 255) for white.
--num <int>     The number of the most used colors is considered only.
                Default is 10.
"""

from PIL import Image
from collections import Counter
import math, os, re, sys


def dieNice(errMsg: str):
    """
    Die nice with an error message and return exit code 1.
    """
    print(f"Error: {errMsg}\nSee --help for more details.")
    sys.exit(1)

def colorFromArg(arg: str):
    """
    Convert a color argument to an RGB tuple. The color can
    be provided as a hex number e.g. #ffffff or as 255,255,255 for white.
    """

    if arg[0:1] == '#':
        color = tuple(int(arg[i:i+2], 16) for i in (1, 3, 5))
    else:
        color = tuple(map(int, re.sub(r"[^\d\,]", '', arg).split(',')))
        if len(color) != 3:
            raise ValueError
    for c in color:
        if c < 0 or c > 255:
            raise ValueError
    return color

def rgba2rgb(rgba, background=(255, 255, 255)):
    """
    Convert an RGBA color to RGB by compositing it onto a background color.
    The background color is provided as an RGB tuple.
    """
    if not rgba:
        return rgba
    channels = len(rgba)

    if channels == 3:
        return rgba

    assert channels == 4, "RGBA image must have 4 channels."

    r, g, b, a = rgba
    alpha = a / 255.0
    R, G, B = background
    return (
        int(r * alpha + (1 - alpha) * R),
        int(g * alpha + (1 - alpha) * G),
        int(b * alpha + (1 - alpha) * B),
    )

def getUsedColors(file: str, numColors: int, **kwargs):
    """
    Get an associative array of colors and their proportions from an image file.
    The key is the RGB color as a string "R,G,B", and the value is the proportion
    of that color in the image.
    Every pixel is counted at the image's native resolution (no downscaling / interpolation,
    so no colors are invented and no counts are smeared across neighbours), and transparency is
    decided per pixel from the alpha channel. Pixels whose alpha is at or below
    $alphaThreshold are dropped; the rest are composited onto $bgColor.
    """
    # Load image as RGBA so transparency lives in the alpha channel. 
    img = Image.open(file).convert("RGBA")

    bgColor = kwargs.get('bgColor', (255, 255, 255)) # Default background color (white)
    # Pixels with alpha at or below this value are treated as transparent and
    # dropped. Default 0 = drop only fully transparent pixels.
    alphaThreshold = kwargs.get('alphaThreshold', 0)

    # Decide transparency per pixel via its alpha value, never by guessing a
    # single "transparent color" (an RGBA image does not have one). Drop the
    # transparent pixels, then composite the rest onto the background.
    colorCounts = Counter(
        rgba2rgb(p, background=bgColor)
        for p in img.get_flattened_data()
        if p[3] > alphaThreshold
    )

    # Remove the color to ignore.
    ignore = kwargs.get('ignore')
    if ignore is not None:
        del colorCounts[ignore]
    
    # Sort by most used colors.
    sortedColors = sorted(colorCounts.items(), key=lambda x: x[1], reverse=True)

    # Keep only the top $numColors colors and drop the rest lesser used colors.
    if numColors < len(sortedColors):
        sortedColors = sortedColors[0:numColors]
    
    # Total pixels, needed to convert counts to proportions (but do not count pixels
    # used the dropped colors).
    total = 0
    for color in sortedColors:
        total += color[1]

    # Convert to proportions.
    for key, item in enumerate(sortedColors):
        sortedColors[key] = (item[0], item[1] / total)

    return sortedColors

def writeBarSvg(width: int, height: int, colors: list, borderColor: int = None, borderWidth: int = 0) -> str:
    """
    Write a bar chart as SVG. Input is the width and height and a list of colors (r, g, b) and their proportions (0-1).
    Also a border color and width can be defined. The result is the SVG as a string.
    """

    svgParts = []
    svgParts.append(f'<svg xmlns="http://www.w3.org/2000/svg" width="{width}" height="{height}">')

    offset = 0

    for (r, g, b), count in colors:
        if (width < height):
            barPart = float(count) * height
            svgParts.append(
                f'<rect x="0" y="{offset}" width="{width}" height="{barPart}" '
                f'fill="rgb({r},{g},{b})" />'
            )
            offset += barPart
        else:
            barPart = float(count) * width
            svgParts.append(
                f'<rect x="{offset}" y="0" width="{barPart}" height="{height}" '
                f'fill="rgb({r},{g},{b})" />'
            )
            offset += barPart

    # Add border if specified
    if borderColor is not None:
        fillStr = f'rgb({borderColor[0]},{borderColor[1]},{borderColor[2]})'
        svgParts.append(
            f'<rect x="0" y="0" width="{width}" height="{borderWidth}" fill="{fillStr}" />'
        )
        svgParts.append(
            f'<rect x="0" y="{height - borderWidth}" width="{width}" height="{borderWidth}" fill="{fillStr}" />'
        )
        svgParts.append(
            f'<rect x="0" y="0" width="{borderWidth}" height="{height}" fill="{fillStr}" />'
        )
        svgParts.append(
            f'<rect x="{width - borderWidth}" y="0" width="{borderWidth}" height="{height}" fill="{fillStr}" />'
        )

    svgParts.append('</svg>')

    return "\n".join(svgParts)

def writePieSvg(radius: int, colors: list, borderColor: int = None, borderWidth: int = 0) -> str:
    """
    Write a pie chart as SVG. Input is the radius and a list of colors (r, g, b) and their proportions (0-1).
    Also a border color and width can be defined. The result is the SVG as a string.
    """

    svgParts = []
    svgParts.append(f'<svg xmlns="http://www.w3.org/2000/svg" width="{radius*2}" height="{radius*2}">')

    cx, cy = radius, radius
    startAngle = 0

    for (r, g, b), count in colors:
        endAngle = startAngle + count * 360
        largeArcFlag = 1 if endAngle - startAngle > 180 else 0

        x1 = cx + radius * math.cos(math.radians(startAngle))
        y1 = cy + radius * math.sin(math.radians(startAngle))
        x2 = cx + radius * math.cos(math.radians(endAngle))
        y2 = cy + radius * math.sin(math.radians(endAngle))

        pathData = f'M {cx},{cy} L {x1},{y1} A {radius},{radius} 0 {largeArcFlag},1 {x2},{y2} Z'
        svgParts.append(f'<path d="{pathData}" fill="rgb({r},{g},{b})" />')

        startAngle = endAngle

    # Add border if specified
    if borderColor is not None and borderWidth > 0:
        fillStr = f'rgb({borderColor[0]},{borderColor[1]},{borderColor[2]})'
        radius = radius - borderWidth / 2
        svgParts.append(
            f'<circle cx="{cx}" cy="{cy}" r="{radius}" fill="none" stroke="{fillStr}" stroke-width="{borderWidth}" />'
        )

    svgParts.append('</svg>')

    return "\n".join(svgParts)

def main():
    """Evaluate command line arguments, build up worklog and start
    processing the wiki articles"""

    # available options that can be changed via the command line
    options = ['ignore', 'num', 'dim', 'help', 'dump', 'border', 'bgcolor', 'alpha']
    inFile = ''
    outFile = ''
    numColors = 10
    dim = []
    dump = False
    ignore = None
    borderColor = None
    borderWidth = 0
    bgColor = (255, 255, 255)
    alphaThreshold = 0

    # try to fetch the command line args
    currentCmd = ''
    for i in range(len(sys.argv)):
        if i == 0:
            continue
        arg = sys.argv[i]
        # we have a command identified by -- remember it in currentCmd
        # in case this command needs an argument, or just set the
        # appropriate parameter in the worklog or execute some action
        if arg[0:2] == '--':
            currentCmd = arg[2:]
            if not(currentCmd in options):
                dieNice("Invalid argument %s" % currentCmd)
            if currentCmd == 'help':
                print(__doc__)
                sys.exit(0)
            if currentCmd == 'dump':
                dump = True
                currentCmd = ''
        # we have an argument, what was the previous command, do this
        # action in the worklog.
        elif len(currentCmd) > 0:
            if currentCmd == 'num':
                try:
                    numColors = int(arg)
                except:
                    dieNice('Invalid value for --num')
            elif currentCmd == 'alpha':
                try:
                    alphaThreshold = int(arg)
                    if alphaThreshold < 0 or alphaThreshold > 255:
                        raise ValueError
                except:
                    dieNice('Invalid value for --alpha, expected 0-255')
            elif currentCmd == 'dim':
                try:
                    dim = list(map(int, arg.split("x")))
                    if len(dim) < 1 or len(dim) > 2:
                        raise ValueError
                except:
                    dieNice('dimension value incorrect')
            elif currentCmd == 'border':
                try:
                    borderParts = arg.split(':')
                    if len(borderParts) != 2:
                        raise ValueError
                    borderColor, borderWidth = borderParts
                    borderWidth = int(borderWidth)
                    borderColor = colorFromArg(borderColor)
                except:
                    dieNice('Invalid value for --border')        
            elif currentCmd == 'ignore':
                try:
                    ignore = colorFromArg(arg)
                except:
                    dieNice('Invalid value for color to ignore')
            currentCmd = ''
        else:
            # No argument prefix hence it must be a file name.
            if inFile == '':
                inFile = arg if arg[0:1] == '/' else os.getcwd() + '/' + arg
                # Check if current arg is a file that actually exists.
                if not os.path.isfile(inFile):
                    dieNice(f'File {arg} does not exist.')
            elif outFile == '':
                outFile = arg if arg[0:1] == '/' else os.getcwd() + '/' + arg
            else:
                dieNice('Invalid argument')

    if inFile == '':
        dieNice('Image file not given')
    # process the data now
    colors = getUsedColors(inFile, numColors, ignore=ignore, bgColor=bgColor, alphaThreshold=alphaThreshold)

    if dump == True:
        # Print top colors.
        for color, proportion in colors:
            print(color, proportion)
        print(sum(value for _, value in colors))
        sys.exit(0)

    chartFormat = 'bar'
    if len(dim) == 0:
        dim = [200,300]
    elif len(dim) == 1:
        chartFormat = 'pie'

    if outFile == '':
        if chartFormat == 'bar':
            print(writeBarSvg(dim[0], dim[1], colors, borderColor=borderColor, borderWidth=borderWidth))
        else:
            print(writePieSvg(dim[0], colors, borderColor=borderColor, borderWidth=borderWidth))
    else:
        with open(outFile, "w") as f:
            if chartFormat == 'bar':
                f.write(writeBarSvg(dim[0], dim[1], colors, borderColor=borderColor, borderWidth=borderWidth))
            else:
                f.write(writePieSvg(dim[0], colors, borderColor=borderColor, borderWidth=borderWidth))

if __name__ == "__main__":
    main()
