#include "operation.h"

#include <iterator>
#include <limits>
#include <stdexcept>
#include <string>
#include <vector>

namespace
{
[[noreturn]] void fail(const std::string& message)
{
    throw std::runtime_error(message);
}

void require(bool condition, const std::string& message)
{
    if (!condition)
        fail(message);
}

std::size_t asSize(std::ptrdiff_t value)
{
    require(value >= 0, "negative document coordinate");
    return static_cast<std::size_t>(value);
}

void createRandomExtractions(Ranges& ranges, const std::string& text, std::mt19937& random)
{
    std::bernoulli_distribution extract(0.5);
    Range current;
    for (std::ptrdiff_t i = 0; i < static_cast<std::ptrdiff_t>(text.size()); ++i)
    {
        if (extract(random))
            current.s += text[asSize(i)];
        else
        {
            if (!current.s.empty())
                ranges.push_back(current);
            current = Range{i + 1, {}};
        }
    }
    if (!current.s.empty())
        ranges.push_back(current);
}

void createRandomInsertions(Ranges& ranges, std::ptrdiff_t size, char& character, std::mt19937& random)
{
    std::uniform_int_distribution<int> length(-1, 3);
    std::ptrdiff_t shift = 0;
    for (std::ptrdiff_t i = 0; i <= size; ++i)
    {
        const int count = length(random);
        if (count <= 0)
            continue;
        Range range{shift + i, {}};
        for (int j = 0; j < count; ++j)
        {
            range.s += character;
            character = character == 'Z' ? 'A' : static_cast<char>(character + 1);
        }
        shift += range.size();
        ranges.push_back(std::move(range));
    }
}

void applyInsertions(std::string& text, const Ranges& ranges)
{
    for (const auto& range : ranges)
    {
        require(range.p <= static_cast<std::ptrdiff_t>(text.size()), "insertion position is outside the document");
        text.insert(asSize(range.p), range.s);
    }
}

void applyExtractions(std::string& text, const Ranges& ranges)
{
    for (auto i = ranges.rbegin(); i != ranges.rend(); ++i)
    {
        require(i->p + i->size() <= static_cast<std::ptrdiff_t>(text.size()), "extraction range is outside the document");
        require(text.substr(asSize(i->p), asSize(i->size())) == i->s, "extraction content does not match the document");
        text.erase(asSize(i->p), asSize(i->size()));
    }
}

void validateRanges(const Ranges& ranges, const std::string& text)
{
    for (auto i = ranges.begin(); i != ranges.end(); ++i)
    {
        require(i->p + i->size() <= static_cast<std::ptrdiff_t>(text.size()), "range is outside the document");
        require(text.substr(asSize(i->p), asSize(i->size())) == i->s, "range content does not match the document");
        const auto next = std::next(i);
        if (next != ranges.end())
            require(i->p + i->size() < next->p, "ranges overlap or have not been coalesced");
    }
}

void coalesceRanges(Ranges& ranges)
{
    for (auto i = ranges.begin(); i != ranges.end(); ++i)
    {
        auto next = std::next(i);
        while (next != ranges.end() && i->p + i->size() == next->p)
        {
            i->s += next->s;
            next = ranges.erase(next);
        }
    }
}

void splitRange(Ranges& ranges, Ranges::iterator i, std::ptrdiff_t leftSize)
{
    require(leftSize > 0 && leftSize < i->size(), "invalid range split");
    const auto right = ranges.insert(std::next(i), Range{i->p + leftSize, i->s.substr(asSize(leftSize))});
    (void)right;
    i->s.erase(asSize(leftSize));
}

std::string makeDocument(std::ptrdiff_t size)
{
    std::string result;
    for (std::ptrdiff_t i = 0; i < size; ++i)
        result += static_cast<char>('a' + i % 26);
    return result;
}

Ranges extractionsFromMask(const std::string& text, std::uint64_t mask)
{
    Ranges result;
    Range current;
    for (std::ptrdiff_t i = 0; i < static_cast<std::ptrdiff_t>(text.size()); ++i)
    {
        if ((mask & (std::uint64_t{1} << i)) != 0)
        {
            if (current.s.empty())
                current.p = i;
            current.s += text[asSize(i)];
        }
        else if (!current.s.empty())
        {
            result.push_back(std::move(current));
            current = {};
        }
    }
    if (!current.s.empty())
        result.push_back(std::move(current));
    return result;
}

Ranges insertionsFromMask(std::ptrdiff_t documentSize, std::uint64_t mask, char firstCharacter)
{
    Ranges result;
    std::ptrdiff_t shift = 0;
    char character = firstCharacter;
    for (std::ptrdiff_t p = 0; p <= documentSize; ++p)
    {
        if ((mask & (std::uint64_t{1} << p)) != 0)
        {
            result.push_back(Range{p + shift, std::string(1, character)});
            ++shift;
            character = character == 'Z' ? 'A' : static_cast<char>(character + 1);
        }
    }
    return result;
}

void authenticateTranspose(const std::string& original, Ranges insertions, Ranges extractions)
{
    std::string sequential = original;
    applyInsertions(sequential, insertions);
    applyExtractions(sequential, extractions);

    transposeInsertExtraction(insertions, extractions);
    std::string transposed = original;
    validateRanges(extractions, transposed);
    applyExtractions(transposed, extractions);
    applyInsertions(transposed, insertions);
    validateRanges(insertions, transposed);
    require(sequential == transposed, "insertion/extraction transpose changed the result");
}

void authenticateExtractionMerge(const std::string& original, Ranges first, const Ranges& second)
{
    std::string sequential = original;
    applyExtractions(sequential, first);
    applyExtractions(sequential, second);

    mergeExtractions(first, second);
    std::string merged = original;
    validateRanges(first, merged);
    applyExtractions(merged, first);
    require(sequential == merged, "merged extractions changed the result");
}

void authenticateInsertionMerge(const std::string& original, Ranges first, const Ranges& second)
{
    std::string sequential = original;
    applyInsertions(sequential, first);
    applyInsertions(sequential, second);

    mergeInsertions(first, second);
    std::string merged = original;
    applyInsertions(merged, first);
    validateRanges(first, merged);
    require(sequential == merged, "merged insertions changed the result");
}

void runFocusedTests()
{
    authenticateTranspose("", Ranges{{0, "ABC"}}, Ranges{{0, "ABC"}});
    authenticateTranspose("abcd", Ranges{{0, "L"}, {3, "M"}, {6, "R"}},
                          Ranges{{0, "L"}, {2, "bM"}, {5, "d"}});
    authenticateExtractionMerge("abcdef", Ranges{{0, "ab"}, {4, "e"}}, Ranges{{0, "c"}, {2, "f"}});
    authenticateExtractionMerge("abcdef", Ranges{{0, "abcdef"}}, {});
    authenticateInsertionMerge("", Ranges{{0, "abc"}}, Ranges{{0, "L"}, {4, "R"}});
    authenticateInsertionMerge("abcd", Ranges{{0, "L"}, {3, "M"}, {6, "R"}},
                               Ranges{{0, "A"}, {3, "B"}, {7, "C"}, {10, "D"}});
}

void runExhaustiveSmallStateTests()
{
    // Every insertion subset is crossed with every subsequent extraction subset.
    for (std::ptrdiff_t size = 0; size <= 5; ++size)
    {
        const std::string original = makeDocument(size);
        const auto insertionCount = std::uint64_t{1} << (size + 1);
        for (std::uint64_t insertionMask = 0; insertionMask < insertionCount; ++insertionMask)
        {
            auto insertions = insertionsFromMask(size, insertionMask, 'A');
            std::string intermediate = original;
            applyInsertions(intermediate, insertions);
            const auto extractionCount = std::uint64_t{1} << intermediate.size();
            for (std::uint64_t extractionMask = 0; extractionMask < extractionCount; ++extractionMask)
                authenticateTranspose(original, insertions, extractionsFromMask(intermediate, extractionMask));
        }
    }

    // Every pair of contextually serialised extraction sets on documents up to seven characters.
    for (std::ptrdiff_t size = 0; size <= 7; ++size)
    {
        const std::string original = makeDocument(size);
        const auto firstCount = std::uint64_t{1} << size;
        for (std::uint64_t firstMask = 0; firstMask < firstCount; ++firstMask)
        {
            auto first = extractionsFromMask(original, firstMask);
            std::string intermediate = original;
            applyExtractions(intermediate, first);
            const auto secondCount = std::uint64_t{1} << intermediate.size();
            for (std::uint64_t secondMask = 0; secondMask < secondCount; ++secondMask)
                authenticateExtractionMerge(original, first, extractionsFromMask(intermediate, secondMask));
        }
    }

    // Every pair of single-character insertion subsets on documents up to five characters.
    for (std::ptrdiff_t size = 0; size <= 5; ++size)
    {
        const std::string original = makeDocument(size);
        const auto firstCount = std::uint64_t{1} << (size + 1);
        for (std::uint64_t firstMask = 0; firstMask < firstCount; ++firstMask)
        {
            auto first = insertionsFromMask(size, firstMask, 'A');
            std::string intermediate = original;
            applyInsertions(intermediate, first);
            const auto secondCount = std::uint64_t{1} << (intermediate.size() + 1);
            for (std::uint64_t secondMask = 0; secondMask < secondCount; ++secondMask)
                authenticateInsertionMerge(original, first,
                    insertionsFromMask(static_cast<std::ptrdiff_t>(intermediate.size()), secondMask, 'N'));
        }
    }
}

void runRandomComponentTests(std::size_t count, std::mt19937& random)
{
    std::uniform_int_distribution<std::ptrdiff_t> documentLength(0, 64);
    for (std::size_t iteration = 0; iteration < count; ++iteration)
    {
        const std::string original = makeDocument(documentLength(random));
        char character = 'A';

        Ranges insertions;
        createRandomInsertions(insertions, static_cast<std::ptrdiff_t>(original.size()), character, random);
        std::string afterInsertions = original;
        applyInsertions(afterInsertions, insertions);
        Ranges extractions;
        createRandomExtractions(extractions, afterInsertions, random);
        authenticateTranspose(original, insertions, extractions);

        Ranges firstExtractions;
        createRandomExtractions(firstExtractions, original, random);
        std::string afterFirstExtractions = original;
        applyExtractions(afterFirstExtractions, firstExtractions);
        Ranges secondExtractions;
        createRandomExtractions(secondExtractions, afterFirstExtractions, random);
        authenticateExtractionMerge(original, firstExtractions, secondExtractions);

        Ranges firstInsertions;
        createRandomInsertions(firstInsertions, static_cast<std::ptrdiff_t>(original.size()), character, random);
        std::string afterFirstInsertions = original;
        applyInsertions(afterFirstInsertions, firstInsertions);
        Ranges secondInsertions;
        createRandomInsertions(secondInsertions, static_cast<std::ptrdiff_t>(afterFirstInsertions.size()), character, random);
        authenticateInsertionMerge(original, firstInsertions, secondInsertions);
    }
}
}

