mruby-numeric-ext (int_sqrt): support bigint

This commit is contained in:
Yukihiro "Matz" Matsumoto
2025-04-16 23:08:23 +09:00
parent dee661913f
commit ddfe65763e
3 changed files with 91 additions and 7 deletions
+1
View File
@@ -248,6 +248,7 @@ void mrb_gc_free_bint(mrb_state *mrb, struct RBasic *x);
void mrb_bint_copy(mrb_state *mrb, mrb_value x, mrb_value y);
size_t mrb_bint_memsize(mrb_value x);
mrb_value mrb_bint_hash(mrb_state *mrb, mrb_value x);
mrb_value mrb_bint_sqrt(mrb_state *mrb, mrb_value x);
#endif
#endif /* MRUBY_INTERNAL_H */
+74
View File
@@ -1246,6 +1246,63 @@ mpz_gcd(mrb_state *mrb, mpz_t *gg, mpz_t *aa, mpz_t *bb)
}
#endif
static size_t
mpz_bit_length(const mpz_t *x)
{
if (x->sz == 0 || x->sn == 0) return 0;
// 最上位 limb を取得(リトルエンディアン: 末尾)
size_t i = x->sz - 1;
mp_limb high = x->p[i];
// high の中で最上位のビットを探す
size_t bits = 0;
while (high != 0) {
high >>= 1;
bits++;
}
return i * (sizeof(mp_limb) * 8) + bits;
}
static void
mpz_sqrt(mrb_state *mrb, mpz_t *z, mpz_t *x)
{
mrb_assert(x->sn >= 0);
if (x->sz == 0) {
// sqrt(0) = 0
z->sn = 0;
z->sz = 0;
return;
}
// 初期値の設定: bit-length の半分の位置に 1 を立てる
size_t xbits = mpz_bit_length(x);
size_t sbit = (xbits + 1) / 2;
mpz_t s, t;
mpz_init_set_int(mrb, &s, 1);
mpz_mul_2exp(mrb, &s, &s, sbit);
mpz_init(mrb, &t);
// ループ: s = (s + x / s) / 2
for (;;) {
mpz_mdiv(mrb, &t, x, &s); // t = x / s
mpz_add(mrb, &t, &t, &s); // t = s + x/s
mpz_div_2exp(mrb, &t, &t, 1); // t = (s + x/s) / 2
if (mpz_cmp(mrb, &t, &s) >= 0) {
break; // 収束
}
mpz_set(mrb, &s, &t);
}
mpz_move(mrb, z, &s);
mpz_clear(mrb, &t);
}
/* --- mruby functions --- */
/* initialize mpz_t from RBigint (not need to clear) */
static void
@@ -2014,6 +2071,23 @@ mrb_bint_memsize(mrb_value x)
return z.sz * sizeof(mp_limb);
}
mrb_value
mrb_bint_sqrt(mrb_state *mrb, mrb_value x)
{
mpz_t a;
bint_as_mpz(RBIGINT(x), &a);
if (a.sn < 0) {
mrb_raise(mrb, E_ARGUMENT_ERROR, "square root of negative number");
}
mpz_t z;
mpz_init(mrb, &z);
mpz_sqrt(mrb, &z, &a);
return bint_norm(mrb, bint_new(mrb, &z));
}
mrb_value
mrb_bint_hash(mrb_state *mrb, mrb_value x)
{
+16 -7
View File
@@ -283,14 +283,23 @@ isqrt(mrb_int n)
static mrb_value
int_sqrt(mrb_state *mrb, mrb_value self)
{
mrb_int n;
mrb_get_args(mrb, "i", &n);
if (n < 0) {
mrb_raise(mrb, E_ARGUMENT_ERROR, "non-negative integer required");
}
mrb_value arg = mrb_get_arg1(mrb);
mrb_int result = isqrt(n);
return mrb_int_value(mrb, result);
if (mrb_integer_p(arg)) {
mrb_int n = mrb_integer(arg);
if (n < 0) {
mrb_raise(mrb, E_ARGUMENT_ERROR, "non-negative integer required");
}
return mrb_int_value(mrb, isqrt(n));
}
#ifdef MRB_USE_BIGINT
else if (mrb_bigint_p(arg)) {
return mrb_bint_sqrt(mrb, arg);
}
#endif
else {
mrb_raise(mrb, E_TYPE_ERROR, "expected Integer");
}
}
void