bignum.cpp raw
1 // Copyright (c) 2025 The Limenka developers
2 // Distributed under the MIT software license, see the accompanying
3 // file COPYING or http://www.opensource.org/licenses/mit-license.php.
4
5 #include <crypto/bignum.h>
6
7 #include <algorithm>
8 #include <cassert>
9 #include <cstring>
10
11 void BigNum::trim()
12 {
13 while (m_limbs.size() > 1 && m_limbs.back() == 0) {
14 m_limbs.pop_back();
15 }
16 }
17
18 BigNum::BigNum(const std::vector<uint8_t>& bytes, bool big_endian)
19 {
20 if (bytes.empty()) { m_limbs = {0}; return; }
21 size_t nlimbs = (bytes.size() + 3) / 4;
22 m_limbs.resize(nlimbs, 0);
23 if (big_endian) {
24 for (size_t i = 0; i < bytes.size(); i++) {
25 size_t limb_pos = bytes.size() - 1 - i; // position of byte i from the LSB
26 size_t li = limb_pos / 4;
27 size_t shift = 8 * (limb_pos % 4);
28 m_limbs[li] |= uint32_t(bytes[i]) << shift;
29 }
30 } else {
31 std::memcpy(m_limbs.data(), bytes.data(), bytes.size());
32 }
33 trim();
34 }
35
36 BigNum::BigNum(uint32_t val)
37 {
38 m_limbs = {val};
39 trim();
40 }
41
42 std::vector<uint8_t> BigNum::to_bytes(size_t width) const
43 {
44 size_t nbytes = m_limbs.size() * 4;
45 if (width > 0 && width > nbytes) nbytes = width;
46 std::vector<uint8_t> out(nbytes, 0);
47 for (size_t i = 0; i < m_limbs.size(); i++) {
48 uint32_t v = m_limbs[i];
49 out[i * 4 + 0] = uint8_t(v);
50 out[i * 4 + 1] = uint8_t(v >> 8);
51 out[i * 4 + 2] = uint8_t(v >> 16);
52 out[i * 4 + 3] = uint8_t(v >> 24);
53 }
54 return out;
55 }
56
57 bool BigNum::is_zero() const { return m_limbs.size() == 1 && m_limbs[0] == 0; }
58 bool BigNum::is_one() const { return m_limbs.size() == 1 && m_limbs[0] == 1; }
59
60 size_t BigNum::bit_length() const
61 {
62 if (is_zero()) return 0;
63 uint32_t top = m_limbs.back();
64 size_t bits = (m_limbs.size() - 1) * 32;
65 while (top) { bits++; top >>= 1; }
66 return bits;
67 }
68
69 int BigNum::compare(const BigNum& other) const
70 {
71 if (m_limbs.size() != other.m_limbs.size())
72 return m_limbs.size() < other.m_limbs.size() ? -1 : 1;
73 for (size_t i = m_limbs.size(); i-- > 0; )
74 if (m_limbs[i] != other.m_limbs[i])
75 return m_limbs[i] < other.m_limbs[i] ? -1 : 1;
76 return 0;
77 }
78
79 BigNum& BigNum::operator+=(const BigNum& other)
80 {
81 size_t n = std::max(m_limbs.size(), other.m_limbs.size());
82 m_limbs.resize(n + 1, 0);
83 uint64_t carry = 0;
84 for (size_t i = 0; i < n; i++) {
85 uint64_t sum = uint64_t(m_limbs[i]) + uint64_t(other.limb(i)) + carry;
86 m_limbs[i] = uint32_t(sum);
87 carry = sum >> 32;
88 }
89 m_limbs[n] = uint32_t(carry);
90 trim();
91 return *this;
92 }
93
94 BigNum& BigNum::operator-=(const BigNum& other)
95 {
96 assert(compare(other) >= 0);
97 int64_t borrow = 0;
98 for (size_t i = 0; i < m_limbs.size(); i++) {
99 int64_t diff = int64_t(m_limbs[i]) - int64_t(other.limb(i)) - borrow;
100 if (diff < 0) { diff += (uint64_t(1) << 32); borrow = 1; }
101 else { borrow = 0; }
102 m_limbs[i] = uint32_t(diff);
103 }
104 trim();
105 return *this;
106 }
107
108 BigNum& BigNum::operator*=(const BigNum& other)
109 {
110 size_t n = m_limbs.size(), m = other.m_limbs.size();
111 std::vector<uint32_t> result(n + m, 0);
112 for (size_t i = 0; i < n; i++) {
113 uint64_t carry = 0;
114 for (size_t j = 0; j < m; j++) {
115 uint64_t prod = uint64_t(m_limbs[i]) * uint64_t(other.m_limbs[j])
116 + uint64_t(result[i + j]) + carry;
117 result[i + j] = uint32_t(prod);
118 carry = prod >> 32;
119 }
120 result[i + m] = uint32_t(carry);
121 }
122 m_limbs = std::move(result);
123 trim();
124 return *this;
125 }
126
127 void BigNum::div_rem(const BigNum& num, const BigNum& den,
128 BigNum& quot, BigNum& rem)
129 {
130 assert(!den.is_zero());
131 if (num.compare(den) < 0) { quot = BigNum(0u); rem = num; return; }
132
133 size_t n = num.m_limbs.size();
134 size_t m = den.m_limbs.size();
135
136 // Knuth algorithm D (TAOCP 4.3.1). Normalize so the divisor's top bit
137 // is set, keeping the quotient estimate accurate to within one.
138 uint32_t top = den.m_limbs[m - 1];
139 size_t s = 0;
140 while ((top & 0x80000000u) == 0) { top <<= 1; s++; }
141
142 // v = den << s (m limbs); u = num << s (n + 1 limbs).
143 std::vector<uint32_t> v(m, 0);
144 for (size_t i = 0; i < m; i++) {
145 uint64_t cur = uint64_t(den.limb(i)) << s;
146 if (i > 0) cur |= uint64_t(den.limb(i - 1)) >> (32 - s);
147 v[i] = uint32_t(cur);
148 }
149 std::vector<uint32_t> u(n + 1, 0);
150 for (size_t i = 0; i < n; i++) {
151 uint64_t cur = uint64_t(num.limb(i)) << s;
152 if (i > 0) cur |= uint64_t(num.limb(i - 1)) >> (32 - s);
153 u[i] = uint32_t(cur);
154 }
155 u[n] = uint32_t(uint64_t(num.limb(n - 1)) >> (32 - s));
156
157 std::vector<uint32_t> q(n - m + 1, 0);
158
159 for (size_t j = n - m + 1; j-- > 0; ) {
160 // D3: estimate qhat from the top two limbs.
161 uint64_t num2 = (uint64_t(u[j + m]) << 32) | u[j + m - 1];
162 uint64_t qhat = num2 / v[m - 1];
163 uint64_t rhat = num2 % v[m - 1];
164 if (qhat >= 0x100000000ULL) {
165 qhat = 0xFFFFFFFFULL;
166 rhat = num2 - qhat * uint64_t(v[m - 1]);
167 }
168 if (m >= 2) {
169 uint64_t lhs = qhat * uint64_t(v[m - 2]);
170 uint64_t rhs = (rhat << 32) | u[j + m - 2];
171 while (lhs > rhs) {
172 qhat--;
173 rhat += v[m - 1];
174 if (rhat >= 0x100000000ULL) break;
175 lhs = qhat * uint64_t(v[m - 2]);
176 rhs = (rhat << 32) | u[j + m - 2];
177 }
178 }
179
180 // D4: u[j..j+m] -= qhat * v.
181 uint64_t borrow = 0;
182 uint64_t carry = 0;
183 for (size_t i = 0; i < m; i++) {
184 uint64_t p = qhat * uint64_t(v[i]) + carry;
185 carry = p >> 32;
186 uint64_t res = uint64_t(u[j + i]) - uint32_t(p) - borrow;
187 u[j + i] = uint32_t(res);
188 borrow = res >> 63;
189 }
190 uint64_t res_top = uint64_t(u[j + m]) - carry - borrow;
191 u[j + m] = uint32_t(res_top);
192 bool negative = (res_top >> 63) != 0;
193
194 // D5: if the remainder went negative, qhat was one too large.
195 if (negative) {
196 qhat--;
197 uint64_t add_carry = 0;
198 for (size_t i = 0; i < m; i++) {
199 uint64_t sum = uint64_t(u[j + i]) + uint64_t(v[i]) + add_carry;
200 u[j + i] = uint32_t(sum);
201 add_carry = sum >> 32;
202 }
203 u[j + m] = uint32_t(uint64_t(u[j + m]) + add_carry);
204 }
205 q[j] = uint32_t(qhat);
206 }
207
208 // D8: unnormalize the remainder r = u[0..m-1] >> s.
209 std::vector<uint32_t> r(m, 0);
210 for (size_t i = 0; i < m; i++) {
211 uint64_t cur = uint64_t(u[i]) >> s;
212 if (s > 0 && i + 1 < n + 1) cur |= uint64_t(u[i + 1]) << (32 - s);
213 r[i] = uint32_t(cur);
214 }
215
216 quot.m_limbs = std::move(q);
217 quot.trim();
218 rem.m_limbs = std::move(r);
219 rem.trim();
220 }
221
222 BigNum& BigNum::operator%=(const BigNum& mod)
223 {
224 BigNum q, r;
225 div_rem(*this, mod, q, r);
226 *this = std::move(r);
227 return *this;
228 }
229
230 uint32_t BigNum::mont_mp(const BigNum& m)
231 {
232 uint32_t m0 = m.limb(0);
233 uint32_t x = 1;
234 for (int i = 0; i < 5; i++) x *= 2 - m0 * x;
235 return -x;
236 }
237
238 BigNum BigNum::mont_r2(const BigNum& m)
239 {
240 // R = 2^(32 * W) where W = m.num_limbs(). R^2 = 2^(64*W).
241 // Start with 1, double 64*W times, reduce mod m each step.
242 BigNum r(1u);
243 BigNum two(2u);
244 size_t steps = 64 * m.m_limbs.size();
245 for (size_t i = 0; i < steps; i++) {
246 r = r + r;
247 while (r.compare(m) >= 0) r -= m;
248 }
249 return r;
250 }
251
252 BigNum BigNum::mont_mul(const BigNum& a, const BigNum& b,
253 const BigNum& m, uint32_t m_prime)
254 {
255 // Returns a * b * R^-1 mod m, where R = 2^(32*n). Requires a, b < m and
256 // m odd. Implemented as full product followed by REDC (Montgomery
257 // reduction): for each limb i, add u*m to make T[i] zero, then shift.
258 size_t n = m.m_limbs.size();
259
260 // T = a * b (2n limbs), n + 1 for the reduction carry.
261 std::vector<uint32_t> T(2 * n + 1, 0);
262 for (size_t i = 0; i < n; i++) {
263 uint64_t carry = 0;
264 for (size_t j = 0; j < n; j++) {
265 uint64_t x = uint64_t(T[i + j]) + uint64_t(a.limb(i)) * uint64_t(b.limb(j)) + carry;
266 T[i + j] = uint32_t(x);
267 carry = x >> 32;
268 }
269 T[i + n] = uint32_t(carry);
270 }
271
272 // REDC: zero out the low n limbs.
273 for (size_t i = 0; i < n; i++) {
274 uint32_t u = uint32_t(uint64_t(T[i]) * uint64_t(m_prime));
275 uint64_t carry = 0;
276 for (size_t j = 0; j < n; j++) {
277 uint64_t x = uint64_t(T[i + j]) + uint64_t(u) * uint64_t(m.limb(j)) + carry;
278 T[i + j] = uint32_t(x);
279 carry = x >> 32;
280 }
281 size_t k = i + n;
282 while (carry) {
283 uint64_t x = uint64_t(T[k]) + carry;
284 T[k] = uint32_t(x);
285 carry = x >> 32;
286 k++;
287 }
288 }
289
290 BigNum result(0u);
291 result.m_limbs.assign(T.begin() + n, T.begin() + 2 * n + 1);
292 result.trim();
293 if (result.compare(m) >= 0) result -= m;
294 return result;
295 }
296
297 BigNum BigNum::mod_pow(BigNum base, BigNum exp, const BigNum& m)
298 {
299 if (exp.is_zero()) return BigNum(1u);
300 if (m.is_one()) return BigNum(0u);
301
302 // For moduli up to 256 bits, use naive arithmetic (faster
303 // because Montgomery r2 computation overhead dominates).
304 if (m.num_limbs() <= 8) {
305 BigNum result(1u);
306 base = base % m;
307 size_t bits = exp.bit_length();
308 for (size_t i = 0; i < bits; i++) {
309 size_t li = i / 32, bi = i % 32;
310 if (exp.limb(li) & (1u << bi)) {
311 result = (result * base) % m;
312 }
313 base = (base * base) % m;
314 }
315 return result;
316 }
317
318 // Montgomery modular exponentiation for large moduli.
319 uint32_t mp = mont_mp(m);
320 BigNum r2 = mont_r2(m);
321 BigNum one(1u);
322
323 // Convert base to Montgomery form: base_mont = base * R mod m
324 BigNum base_mont = mont_mul(base % m, r2, m, mp);
325
326 // Result starts at 1 in Montgomery form: R mod m
327 BigNum result_mont = mont_mul(one, r2, m, mp);
328
329 size_t bits = exp.bit_length();
330 for (size_t i = bits; i-- > 0; ) {
331 result_mont = mont_mul(result_mont, result_mont, m, mp);
332 size_t li = i / 32, bi = i % 32;
333 if (exp.limb(li) & (1u << bi)) {
334 result_mont = mont_mul(result_mont, base_mont, m, mp);
335 }
336 }
337
338 // Convert back: result = result_mont * 1 * R^-1 mod m
339 return mont_mul(result_mont, one, m, mp);
340 }
341
342 BigNum BigNum::sqr_chain(const BigNum& base, uint32_t D, const BigNum& m,
343 std::vector<BigNum>* intermediates)
344 {
345 // For small moduli, use naive squaring.
346 if (m.num_limbs() <= 8) {
347 BigNum x = base % m;
348 if (intermediates) {
349 intermediates->clear();
350 intermediates->reserve(D / 16 + 2);
351 intermediates->push_back(x);
352 }
353 for (uint32_t i = 0; i < D; i++) {
354 x = (x * x) % m;
355 if (intermediates && (i % 16 == 0)) {
356 intermediates->push_back(x);
357 }
358 }
359 return x;
360 }
361
362 // Montgomery squaring chain for large moduli.
363 uint32_t mp = mont_mp(m);
364 BigNum r2 = mont_r2(m);
365 BigNum one(1u);
366
367 // Convert base to Montgomery form: x = base * R mod m
368 BigNum x = mont_mul(base % m, r2, m, mp);
369
370 if (intermediates) {
371 intermediates->clear();
372 intermediates->reserve(D / 16 + 2);
373 intermediates->push_back(base % m);
374 }
375
376 for (uint32_t i = 0; i < D; i++) {
377 x = mont_mul(x, x, m, mp);
378 if (intermediates && (i % 16 == 0)) {
379 BigNum plain = mont_mul(x, one, m, mp);
380 intermediates->push_back(plain);
381 }
382 }
383
384 return mont_mul(x, one, m, mp);
385 }
386
387 BigNum BigNum::gcd(const BigNum& a, const BigNum& b)
388 {
389 BigNum x(a), y(b);
390 for (int i = 0; i < 1000 && !y.is_zero(); i++) {
391 BigNum r = x % y;
392 x = std::move(y);
393 y = std::move(r);
394 }
395 // Safety: if b != 0, gcd cannot exceed 1000 steps for
396 // any well-formed input. Return 1 (co-prime) as fallback.
397 if (!y.is_zero()) return BigNum(1u);
398 return x;
399 }
400
401 std::string BigNumToHex(const BigNum& n)
402 {
403 auto bytes = n.to_bytes();
404 if (bytes.empty()) return "00";
405 std::string hex;
406 for (size_t i = bytes.size(); i-- > 0; ) {
407 char buf[3];
408 snprintf(buf, sizeof(buf), "%02x", bytes[i]);
409 hex += buf;
410 }
411 return hex;
412 }
413