#define ACCENT_SCORE_IMPLEMENTATION
#include "../vendor/accent_score.h"

#include "candidates.h"

#include <math.h>
#include <stdlib.h>
#include <string.h>

#define DEFAULT_MERGE_DISTANCE 0.035
#define CHROMATIC_POOL_FRACTION 2U
#define MIN_CHROMA 0.04

void candidates_init(CandidateSet *set) {
    set->items = NULL;
    set->count = 0;
    set->capacity = 0;
    set->merge_distance = DEFAULT_MERGE_DISTANCE;
    arena_init(&set->arena);
}

void candidates_destroy(CandidateSet *set) {
    arena_destroy(&set->arena);
    set->items = NULL;
    set->count = 0;
    set->capacity = 0;
    set->merge_distance = DEFAULT_MERGE_DISTANCE;
}

static bool grow(CandidateSet *set) {
    size_t new_capacity;
    ColourCluster *items;
    if (set->count < set->capacity) {
        return true;
    }
    new_capacity = set->capacity == 0 ? 32 : set->capacity * 2;
    if (new_capacity < set->capacity ||
        new_capacity > SIZE_MAX / sizeof(*set->items)) {
        return false;
    }
    items = arena_allocate(&set->arena, new_capacity, sizeof(*items));
    if (items == NULL) {
        return false;
    }
    if (set->items != NULL) {
        memcpy(items, set->items, set->count * sizeof(*items));
    }
    set->items = items;
    set->capacity = new_capacity;
    return true;
}

bool candidates_add(CandidateSet *set, RGB colour, double prevalence) {
    const Oklab lab = colour_to_oklab(colour);
    size_t nearest = SIZE_MAX;
    double nearest_distance = set->merge_distance;
    size_t index;

    if (!(prevalence > 0.0) || !isfinite(prevalence)) {
        return true;
    }
    for (index = 0; index < set->count; ++index) {
        const double distance = colour_oklab_distance(lab, set->items[index].lab);
        if (distance < nearest_distance) {
            nearest = index;
            nearest_distance = distance;
        }
    }
    if (nearest != SIZE_MAX) {
        ColourCluster *cluster = &set->items[nearest];
        const double red = colour_srgb_to_linear((double)colour.r / 255.0);
        const double green = colour_srgb_to_linear((double)colour.g / 255.0);
        const double blue = colour_srgb_to_linear((double)colour.b / 255.0);
        cluster->red_sum += red * prevalence;
        cluster->green_sum += green * prevalence;
        cluster->blue_sum += blue * prevalence;
        cluster->prevalence += prevalence;
        cluster->colour = colour_from_linear(
            cluster->red_sum / cluster->prevalence,
            cluster->green_sum / cluster->prevalence,
            cluster->blue_sum / cluster->prevalence);
        cluster->lab = colour_to_oklab(cluster->colour);
        return true;
    }
    if (!grow(set)) {
        return false;
    }
    set->items[set->count].colour = colour;
    set->items[set->count].lab = lab;
    set->items[set->count].red_sum =
        colour_srgb_to_linear((double)colour.r / 255.0) * prevalence;
    set->items[set->count].green_sum =
        colour_srgb_to_linear((double)colour.g / 255.0) * prevalence;
    set->items[set->count].blue_sum =
        colour_srgb_to_linear((double)colour.b / 255.0) * prevalence;
    set->items[set->count].prevalence = prevalence;
    ++set->count;
    return true;
}

static int compare_prevalence(const void *left, const void *right) {
    const ScoredCandidate *a = left;
    const ScoredCandidate *b = right;
    if (a->prevalence > b->prevalence) return -1;
    if (a->prevalence < b->prevalence) return 1;
    if (a->colour.r != b->colour.r) return (int)a->colour.r - (int)b->colour.r;
    if (a->colour.g != b->colour.g) return (int)a->colour.g - (int)b->colour.g;
    return (int)a->colour.b - (int)b->colour.b;
}

static int compare_score(const void *left, const void *right) {
    const ScoredCandidate *a = left;
    const ScoredCandidate *b = right;
    if (a->score > b->score) return -1;
    if (a->score < b->score) return 1;
    return compare_prevalence(left, right);
}

static double candidate_chroma(const ScoredCandidate *candidate) {
    const Oklab lab = colour_to_oklab(candidate->colour);
    return sqrt(lab.a * lab.a + lab.b * lab.b);
}

