From 4dcdf78f8ae834db2e81b274aa432913f37885f5 Mon Sep 17 00:00:00 2001 From: "Yukihiro \"Matz\" Matsumoto" Date: Thu, 17 Mar 2022 23:25:50 +0900 Subject: [PATCH] numeric_ext.c: add `Integer#pow()` method. Which takes optional second argument of modulo. --- mrbgems/mruby-numeric-ext/src/numeric_ext.c | 58 ++++++++++++++++++--- mrbgems/mruby-numeric-ext/test/numeric.rb | 6 +++ src/numeric.c | 49 ++++++++--------- 3 files changed, 81 insertions(+), 32 deletions(-) diff --git a/mrbgems/mruby-numeric-ext/src/numeric_ext.c b/mrbgems/mruby-numeric-ext/src/numeric_ext.c index 2f2a6c755..4220f9207 100644 --- a/mrbgems/mruby-numeric-ext/src/numeric_ext.c +++ b/mrbgems/mruby-numeric-ext/src/numeric_ext.c @@ -2,6 +2,9 @@ #include #include +void mrb_int_zerodiv(mrb_state *mrb); +void mrb_int_overflow(mrb_state *mrb, const char *reason); + /* * call-seq: * int.allbits?(mask) -> true or false @@ -50,12 +53,6 @@ int_nobits(mrb_state *mrb, mrb_value self) return mrb_bool_value((n & m) == 0); } -static void -zerodiv(mrb_state *mrb) -{ - mrb_raise(mrb, E_ZERODIV_ERROR, "divided by 0"); -} - /* * call-seq: * num.remainder(numeric) -> real @@ -73,7 +70,7 @@ int_remainder(mrb_state *mrb, mrb_value x) a = mrb_integer(x); if (mrb_integer_p(y)) { b = mrb_integer(y); - if (b == 0) zerodiv(mrb); + if (b == 0) mrb_int_zerodiv(mrb); if (a == MRB_INT_MIN && b == -1) return mrb_fixnum_value(0); return mrb_int_value(mrb, a % b); } @@ -88,6 +85,49 @@ int_remainder(mrb_state *mrb, mrb_value x) #endif } +mrb_value mrb_int_pow(mrb_state *mrb, mrb_value x); + +/* + * call-seq: + * integer.pow(numeric) -> numeric + * integer.pow(integer, integer) -> integer + * + * Returns (modular) exponentiation as: + * + * a.pow(b) #=> same as a**b + * a.pow(b, m) #=> same as (a**b) % m, but avoids huge temporary values + */ +static mrb_value +int_powm(mrb_state *mrb, mrb_value x) +{ + mrb_int base, exp, mod, result = 1; + + if (mrb_get_argc(mrb) == 1) { + return mrb_int_pow(mrb, x); + } + mrb_get_args(mrb, "ii", &exp, &mod); + if (exp < 0) mrb_raise(mrb, E_TYPE_ERROR, "int.pow(n,m): n must be positive"); + if (mod < 0) mrb_raise(mrb, E_TYPE_ERROR, "int.pow(n,m): m must be positive"); + if (mod == 0) mrb_int_zerodiv(mrb); + if (mod == 1) return mrb_fixnum_value(0); + base = mrb_integer(x); + for (;;) { + if (exp & 1) { + if (mrb_int_mul_overflow(result, base, &result)) { + mrb_int_overflow(mrb, "pow"); + } + result %= mod; + } + exp >>= 1; + if (exp == 0) break; + if (mrb_int_mul_overflow(base, base, &base)) { + mrb_int_overflow(mrb, "pow"); + } + base %= mod; + } + return mrb_int_value(mrb, result); +} + #ifndef MRB_NO_FLOAT static mrb_value flo_remainder(mrb_state *mrb, mrb_value self) @@ -96,7 +136,7 @@ flo_remainder(mrb_state *mrb, mrb_value self) a = mrb_float(self); mrb_get_args(mrb, "f", &b); - if (b == 0) zerodiv(mrb); + if (b == 0) mrb_int_zerodiv(mrb); if (isinf(b)) return mrb_float_value(mrb, a); return mrb_float_value(mrb, a-b*trunc(a/b)); } @@ -114,6 +154,8 @@ mrb_mruby_numeric_ext_gem_init(mrb_state* mrb) mrb_define_alias(mrb, i, "modulo", "%"); mrb_define_method(mrb, i, "remainder", int_remainder, MRB_ARGS_REQ(1)); + mrb_define_method_id(mrb, i, MRB_SYM(pow), int_powm, MRB_ARGS_ARG(1,1)); + #ifndef MRB_NO_FLOAT struct RClass *f = mrb_class_get(mrb, "Float"); diff --git a/mrbgems/mruby-numeric-ext/test/numeric.rb b/mrbgems/mruby-numeric-ext/test/numeric.rb index efdf25a34..7ff6bdb97 100644 --- a/mrbgems/mruby-numeric-ext/test/numeric.rb +++ b/mrbgems/mruby-numeric-ext/test/numeric.rb @@ -19,3 +19,9 @@ assert('Integer#nonzero?') do assert_equal nil, 0.nonzero? assert_equal 1000, 1000.nonzero? end + +assert('Integer#pow') do + assert_equal(8, 2.pow(3)) + assert_equal(-8, (-2).pow(3)) + assert_equal(361, 9.pow(1024,1000)) +end diff --git a/src/numeric.c b/src/numeric.c index 109833bb9..aa07853ac 100644 --- a/src/numeric.c +++ b/src/numeric.c @@ -32,14 +32,14 @@ mrb_value mrb_rational_add(mrb_state *mrb, mrb_value x, mrb_value y); mrb_value mrb_rational_sub(mrb_state *mrb, mrb_value x, mrb_value y); mrb_value mrb_rational_mul(mrb_state *mrb, mrb_value x, mrb_value y); -static void -int_overflow(mrb_state *mrb, const char *reason) +void +mrb_int_overflow(mrb_state *mrb, const char *reason) { mrb_raisef(mrb, E_RANGE_ERROR, "integer overflow in %s", reason); } -static void -int_zerodiv(mrb_state *mrb) +void +mrb_int_zerodiv(mrb_state *mrb) { mrb_raise(mrb, E_ZERODIV_ERROR, "divided by 0"); } @@ -53,8 +53,8 @@ int_zerodiv(mrb_state *mrb) * * 2.0**3 #=> 8.0 */ -static mrb_value -int_pow(mrb_state *mrb, mrb_value x) +mrb_value +mrb_int_pow(mrb_state *mrb, mrb_value x) { mrb_int base = mrb_integer(x); mrb_int result = 1; @@ -78,32 +78,33 @@ int_pow(mrb_state *mrb, mrb_value x) #ifndef MRB_NO_FLOAT return mrb_float_value(mrb, pow((double)base, (double)exp)); #else - int_overflow(mrb, "negative power"); + mrb_int_overflow(mrb, "negative power"); #endif } for (;;) { if (exp & 1) { if (mrb_int_mul_overflow(result, base, &result)) { - int_overflow(mrb, "power"); + mrb_int_overflow(mrb, "power"); } } exp >>= 1; if (exp == 0) break; if (mrb_int_mul_overflow(base, base, &base)) { - int_overflow(mrb, "power"); + mrb_int_overflow(mrb, "power"); } } return mrb_int_value(mrb, result); } +#define int_pow mrb_int_pow mrb_int mrb_div_int(mrb_state *mrb, mrb_int x, mrb_int y) { if (y == 0) { - int_zerodiv(mrb); + mrb_int_zerodiv(mrb); } else if(x == MRB_INT_MIN && y == -1) { - int_overflow(mrb, "division"); + mrb_int_overflow(mrb, "division"); } else { mrb_int div = x / y; @@ -165,7 +166,7 @@ int_idiv(mrb_state *mrb, mrb_value x) mrb_get_args(mrb, "i", &y); if (y == 0) { - int_zerodiv(mrb); + mrb_int_zerodiv(mrb); } return mrb_int_value(mrb, mrb_integer(x) / y); } @@ -180,7 +181,7 @@ int_quo(mrb_state *mrb, mrb_value xv) mrb_get_args(mrb, "f", &y); if (y == 0) { - int_zerodiv(mrb); + mrb_int_zerodiv(mrb); } return mrb_float_value(mrb, mrb_integer(xv) / y); #endif @@ -412,7 +413,7 @@ flodivmod(mrb_state *mrb, double x, double y, mrb_float *divp, mrb_float *modp) goto exit; } if (y == 0.0) { - int_zerodiv(mrb); + mrb_int_zerodiv(mrb); } if (isinf(y) && !isinf(x)) { mod = x; @@ -553,7 +554,7 @@ static mrb_value int64_value(mrb_state *mrb, int64_t v) { if (!TYPED_FIXABLE(v,int64_t)) { - int_overflow(mrb, "bit operation"); + mrb_int_overflow(mrb, "bit operation"); } return mrb_fixnum_value((mrb_int)v); } @@ -1014,7 +1015,7 @@ mrb_int_mul(mrb_state *mrb, mrb_value x, mrb_value y) if (a == 0) return x; b = mrb_integer(y); if (mrb_int_mul_overflow(a, b, &c)) { - int_overflow(mrb, "multiplication"); + mrb_int_overflow(mrb, "multiplication"); } return mrb_int_value(mrb, c); } @@ -1058,10 +1059,10 @@ static void intdivmod(mrb_state *mrb, mrb_int x, mrb_int y, mrb_int *divp, mrb_int *modp) { if (y == 0) { - int_zerodiv(mrb); + mrb_int_zerodiv(mrb); } else if(x == MRB_INT_MIN && y == -1) { - int_overflow(mrb, "division"); + mrb_int_overflow(mrb, "division"); } else { mrb_int div = x / y; @@ -1094,7 +1095,7 @@ int_mod(mrb_state *mrb, mrb_value x) a = mrb_integer(x); if (mrb_integer_p(y)) { b = mrb_integer(y); - if (b == 0) int_zerodiv(mrb); + if (b == 0) mrb_int_zerodiv(mrb); if (a == MRB_INT_MIN && b == -1) return mrb_fixnum_value(0); mrb_int mod = a % b; if ((a < 0) != (b < 0) && mod != 0) { @@ -1338,7 +1339,7 @@ int_lshift(mrb_state *mrb, mrb_value x) val = mrb_integer(x); if (val == 0) return x; if (!mrb_num_shift(mrb, val, width, &val)) { - int_overflow(mrb, "bit shift"); + mrb_int_overflow(mrb, "bit shift"); } return mrb_int_value(mrb, val); } @@ -1362,9 +1363,9 @@ int_rshift(mrb_state *mrb, mrb_value x) } val = mrb_integer(x); if (val == 0) return x; - if (width == MRB_INT_MIN) int_overflow(mrb, "bit shift"); + if (width == MRB_INT_MIN) mrb_int_overflow(mrb, "bit shift"); if (!mrb_num_shift(mrb, val, -width, &val)) { - int_overflow(mrb, "bit shift"); + mrb_int_overflow(mrb, "bit shift"); } return mrb_int_value(mrb, val); } @@ -1434,7 +1435,7 @@ mrb_int_add(mrb_state *mrb, mrb_value x, mrb_value y) if (a == 0) return y; b = mrb_integer(y); if (mrb_int_add_overflow(a, b, &c)) { - int_overflow(mrb, "addition"); + mrb_int_overflow(mrb, "addition"); } return mrb_int_value(mrb, c); } @@ -1484,7 +1485,7 @@ mrb_int_sub(mrb_state *mrb, mrb_value x, mrb_value y) b = mrb_integer(y); if (mrb_int_sub_overflow(a, b, &c)) { - int_overflow(mrb, "subtraction"); + mrb_int_overflow(mrb, "subtraction"); } return mrb_int_value(mrb, c); }