#include "qfa_bundle.h"
#include "objective.h"
#include "dickson_hybrid_topology.h"
#include <symengine/lambda_double.h>
#include <symengine/symbol.h>
#include <symengine/integer.h>
#include <symengine/visitor.h>
#include <cassert>
#include <cmath>
#include <vector>
using namespace SymEngine;

static RCP<const Basic> objective(const Topology &t, const std::vector<RCP<const Basic>> &x,
                                  bool fsl) {
    map_basic_basic fixed;
    fixed[symbol("D")] = real_double(.75);
    fixed[t.frequency] = real_double(1);
    for (auto &e : t.capacitor_esr) fixed[e] = real_double(0);
    for (unsigned i = 0; i < t.capacitances.size(); ++i)
        fixed[t.capacitances[i]] = fsl ? RCP<const Basic>(real_double(1)) : RCP<const Basic>(x[i % x.size()]);
    for (unsigned i = 0; i < t.switch_resistances.size(); ++i)
        fixed[t.switch_resistances[i]] = fsl ? div(integer(1), x[i % x.size()]) : RCP<const Basic>(real_double(1));
    const auto &z = fsl ? t.ZFSL : t.ZSSL;
    RCP<const Basic> sum = integer(0);
    for (unsigned i=0; i<z.nrows(); ++i) for (unsigned j=0; j<z.ncols(); ++j)
        sum = add(sum, z.get(i,j)->subs(fixed));
    return sum;
}
static void check(bool fsl) {
    const unsigned n = fsl ? 9 : 5;
    Topology t = dickson_hybrid_topology(5, Expression(symbol("D")), {symbol("D")}, true, false);
    complete_qfa(t, dickson_arch(5));
    std::vector<RCP<const Basic>> x;
    for (unsigned i=0; i<n; ++i) x.push_back(symbol("x" + std::to_string(i)));
    auto expr = objective(t, x, fsl);
    auto free = free_symbols(*expr);
    assert(free.size() == n);
    for (auto &v : x) assert(free.count(v));
    CompiledObjective compiled(t, fsl ? "fsl" : "ssl", -1, .75, std::vector<double>(7, 1.0));
    LambdaDoubleVisitor<double> lambda;
    lambda.init(x, *expr, false);
    for (double scale : {0.7, 1.3}) {
        std::vector<double> values(n, scale);
        const double got = lambda.call(values);
        assert(std::isfinite(got));
        assert(std::isfinite(compiled.call(values)));
    }
}
int main() { check(false); check(true); }
