#include <systemc>

using namespace sc_core;
using namespace sc_dt;
using namespace std;

SC_MODULE(Dff) {
    sc_in_clk clk;
    sc_in<bool> d, reset;
    sc_out<bool> q;
    SC_CTOR(Dff) {
        SC_METHOD(store);
        sensitive << clk.pos() << reset.pos();
    }
private:
    void store() {
        if (reset.read()) {
            q.write(false);
        }
        else {
            q.write(d.read());
        }
    }
};

SC_MODULE(Xor) {
    sc_in<bool> i0, i1;
    sc_out<bool> o;
    SC_CTOR(Xor) {
        SC_METHOD(run);
        sensitive << i0 << i1;
    }
private:
    void run() {
        o.write(i0.read() ^ i1.read());
    }
};

SC_MODULE(And) {
    sc_in<bool> i0, i1;
    sc_out<bool> o;
    SC_CTOR(And) {
        SC_METHOD(run);
        sensitive << i0 << i1;
    }
private:
    void run() {
        o.write(i0.read() && i1.read());
    }
};

template<int N>
SC_MODULE(Collect) {
    sc_in<bool> i[N];
    sc_out<sc_uint<N>> o;
    SC_CTOR(Collect) {
        SC_METHOD(run);
        for (int n = 0; n < N; ++n) {
            sensitive << i[n];
        }
    }
private:
    void run() {
        sc_uint<N> out = 0;
        for (int n = N - 1; n >= 0; --n) {
            out = out * 2 + i[n].read();
        }
        o.write(out);
    }
};

template<int N>
SC_MODULE(counterNbit) {
    sc_in_clk clk;
    sc_in<bool> reset, enable;
    sc_out<sc_uint<N>> counter_out;
    SC_CTOR(counterNbit): collect("collect") {
        for (int n = 0; n < N; ++n) {
            dff[n] = new Dff("");
            dff[n]->clk(clk);
            dff[n]->reset(reset);
            dff[n]->d(d[n]);
            dff[n]->q(q[n]);
            xor[n] = new Xor("");
            xor[n]->i0(q[n]);
            if (n == 0) {
                xor[n]->i1(enable);
            } else {
                xor[n]->i1(a[n-1]);
            }
            xor[n]->o(d[n]);
            if (n < N - 1) {
                and[n] = new And("");
                and[n]->i0(q[n]);
                if (n == 0) {
                    and[n]->i1(enable);
                }
                else {
                    and[n]->i1(a[n-1]);
                }
                and[n]->o(a[n]);
            }
            collect.i[n](q[n]);
        }
        collect.o(counter_out);
    }
    static const int number_of_bits = N;
private:
    Dff* dff[N];
    Xor* xor[N];
    And* and[N - 1];
    Collect<N> collect;
    sc_signal<bool> d[N], q[N], a[N - 1];
};

SC_MODULE(Testbench) {
    sc_in_clk clk;
    sc_out<bool> reset, enable;
    SC_CTOR(Testbench){
        SC_THREAD(testprocess);
        sensitive << clk.pos();
    }
private:
    void testprocess() {
        enable.write(true);
        reset.write(true);
        wait(25, SC_NS);
        reset.write(false);
        wait(200, SC_NS);
        enable.write(false);
        wait(50, SC_NS);
        reset.write(true);
        wait(50, SC_NS);
        reset.write(false);
    }
};

int sc_main(int argc, char *argv[]) {
    sc_clock clock("clock", 20, SC_NS);
    sc_signal<bool> reset, enable;
    Testbench testbench("testbench");
    testbench.reset(reset);
    testbench.enable(enable);
    testbench.clk(clock);

    counterNbit<2> counter4("counter4");
    sc_signal<sc_uint<counter4.number_of_bits>> count4;
    counter4.clk(clock);
    counter4.reset(reset);
    counter4.enable(enable);
    counter4.counter_out(count4);

    counterNbit<3> counter8("counter8");
    sc_signal<sc_uint<counter8.number_of_bits>> count8;
    counter8.clk(clock);
    counter8.reset(reset);
    counter8.enable(enable);
    counter8.counter_out(count8);
    
    // Record (trace) signals for verification
    auto tf = sc_create_vcd_trace_file("trace");
    tf->set_time_unit(1, SC_NS);
    sc_trace(tf, clock, "clock");
    sc_trace(tf, reset, "reset");
    sc_trace(tf, enable, "enable");
    sc_trace(tf, count4, "count4");
    sc_trace(tf, count8, "count8");

    // Start the simulation for 200ns 
    sc_start(500, SC_NS);

    sc_close_vcd_trace_file(tf);
    cin.get();
    return 0;
}
