diff --git a/include/mruby/internal.h b/include/mruby/internal.h index 83738506f..1d752d131 100644 --- a/include/mruby/internal.h +++ b/include/mruby/internal.h @@ -261,6 +261,9 @@ mrb_value mrb_bint_sqrt(mrb_state *mrb, mrb_value x); mrb_int mrb_bint_size(mrb_state *mrb, mrb_value bint); mrb_value mrb_bint_from_bytes(mrb_state *mrb, const uint8_t *bytes, mrb_int len); mrb_int mrb_bint_sign(mrb_state *mrb, mrb_value bint); +mrb_value mrb_bint_gcd(mrb_state *mrb, mrb_value x, mrb_value y); +mrb_value mrb_bint_lcm(mrb_state *mrb, mrb_value x, mrb_value y); +mrb_value mrb_bint_abs(mrb_state *mrb, mrb_value x); #endif #endif /* MRUBY_INTERNAL_H */ diff --git a/mrbgems/mruby-bigint/core/bigint.c b/mrbgems/mruby-bigint/core/bigint.c index faf725722..5354cd7fd 100644 --- a/mrbgems/mruby-bigint/core/bigint.c +++ b/mrbgems/mruby-bigint/core/bigint.c @@ -1515,7 +1515,6 @@ mpz_powm_i(mrb_state *mrb, mpz_t *zz, mpz_t *x, mrb_int ex, mpz_t *n) mpz_clear(mrb, &b); } -#ifdef MRB_USE_RATIONAL static void mpz_abs(mrb_state *mrb, mpz_t *x, mpz_t *y) { @@ -1833,7 +1832,6 @@ mpz_gcd(mrb_state *mrb, mpz_t *gg, mpz_t *aa, mpz_t *bb) mpz_move(mrb, gg, &b); mpz_clear(mrb, &a); } -#endif static size_t mpz_bits(const mpz_t *x) @@ -2844,3 +2842,69 @@ mrb_bint_reduce(mrb_state *mrb, mrb_value *xp, mrb_value *yp) *yp = mrb_obj_value(b2); } #endif + +mrb_value +mrb_bint_gcd(mrb_state *mrb, mrb_value x, mrb_value y) +{ + mpz_t r, a, b; + + mpz_init(mrb, &r); + bint_as_mpz(RBIGINT(x), &a); + bint_as_mpz(RBIGINT(y), &b); + + mpz_gcd(mrb, &r, &a, &b); + + struct RBigint *result = bint_new(mrb, &r); + mpz_clear(mrb, &r); + return bint_norm(mrb, result); +} + +mrb_value +mrb_bint_lcm(mrb_state *mrb, mrb_value x, mrb_value y) +{ + mpz_t gcd_val, x_mpz, y_mpz, abs_x, abs_y, product, result_mpz; + mrb_value zero = mrb_bint_new_int(mrb, 0); + + if (mrb_bint_cmp(mrb, x, zero) == 0 || mrb_bint_cmp(mrb, y, zero) == 0) { + return zero; + } + + mpz_init(mrb, &gcd_val); + mpz_init(mrb, &abs_x); + mpz_init(mrb, &abs_y); + mpz_init(mrb, &product); + mpz_init(mrb, &result_mpz); + + bint_as_mpz(RBIGINT(x), &x_mpz); + bint_as_mpz(RBIGINT(y), &y_mpz); + + mpz_abs(mrb, &abs_x, &x_mpz); + mpz_abs(mrb, &abs_y, &y_mpz); + + mpz_gcd(mrb, &gcd_val, &abs_x, &abs_y); + mpz_mul(mrb, &product, &abs_x, &abs_y); + mpz_mdiv(mrb, &result_mpz, &product, &gcd_val); + + mpz_clear(mrb, &gcd_val); + mpz_clear(mrb, &abs_x); + mpz_clear(mrb, &abs_y); + mpz_clear(mrb, &product); + + struct RBigint *result = bint_new(mrb, &result_mpz); + mpz_clear(mrb, &result_mpz); + return mrb_obj_value(result); +} + +mrb_value +mrb_bint_abs(mrb_state *mrb, mrb_value x) +{ + mpz_t a, result_mpz; + + mpz_init(mrb, &result_mpz); + bint_as_mpz(RBIGINT(x), &a); + mpz_abs(mrb, &result_mpz, &a); + + struct RBigint *result = bint_new(mrb, &result_mpz); + mpz_clear(mrb, &result_mpz); + return mrb_obj_value(result); +} diff --git a/mrbgems/mruby-numeric-ext/src/numeric_ext.c b/mrbgems/mruby-numeric-ext/src/numeric_ext.c index 445464438..5acc29ded 100644 --- a/mrbgems/mruby-numeric-ext/src/numeric_ext.c +++ b/mrbgems/mruby-numeric-ext/src/numeric_ext.c @@ -46,6 +46,90 @@ int_remainder(mrb_state *mrb, mrb_value x) mrb_value mrb_int_pow(mrb_state *mrb, mrb_value x, mrb_value y); +static mrb_int +mrb_int_gcd(mrb_int x, mrb_int y) +{ + if (x < 0) x = -x; + if (y < 0) y = -y; + + while (y != 0) { + mrb_int temp = y; + y = x % y; + x = temp; + } + + return x; +} + +/* + * call-seq: + * int.gcd(other_int) -> integer + * + * Returns the greatest common divisor of the two integers. + * The result is always positive. + */ +static mrb_value +int_gcd(mrb_state *mrb, mrb_value x) +{ + mrb_value y = mrb_get_arg1(mrb); + +#ifdef MRB_USE_BIGINT + if (mrb_bigint_p(x) || mrb_bigint_p(y)) { + if (!mrb_integer_p(y) && !mrb_bigint_p(y)) { + mrb_raisef(mrb, E_TYPE_ERROR, "can't convert %Y into Integer", y); + } + if (!mrb_bigint_p(x)) x = mrb_bint_new_int(mrb, mrb_integer(x)); + if (!mrb_bigint_p(y)) y = mrb_bint_new_int(mrb, mrb_integer(y)); + return mrb_bint_gcd(mrb, x, y); + } +#endif + + if (!mrb_integer_p(y)) { + mrb_raisef(mrb, E_TYPE_ERROR, "can't convert %Y into Integer", y); + } + return mrb_int_value(mrb, mrb_int_gcd(mrb_integer(x), mrb_integer(y))); +} + +/* + * call-seq: + * int.lcm(other_int) -> integer + * + * Returns the least common multiple of the two integers. + * The result is always positive. + */ +static mrb_value +int_lcm(mrb_state *mrb, mrb_value x) +{ + mrb_value y = mrb_get_arg1(mrb); + mrb_int a, b, gcd_val; + +#ifdef MRB_USE_BIGINT + if (mrb_bigint_p(x) || mrb_bigint_p(y)) { + if (!mrb_integer_p(y) && !mrb_bigint_p(y)) { + mrb_raisef(mrb, E_TYPE_ERROR, "can't convert %Y into Integer", y); + } + if (!mrb_bigint_p(x)) x = mrb_bint_new_int(mrb, mrb_integer(x)); + if (!mrb_bigint_p(y)) y = mrb_bint_new_int(mrb, mrb_integer(y)); + return mrb_bint_lcm(mrb, x, y); + } +#endif + + if (!mrb_integer_p(y)) { + mrb_raisef(mrb, E_TYPE_ERROR, "can't convert %Y into Integer", y); + } + + a = mrb_integer(x); + b = mrb_integer(y); + + if (a == 0 || b == 0) return mrb_int_value(mrb, 0); + + gcd_val = mrb_int_gcd(a, b); + if (a < 0) a = -a; + if (b < 0) b = -b; + + return mrb_int_value(mrb, (a / gcd_val) * b); +} + /* * call-seq: * integer.pow(numeric) -> numeric @@ -353,6 +437,8 @@ mrb_mruby_numeric_ext_gem_init(mrb_state* mrb) mrb_define_method_id(mrb, ic, MRB_SYM(size), int_size, MRB_ARGS_NONE()); mrb_define_method_id(mrb, ic, MRB_SYM_Q(odd), int_odd, MRB_ARGS_NONE()); mrb_define_method_id(mrb, ic, MRB_SYM_Q(even), int_even, MRB_ARGS_NONE()); + mrb_define_method_id(mrb, ic, MRB_SYM(gcd), int_gcd, MRB_ARGS_REQ(1)); + mrb_define_method_id(mrb, ic, MRB_SYM(lcm), int_lcm, MRB_ARGS_REQ(1)); mrb_define_class_method_id(mrb, ic, MRB_SYM(sqrt), int_sqrt, MRB_ARGS_REQ(1)); #ifndef MRB_NO_FLOAT diff --git a/mrbgems/mruby-numeric-ext/test/numeric.rb b/mrbgems/mruby-numeric-ext/test/numeric.rb index bb4c461d6..4578afd26 100644 --- a/mrbgems/mruby-numeric-ext/test/numeric.rb +++ b/mrbgems/mruby-numeric-ext/test/numeric.rb @@ -26,6 +26,30 @@ assert('Integer#pow') do assert_equal(361, 9.pow(1024,1000)) end +assert('Integer#gcd') do + assert_equal(1, 2.gcd(3)) + assert_equal(5, 10.gcd(15)) + assert_equal(6, 24.gcd(18)) + assert_equal(7, 7.gcd(0)) + assert_equal(7, 0.gcd(7)) + assert_equal(0, 0.gcd(0)) + assert_equal(5, (-10).gcd(15)) + assert_equal(5, 10.gcd(-15)) + assert_equal(5, (-10).gcd(-15)) +end + +assert('Integer#lcm') do + assert_equal(6, 2.lcm(3)) + assert_equal(30, 10.lcm(15)) + assert_equal(72, 24.lcm(18)) + assert_equal(0, 7.lcm(0)) + assert_equal(0, 0.lcm(7)) + assert_equal(0, 0.lcm(0)) + assert_equal(30, (-10).lcm(15)) + assert_equal(30, 10.lcm(-15)) + assert_equal(30, (-10).lcm(-15)) +end + assert('Integer#ceildiv') do assert_equal(0, 0.ceildiv(3)) assert_equal(1, 1.ceildiv(3))