#include "Simulation.h"

#include "EffectsDocument.h"
#include "Operation.h"
#include "RandomOperation.h"

#include <deque>
#include <numeric>
#include <random>
#include <sstream>
#include <stdexcept>
#include <utility>
#include <vector>

namespace
{
using VectorTime = std::vector<int>;

bool contains(const VectorTime& time, const Operation& operation)
{
    return time.at(static_cast<std::size_t>(operation.opid.id)) > operation.opid.t;
}

EffectsDocumentSet initialDocument(std::size_t documents, std::size_t charactersPerDocument)
{
    if (documents == 0 || charactersPerDocument == 0 || documents * charactersPerDocument > 90)
        throw std::runtime_error("initial document dimensions must contain between 1 and 90 characters");
    EffectsDocumentSet result(static_cast<int>(documents));
    char value = '!';
    for (auto& document : result)
    {
        document.SetState(value, static_cast<char>(value + charactersPerDocument));
        value = static_cast<char>(value + charactersPerDocument);
    }
    return result;
}

struct Site
{
    Site(int id, std::size_t siteCount, const EffectsDocumentSet& initial)
        : id(id), documents(initial), time(siteCount, 0) {}
    int id;
    EffectsDocumentSet documents;
    VectorTime time;
    std::vector<Operation> history;
};

std::string describe(const VectorTime& time)
{
    std::ostringstream out;
    out << '[';
    for (std::size_t i = 0; i < time.size(); ++i) out << (i ? "," : "") << time[i];
    return out.str() + ']';
}

class Simulation
{
public:
    Simulation(std::uint32_t seed, std::size_t siteCount, const SimulationSettings& settings)
        : random_(seed), settings_(settings)
    {
        setOperationRandomSeed(seed ^ 0x9e3779b9u);
        const auto initial = initialDocument(settings.documents, settings.initialCharactersPerDocument);
        for (std::size_t i = 0; i < siteCount; ++i)
            sites_.emplace_back(static_cast<int>(i), siteCount, initial);
    }

    void run(std::size_t eventCount)
    {
        for (std::size_t event = 0; event < eventCount; ++event)
        {
            const int choice = integer(0, 10);
            if (choice < 5) generate();
            else if (choice < 9) sendRandom();
            else synchroniseRandomPair();
            validateSites();
        }
        drainAndValidateAll();
    }
    const std::deque<std::string>& trace() const { return trace_; }

private:
    int integer(int first, int pastLast)
    {
        return std::uniform_int_distribution<int>(first, pastLast - 1)(random_);
    }

    std::pair<std::size_t, std::size_t> differentSites()
    {
        const auto first = static_cast<std::size_t>(integer(0, sites_.size()));
        std::size_t second;
        do second = static_cast<std::size_t>(integer(0, sites_.size())); while (first == second);
        return {first, second};
    }

    void record(std::string event)
    {
        trace_.push_back(std::move(event));
        if (trace_.size() > 1000) trace_.pop_front();
    }

    void generate()
    {
        Site& site = sites_[static_cast<std::size_t>(integer(0, sites_.size()))];
        Operation operation;
        RandomOperationSettings settings;
        settings.minAtomic = static_cast<int>(settings_.minAtomic);
        settings.maxAtomic = static_cast<int>(settings_.maxAtomic);
        SetRandom(operation, settings, site.documents,
                  OpId{site.id, site.time[static_cast<std::size_t>(site.id)]});
        std::ostringstream description;
        description << operation;
        record("S" + std::to_string(site.id) + " generate " + description.str());
        operation.Apply(site.documents);
        site.history.push_back(operation);
        ++site.time[static_cast<std::size_t>(site.id)];
    }

    static std::size_t separateHistory(std::vector<Operation>& history,
                                       const VectorTime& context)
    {
        std::size_t prefix = 0;
        while (prefix < history.size() && contains(context, history[prefix])) ++prefix;
        for (std::size_t candidate = prefix + 1; candidate < history.size(); ++candidate)
        {
            if (!contains(context, history[candidate])) continue;
            for (std::size_t position = candidate; position > prefix; --position)
            {
                Transpose(history[position - 1], history[position]);
                std::swap(history[position - 1], history[position]);
            }
            ++prefix;
        }
        return prefix;
    }

