GitHub

@@ -942,7 +942,41 @@ mpz_pow(mrb_state *mrb, mpz_t *zz, mpz_t *x, mrb_int e)

942942

}

943943944944

static void

945-

mpz_powm(mrb_state *mrb, mpz_t *zz, mpz_t *x, mrb_int ex, mpz_t *n)

945+

mpz_powm(mrb_state *mrb, mpz_t *zz, mpz_t *x, mpz_t *ex, mpz_t *n)

946+

{

947+

mpz_t t, b;

948+949+

if (uzero(ex)) {

950+

mpz_set_int(mrb, zz, 1);

951+

return;

952+

}

953+954+

if (ex->sn < 0) {

955+

return;

956+

}

957+958+

mpz_init_set_int(mrb, &t, 1);

959+

mpz_init_set(mrb, &b, x);

960+961+

size_t len = digits(ex);

962+

for (size_t i=0; i<len; i++) {

963+

mp_limb e = ex->p[i];

964+

for (size_t j=0; j<sizeof(mp_limb)*8; j++) {

965+

if ((e & 1) == 1) {

966+

mpz_mul(mrb, &t, &t, &b);

967+

mpz_mod(mrb, &t, &t, n);

968+

}

969+

e >>= 1;

970+

mpz_mul(mrb, &b, &b, &b);

971+

mpz_mod(mrb, &b, &b, n);

972+

}

973+

}

974+

mpz_move(mrb, zz, &t);

975+

mpz_clear(mrb, &b);

976+

}

977+978+

static void

979+

mpz_powm_i(mrb_state *mrb, mpz_t *zz, mpz_t *x, mrb_int ex, mpz_t *n)

946980

{

947981

mpz_t t, b;

948982

@@ -1429,31 +1463,36 @@ mrb_bint_pow(mrb_state *mrb, mrb_value x, mrb_value y)

14291463

}

1430146414311465

mrb_value

1432-

mrb_bint_powm(mrb_state *mrb, mrb_value x, mrb_int exp, mrb_value mod)

1466+

mrb_bint_powm(mrb_state *mrb, mrb_value x, mrb_value exp, mrb_value mod)

14331467

{

14341468

struct RBigint *b = RBIGINT(x);

1435-

switch (mrb_type(mod)) {

1436-

case MRB_TT_INTEGER:

1437-

{

1438-

mrb_int m = mrb_integer(mod);

1439-

if (m == 0) mrb_int_zerodiv(mrb);

1440-

struct RBigint *b2 = bint_new_int(mrb, m);

1441-

struct RBigint *b3 = bint_new(mrb);

1442-

mpz_powm(mrb, &b3->mp, &b->mp, exp, &b2->mp);

1443-

return bint_norm(mrb, b3);

1469+

struct RBigint *b2, *b3;

1470+1471+

if (mrb_bigint_p(mod)) {

1472+

b2 = RBIGINT(mod);

1473+

if (uzero(&b2->mp)) mrb_int_zerodiv(mrb);

1474+

}

1475+

else {

1476+

mrb_int m = mrb_integer(mod);

1477+

if (m == 0) mrb_int_zerodiv(mrb);

1478+

b2 = bint_new_int(mrb, m);

1479+

}

1480+

b3 = bint_new(mrb);

1481+

if (mrb_bigint_p(exp)) {

1482+

struct RBigint *be = RBIGINT(exp);

1483+

if (be->mp.sn < 0) {

1484+

mrb_raise(mrb, E_ARGUMENT_ERROR, "int.pow(n,m): n must be positive");

14441485

}

1445-

case MRB_TT_BIGINT:

1446-

{

1447-

struct RBigint *b2 = RBIGINT(mod);

1448-

struct RBigint *b3 = bint_new(mrb);

1449-

if (uzero(&b2->mp)) mrb_int_zerodiv(mrb);

1450-

mpz_powm(mrb, &b3->mp, &b->mp, exp, &b2->mp);

1451-

return bint_norm(mrb, b3);

1486+

mpz_powm(mrb, &b3->mp, &b->mp, &be->mp, &b2->mp);

1487+

}

1488+

else {

1489+

mrb_int e = mrb_integer(exp);

1490+

if (e < 0) {

1491+

mrb_raise(mrb, E_ARGUMENT_ERROR, "int.pow(n,m): n must be positive");

14521492

}

1453-

default:

1454-

mrb_raisef(mrb, E_TYPE_ERROR, "%v cannot be convert to integer", mod);

1493+

mpz_powm_i(mrb, &b3->mp, &b->mp, e, &b2->mp);

14551494

}

1456-

return mrb_nil_value();

1495+

return bint_norm(mrb, b3);

14571496

}

1458149714591498

mrb_value

Read the original on github.com ↗