// Copyright (c) 2025 The Limenka developers // Distributed under the MIT software license, see the accompanying // file COPYING or http://www.opensource.org/licenses/mit-license.php. #include #include #include #include void BigNum::trim() { while (m_limbs.size() > 1 && m_limbs.back() == 0) { m_limbs.pop_back(); } } BigNum::BigNum(const std::vector& bytes, bool big_endian) { if (bytes.empty()) { m_limbs = {0}; return; } size_t nlimbs = (bytes.size() + 3) / 4; m_limbs.resize(nlimbs, 0); if (big_endian) { for (size_t i = 0; i < bytes.size(); i++) { size_t limb_pos = bytes.size() - 1 - i; // position of byte i from the LSB size_t li = limb_pos / 4; size_t shift = 8 * (limb_pos % 4); m_limbs[li] |= uint32_t(bytes[i]) << shift; } } else { std::memcpy(m_limbs.data(), bytes.data(), bytes.size()); } trim(); } BigNum::BigNum(uint32_t val) { m_limbs = {val}; trim(); } std::vector BigNum::to_bytes(size_t width) const { size_t nbytes = m_limbs.size() * 4; if (width > 0 && width > nbytes) nbytes = width; std::vector out(nbytes, 0); for (size_t i = 0; i < m_limbs.size(); i++) { uint32_t v = m_limbs[i]; out[i * 4 + 0] = uint8_t(v); out[i * 4 + 1] = uint8_t(v >> 8); out[i * 4 + 2] = uint8_t(v >> 16); out[i * 4 + 3] = uint8_t(v >> 24); } return out; } bool BigNum::is_zero() const { return m_limbs.size() == 1 && m_limbs[0] == 0; } bool BigNum::is_one() const { return m_limbs.size() == 1 && m_limbs[0] == 1; } size_t BigNum::bit_length() const { if (is_zero()) return 0; uint32_t top = m_limbs.back(); size_t bits = (m_limbs.size() - 1) * 32; while (top) { bits++; top >>= 1; } return bits; } int BigNum::compare(const BigNum& other) const { if (m_limbs.size() != other.m_limbs.size()) return m_limbs.size() < other.m_limbs.size() ? -1 : 1; for (size_t i = m_limbs.size(); i-- > 0; ) if (m_limbs[i] != other.m_limbs[i]) return m_limbs[i] < other.m_limbs[i] ? -1 : 1; return 0; } BigNum& BigNum::operator+=(const BigNum& other) { size_t n = std::max(m_limbs.size(), other.m_limbs.size()); m_limbs.resize(n + 1, 0); uint64_t carry = 0; for (size_t i = 0; i < n; i++) { uint64_t sum = uint64_t(m_limbs[i]) + uint64_t(other.limb(i)) + carry; m_limbs[i] = uint32_t(sum); carry = sum >> 32; } m_limbs[n] = uint32_t(carry); trim(); return *this; } BigNum& BigNum::operator-=(const BigNum& other) { assert(compare(other) >= 0); int64_t borrow = 0; for (size_t i = 0; i < m_limbs.size(); i++) { int64_t diff = int64_t(m_limbs[i]) - int64_t(other.limb(i)) - borrow; if (diff < 0) { diff += (uint64_t(1) << 32); borrow = 1; } else { borrow = 0; } m_limbs[i] = uint32_t(diff); } trim(); return *this; } BigNum& BigNum::operator*=(const BigNum& other) { size_t n = m_limbs.size(), m = other.m_limbs.size(); std::vector result(n + m, 0); for (size_t i = 0; i < n; i++) { uint64_t carry = 0; for (size_t j = 0; j < m; j++) { uint64_t prod = uint64_t(m_limbs[i]) * uint64_t(other.m_limbs[j]) + uint64_t(result[i + j]) + carry; result[i + j] = uint32_t(prod); carry = prod >> 32; } result[i + m] = uint32_t(carry); } m_limbs = std::move(result); trim(); return *this; } void BigNum::div_rem(const BigNum& num, const BigNum& den, BigNum& quot, BigNum& rem) { assert(!den.is_zero()); if (num.compare(den) < 0) { quot = BigNum(0u); rem = num; return; } size_t n = num.m_limbs.size(); size_t m = den.m_limbs.size(); // Knuth algorithm D (TAOCP 4.3.1). Normalize so the divisor's top bit // is set, keeping the quotient estimate accurate to within one. uint32_t top = den.m_limbs[m - 1]; size_t s = 0; while ((top & 0x80000000u) == 0) { top <<= 1; s++; } // v = den << s (m limbs); u = num << s (n + 1 limbs). std::vector v(m, 0); for (size_t i = 0; i < m; i++) { uint64_t cur = uint64_t(den.limb(i)) << s; if (i > 0) cur |= uint64_t(den.limb(i - 1)) >> (32 - s); v[i] = uint32_t(cur); } std::vector u(n + 1, 0); for (size_t i = 0; i < n; i++) { uint64_t cur = uint64_t(num.limb(i)) << s; if (i > 0) cur |= uint64_t(num.limb(i - 1)) >> (32 - s); u[i] = uint32_t(cur); } u[n] = uint32_t(uint64_t(num.limb(n - 1)) >> (32 - s)); std::vector q(n - m + 1, 0); for (size_t j = n - m + 1; j-- > 0; ) { // D3: estimate qhat from the top two limbs. uint64_t num2 = (uint64_t(u[j + m]) << 32) | u[j + m - 1]; uint64_t qhat = num2 / v[m - 1]; uint64_t rhat = num2 % v[m - 1]; if (qhat >= 0x100000000ULL) { qhat = 0xFFFFFFFFULL; rhat = num2 - qhat * uint64_t(v[m - 1]); } if (m >= 2) { uint64_t lhs = qhat * uint64_t(v[m - 2]); uint64_t rhs = (rhat << 32) | u[j + m - 2]; while (lhs > rhs) { qhat--; rhat += v[m - 1]; if (rhat >= 0x100000000ULL) break; lhs = qhat * uint64_t(v[m - 2]); rhs = (rhat << 32) | u[j + m - 2]; } } // D4: u[j..j+m] -= qhat * v. uint64_t borrow = 0; uint64_t carry = 0; for (size_t i = 0; i < m; i++) { uint64_t p = qhat * uint64_t(v[i]) + carry; carry = p >> 32; uint64_t res = uint64_t(u[j + i]) - uint32_t(p) - borrow; u[j + i] = uint32_t(res); borrow = res >> 63; } uint64_t res_top = uint64_t(u[j + m]) - carry - borrow; u[j + m] = uint32_t(res_top); bool negative = (res_top >> 63) != 0; // D5: if the remainder went negative, qhat was one too large. if (negative) { qhat--; uint64_t add_carry = 0; for (size_t i = 0; i < m; i++) { uint64_t sum = uint64_t(u[j + i]) + uint64_t(v[i]) + add_carry; u[j + i] = uint32_t(sum); add_carry = sum >> 32; } u[j + m] = uint32_t(uint64_t(u[j + m]) + add_carry); } q[j] = uint32_t(qhat); } // D8: unnormalize the remainder r = u[0..m-1] >> s. std::vector r(m, 0); for (size_t i = 0; i < m; i++) { uint64_t cur = uint64_t(u[i]) >> s; if (s > 0 && i + 1 < n + 1) cur |= uint64_t(u[i + 1]) << (32 - s); r[i] = uint32_t(cur); } quot.m_limbs = std::move(q); quot.trim(); rem.m_limbs = std::move(r); rem.trim(); } BigNum& BigNum::operator%=(const BigNum& mod) { BigNum q, r; div_rem(*this, mod, q, r); *this = std::move(r); return *this; } uint32_t BigNum::mont_mp(const BigNum& m) { uint32_t m0 = m.limb(0); uint32_t x = 1; for (int i = 0; i < 5; i++) x *= 2 - m0 * x; return -x; } BigNum BigNum::mont_r2(const BigNum& m) { // R = 2^(32 * W) where W = m.num_limbs(). R^2 = 2^(64*W). // Start with 1, double 64*W times, reduce mod m each step. BigNum r(1u); BigNum two(2u); size_t steps = 64 * m.m_limbs.size(); for (size_t i = 0; i < steps; i++) { r = r + r; while (r.compare(m) >= 0) r -= m; } return r; } BigNum BigNum::mont_mul(const BigNum& a, const BigNum& b, const BigNum& m, uint32_t m_prime) { // Returns a * b * R^-1 mod m, where R = 2^(32*n). Requires a, b < m and // m odd. Implemented as full product followed by REDC (Montgomery // reduction): for each limb i, add u*m to make T[i] zero, then shift. size_t n = m.m_limbs.size(); // T = a * b (2n limbs), n + 1 for the reduction carry. std::vector T(2 * n + 1, 0); for (size_t i = 0; i < n; i++) { uint64_t carry = 0; for (size_t j = 0; j < n; j++) { uint64_t x = uint64_t(T[i + j]) + uint64_t(a.limb(i)) * uint64_t(b.limb(j)) + carry; T[i + j] = uint32_t(x); carry = x >> 32; } T[i + n] = uint32_t(carry); } // REDC: zero out the low n limbs. for (size_t i = 0; i < n; i++) { uint32_t u = uint32_t(uint64_t(T[i]) * uint64_t(m_prime)); uint64_t carry = 0; for (size_t j = 0; j < n; j++) { uint64_t x = uint64_t(T[i + j]) + uint64_t(u) * uint64_t(m.limb(j)) + carry; T[i + j] = uint32_t(x); carry = x >> 32; } size_t k = i + n; while (carry) { uint64_t x = uint64_t(T[k]) + carry; T[k] = uint32_t(x); carry = x >> 32; k++; } } BigNum result(0u); result.m_limbs.assign(T.begin() + n, T.begin() + 2 * n + 1); result.trim(); if (result.compare(m) >= 0) result -= m; return result; } BigNum BigNum::mod_pow(BigNum base, BigNum exp, const BigNum& m) { if (exp.is_zero()) return BigNum(1u); if (m.is_one()) return BigNum(0u); // For moduli up to 256 bits, use naive arithmetic (faster // because Montgomery r2 computation overhead dominates). if (m.num_limbs() <= 8) { BigNum result(1u); base = base % m; size_t bits = exp.bit_length(); for (size_t i = 0; i < bits; i++) { size_t li = i / 32, bi = i % 32; if (exp.limb(li) & (1u << bi)) { result = (result * base) % m; } base = (base * base) % m; } return result; } // Montgomery modular exponentiation for large moduli. uint32_t mp = mont_mp(m); BigNum r2 = mont_r2(m); BigNum one(1u); // Convert base to Montgomery form: base_mont = base * R mod m BigNum base_mont = mont_mul(base % m, r2, m, mp); // Result starts at 1 in Montgomery form: R mod m BigNum result_mont = mont_mul(one, r2, m, mp); size_t bits = exp.bit_length(); for (size_t i = bits; i-- > 0; ) { result_mont = mont_mul(result_mont, result_mont, m, mp); size_t li = i / 32, bi = i % 32; if (exp.limb(li) & (1u << bi)) { result_mont = mont_mul(result_mont, base_mont, m, mp); } } // Convert back: result = result_mont * 1 * R^-1 mod m return mont_mul(result_mont, one, m, mp); } BigNum BigNum::sqr_chain(const BigNum& base, uint32_t D, const BigNum& m, std::vector* intermediates) { // For small moduli, use naive squaring. if (m.num_limbs() <= 8) { BigNum x = base % m; if (intermediates) { intermediates->clear(); intermediates->reserve(D / 16 + 2); intermediates->push_back(x); } for (uint32_t i = 0; i < D; i++) { x = (x * x) % m; if (intermediates && (i % 16 == 0)) { intermediates->push_back(x); } } return x; } // Montgomery squaring chain for large moduli. uint32_t mp = mont_mp(m); BigNum r2 = mont_r2(m); BigNum one(1u); // Convert base to Montgomery form: x = base * R mod m BigNum x = mont_mul(base % m, r2, m, mp); if (intermediates) { intermediates->clear(); intermediates->reserve(D / 16 + 2); intermediates->push_back(base % m); } for (uint32_t i = 0; i < D; i++) { x = mont_mul(x, x, m, mp); if (intermediates && (i % 16 == 0)) { BigNum plain = mont_mul(x, one, m, mp); intermediates->push_back(plain); } } return mont_mul(x, one, m, mp); } BigNum BigNum::gcd(const BigNum& a, const BigNum& b) { BigNum x(a), y(b); for (int i = 0; i < 1000 && !y.is_zero(); i++) { BigNum r = x % y; x = std::move(y); y = std::move(r); } // Safety: if b != 0, gcd cannot exceed 1000 steps for // any well-formed input. Return 1 (co-prime) as fallback. if (!y.is_zero()) return BigNum(1u); return x; } std::string BigNumToHex(const BigNum& n) { auto bytes = n.to_bytes(); if (bytes.empty()) return "00"; std::string hex; for (size_t i = bytes.size(); i-- > 0; ) { char buf[3]; snprintf(buf, sizeof(buf), "%02x", bytes[i]); hex += buf; } return hex; }