void transposeInsertExtraction(Ranges& insertions, Ranges& extractions)
{
    std::ptrdiff_t insertionShift = 0;
    std::ptrdiff_t extractionShift = 0;
    auto insertion = insertions.begin();
    auto extraction = extractions.begin();

    while (insertion != insertions.end() && extraction != extractions.end())
    {
        if (insertion->p + insertion->size() <= extraction->p)
        {
            insertion->p -= extractionShift;
            insertionShift += insertion->size();
            ++insertion;
        }
        else if (extraction->p + extraction->size() <= insertion->p)
        {
            extraction->p -= insertionShift;
            extractionShift += extraction->size();
            ++extraction;
        }
        else
        {
            if (insertion->p < extraction->p)
            {
                splitRange(insertions, insertion, extraction->p - insertion->p);
                insertion->p -= extractionShift;
                insertionShift += insertion->size();
                ++insertion;
            }
            else if (extraction->p < insertion->p)
            {
                splitRange(extractions, extraction, insertion->p - extraction->p);
                extraction->p -= insertionShift;
                extractionShift += extraction->size();
                ++extraction;
            }

            require(extraction->p == insertion->p, "overlapping ranges did not align");
            if (extraction->size() < insertion->size())
                splitRange(insertions, insertion, extraction->size());
            else if (insertion->size() < extraction->size())
                splitRange(extractions, extraction, insertion->size());

            extractionShift += extraction->size();
            insertionShift += insertion->size();
            insertion = insertions.erase(insertion);
            extraction = extractions.erase(extraction);
        }
    }

    for (; extraction != extractions.end(); ++extraction)
        extraction->p -= insertionShift;
    for (; insertion != insertions.end(); ++insertion)
        insertion->p -= extractionShift;
    coalesceRanges(insertions);
    coalesceRanges(extractions);
}