static bool is_chromatic(const ScoredCandidate *candidate) {
    return candidate_chroma(candidate) >= MIN_CHROMA;
}

static bool candidate_is_in_pool(const ScoredCandidate *candidate,
                                 const ScoredCandidate *ranked,
                                 size_t pool_count) {
    size_t index;
    for (index = 0; index < pool_count; ++index) {
        if (candidate->colour.r == ranked[index].colour.r &&
            candidate->colour.g == ranked[index].colour.g &&
            candidate->colour.b == ranked[index].colour.b) {
            return true;
        }
    }
    return false;
}

static size_t least_prevalent_neutral(const ScoredCandidate *ranked,
                                      size_t pool_count) {
    size_t index;
    for (index = pool_count; index > 0; --index) {
        if (!is_chromatic(&ranked[index - 1U])) {
            return index - 1U;
        }
    }
    return SIZE_MAX;
}

static void ensure_chromatic_representation(ScoredCandidate *ranked,
                                            size_t ranked_count,
                                            size_t pool_count) {
    const size_t target = (pool_count + CHROMATIC_POOL_FRACTION - 1U) /
                          CHROMATIC_POOL_FRACTION;
    const size_t prevalence_target = (target + 1U) / 2U;
    size_t chromatic_count = 0;
    size_t scan;
    size_t index;

    for (index = 0; index < pool_count; ++index) {
        if (is_chromatic(&ranked[index])) {
            ++chromatic_count;
        }
    }
    for (scan = pool_count;
         scan < ranked_count && chromatic_count < prevalence_target;
         ++scan) {
        size_t replace;
        if (!is_chromatic(&ranked[scan])) {
            continue;
        }
        replace = least_prevalent_neutral(ranked, pool_count);
        if (replace == SIZE_MAX) {
            break;
        }
        ranked[replace] = ranked[scan];
        ++chromatic_count;
    }
    while (chromatic_count < target) {
        size_t best = SIZE_MAX;
        double best_chroma = 0.0;
        const size_t replace = least_prevalent_neutral(ranked, pool_count);
        if (replace == SIZE_MAX) {
            break;
        }
        for (scan = pool_count; scan < ranked_count; ++scan) {
            const double chroma = candidate_chroma(&ranked[scan]);
            if (chroma >= MIN_CHROMA && chroma > best_chroma &&
                !candidate_is_in_pool(&ranked[scan], ranked, pool_count)) {
                best = scan;
                best_chroma = chroma;
            }
        }
        if (best == SIZE_MAX) {
            break;
        }
        ranked[replace] = ranked[best];
        ++chromatic_count;
    }
}

bool candidates_rank(CandidateSet *set, size_t limit, RGB background,
                     float minimum_score, ScoredCandidate **results,
                     size_t *result_count) {
    ScoredCandidate *ranked;
    size_t pool_count;
    size_t index;
    size_t kept = 0;

    *results = NULL;
    *result_count = 0;
    if (set->count == 0 || limit == 0) {
        return true;
    }
    ranked = arena_allocate(&set->arena, set->count, sizeof(*ranked));
    if (ranked == NULL) {
        return false;
    }
    for (index = 0; index < set->count; ++index) {
        ranked[index].colour = set->items[index].colour;
        ranked[index].prevalence = set->items[index].prevalence;
        ranked[index].score = 0.0f;
    }
    qsort(ranked, set->count, sizeof(*ranked), compare_prevalence);
    pool_count = set->count < limit ? set->count : limit;
    ensure_chromatic_representation(ranked, set->count, pool_count);
    for (index = 0; index < pool_count; ++index) {
        const float score = accent_score(ranked[index].colour, background);
        if (score >= minimum_score) {
            ranked[kept] = ranked[index];
            ranked[kept].score = score;
            ++kept;
        }
    }
    qsort(ranked, kept, sizeof(*ranked), compare_score);
    if (kept == 0) {
        return true;
    }
    *results = ranked;
    *result_count = kept;
    return true;
}

void candidates_tweak(ScoredCandidate *results, size_t result_count,
                      RGB background) {
    size_t index;
    for (index = 0; index < result_count; ++index) {
        const RGB tweaked = colour_tweak_accent(results[index].colour,
                                                background);
        const float tweaked_score = accent_score(tweaked, background);
        if (tweaked_score >= results[index].score) {
            results[index].colour = tweaked;
            results[index].score = tweaked_score;
        }
    }
    if (result_count > 1) {
        qsort(results, result_count, sizeof(*results), compare_score);
    }
}