#include "utilities/Sort.h"
#include "unittest/TestSuite.h"
#include "stem_core.h"
#include <stdint.h>

static int compareUInt32Ascending(const void * lhsUntyped, const void * rhsUntyped, void * context) {
	unsigned int * callCount = context;
	++*callCount;
	const uint32_t * lhs = lhsUntyped, * rhs = rhsUntyped;
	return (*lhs > *rhs) * 2 - 1;
}

static int compareUInt32Descending(const void * lhsUntyped, const void * rhsUntyped, void * context) {
	unsigned int * callCount = context;
	++*callCount;
	const uint32_t * lhs = lhsUntyped, * rhs = rhsUntyped;
	return (*lhs < *rhs) * 2 - 1;
}

static int compareUInt64Ascending(const void * lhsUntyped, const void * rhsUntyped, void * context) {
	unsigned int * callCount = context;
	++*callCount;
	const uint64_t * lhs = lhsUntyped, * rhs = rhsUntyped;
	return (*lhs > *rhs) * 2 - 1;
}

struct largeItem_align64 {
	uint64_t index;
	uint64_t extra1;
	uint64_t extra2;
	uint64_t extra3;
	uint64_t extra4;
};

static int compareLargeItemAlign64Descending(const void * lhsUntyped, const void * rhsUntyped, void * context) {
	const struct largeItem_align64 * lhs = lhsUntyped, * rhs = rhsUntyped;
	return (lhs->index < rhs->index) * 2 - 1;
}

