#include "simulation.h"
#include "repository_field.h"

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

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

struct GeneratedOp
{
    VectorTime v;
    SingleAssignmentOp op;
};

Value OracleFieldValue(
    const ReposAssignableField& canonicalOrder, const VectorTime& time)
{
    for (const auto& assignment : canonicalOrder)
        if (Contains(time, assignment)) return assignment.value;
    return Value();
}

void CheckQueries(const ReposAssignableField& repository,
                  const ReposAssignableField& canonicalOrder,
                  const std::vector<VectorTime>& times, bool exhaustiveDiffs)
{
    for (const auto& time : times)
    {
        const auto actual = GetFieldValue(repository, time);
        const auto expected = OracleFieldValue(canonicalOrder, time);
        if (actual != expected)
        {
            std::ostringstream message;
            message << "state mismatch at " << time;
            throw std::runtime_error(message.str());
        }
    }
    const auto checkDiff = [&](const VectorTime& v1, const VectorTime& v2)
    {
        MultiAssignmentOp expected;
        expected.prevValue = OracleFieldValue(canonicalOrder, v1);
        int q = 0;
        for (const auto& r : canonicalOrder)
        {
            if (!Contains(v2, r)) continue;
            if (!Contains(v1, r)) expected.L.push_back({r.s, r.t, r.value, q});
            ++q;
        }
        if (GetDiff(repository, v1, v2) != expected)
        {
            std::ostringstream message;
            message << "diff mismatch from " << v1 << " to " << v2;
            throw std::runtime_error(message.str());
        }
    };
    if (exhaustiveDiffs)
    {
        for (const auto& v1 : times)
            for (const auto& v2 : times)
                if (std::ranges::equal(v1, v2,
                        [](TimeIndex a, TimeIndex b) { return a <= b; })) checkDiff(v1, v2);
    }
    else
    {
        for (const auto& v : times)
        {
            checkDiff(times.front(), v);
            checkDiff(v, times.back());
        }
    }
}

struct History
{
    std::vector<GeneratedOp> ops;
    std::vector<VectorTime> times;
};

History GenerateHistory(std::mt19937& random, std::size_t siteCount, std::size_t count)
{
    std::vector<VectorTime> knowledge(siteCount, VectorTime(siteCount, 0));
    History history;
    history.times.push_back(VectorTime(siteCount, 0));
    std::uniform_int_distribution<std::size_t> site(0, siteCount - 1);
    std::bernoulli_distribution synchronise(0.40);

    while (history.ops.size() < count)
    {
        if (synchronise(random))
        {
            const auto sender = site(random);
            const auto receiver = site(random);
            if (sender != receiver)
            {
                for (std::size_t i = 0; i < siteCount; ++i)
                    knowledge[receiver][i] = std::max(knowledge[receiver][i], knowledge[sender][i]);
                history.times.push_back(knowledge[receiver]);
            }
            continue;
        }
        const auto origin = site(random);
        GeneratedOp generated;
        generated.v = knowledge[origin];
        generated.op.s = static_cast<SiteId>(origin);
        generated.op.t = knowledge[origin][origin];
        generated.op.value =
            static_cast<Value>(1000000 + origin * count + generated.op.t);
        // A user assignment wins the field in its generation context, so it is generated at
        // q = 0.  Nonzero q-positions occur when GetDiff packages historical assignments.
        generated.op.q = 0;
        history.ops.push_back(generated);
        ++knowledge[origin][origin];
        history.times.push_back(knowledge[origin]);
    }
    VectorTime all(siteCount, 0);
    for (const auto& generated : history.ops)
        ++all[static_cast<std::size_t>(generated.op.s)];
    history.times.push_back(std::move(all));
    return history;
}

ReposAssignableField ImportInGenerationOrder(const std::vector<GeneratedOp>& ops)
{
    ReposAssignableField repository;
    for (const auto& generated : ops) Checkin(repository, generated.v, generated.op);
    Validate(repository);
    return repository;
}

