/* C port of orlp/pdqsort (commit b1ef26a, 2021-03-14) for int keys.
 * Only the branchy partition_right is ported; the comparison sequence is identical to
 * pdqsort.h with a non-default comparator, which bench.cpp checks.
 * The heapsort fallback mirrors libstdc++'s make_heap + sort_heap for the same reason. */
#include <stdbool.h>
#include <string.h>
#include "pdq.h"

enum {
    INSERTION_SORT_THRESHOLD = 24,
    NINTHER_THRESHOLD = 128,
    PARTIAL_INSERTION_SORT_LIMIT = 8,
};

static pdq_stats S;

static inline bool lt(int a, int b)
{
    S.comparisons++;
    return a < b;
}

static inline void swap(int *a, int *b)
{
    int t = *a;
    *a = *b;
    *b = t;
}

static int log2_floor(size_t n)
{
    int log = 0;
    while (n >>= 1) ++log;
    return log;
}

static void insertion_sort(int *begin, int *end)
{
    if (begin == end) return;
    for (int *cur = begin + 1; cur != end; ++cur) {
        int *sift = cur, *sift_1 = cur - 1;
        if (lt(*sift, *sift_1)) {
            int tmp = *sift;
            do { *sift-- = *sift_1; } while (sift != begin && lt(tmp, *--sift_1));
            *sift = tmp;
        }
    }
}

/* Requires begin[-1] <= every element of [begin, end): it acts as the sentinel. */
static void unguarded_insertion_sort(int *begin, int *end)
{
    if (begin == end) return;
    for (int *cur = begin + 1; cur != end; ++cur) {
        int *sift = cur, *sift_1 = cur - 1;
        if (lt(*sift, *sift_1)) {
            int tmp = *sift;
            do { *sift-- = *sift_1; } while (lt(tmp, *--sift_1));
            *sift = tmp;
        }
    }
}

/* Insertion sort that gives up after moving more than PARTIAL_INSERTION_SORT_LIMIT elements. */
static bool partial_insertion_sort(int *begin, int *end)
{
    if (begin == end) return true;
    size_t moves = 0;
    for (int *cur = begin + 1; cur != end; ++cur) {
        int *sift = cur, *sift_1 = cur - 1;
        if (lt(*sift, *sift_1)) {
            int tmp = *sift;
            do { *sift-- = *sift_1; } while (sift != begin && lt(tmp, *--sift_1));
            *sift = tmp;
            moves += (size_t)(cur - sift);
        }
        if (moves > PARTIAL_INSERTION_SORT_LIMIT) return false;
    }
    return true;
}

static inline void sort2(int *a, int *b)
{
    if (lt(*b, *a)) swap(a, b);
}

static inline void sort3(int *a, int *b, int *c)
{
    sort2(a, b);
    sort2(b, c);
    sort2(a, b);
}

/* Pivot is *begin. Elements equal to the pivot go right. Returns the pivot's final position;
 * *already is set when no pair had to be swapped. */
static int *partition_right(int *begin, int *end, bool *already)
{
    int pivot = *begin;
    int *first = begin, *last = end;

    while (lt(*++first, pivot));
    if (first - 1 == begin)
        while (first < last && !lt(*--last, pivot));
    else
        while (!lt(*--last, pivot));

    *already = first >= last;
    while (first < last) {
        swap(first, last);
        while (lt(*++first, pivot));
        while (!lt(*--last, pivot));
    }

    int *pivot_pos = first - 1;
    *begin = *pivot_pos;
    *pivot_pos = pivot;
    return pivot_pos;
}

/* Pivot is *begin. Elements equal to the pivot go left. */
static int *partition_left(int *begin, int *end)
{
    int pivot = *begin;
    int *first = begin, *last = end;

    while (lt(pivot, *--last));
    if (last + 1 == end)
        while (first < last && !lt(pivot, *++first));
    else
        while (!lt(pivot, *++first));

    while (first < last) {
        swap(first, last);
        while (lt(pivot, *--last));
        while (!lt(pivot, *++first));
    }

    int *pivot_pos = last;
    *begin = *pivot_pos;
    *pivot_pos = pivot;
    return pivot_pos;
}

static void adjust_heap(int *a, ptrdiff_t hole, ptrdiff_t len, int value)
{
    const ptrdiff_t top = hole;
    ptrdiff_t child = hole;
    while (child < (len - 1) / 2) {
        child = 2 * (child + 1);
        if (lt(a[child], a[child - 1])) child--;
        a[hole] = a[child];
        hole = child;
    }
    if ((len & 1) == 0 && child == (len - 2) / 2) {
        child = 2 * (child + 1);
        a[hole] = a[child - 1];
        hole = child - 1;
    }
    ptrdiff_t parent = (hole - 1) / 2;
    while (hole > top && lt(a[parent], value)) {
        a[hole] = a[parent];
        hole = parent;
        parent = (hole - 1) / 2;
    }
    a[hole] = value;
}

