diff --git a/crypto/modexp/gmp/gmp.go b/crypto/modexp/gmp/gmp.go index 76bb83f533..849c7ee651 100644 --- a/crypto/modexp/gmp/gmp.go +++ b/crypto/modexp/gmp/gmp.go @@ -20,10 +20,7 @@ package gmp // } // // // Initialize GMP integers -// mpz_init(base_mpz); -// mpz_init(exp_mpz); -// mpz_init(mod_mpz); -// mpz_init(result_mpz); +// mpz_inits(base_mpz, exp_mpz, mod_mpz, result_mpz, NULL); // // // Import big-endian byte arrays into GMP integers // // Handle empty arrays specially - GMP treats NULL with size 0 as 0 @@ -37,60 +34,32 @@ package gmp // mpz_import(mod_mpz, mod_len, 1, 1, 0, 0, mod); // } // -// // Special case: modulus is zero - return empty result (EVM behavior) -// if (mpz_cmp_ui(mod_mpz, 0) == 0) { -// *result_len = 0; -// mpz_clear(base_mpz); -// mpz_clear(exp_mpz); -// mpz_clear(mod_mpz); -// mpz_clear(result_mpz); -// return 0; -// } +// // Perform modular exponentiation +// mpz_powm(result_mpz, base_mpz, exp_mpz, mod_mpz); // -// // Special case: base has bit length 1 (base == 1) -// // Just return base % mod -// if (mpz_sizeinbase(base_mpz, 2) == 1) { -// mpz_mod(result_mpz, base_mpz, mod_mpz); -// } else { -// // Normal case: perform modular exponentiation -// mpz_powm(result_mpz, base_mpz, exp_mpz, mod_mpz); -// } -// -// // Get size needed for result -// size_t needed = (mpz_sizeinbase(result_mpz, 2) + 7) / 8; -// if (needed == 0) needed = 1; // For zero result +// // Get exact size needed for result +// size_t needed = 0; +// mpz_export(NULL, &needed, 1, 1, 0, 0, result_mpz); // // // Check if result buffer is large enough // if (*result_len < needed) { -// mpz_clear(base_mpz); -// mpz_clear(exp_mpz); -// mpz_clear(mod_mpz); -// mpz_clear(result_mpz); +// mpz_clears(base_mpz, exp_mpz, mod_mpz, result_mpz, NULL); // return -2; // } // // // Export result to big-endian byte array -// size_t count; -// mpz_export(result, &count, 1, 1, 0, 0, result_mpz); -// *result_len = count; -// -// // Handle zero result specially -// if (count == 0) { -// result[0] = 0; -// *result_len = 1; -// } +// mpz_export(result, &needed, 1, 1, 0, 0, result_mpz); +// *result_len = needed; // // // Clean up -// mpz_clear(base_mpz); -// mpz_clear(exp_mpz); -// mpz_clear(mod_mpz); -// mpz_clear(result_mpz); +// mpz_clears(base_mpz, exp_mpz, mod_mpz, result_mpz, NULL); // // return 0; // } import "C" import ( "errors" + "runtime" "unsafe" ) @@ -105,12 +74,55 @@ func ModExp(base, exp, mod []byte) ([]byte, error) { return []byte{}, nil } + // Special case: zero modulus + // TODO: Check to see if theres a cleaner way to do this + allZero := true + for _, b := range mod { + if b != 0 { + allZero = false + break + } + } + if allZero { + return []byte{}, nil + } + + // // Special case: base == 1 + // // Check if base is 1 (only one byte with value 1, or leading zeros followed by 1) + // baseIsOne := false + // if len(base) == 0 { + // baseIsOne = false + // } else if len(base) == 1 && base[0] == 1 { + // baseIsOne = true + // } else { + // // Check for leading zeros followed by 1 + // allZeroExceptLast := true + // for i := 0; i < len(base)-1; i++ { + // if base[i] != 0 { + // allZeroExceptLast = false + // break + // } + // } + // if allZeroExceptLast && base[len(base)-1] == 1 { + // baseIsOne = true + // } + // } + + // if baseIsOne { + // // base^exp mod mod = 1 mod mod = 1 (if mod > 1), 0 (if mod == 1) + // // Just return base % mod which is 1 % mod + // if len(mod) == 1 && mod[0] == 1 { + // return []byte{}, nil // 1 % 1 = 0 + // } + // return []byte{1}, nil // 1 % mod = 1 for mod > 1 + // } + // Allocate result buffer (size of modulus is the max possible result) result := make([]byte, len(mod)) resultLen := C.size_t(len(result)) // Handle empty slices - pass a dummy non-nil pointer with length 0 - // This avoids UB when the length is zero + // This avoids UB when the length is zero. dummy := C.uint8_t(0) var basePtr, expPtr, modPtr *C.uint8_t = &dummy, &dummy, &dummy @@ -132,6 +144,12 @@ func ModExp(base, exp, mod []byte) ([]byte, error) { (*C.uint8_t)(unsafe.Pointer(&result[0])), &resultLen, ) + // Keep the slices alive until after the C call completes + runtime.KeepAlive(base) + runtime.KeepAlive(exp) + runtime.KeepAlive(mod) + runtime.KeepAlive(result) + // Check for errors switch ret { case 0: