#include "dickson_hybrid_topology.h"
#include "graph_primitives.h"
#include "solve_charge_vectors.h"
#include "qfa_bundle.h"
#include <symengine/integer.h>
#include <symengine/symbol.h>
#include <symengine/visitor.h>
#include <stdexcept>
#include <cmath>
#include <symengine/pow.h>
#include <symengine/eval_double.h>
using SymEngine::DenseMatrix;

namespace {
DenseMatrix select_columns(const DenseMatrix &matrix, const std::vector<unsigned> &columns) {
    DenseMatrix selected(matrix.nrows(), columns.size());
    for (unsigned row = 0; row < matrix.nrows(); ++row)
        for (unsigned column = 0; column < columns.size(); ++column)
            selected.set(row, column, matrix.get(row, columns[column]));
    return selected;
}
DenseMatrix select_rows(const DenseMatrix &matrix, const std::vector<unsigned> &rows) {
    DenseMatrix selected(rows.size(), matrix.ncols());
    for (unsigned row = 0; row < rows.size(); ++row)
        for (unsigned column = 0; column < matrix.ncols(); ++column)
            selected.set(row, column, matrix.get(rows[row], column));
    return selected;
}
DenseMatrix select_square(const DenseMatrix &matrix, const std::vector<unsigned> &columns) {
    DenseMatrix selected(columns.size(), columns.size());
    for (unsigned row = 0; row < columns.size(); ++row)
        for (unsigned column = 0; column < columns.size(); ++column)
            selected.set(row, column, matrix.get(columns[row], columns[column]));
    return selected;
}
}

Topology dickson_hybrid_topology(int n_caps, const SymEngine::Expression &duty,
                                 const std::vector<SymEngine::RCP<const SymEngine::Symbol>> &symbols,
                                 bool dc_out, bool half_point) {
    if (dc_out == false || half_point) throw std::invalid_argument("dickson_hybrid_topology: unsupported option");
    auto free = SymEngine::free_symbols(*duty.get_basic());
    if (free.empty()) {
        double value=SymEngine::eval_double(*duty.get_basic());
        if (!std::isfinite(value) || value<=0 || value>=1)
            throw std::invalid_argument("dickson_hybrid_topology: duty must be in (0,1)");
    }
    if (free.size() != symbols.size()) throw std::invalid_argument("dickson_hybrid_topology: symbols do not match duty");
    for (const auto &symbol : symbols)
        if (free.find(symbol) == free.end()) throw std::invalid_argument("dickson_hybrid_topology: symbols do not match duty");
    ArchDef arch = dickson_arch(n_caps);
    Topology top;
    top.ordered_symbols = symbols;
    top.N_caps = n_caps; top.N_sw = arch.Asw.ncols(); top.vo_swing = 1.0 / n_caps;
    top.duty = SymEngine::DenseMatrix(1, 2);
    top.duty.set(0, 0, duty.get_basic());
    top.duty.set(0, 1, SymEngine::sub(SymEngine::integer(1), duty.get_basic()));
    DenseMatrix loads(arch.Acaps.nrows(), arch.Acaps.nrows()-1), supply(arch.Acaps.nrows(),1);
    for (unsigned r=0;r<loads.nrows();++r) for (unsigned c=0;c<loads.ncols();++c) loads.set(r,c,SymEngine::integer(r==c+1));
    for (unsigned r=0;r<supply.nrows();++r) supply.set(r,0,SymEngine::integer(r==0));
    for (unsigned p=0;p<2;++p) {
        DenseMatrix on(arch.Asw.nrows(),0), off(arch.Asw.nrows(),0); std::vector<unsigned> indexes;
        for (unsigned c=0;c<arch.Asw.ncols();++c) {
            DenseMatrix one(arch.Asw.nrows(),1); for (unsigned r=0;r<one.nrows();++r) one.set(r,0,arch.Asw.get(r,c));
            bool active = SymEngine::eq(*arch.Asw_act.get(p,c), *SymEngine::integer(1));
            if (active) { DenseMatrix x(on.nrows(),on.ncols()+1); for(unsigned r=0;r<x.nrows();++r){for(unsigned j=0;j<on.ncols();++j)x.set(r,j,on.get(r,j));x.set(r,on.ncols(),one.get(r,0));} on=x; indexes.push_back(c+1); }
            else { DenseMatrix x(off.nrows(),off.ncols()+1); for(unsigned r=0;r<x.nrows();++r){for(unsigned j=0;j<off.ncols();++j)x.set(r,j,off.get(r,j));x.set(r,off.ncols(),one.get(r,0));} off=x; }
        }
        DenseMatrix phase_loads = loads;
        auto phase_duty = p == 0 ? duty.get_basic() : SymEngine::sub(SymEngine::integer(1), duty.get_basic());
        for (unsigned c = 0; c < phase_loads.ncols(); ++c)
            phase_loads.set(c + 1, c, phase_duty);
        top.phase.emplace_back(on,arch.Acaps,off,phase_loads,n_caps,indexes,supply);
        top.phase.back().duty = p == 0 ? duty : SymEngine::Expression(SymEngine::sub(SymEngine::integer(1), duty.get_basic()));
        top.phase.back().symbols = symbols;
        top.phase.back().graph = top.phase.back().get_on_no_sw();
        auto graph=top.phase.back().graph;
        top.phase.back().tree=build_tree(graph,0);
        if (top.phase.back().tree.size()==1 && top.phase.back().tree[0]==static_cast<unsigned>(-1)) throw std::runtime_error("dickson_hybrid_topology: singular graph");
        top.phase.back().cutset=fun_cutset(graph,top.phase.back().tree);
        top.phase.back().loop=fun_loop(graph,top.phase.back().tree);
    }
    std::vector<DenseMatrix> cutsets;
    for (const auto &phase : top.phase) cutsets.push_back(phase.cutset);
    ChargeSolution charge = solve_charge_vectors(cutsets, n_caps, top.duty, symbols);
    top.m_ratios = charge.m;
    top.ratio = charge.m;
    for (unsigned p = 0; p < top.phase.size(); ++p) top.phase[p].set_a_vector(charge.a[p]);
    complete_qfa(top, arch);
    return top;
}