static void heap_sort(int *begin, int *end)
{
    ptrdiff_t len = end - begin;
    if (len < 2) return;
    for (ptrdiff_t parent = (len - 2) / 2;; parent--) {
        adjust_heap(begin, parent, len, begin[parent]);
        if (parent == 0) break;
    }
    for (ptrdiff_t n = len - 1; n > 0; n--) {
        int value = begin[n];
        begin[n] = begin[0];
        adjust_heap(begin, 0, n, value);
    }
}

static void pdq_loop(int *begin, int *end, int bad_allowed, bool leftmost,
                     unsigned long long depth)
{
    if (depth > S.max_depth) S.max_depth = depth;

    for (;;) {
        ptrdiff_t size = end - begin;

        if (size < INSERTION_SORT_THRESHOLD) {
            if (leftmost) insertion_sort(begin, end);
            else unguarded_insertion_sort(begin, end);
            return;
        }

        /* median of 3, or Tukey's ninther; the pivot ends up in *begin */
        ptrdiff_t s2 = size / 2;
        if (size > NINTHER_THRESHOLD) {
            sort3(begin, begin + s2, end - 1);
            sort3(begin + 1, begin + (s2 - 1), end - 2);
            sort3(begin + 2, begin + (s2 + 1), end - 3);
            sort3(begin + (s2 - 1), begin + s2, begin + (s2 + 1));
            swap(begin, begin + s2);
        } else {
            sort3(begin + s2, begin, end - 1);
        }

        /* begin[-1] is the pivot of an ancestor and <= everything here. If the new pivot
         * equals it, [begin, pivot] is all equal: skip it without recursing. */
        if (!leftmost && !lt(begin[-1], *begin)) {
            S.partition_left++;
            begin = partition_left(begin, end) + 1;
            continue;
        }

        bool already_partitioned;
        int *pivot_pos = partition_right(begin, end, &already_partitioned);
        S.partitions++;

        ptrdiff_t l_size = pivot_pos - begin;
        ptrdiff_t r_size = end - (pivot_pos + 1);
        bool highly_unbalanced = l_size < size / 8 || r_size < size / 8;

        if (highly_unbalanced) {
            S.bad_partitions++;
            if (--bad_allowed == 0) {
                S.heapsort_calls++;
                S.heapsort_elems += (unsigned long long)size;
                heap_sort(begin, end);
                return;
            }
            /* deterministic pattern breaking: swap pivot candidates with elements
             * around the 25% and 75% positions of each side */
            if (l_size >= INSERTION_SORT_THRESHOLD) {
                swap(begin, begin + l_size / 4);
                swap(pivot_pos - 1, pivot_pos - l_size / 4);
                if (l_size > NINTHER_THRESHOLD) {
                    swap(begin + 1, begin + (l_size / 4 + 1));
                    swap(begin + 2, begin + (l_size / 4 + 2));
                    swap(pivot_pos - 2, pivot_pos - (l_size / 4 + 1));
                    swap(pivot_pos - 3, pivot_pos - (l_size / 4 + 2));
                }
            }
            if (r_size >= INSERTION_SORT_THRESHOLD) {
                swap(pivot_pos + 1, pivot_pos + (1 + r_size / 4));
                swap(end - 1, end - r_size / 4);
                if (r_size > NINTHER_THRESHOLD) {
                    swap(pivot_pos + 2, pivot_pos + (2 + r_size / 4));
                    swap(pivot_pos + 3, pivot_pos + (3 + r_size / 4));
                    swap(end - 2, end - (1 + r_size / 4));
                    swap(end - 3, end - (2 + r_size / 4));
                }
            }
        } else if (already_partitioned
                   && partial_insertion_sort(begin, pivot_pos)
                   && partial_insertion_sort(pivot_pos + 1, end)) {
            S.partial_ok++;
            return;
        }

        /* recurse on the left part, loop on the right part */
        pdq_loop(begin, pivot_pos, bad_allowed, leftmost, depth + 1);
        begin = pivot_pos + 1;
        leftmost = false;
    }
}

void pdq_sort(int *a, size_t n, pdq_stats *stats)
{
    memset(&S, 0, sizeof S);
    if (n > 0) pdq_loop(a, a + n, log2_floor(n), true, 1);
    if (stats) *stats = S;
}
