#!/usr/bin/env python3

import argparse
import sys

from src import pyjpeg


def print_data_unit(data_unit: list[int]) -> None:
    values = pyjpeg.unzig_zag(data_unit)

    cols = []
    for x in range(8):
        col = []
        for y in range(8):
            col.append("%d" % values[y * 8 + x])
        cols.append(col)

    col_widths = []
    for x in range(8):
        width = 0
        for y in range(8):
            width = max(width, len(cols[x][y]))
        col_widths.append(width)

    for y in range(8):
        row = []
        for x in range(8):
            row.append(cols[x][y].rjust(col_widths[x]))
        print("  %s" % " ".join(row))


parser = argparse.ArgumentParser(
    prog="analyze", description="A tool for analyze JPEG images."
)
parser.add_argument("input", help="Path to the input JPEG file")
args = parser.parse_args(sys.argv[1:])

reader = pyjpeg.FileReader(open(args.input, "rb"))
stream = pyjpeg.Stream.read(reader)

is_lossless = False
is_ls = False
for segment in stream.segments:
    if isinstance(segment, pyjpeg.StartOfImage):
        print("SOI Start of Image")
    elif isinstance(segment, pyjpeg.JfifHeader):
        print("APP%d JFIF" % segment.n)
        print(" Version: %d.%d" % (segment.version[0], segment.version[1]))
        if segment.density.unit == pyjpeg.JfifDensityUnit.ASPECT_RATIO:
            print(" Aspect Ratio: %dx%d" % (segment.density.x, segment.density.y))
        elif segment.density.unit == pyjpeg.JfifDensityUnit.DPI:
            print(" Density: %dx%ddpi" % (segment.density.x, segment.density.y))
        elif segment.density.unit == pyjpeg.JfifDensityUnit.DPCM:
            print(" Density: %dx%ddpcm" % (segment.density.x, segment.density.y))
        if len(segment.thumbnail_data) > 0:
            # FIXME: Support RGB thumbnails
            s = " Thumbnail %dx%d:" % (
                segment.thumbnail_size[0],
                segment.thumbnail_size[1],
            )
            for i in range(0, len(segment.thumbnail_data), 3):
                if i % (segment.thumbnail_size[0] * 3) == 0:
                    s += "\n "
                s += " %d,%d,%d" % (
                    segment.thumbnail_data[i],
                    segment.thumbnail_data[i + 1],
                    segment.thumbnail_data[i + 2],
                )
            print(s)
    elif isinstance(segment, pyjpeg.JfifJpegThumbnail):
        print("APP%d JPEG Thumbnail" % segment.n)
        print(" Data: %s" % repr(segment.data))
    elif isinstance(segment, pyjpeg.JfifPalletizedThumbnail):
        print("APP%d Palletized Thumbnail" % segment.n)
        print(" Width: %d" % segment.width)
        print(" Height: %d" % segment.height)
        print(" Data: %s" % segment.data)
    elif isinstance(segment, pyjpeg.JfifRgbThumbnail):
        print("APP%d RGB Thumbnail" % segment.n)
        print(" Width: %d" % segment.width)
        print(" Height: %d" % segment.height)
        print(" Data: %s" % segment.data)
    elif isinstance(segment, pyjpeg.SpiffHeader):
        print("APP%d SPIFF" % segment.n)
        print(" Version: %d.%d" % (segment.version[0], segment.version[1]))
        print(" Profile: %d" % segment.profile)
        print(" Number of Components: %d" % segment.number_of_components)
        print(" Height: %d" % segment.height)
        print(" Width: %d" % segment.width)
        print(" Color Space: %d" % segment.color_space)
        print(" Bits per Sample: %d" % segment.bits_per_sample)
        print(" Compression Type: %d" % segment.compression_type)
        print(" Resolution Units: %d" % segment.resolution_units)
        print(" Vertical Resolution: %d" % segment.vertical_resoution)
        print(" Horizontal Resolution: %d" % segment.horizontal_resolution)
    elif isinstance(segment, pyjpeg.ExifHeader):
        print("APP%d EXIF" % segment.n)
        print(" Data: %r" % segment.data)
    elif isinstance(segment, pyjpeg.AdobeHeader):
        print("APP%d Adobe" % segment.n)
        print(" Version: %d" % segment.version)
        print(" Flags 0: %04x" % segment.flags0)
        print(" Flags 1: %04x" % segment.flags1)
        print(
            " Colorspace: %s"
            % {
                pyjpeg.AdobeColorSpace.RGB_OR_CMYK: "RGB or CMYK",
                pyjpeg.AdobeColorSpace.Y_CB_CR: "YCbCr",
                pyjpeg.AdobeColorSpace.Y_CB_CR_K: "YCbCrK",
            }.get(segment.color_space, "%d" % segment.color_space)
        )
    elif isinstance(segment, pyjpeg.UnknownApplicationSpecificData):
        print("APP%d Application Specific Data" % segment.n)
        s = " Data: "
        for d in segment.data:
            s += "%02X" % d
        print(s)
    elif isinstance(segment, pyjpeg.Comment):
        print("COM Comment")
        print(" Data: %s" % repr(segment.data))
    elif isinstance(segment, pyjpeg.DefineQuantizationTables):
        print("DQT Define Quantization Tables")
        for quantization_table in segment.tables:
            print(" Table %d:" % quantization_table.destination)
            print("  Precision: %d bits" % quantization_table.precision)
            print_data_unit(quantization_table.values)
    elif isinstance(segment, pyjpeg.DefineHuffmanTables):
        print("DHT Define Huffman Tables")
        for huffman_table in segment.tables:
            print(
                " %s Table %d:"
                % (
                    {
                        pyjpeg.HuffmanTableClass.DC: "DC",
                        pyjpeg.HuffmanTableClass.AC: "AC",
                    }[huffman_table.table_class],
                    huffman_table.destination,
                )
            )
            for i, symbols in enumerate(huffman_table.table):
                if len(symbols) > 0:
                    s = "  Symbols of length %d:" % (i + 1)
                    for symbol in symbols:
                        s += " %02x" % symbol
                    print(s)
    elif isinstance(segment, pyjpeg.DefineArithmeticConditioning):
        print("DAC Define Arithmetic Conditioning")
        for conditioning in segment.tables:
            print(
                " %s Table %d: %s"
                % (
                    {
                        pyjpeg.ArithmeticConditioningTableClass.DC: "DC",
                        pyjpeg.ArithmeticConditioningTableClass.AC: "AC",
                    }[conditioning.table_class],
                    conditioning.destination,
                    repr(conditioning.value),
                )
            )
    elif isinstance(segment, pyjpeg.DefineRestartInterval):
        print("DRI Define Restart Interval")
        print(" Restart interval: %d" % segment.restart_interval)
    elif isinstance(segment, pyjpeg.ExpandReferenceComponents):
        print("EXP Expand Reference Components")
        print(
            " Expand Horizontal: %s"
            % {False: "No", True: "Yes"}[segment.expand_horizontal != 0]
        )
        print(
            " Expand Vertical: %s"
            % {False: "No", True: "Yes"}[segment.expand_vertical != 0]
        )
    elif isinstance(segment, pyjpeg.StartOfFrame):
        is_lossless = segment.n in (3, 7, 11, 15)
        is_ls = segment.n == 55
        print(
            "SOF%d Start of Frame, %s"
            % (
                segment.n,
                {
                    pyjpeg.FrameType.BASELINE: "Baseline DCT",
                    pyjpeg.FrameType.EXTENDED_HUFFMAN: "Extended sequential DCT, Huffman coding",
                    pyjpeg.FrameType.PROGRESSIVE_HUFFMAN: "Progressive DCT, Huffman coding",
                    pyjpeg.FrameType.LOSSLESS_HUFFMAN: "Lossless (sequential), Huffman coding",
                    pyjpeg.FrameType.DIFFERENTIAL_SEQUENTIAL_HUFFMAN: "Differential sequential DCT, Huffman coding",
                    pyjpeg.FrameType.DIFFERENTIAL_PROGRESSIVE_HUFFMAN: "Differential progressive DCT, Huffman coding",
                    pyjpeg.FrameType.DIFFERENTIAL_LOSSLESS_HUFFMAN: "Differential lossless (sequential), Huffman coding",
                    pyjpeg.FrameType.EXTENDED_ARITHMETIC: "Extended sequential DCT, Arithmetic coding",
                    pyjpeg.FrameType.PROGRESSIVE_ARITHMETIC: "Progressive DCT, Arithmetic coding",
                    pyjpeg.FrameType.LOSSLESS_ARITHMETIC: "Lossless (sequential), Arithmetic coding",
                    pyjpeg.FrameType.DIFFERENTIAL_SEQUENTIAL_ARITHMETIC: "Differential sequential DCT, Arithmetic coding",
                    pyjpeg.FrameType.DIFFERENTIAL_PROGRESSIVE_ARITHMETIC: "Differential progressive DCT, Arithmetic coding",
                    pyjpeg.FrameType.DIFFERENTIAL_LOSSLESS_ARITHMETIC: "Differential lossless (sequential), Arithmetic coding",
                    pyjpeg.FrameType.LS: "JPEG-LS",
                }[segment.n],
            )
        )
        print(" Precision: %d bits" % segment.precision)
        print(
            " Number of lines: %d" % segment.number_of_lines
        )  # FIXME: Note if zero defined later
        print(" Number of samples per line: %d" % segment.samples_per_line)
        for frame_component in segment.components:
            print(" Component:")
            print("  Id: %d" % frame_component.id)
            print(
                "  Sampling Factor: %dx%d"
                % (
                    frame_component.sampling_factor[0],
                    frame_component.sampling_factor[1],
                )
            )
            if not is_lossless and not is_ls:
                print(
                    "  Quantization Table: %d"
                    % frame_component.quantization_table_index
                )
    elif isinstance(segment, pyjpeg.StartOfScan):
        print("SOS Start of Scan")
        for scan_component in segment.components:
            print(" Component:")
            print("  Id: %d" % scan_component.component_selector)
            if is_ls:
                print("  Mapping table: %d" % scan_component.get_mapping_table())
            else:
                print("  DC Table: %d" % scan_component.dc_table)
                if not is_lossless:
                    print("  AC Table: %d" % scan_component.ac_table)
        if is_lossless:
            print(" Predictor: %d" % segment.spectral_selection[0])
        elif is_ls:
            print(" Near: %d" % segment.spectral_selection[0])
            print(
                " Interleave Mode: %s"
                % {
                    pyjpeg.LSInterleaveMode.NONE: "None",
                    pyjpeg.LSInterleaveMode.LINE: "Line",
                    pyjpeg.LSInterleaveMode.SAMPLE: "Sample",
                }.get(
                    segment.spectral_selection[1], "%d" % segment.spectral_selection[1]
                )
            )
        else:
            print(
                " Spectral Selection: %d-%d"
                % (segment.spectral_selection[0], segment.spectral_selection[1])
            )

        if (segment.point_transform & 0xF0) != 0:
            print(" Previous Point Transform: %d" % (segment.point_transform >> 4))
        print(" Point Transform: %d" % (segment.point_transform & 0xF))
    elif isinstance(segment, pyjpeg.HuffmanDCTScan) or isinstance(
        segment, pyjpeg.ArithmeticDCTScan
    ):
        for data_unit in segment.data_units:
            print_data_unit(data_unit)
    elif (
        isinstance(segment, pyjpeg.HuffmanLosslessScan)
        or isinstance(segment, pyjpeg.ArithmeticLosslessScan)
        or isinstance(segment, pyjpeg.LSScan)
    ):
        s = " Samples:"
        for sample in segment.samples:
            s += " %d" % sample
        print(s)
    elif isinstance(segment, pyjpeg.Restart):
        print("RST%d Restart" % segment.index)
    elif isinstance(segment, pyjpeg.DefineNumberOfLines):
        print("DNL Define Number of Lines")
        print(" Number of lines: %d" % segment.number_of_lines)
    elif isinstance(segment, pyjpeg.LSCodingParameters):
        print("LSE Coding Parameters")
        print(" Maximum value: %d" % segment.maxval)
        print(
            " Gradient thresholds: %d, %d, %d"
            % (
                segment.gradient_thresholds[0],
                segment.gradient_thresholds[1],
                segment.gradient_thresholds[2],
            )
        )
        print(" Reset: %d" % segment.reset)
    elif isinstance(segment, pyjpeg.LSMappingTable):
        print("LSE Mapping Table")
        print(" Table ID: %d" % segment.table_id)
        print(" Weight: %d" % segment.weight)
        print(" Table Data: %s" % segment.table.hex())
    elif isinstance(segment, pyjpeg.LSMappingTableContinuation):
        print("LSE Mapping Table Continuation")
        print(" Table ID: %d" % segment.table_id)
        print(" Weight: %d" % segment.weight)
        print(" Table Data: %s" % segment.table.hex())
    elif isinstance(segment, pyjpeg.LSOversizeImageDimensions):
        print("LSE Oversize Image Dimensions")
        print(" Number of lines: %d" % segment.number_of_lines)
        print(" Number of samples per line: %d" % segment.samples_per_line)
    elif isinstance(segment, pyjpeg.LSUnknownPresetParameters):
        print("LSE Preset Parameters")
        print(" ID: %d" % segment.id)
        print(" Data: %s" % segment.data.hex())
    elif isinstance(segment, pyjpeg.EndOfImage):
        print("EOI End of Image")
    else:
        print(segment)
