📄 Colors/src/MmcqQuantizer.cs
using System;
using System.Collections.Generic;
using System.Linq;
using SixLabors.ImageSharp;
using SixLabors.ImageSharp.PixelFormats;

namespace Mosaic.Colors;

internal static class MmcqQuantizer
{
    internal const int SignificantBits = 5;
    internal const int Shift = 8 - SignificantBits;
    internal const int HistogramSize = 1 << (3 * SignificantBits);
    internal const int VboxLength = 1 << SignificantBits;
    internal const int MaxValue = 256 >> Shift;
    internal const int Multiplier = 1 << Shift;
    internal const int MaxIterations = 1000;

    extension(Image<Rgba32> image)
    {
        internal Dictionary<Rgba32, int> Quantize(int maxColorCount)
        {
            if (maxColorCount < 2 || 256 < maxColorCount)
            {
                throw new ArgumentOutOfRangeException(
                    paramName: nameof(maxColorCount),
                    message: "Should be between 2 and 256"
                );
            }

            var histogram = image.GetHistogram();

            var vBox = image.GetVbox(histogram);
            var queue = new PriorityQueue<VBox, VBox>(VBox.CountComparer);
            queue.Enqueue(vBox, vBox);

            Iterate(queue, (int)float.Ceiling(maxColorCount * 0.75f), histogram);

            var queue2 = new PriorityQueue<VBox, VBox>(queue.UnorderedItems, VBox.CountVolumeComparer);
            Iterate(queue2, maxColorCount, histogram);

            return queue2.UnorderedItems.ToDictionary(p => p.Element.Average, p => p.Element.Count);
        }

        internal Memory<int> GetHistogram()
        {
            var builder = new HistogramBuilder();
            image.ProcessPixelRows(builder.Populate);
            return builder.Histogram;
        }

        internal VBox GetVbox(ReadOnlyMemory<int> histogram)
        {
            var builder = new VboxBuilder();
            image.ProcessPixelRows(builder.Populate);
            return builder.Build(histogram);
        }
    }

    internal static void Iterate(PriorityQueue<VBox, VBox> queue, int target, ReadOnlyMemory<int> histogram)
    {
        for (var i = 0; i < MaxIterations; i++)
        {
            if (queue.Count >= target)
            {
                return;
            }

            var vbox = queue.Dequeue();
            if (vbox.Count is 0)
            {
                queue.Enqueue(vbox, vbox);
                continue;
            }

            var result = MedianCutApply(histogram, vbox);
            if (result is not [var next, ..])
            {
                queue.Enqueue(vbox, vbox);
                return;
            }

            queue.Enqueue(next, next);
            if (result is [_, var nextNext, ..])
            {
                queue.Enqueue(nextNext, nextNext);
            }
        }
    }

    internal static VBox[] MedianCutApply(ReadOnlyMemory<int> histogram, VBox vBox)
    {
        if (vBox.Count is 0)
        {
            return [];
        }

        if (vBox.Count is 1)
        {
            return [vBox];
        }

        var rSize = vBox.RMax - vBox.RMin + 1;
        var gSize = vBox.GMax - vBox.GMin + 1;
        var bSize = vBox.BMax - vBox.BMin + 1;
        var maxSize = int.Max(rSize, int.Max(gSize, bSize));

        var total = 0;
        var partialSum = new int[VboxLength];
        var lookAheadSum = new int[VboxLength];
        for (var i = 0; i < VboxLength; i++)
        {
            partialSum[i] = -1;
            lookAheadSum[i] = -1;
        }

        if (maxSize == rSize)
        {
            var span = histogram.Span;
            for (var r = vBox.RMin; r <= vBox.RMax; r++)
            {
                var sum = 0;
                for (var g = vBox.GMin; g <= vBox.GMax; g++)
                for (var b = vBox.BMin; b <= vBox.BMax; b++)
                {
                    sum += span[GetColorIndex(r, g, b)];
                }
                total += sum;
                partialSum[r] = total;
            }
        }
        else if (maxSize == gSize)
        {
            var span = histogram.Span;
            for (var g = vBox.GMin; g <= vBox.GMax; g++)
            {
                var sum = 0;
                for (var r = vBox.RMin; r <= vBox.RMax; r++)
                for (var b = vBox.BMin; b <= vBox.BMax; b++)
                {
                    sum += span[GetColorIndex(r, g, b)];
                }
                total += sum;
                partialSum[g] = total;
            }
        }
        else
        {
            var span = histogram.Span;
            for (var b = vBox.BMin; b <= vBox.BMax; b++)
            {
                var sum = 0;
                for (var r = vBox.RMin; r <= vBox.RMax; r++)
                for (var g = vBox.GMin; g <= vBox.GMax; g++)
                {
                    sum += span[GetColorIndex(r, g, b)];
                }
                total += sum;
                partialSum[b] = total;
            }
        }

        for (var i = 0; i < VboxLength; i++)
        {
            if (partialSum[i] >= 0)
            {
                lookAheadSum[i] = total - partialSum[i];
            }
        }

        if (maxSize == rSize)
        {
            return Cut(vBox.RMin, vBox.RMax) is (var left, var right)
                ?
                [
                    new(vBox.RMin, left, vBox.GMin, vBox.GMax, vBox.BMin, vBox.BMax, vBox.Histogram),
                    new(right, vBox.RMax, vBox.GMin, vBox.GMax, vBox.BMin, vBox.BMax, vBox.Histogram),
                ]
                : [];
        }
        else if (maxSize == gSize)
        {
            return Cut(vBox.GMin, vBox.GMax) is (var left, var right)
                ?
                [
                    new(vBox.RMin, vBox.RMax, vBox.GMin, left, vBox.BMin, vBox.BMax, vBox.Histogram),
                    new(vBox.RMin, vBox.RMax, right, vBox.GMax, vBox.BMin, vBox.BMax, vBox.Histogram),
                ]
                : [];
        }
        else
        {
            return Cut(vBox.BMin, vBox.BMax) is (var left, var right)
                ?
                [
                    new(vBox.RMin, vBox.RMax, vBox.GMin, vBox.GMax, vBox.BMin, left, vBox.Histogram),
                    new(vBox.RMin, vBox.RMax, vBox.GMin, vBox.GMax, right, vBox.BMax, vBox.Histogram),
                ]
                : [];
        }

        (int, int)? Cut(int min, int max)
        {
            for (var i = min; i <= max; i++)
            {
                if (partialSum[i] < total / 2)
                {
                    continue;
                }

                var left = i - min;
                var right = max - i;
                var d2 = left <= right ? int.Min(max - 1, i + right / 2) : int.Max(min, i - 1 - left / 2);

                while (d2 < 0 || partialSum[d2] <= 0)
                {
                    d2++;
                }
                for (
                    var count2 = lookAheadSum[d2];
                    count2 <= 0 && d2 > min && partialSum[d2 - 1] > 0;
                    count2 = lookAheadSum[d2]
                )
                {
                    d2--;
                }

                return (d2, d2 + 1);
            }

            return null;
        }
    }

    internal static int GetColorIndex(int r, int g, int b) => (r << (2 * SignificantBits)) + (g << SignificantBits) + b;
}