Topology select_outputs(const Topology &topology, const std::vector<unsigned> &outputs) {
    for (unsigned output : outputs)
        if (output >= topology.m_ratios.nrows()) throw std::invalid_argument("select_outputs: output index out of range");
    Topology selected = topology;
    selected.m_ratios = select_rows(topology.m_ratios, outputs);
    selected.ratio = selected.m_ratios;
    selected.ZSSL = select_square(topology.ZSSL, outputs);
    selected.ZFSL = select_square(topology.ZFSL, outputs);
    selected.ZESR = select_square(topology.ZESR, outputs);
    selected.ZSCC = select_square(topology.ZSCC, outputs);
    selected.is = select_columns(topology.is, outputs);
    for (unsigned phase = 0; phase < selected.phase.size(); ++phase) {
        selected.phase[phase].a_vector = select_columns(topology.phase[phase].a_vector, outputs);
        selected.phase[phase].b_vector = select_columns(topology.phase[phase].b_vector, outputs);
        selected.phase[phase].r_vector = select_columns(topology.phase[phase].r_vector, outputs);
        selected.phase[phase].ar_vector = select_columns(topology.phase[phase].ar_vector, outputs);
    }
    return selected;
}

Topology dickson_hybrid_topology(int n_caps, double duty, bool dc_out, bool half_point) {
    if (!std::isfinite(duty) || duty<=0 || duty>=1)
        throw std::invalid_argument("dickson_hybrid_topology: duty must be in (0,1)");
    int exponent;
    double mantissa=std::frexp(duty,&exponent);
    auto exact=SymEngine::mul(SymEngine::integer(static_cast<long long>(std::ldexp(mantissa,53))),
                             SymEngine::pow(SymEngine::integer(2),SymEngine::integer(exponent-53)));
    return dickson_hybrid_topology(n_caps, SymEngine::Expression(exact), {}, dc_out, half_point);
}