void mergeExtractions(Ranges& first, const Ranges& second)
{
    std::ptrdiff_t shift = 0;
    auto firstRange = first.begin();
    for (const auto& secondRange : second)
    {
        while (firstRange != first.end() && firstRange->p < shift + secondRange.p)
        {
            shift += firstRange->size();
            ++firstRange;
        }

        Range combined{shift + secondRange.p, {}};
        auto source = secondRange.s.begin();
        while (firstRange != first.end() && firstRange->p <= shift + secondRange.p + secondRange.size())
        {
            const auto boundary = secondRange.s.begin() + firstRange->p - (shift + secondRange.p);
            combined.s.append(source, boundary);
            source = boundary;
            combined.s += firstRange->s;
            shift += firstRange->size();
            firstRange = first.erase(firstRange);
        }
        combined.s.append(source, secondRange.s.end());
        first.insert(firstRange, std::move(combined));
    }
}

void mergeInsertions(Ranges& first, const Ranges& second)
{
    std::ptrdiff_t shift = 0;
    auto firstRange = first.begin();
    auto secondRange = second.begin();
    while (firstRange != first.end())
    {
        while (secondRange != second.end() && secondRange->p < shift + firstRange->p)
        {
            first.insert(firstRange, *secondRange);
            shift += secondRange->size();
            ++secondRange;
        }
        if (secondRange == second.end())
        {
            for (; firstRange != first.end(); ++firstRange)
                firstRange->p += shift;
            return;
        }

        firstRange->p += shift;
        while (secondRange != second.end() && secondRange->p <= firstRange->p + firstRange->size())
        {
            firstRange->s.insert(asSize(secondRange->p - firstRange->p), secondRange->s);
            shift += secondRange->size();
            ++secondRange;
        }
        ++firstRange;
    }
    first.insert(first.end(), secondRange, second.end());
}

