#include <systemc>

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

template <typename T>
SC_MODULE(gcd) {
    sc_in_clk clk;
    sc_in<bool> go_i;
    sc_in<T> x_i, y_i;
    sc_out<bool> done_o;
    sc_out<T> r_o;
    SC_CTOR(gcd) {
        SC_THREAD(run);
        sensitive << clk.pos();
    }
private:
    void run() {
        wait();
        while(1) {
            do {
                wait();
            } while (!go_i.read());
            T x = x_i.read();
            T y = y_i.read();
            wait();
            while (go_i.read() && x != y) {
                if (x > y) {
                    x -= y;
                } else {
                    y -= x;
                }
                wait();
            }
            if (go_i.read()) {
                r_o.write(x);
                done_o.write(true);
            }
            do {
                wait();
            } while (go_i.read());
            done_o.write(false);
        }
    }
};

template <typename T>
SC_MODULE(tb_gcd) {
    sc_in<bool> done_i;
    sc_in<T> r_i;
    sc_out<bool> go_o; 
    sc_out<T> x_o, y_o;
    SC_CTOR(tb_gcd) {
        SC_THREAD(run);
        sensitive << done_i;
    }
private:
    void check(const T& x, const T& y, const T& r) {
        auto start_time_stamp = sc_time_stamp();
        x_o.write(x);
        y_o.write(y);
        go_o.write(true);
        wait();
        assert(r_i.read() == r);
        auto end_time_stamp = sc_time_stamp();
        cout << "@: " << sc_time_stamp() 
             << ": gcd(" << x << "," << y << ") = " << r 
             << " duration: " << end_time_stamp - start_time_stamp << endl;
        go_o.write(false);
        wait();
    }
    void run() {
        check(0, 0, 0);
        check(234, 96, 6);
        check(12345, 67891, 1);
        check(12345, 67890, 15);
        check(12345, 12345, 12345);

        x_o.write(0);
        y_o.write(0);
        sc_stop();
    }
};

int sc_main(int argc, char *argv[]) {
    gcd<unsigned int> sut("sut");
    tb_gcd<unsigned int> tb("tb");

    sc_clock clk("clk", 10, SC_NS);
    sc_signal<bool> go, done;
    sc_buffer<unsigned int> x, y, r;
    
    sut.clk(clk);
    sut.go_i(go);
    sut.x_i(x);
    sut.y_i(y);
    sut.done_o(done);
    sut.r_o(r);

    tb.go_o(go);
    tb.x_o(x);
    tb.y_o(y);
    tb.done_i(done);
    tb.r_i(r);

    auto tf = sc_create_vcd_trace_file("trace");
    tf->set_time_unit(1, SC_NS);
    sc_trace(tf, clk, "clk");
    sc_trace(tf, go, "go");
    sc_trace(tf, x, "x");
    sc_trace(tf, y, "y");
    sc_trace(tf, done, "done");
    sc_trace(tf, r, "r");

    sc_start();

    sc_close_vcd_trace_file(tf);

    cin.get();
    return 0;
}
