/*
 * Basic support for C unit tests.
 *
 * This file is part of Beastie <https://purl.org/nxg/dist/beastie>
 * SPDX-FileCopyrightText: 2024 Norman Gray <https://nxg.me.uk>
 * SPDX-License-Identifier: BSD-2-Clause
 */

#include <stdio.h>
#include <setjmp.h>

#include "c-unit.h"

static jmp_buf jmp_env;
int ntests = 0;
int nsuites = 0;
int total_fails = 0;

#define RED_TEXT(s) "\033[41;37m " s " \033[0m"
#define GREEN_TEXT(s) "\033[42;30m " s " \033[0m"


// compare signed/unsigned ints
void assert_equal_int_fn(const char* label, int actual, int expected)
{
    ntests++;
    if (actual != expected) {
        fprintf(stderr,
                "Test %s " RED_TEXT("failed") "\n  " RED_TEXT("actual")
                ":   %d\n  " GREEN_TEXT("expected") ": %d\n",
                label, actual, expected);
        longjmp(jmp_env, 1);
    }
}

void assert_equal_uint_fn(const char* label, unsigned int actual, unsigned int expected)
{
    ntests++;
    if (actual != expected) {
        fprintf(stderr,
                "Test %s " RED_TEXT("failed") "\n  " RED_TEXT("actual")
                ":   %d\n  " GREEN_TEXT("expected") ": %d\n",
                label, actual, expected);
        longjmp(jmp_env, 1);
    }
}


#define NOTCOMPARE_ARRAYS                                                  \
    ntests++;                                                           \
    for (int i=0; i<n; i++) {                                           \
        if (actual[i] != expected[i]) {                                 \
            fprintf(stderr,                                             \
                    "Test %s " RED_TEXT("failed") " at i=%d\n  " RED_TEXT("actual") \
                    ":   %u\n  " GREEN_TEXT("expected") ": %u\n",       \
                    label, i, actual[i], expected[i]);                  \
            longjmp(jmp_env, 1);                                        \
        }                                                               \
    }

#define COMPARE_ARRAYS                                                  \
    ntests++;                                                           \
    for (int i=0; i<n; i++) {                                           \
        if (actual[i] != expected[i]) {                                 \
            fprintf(stderr,                                             \
                    "Test %s " RED_TEXT("failed") " at i=%d\n",         \
                    label, i);                                          \
            fprintf(stderr, "  " RED_TEXT("actual") ":   [");           \
            for (int ii=0; ii<n; ii++) fprintf(stderr, " %x", actual[ii]); \
            fprintf(stderr, " ]\n  " GREEN_TEXT("expected") ": [");     \
            for (int ii=0; ii<n; ii++) fprintf(stderr, " %x", expected[ii]); \
            fprintf(stderr, " ]\n");                                    \
            longjmp(jmp_env, 1);                                        \
        }                                                               \
    }

void assert_equal_int32_array_fn(const char* label,
                                 const int32_t* actual,
                                 const int32_t* expected,
                                 size_t n)
{
    COMPARE_ARRAYS;
}

void assert_equal_uint32_array_fn(const char* label,
                                  const uint32_t* actual,
                                  const uint32_t* expected,
                                  size_t n)
{
    COMPARE_ARRAYS;
}

void assert_equal_uint16_array_fn(const char* label,
                                  const uint16_t* actual,
                                  const uint16_t* expected,
                                  size_t n)
{
    COMPARE_ARRAYS;
}

void assert_equal_uint8_array_fn(const char* label,
                                 const uint8_t* actual,
                                 const uint8_t* expected,
                                 size_t n)
{
    COMPARE_ARRAYS;
}

// compare a string of specified length with a null-terminated string
void assert_equal_sn_fn(const char* label, // used when reporting
                        const unsigned char* actual, int actual_len, // test-string+length
                        const char* expected) // expected value
{
    ntests++;

    if (actual == NULL || expected == NULL) {
        if (! (actual == NULL && expected == NULL)) {
            fprintf(stderr, "Test %s " RED_TEXT("failed") "\n  " RED_TEXT("actual"), label);
            if (actual == NULL) {
                fprintf(stderr, ":   <NULL>\n  ");
            } else {
                fprintf(stderr, ":   %.*s\n  ", actual_len, actual);
            }
            fprintf(stderr, GREEN_TEXT("expected") ": %s\n",
                    (expected == NULL ? "<NULL>" : expected));

            longjmp(jmp_env, 1);
        }
        return;
    }

    int ok = 1;
    for (int i=0; i<actual_len; i++) {
        //fprintf(stderr, "%d: %x ?= %x\n", i, actual[i], (unsigned char)expected[i]);
        if (actual[i] != (unsigned char)expected[i]) {
            ok = 0;
            break;
        }
    }
    if (expected[actual_len] != '\0') ok = 0;

    if (! ok) {
        fprintf(stderr, "Test %s " RED_TEXT("failed") "\n  " RED_TEXT("actual") ":   ", label);
        for (int i=0; i<actual_len; i++) fprintf(stderr, " %x", actual[i]);
        fprintf(stderr, "\n  " GREEN_TEXT("expected") ": ");
        for (int i=0; i<4; i++) {
            if (expected[i] == '\0') {
                break;
            } else {
                fprintf(stderr, " %x", (unsigned char)expected[i]);
            }
        }
        fprintf(stderr, "\n");
        longjmp(jmp_env, 1);
    }
}

void assert_equal_ptr_fn(const char* label,
                         const void* actual,
                         const void* expected)
{
    ntests++;
    if (actual != expected) {
        fprintf(stderr,
                "Test %s " RED_TEXT("failed") "\n  " RED_TEXT("actual")
                ":   %p\n  " GREEN_TEXT("expected") ": %p\n",
                label, actual, expected);
        longjmp(jmp_env, 1);
    }
}

void assert_true_fn(const char* label, int actual)
{
    ntests++;

    if (! actual) {
        fprintf(stderr, "Test %s " RED_TEXT("failed") "\n  actually "
                RED_TEXT("FALSE") ", expected TRUE\n",
                label);
        longjmp(jmp_env, 1);
    }
}

void assert_false_fn(const char* label, int actual)
{
    ntests++;

    if (actual) {
        fprintf(stderr, "Test %s " RED_TEXT("failed") "\n  actually "
                RED_TEXT("TRUE") ", expected FALSE\n",
                label);
        longjmp(jmp_env, 1);
    }
}

void assert_fail(const char* label)
{
    ntests++;
    fprintf(stderr, "Test %s " RED_TEXT("failed") "\n", label);
    longjmp(jmp_env, 1);
}

void assert_success(const char* label)
{
    ntests++;
}

void run_test_suite(const char* name, void(*fn)(void))
{
    nsuites++;

    int ntests_init = ntests;
    int rval;
    if ((rval = setjmp(jmp_env)) == 0) {
        fn();
    }
    if (rval == 0) {
        printf("  %-60s:%3d tests OK\n", name, ntests-ntests_init);
    } else {
        total_fails++;
        printf("Test %s had some failures\n", name);
    }
}

int report_status(void)
{
    int exit_status = 0;

    if (total_fails == 0) {
        printf("%d tests  :  all %d test-suites pass\n" GREEN_TEXT("OK") "\n", ntests, nsuites);
    } else {
        printf("%d tests  :  %d/%d test-suites failed\n" RED_TEXT("FAILURES") "\n",
               ntests, total_fails, nsuites);
        exit_status = 1;
    }

    return exit_status;
}