void merge(Operation& first, const Operation& second)
{
    auto extractions = second.extractions;
    transposeInsertExtraction(first.insertions, extractions);
    mergeExtractions(first.extractions, extractions);
    mergeInsertions(first.insertions, second.insertions);
}

void runAuthenticationTests(std::size_t count, std::mt19937& random)
{
    runFocusedTests();
    runExhaustiveSmallStateTests();
    runRandomComponentTests(count, random);

    for (std::size_t iteration = 0; iteration < count; ++iteration)
    {
        std::string sequential = "abcdefghijk";
        const std::string original = sequential;
        Operation first;
        Operation second;
        char character = 'A';

        createRandomExtractions(first.extractions, sequential, random);
        validateRanges(first.extractions, sequential);
        applyExtractions(sequential, first.extractions);
        createRandomInsertions(first.insertions, static_cast<std::ptrdiff_t>(sequential.size()), character, random);
        applyInsertions(sequential, first.insertions);
        validateRanges(first.insertions, sequential);

        createRandomExtractions(second.extractions, sequential, random);
        validateRanges(second.extractions, sequential);
        applyExtractions(sequential, second.extractions);
        createRandomInsertions(second.insertions, static_cast<std::ptrdiff_t>(sequential.size()), character, random);
        applyInsertions(sequential, second.insertions);
        validateRanges(second.insertions, sequential);

        merge(first, second);
        std::string compressed = original;
        validateRanges(first.extractions, compressed);
        applyExtractions(compressed, first.extractions);
        applyInsertions(compressed, first.insertions);
        validateRanges(first.insertions, compressed);
        if (sequential != compressed)
            fail("merged operation differs from sequential operations at iteration " + std::to_string(iteration));
    }
}