static void testQuickSort(void) {
	quickSort(NULL, 0, 0, NULL, NULL);
	
	uint32_t testData1[] = {7, 6, 5, 4, 3, 2, 1, 0};
	unsigned int callCount = 0;
	quickSort(testData1, sizeof_count(testData1), sizeof(testData1[0]), compareUInt32Ascending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUIntEqual(testData1[0], 0);
	TestCase_assertUIntEqual(testData1[1], 1);
	TestCase_assertUIntEqual(testData1[2], 2);
	TestCase_assertUIntEqual(testData1[3], 3);
	TestCase_assertUIntEqual(testData1[4], 4);
	TestCase_assertUIntEqual(testData1[5], 5);
	TestCase_assertUIntEqual(testData1[6], 6);
	TestCase_assertUIntEqual(testData1[7], 7);
	
	callCount = 0;
	quickSort(testData1, sizeof_count(testData1), sizeof(testData1[0]), compareUInt32Descending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUIntEqual(testData1[0], 7);
	TestCase_assertUIntEqual(testData1[1], 6);
	TestCase_assertUIntEqual(testData1[2], 5);
	TestCase_assertUIntEqual(testData1[3], 4);
	TestCase_assertUIntEqual(testData1[4], 3);
	TestCase_assertUIntEqual(testData1[5], 2);
	TestCase_assertUIntEqual(testData1[6], 1);
	TestCase_assertUIntEqual(testData1[7], 0);
	
	callCount = 0;
	quickSort(testData1 + 2, 5, sizeof(testData1[0]), compareUInt32Ascending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUIntEqual(testData1[0], 7);
	TestCase_assertUIntEqual(testData1[1], 6);
	TestCase_assertUIntEqual(testData1[2], 1);
	TestCase_assertUIntEqual(testData1[3], 2);
	TestCase_assertUIntEqual(testData1[4], 3);
	TestCase_assertUIntEqual(testData1[5], 4);
	TestCase_assertUIntEqual(testData1[6], 5);
	TestCase_assertUIntEqual(testData1[7], 0);
	
	uint64_t testData2[] = {70, 0, 60, 10, 50, 20, 40, 30};
	quickSort(testData2, sizeof_count(testData2), sizeof(testData2[0]), compareUInt64Ascending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUInt64Equal(testData2[0], 0);
	TestCase_assertUInt64Equal(testData2[1], 10);
	TestCase_assertUInt64Equal(testData2[2], 20);
	TestCase_assertUInt64Equal(testData2[3], 30);
	TestCase_assertUInt64Equal(testData2[4], 40);
	TestCase_assertUInt64Equal(testData2[5], 50);
	TestCase_assertUInt64Equal(testData2[6], 60);
	TestCase_assertUInt64Equal(testData2[7], 70);
	
	struct largeItem_align64 testData3[32];
	for (unsigned int itemIndex = 0; itemIndex < sizeof_count(testData3); itemIndex++) {
		testData3[itemIndex].index = itemIndex;
		testData3[itemIndex].extra1 = itemIndex + 1;
		testData3[itemIndex].extra2 = 12345;
		testData3[itemIndex].extra3 = itemIndex * 2;
		testData3[itemIndex].extra4 = 543210;
	}
	quickSort(testData3, sizeof_count(testData3), sizeof(testData3[0]), compareLargeItemAlign64Descending, NULL);
	for (unsigned int itemIndex = 0; itemIndex < sizeof_count(testData3); itemIndex++) {
		unsigned int reverseItemIndex = sizeof_count(testData3) - itemIndex - 1;
		TestCase_assert(testData3[itemIndex].index == reverseItemIndex,      "Nonmatching index at index %u: Expected %u but got "  UINT64_FORMAT, itemIndex, reverseItemIndex,     testData3[itemIndex].index);
		TestCase_assert(testData3[itemIndex].extra1 == reverseItemIndex + 1, "Nonmatching extra1 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, reverseItemIndex + 1, testData3[itemIndex].extra1);
		TestCase_assert(testData3[itemIndex].extra2 == 12345,                "Nonmatching extra2 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, 12345,                testData3[itemIndex].extra2);
		TestCase_assert(testData3[itemIndex].extra3 == reverseItemIndex * 2, "Nonmatching extra3 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, reverseItemIndex * 2, testData3[itemIndex].extra3);
		TestCase_assert(testData3[itemIndex].extra4 == 543210,               "Nonmatching extra4 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, 543210,               testData3[itemIndex].extra4);
	}
}

static void testHeapSort(void) {
#ifdef INCLUDE_HEAP_SORT
	heapSort(NULL, 0, 0, NULL, NULL);
	
	uint32_t testData1[] = {7, 6, 5, 4, 3, 2, 1, 0};
	unsigned int callCount = 0;
	heapSort(testData1, sizeof_count(testData1), sizeof(testData1[0]), compareUInt32Ascending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUIntEqual(testData1[0], 0);
	TestCase_assertUIntEqual(testData1[1], 1);
	TestCase_assertUIntEqual(testData1[2], 2);
	TestCase_assertUIntEqual(testData1[3], 3);
	TestCase_assertUIntEqual(testData1[4], 4);
	TestCase_assertUIntEqual(testData1[5], 5);
	TestCase_assertUIntEqual(testData1[6], 6);
	TestCase_assertUIntEqual(testData1[7], 7);
	
	callCount = 0;
	heapSort(testData1, sizeof_count(testData1), sizeof(testData1[0]), compareUInt32Descending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUIntEqual(testData1[0], 7);
	TestCase_assertUIntEqual(testData1[1], 6);
	TestCase_assertUIntEqual(testData1[2], 5);
	TestCase_assertUIntEqual(testData1[3], 4);
	TestCase_assertUIntEqual(testData1[4], 3);
	TestCase_assertUIntEqual(testData1[5], 2);
	TestCase_assertUIntEqual(testData1[6], 1);
	TestCase_assertUIntEqual(testData1[7], 0);
	
	callCount = 0;
	heapSort(testData1 + 2, 5, sizeof(testData1[0]), compareUInt32Ascending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUIntEqual(testData1[0], 7);
	TestCase_assertUIntEqual(testData1[1], 6);
	TestCase_assertUIntEqual(testData1[2], 1);
	TestCase_assertUIntEqual(testData1[3], 2);
	TestCase_assertUIntEqual(testData1[4], 3);
	TestCase_assertUIntEqual(testData1[5], 4);
	TestCase_assertUIntEqual(testData1[6], 5);
	TestCase_assertUIntEqual(testData1[7], 0);
	
	uint64_t testData2[] = {70, 0, 60, 10, 50, 20, 40, 30};
	heapSort(testData2, sizeof_count(testData2), sizeof(testData2[0]), compareUInt64Ascending, &callCount);
	TestCase_assert(callCount > 0, "Compare function not called");
	TestCase_assertUInt64Equal(testData2[0], 0);
	TestCase_assertUInt64Equal(testData2[1], 10);
	TestCase_assertUInt64Equal(testData2[2], 20);
	TestCase_assertUInt64Equal(testData2[3], 30);
	TestCase_assertUInt64Equal(testData2[4], 40);
	TestCase_assertUInt64Equal(testData2[5], 50);
	TestCase_assertUInt64Equal(testData2[6], 60);
	TestCase_assertUInt64Equal(testData2[7], 70);
	
	struct largeItem_align64 testData3[32];
	for (unsigned int itemIndex = 0; itemIndex < sizeof_count(testData3); itemIndex++) {
		testData3[itemIndex].index = itemIndex;
		testData3[itemIndex].extra1 = itemIndex + 1;
		testData3[itemIndex].extra2 = 12345;
		testData3[itemIndex].extra3 = itemIndex * 2;
		testData3[itemIndex].extra4 = 543210;
	}
	heapSort(testData3, sizeof_count(testData3), sizeof(testData3[0]), compareLargeItemAlign64Descending, NULL);
	for (unsigned int itemIndex = 0; itemIndex < sizeof_count(testData3); itemIndex++) {
		unsigned int reverseItemIndex = sizeof_count(testData3) - itemIndex - 1;
		TestCase_assert(testData3[itemIndex].index == reverseItemIndex,      "Nonmatching index at index %u: Expected %u but got "  UINT64_FORMAT, itemIndex, reverseItemIndex,     testData3[itemIndex].index);
		TestCase_assert(testData3[itemIndex].extra1 == reverseItemIndex + 1, "Nonmatching extra1 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, reverseItemIndex + 1, testData3[itemIndex].extra1);
		TestCase_assert(testData3[itemIndex].extra2 == 12345,                "Nonmatching extra2 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, 12345,                testData3[itemIndex].extra2);
		TestCase_assert(testData3[itemIndex].extra3 == reverseItemIndex * 2, "Nonmatching extra3 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, reverseItemIndex * 2, testData3[itemIndex].extra3);
		TestCase_assert(testData3[itemIndex].extra4 == 543210,               "Nonmatching extra4 at index %u: Expected %u but got " UINT64_FORMAT, itemIndex, 543210,               testData3[itemIndex].extra4);
	}
#endif
}

TEST_SUITE(SortTest,
           testQuickSort,
           testHeapSort)
