authorgravatar for wander4096@gmail.comtison <wander4096@gmail.com> 2023-11-08 14:39:08+08:00
committergravatar for noreply@github.comGitHub <noreply@github.com> 2023-11-08 08:39:08+02:00
logee47643b6eb175f4414a50c19afa0c22b2d15949
treecee1d8807d87e27487cbe4b301c3b9eb4041fc96
parent77bc8e7b67a1be167f87f19c91a8f606ad5745ad
signaturebadge-question-mark Signed by PGP key 4AEE18F83AFDEB23

std.math.big: fix sqrt with bits > limb_bits

Signed-off-by: tison <wander4096@gmail.com>

2 files changed, 21 insertions(+), 2 deletions(-)

lib/std/math/big/int.zig+2-2
...@@ -69,7 +69,7 @@ pub fn calcPowLimbsBufferLen(a_bit_count: usize, y: usize) usize {...@@ -69,7 +69,7 @@ pub fn calcPowLimbsBufferLen(a_bit_count: usize, y: usize) usize {
69pub fn calcSqrtLimbsBufferLen(a_bit_count: usize) usize {69pub fn calcSqrtLimbsBufferLen(a_bit_count: usize) usize {
70 const a_limb_count = (a_bit_count - 1) / limb_bits + 1;70 const a_limb_count = (a_bit_count - 1) / limb_bits + 1;
71 const shift = (a_bit_count + 1) / 2;71 const shift = (a_bit_count + 1) / 2;
72 const u_s_rem_limb_count = 1 + ((shift - 1) / limb_bits + 1);72 const u_s_rem_limb_count = 1 + ((shift / limb_bits) + 1);
73 return a_limb_count + 3 * u_s_rem_limb_count + calcDivLimbsBufferLen(a_limb_count, u_s_rem_limb_count);73 return a_limb_count + 3 * u_s_rem_limb_count + calcDivLimbsBufferLen(a_limb_count, u_s_rem_limb_count);
74}74}
7575
...@@ -1380,7 +1380,7 @@ pub const Mutable = struct {...@@ -1380,7 +1380,7 @@ pub const Mutable = struct {
1380 var u = b: {1380 var u = b: {
1381 const start = buf_index;1381 const start = buf_index;
1382 const shift = (a.bitCountAbs() + 1) / 2;1382 const shift = (a.bitCountAbs() + 1) / 2;
1383 buf_index += 1 + ((shift - 1) / limb_bits + 1);1383 buf_index += 1 + ((shift / limb_bits) + 1);
1384 var m = Mutable.init(limbs_buffer[start..buf_index], 1);1384 var m = Mutable.init(limbs_buffer[start..buf_index], 1);
1385 m.shiftLeft(m.toConst(), shift); // u must be >= ⌊√a⌋, and should be as small as possible for efficiency1385 m.shiftLeft(m.toConst(), shift); // u must be >= ⌊√a⌋, and should be as small as possible for efficiency
1386 break :b m;1386 break :b m;
lib/std/math/big/int_test.zig+19
...@@ -3161,6 +3161,7 @@ test "big.int.Managed sqrt(0) = 0" {...@@ -3161,6 +3161,7 @@ test "big.int.Managed sqrt(0) = 0" {
3161 try a.setString(10, "0");3161 try a.setString(10, "0");
31623162
3163 try res.sqrt(&a);3163 try res.sqrt(&a);
3164 try testing.expectEqual(@as(i32, 0), try res.to(i32));
3164}3165}
31653166
3166test "big.int.Managed sqrt(-1) = error" {3167test "big.int.Managed sqrt(-1) = error" {
...@@ -3175,3 +3176,21 @@ test "big.int.Managed sqrt(-1) = error" {...@@ -3175,3 +3176,21 @@ test "big.int.Managed sqrt(-1) = error" {
31753176
3176 try testing.expectError(error.SqrtOfNegativeNumber, res.sqrt(&a));3177 try testing.expectError(error.SqrtOfNegativeNumber, res.sqrt(&a));
3177}3178}
3179
3180test "big.int.Managed sqrt(n) succeed with res.bitCountAbs() >= usize bits" {
3181 const allocator = testing.allocator;
3182 var a = try Managed.initSet(allocator, 1);
3183 defer a.deinit();
3184
3185 var res = try Managed.initSet(allocator, 1);
3186 defer res.deinit();
3187
3188 // a.bitCountAbs() = 127 so the first attempt has 64 bits >= usize bits
3189 try a.setString(10, "136036462105870278006290938611834481486");
3190 try res.sqrt(&a);
3191
3192 var expected = try Managed.initSet(allocator, 1);
3193 defer expected.deinit();
3194 try expected.setString(10, "11663466984815033033");
3195 try std.testing.expectEqual(std.math.Order.eq, expected.order(res));
3196}