    bool send(std::size_t senderIndex, std::size_t receiverIndex)
    {
        Site& sender = sites_[senderIndex];
        Site& receiver = sites_[receiverIndex];
        VectorTime context(sites_.size(), 0);
        const Operation* selected = nullptr;
        for (const auto& operation : sender.history)
        {
            if (!contains(receiver.time, operation)) { selected = &operation; break; }
            ++context[static_cast<std::size_t>(operation.opid.id)];
        }
        if (!selected) return false;
        Operation incoming = *selected;
        const auto prefix = separateHistory(receiver.history, context);
        for (std::size_t i = prefix; i < receiver.history.size(); ++i) IT(incoming, receiver.history[i]);
        std::ostringstream description;
        description << incoming;
        record("S" + std::to_string(sender.id) + " -> S" + std::to_string(receiver.id) + ' ' +
               description.str() + " context=" + describe(context));
        incoming.Apply(receiver.documents);
        receiver.history.push_back(incoming);
        ++receiver.time[static_cast<std::size_t>(incoming.opid.id)];
        return true;
    }

    void sendRandom() { const auto [a, b] = differentSites(); send(a, b); }

    void synchronise(std::size_t first, std::size_t second)
    {
        bool changed;
        do { changed = send(first, second); changed = send(second, first) || changed; } while (changed);
        if (sites_[first].time != sites_[second].time ||
            sites_[first].documents != sites_[second].documents)
            throw std::runtime_error("pair failed to converge");
        record("verify S" + std::to_string(sites_[first].id) + " == S" +
               std::to_string(sites_[second].id));
    }

    void synchroniseRandomPair() { const auto [a, b] = differentSites(); synchronise(a, b); }

    void drainAndValidateAll()
    {
        bool changed;
        do
        {
            changed = false;
            for (std::size_t a = 0; a < sites_.size(); ++a)
                for (std::size_t b = 0; b < sites_.size(); ++b)
                    if (a != b) changed = send(a, b) || changed;
        }
        while (changed);
        for (std::size_t i = 1; i < sites_.size(); ++i)
            if (sites_[i].time != sites_[0].time || sites_[i].documents != sites_[0].documents)
                throw std::runtime_error("full exchange did not converge");
    }

    void validateSites() const
    {
        for (const auto& site : sites_)
        {
            const int count = std::accumulate(site.time.begin(), site.time.end(), int{0});
            if (site.history.size() != static_cast<std::size_t>(count))
                throw std::runtime_error("history size and vector time disagree");
            for (const auto& operation : site.history) operation.AssertValid();
        }
    }

    std::mt19937 random_;
    const SimulationSettings& settings_;
    std::vector<Site> sites_;
    std::deque<std::string> trace_;
};
}

void runSimulations(const SimulationSettings& settings)
{
    if (settings.minSites < 2 || settings.minSites > settings.maxSites)
        throw std::runtime_error("site range must satisfy 2 <= min-sites <= max-sites");
    if (settings.minAtomic == 0 || settings.minAtomic > settings.maxAtomic)
        throw std::runtime_error("atomic-operation range must be positive and ordered");
    std::mt19937 seeds(settings.seed);
    for (std::size_t run = 0; run < settings.runs; ++run)
    {
        const auto runSeed = seeds();
        const auto siteCount = settings.minSites + run % (settings.maxSites - settings.minSites + 1);
        Simulation simulation(runSeed, siteCount, settings);
        try { simulation.run(settings.events); }
        catch (const std::exception& error)
        {
            std::ostringstream message;
            message << "run " << run << ", run-seed " << runSeed << ", sites " << siteCount
                    << ": " << error.what() << "\nReplay trace:\n";
            for (const auto& event : simulation.trace()) message << event << '\n';
            throw std::runtime_error(message.str());
        }
    }
}