ReposAssignableField ImportInRandomCausalOrder(
    const std::vector<GeneratedOp>& ops, std::mt19937& random)
{
    ReposAssignableField repository;
    const auto siteCount = ops.front().v.size();
    VectorTime checkedIn(siteCount, 0);
    std::vector<std::vector<const GeneratedOp*>> bySite(siteCount);
    for (const auto& generated : ops)
    {
        const auto s = static_cast<std::size_t>(generated.op.s);
        if (s >= siteCount || generated.op.t != static_cast<TimeIndex>(bySite[s].size()))
            throw std::runtime_error("generated history has a non-contiguous site sequence");
        bySite[s].push_back(&generated);
    }

    for (std::size_t remaining = ops.size(); remaining != 0; --remaining)
    {
        std::vector<std::size_t> readySites;
        for (std::size_t s = 0; s < siteCount; ++s)
        {
            const auto t = static_cast<std::size_t>(checkedIn[s]);
            if (t == bySite[s].size()) continue;
            if (std::ranges::equal(bySite[s][t]->v, checkedIn,
                    [](TimeIndex required, TimeIndex available) { return required <= available; }))
                readySites.push_back(s);
        }
        if (readySites.empty()) throw std::runtime_error("generated history contains a causal cycle");

        std::uniform_int_distribution<std::size_t> choose(0, readySites.size() - 1);
        const auto s = readySites[choose(random)];
        const auto& generated = *bySite[s][static_cast<std::size_t>(checkedIn[s])];
        Checkin(repository, generated.v, generated.op);
        ++checkedIn[s];
    }
    Validate(repository);
    return repository;
}
}

void runDeterministicTests()
{
    ReposAssignableField nonzeroQ;
    Checkin(nonzeroQ, {0,0}, {0,0,10,0});
    Checkin(nonzeroQ, {1,0}, {1,0,20,1});
    Require(nonzeroQ == ReposAssignableField{{0,0,10}, {1,0,20}},
            "nonzero q-position changed");

    const GeneratedOp a{{0, 0, 0}, {0, 0, 10, 0}};
    const GeneratedOp b{{0, 0, 0}, {1, 0, 20, 0}};
    const GeneratedOp c{{0, 0, 0}, {2, 0, 30, 0}};
    const GeneratedOp d{{1, 1, 0}, {1, 1, 40, 0}};
    const GeneratedOp e{{1, 0, 1}, {2, 1, 40, 0}};
    const std::vector<GeneratedOp> ops{a, b, c, d, e};
    const std::vector<VectorTime> times{{0,0,0}, {1,0,0}, {0,1,0}, {0,0,1},
                                        {1,1,1}, {1,2,1}, {1,1,2}, {1,2,2}};

    ReposAssignableField first;
    for (const auto& generated : ops) Checkin(first, generated.v, generated.op);
    const ReposAssignableField expectedOrder{
        {d.op.s,d.op.t,d.op.value}, {e.op.s,e.op.t,e.op.value},
        {a.op.s,a.op.t,a.op.value}, {b.op.s,b.op.t,b.op.value}, {c.op.s,c.op.t,c.op.value}};
    Require(first == expectedOrder, "mixed causal/concurrent ordering changed");
    Validate(first);
    CheckQueries(first, expectedOrder, times, true);

    const MultiAssignmentOp expectedDiff{
        {{d.op.s,d.op.t,d.op.value,0}, {e.op.s,e.op.t,e.op.value,1}}, 10};
    Require(GetDiff(first, {1,1,1}, {1,2,2}) == expectedDiff,
            "information-preserving difference changed");
}

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.assignments == 0) throw std::runtime_error("assignments must be positive");

    std::mt19937 seeds(settings.seed);
    for (std::size_t run = 0; run < settings.runs; ++run)
    {
        const auto runSeed = seeds();
        std::mt19937 random(runSeed);
        const auto siteCount = settings.minSites + run % (settings.maxSites - settings.minSites + 1);
        try
        {
            const auto history = GenerateHistory(random, siteCount, settings.assignments);
            const auto field = ImportInGenerationOrder(history.ops);
            CheckQueries(field, field, history.times, false);

            for (int order = 0; order < 3; ++order)
            {
                const auto reordered = ImportInRandomCausalOrder(history.ops, random);
                Require(reordered == field, "causally valid check-in order changed the repository");
            }
        }
        catch (const std::exception& error)
        {
            std::ostringstream message;
            message << "run " << run << ", run-seed " << runSeed << ", sites " << siteCount
                    << ": " << error.what();
            throw std::runtime_error(message.str());
        }
    }
}
