#include "simulation.h"
#include "lossy_assignment.h"

#include <algorithm>
#include <random>
#include <sstream>
#include <stdexcept>
#include <vector>

namespace
{
void Require(bool condition, const char* message)
{
    if (!condition) throw std::runtime_error(message);
}

bool TrySendOp(HB& sender, HB& receiver)
{
    const int i = sender.GetOpToSend(receiver.hv);
    if (i == -1) return false;
    receiver.ProcessRemoteOp(sender.MakeSentAssignmentOp(i));
    receiver.Validate();
    return true;
}

void Synchronise(HB& a, HB& b)
{
    while (TrySendOp(a, b)) {}
    while (TrySendOp(b, a)) {}
    Require(a.hv == b.hv, "synchronised sites have different vector times");
    Require(a.fieldValues == b.fieldValues, "synchronised sites have different field values");
}

void SynchroniseAll(std::vector<HB>& sites)
{
    bool progress;
    do
    {
        progress = false;
        for (HB& sender : sites)
            for (HB& receiver : sites)
                if (&sender != &receiver) progress = TrySendOp(sender, receiver) || progress;
    } while (progress);

    for (const HB& site : sites)
    {
        Require(site.hv == sites.front().hv, "sites have different final vector times");
        Require(site.fieldValues == sites.front().fieldValues, "sites did not converge");
    }
}

struct Simulation
{
    Simulation(std::mt19937& random_, int numSites, int numFields) : random(random_)
    {
        for (int s = 0; s < numSites; ++s) sites.emplace_back(s, numSites, numFields);
    }

    int PickSite()
    {
        std::uniform_int_distribution<int> distribution(0, static_cast<int>(sites.size()) - 1);
        return distribution(random);
    }

    void PickDifferentSites(int& a, int& b)
    {
        do
        {
            a = PickSite();
            b = PickSite();
        } while (a == b);
    }

    void Generate()
    {
        HB& site = sites[static_cast<std::size_t>(PickSite())];
        std::uniform_int_distribution<int> field(0, static_cast<int>(site.fieldValues.size()) - 1);
        std::uniform_int_distribution<int> value(1, 1000000);
        site.ApplyLocalOp(field(random), value(random));
        site.Validate();
    }

    void Send()
    {
        int sender, receiver;
        PickDifferentSites(sender, receiver);
        TrySendOp(sites[static_cast<std::size_t>(sender)], sites[static_cast<std::size_t>(receiver)]);
    }

    void SynchronisePair()
    {
        int a, b;
        PickDifferentSites(a, b);
        Synchronise(sites[static_cast<std::size_t>(a)], sites[static_cast<std::size_t>(b)]);
    }

    std::mt19937& random;
    std::vector<HB> sites;
};
}

void RunDeterministicTests()
{
    HB a(0, 2, 2);
    HB b(1, 2, 2);

    a.ApplyLocalOp(0, 10);
    a.ApplyLocalOp(1, 20);
    a.ApplyLocalOp(0, 30);
    Require(a.GetExecutionContextOfOp(0) == VectorTime({0, 0}),
            "first execution context was reconstructed incorrectly");
    Require(a.GetExecutionContextOfOp(1) == VectorTime({1, 0}),
            "second execution context was reconstructed incorrectly");
    Require(a.GetExecutionContextOfOp(2) == VectorTime({2, 0}),
            "third execution context was reconstructed incorrectly");
    Require(a.dm[0]->value == 30, "local dominant map was not updated");
    Require(a.dm[0]->prev && a.dm[0]->prev->value == 10, "local backward chain was not updated");
    const CheckinAssignmentOp checkin = a.GetCheckinAssignment(0);
    Require(checkin.s == 0 && checkin.t == 2 && checkin.fid == 0 && checkin.value == 30,
            "check-in compression did not select the winning assignment");
    a.Validate();

    Synchronise(a, b);
    Require(b.fieldValues == a.fieldValues, "initial synchronisation failed");

    a.ApplyLocalOp(0, 40);
    b.ApplyLocalOp(0, 50);
    Synchronise(a, b);
    Require(a.fieldValues[0] == 40, "lower SiteId did not win concurrent assignments");
    Require(std::ranges::any_of(a.list, [](const auto& op) { return !op->enabled; }),
            "losing concurrent assignment was not disabled");

    b.ApplyLocalOp(0, 60);
    Synchronise(a, b);
    Require(a.fieldValues[0] == 60, "causally later assignment did not win");
}

void RunSimulations(const SimulationSettings& settings)
{
    Require(settings.minSites >= 2 && settings.minSites <= settings.maxSites,
            "site range must satisfy 2 <= min-sites <= max-sites");
    Require(settings.fields > 0, "the number of fields must be positive");
    Require(settings.events > 0, "the number of events must be positive");

    std::mt19937 seeds(settings.seed);
    for (std::size_t run = 0; run < settings.runs; ++run)
    {
        const std::uint32_t runSeed = seeds();
        std::mt19937 random(runSeed);
        const int numSites = settings.minSites +
            static_cast<int>(run % static_cast<std::size_t>(settings.maxSites - settings.minSites + 1));
        Simulation simulation(random, numSites, settings.fields);
        std::uniform_int_distribution<int> event(0, 9);

        try
        {
            for (std::size_t i = 0; i < settings.events; ++i)
            {
                const int e = event(random);
                if (e < 5) simulation.Generate();
                else if (e < 9) simulation.Send();
                else simulation.SynchronisePair();
            }
            SynchroniseAll(simulation.sites);
        }
        catch (const std::exception& error)
        {
            std::ostringstream message;
            message << "run " << run << ", run-seed " << runSeed << ", sites " << numSites
                    << ": " << error.what();
            throw std::runtime_error(message.str());
        }
    }
}
