numeric_ext.c: add Integer#pow() method.

Which takes optional second argument of modulo.
This commit is contained in:
Yukihiro "Matz" Matsumoto
2022-03-17 23:25:50 +09:00
parent 6b2f08d933
commit 4dcdf78f8a
3 changed files with 81 additions and 32 deletions
+50 -8
View File
@@ -2,6 +2,9 @@
#include <mruby/numeric.h>
#include <mruby/presym.h>
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");
@@ -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
+25 -24
View File
@@ -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);
}