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