some cleanup and temp commenting out the

This commit is contained in:
Kevaundray Wedderburn 2025-07-08 21:20:47 +01:00
parent c496edfa2e
commit 92adee104b

View file

@ -20,10 +20,7 @@ package gmp
// } // }
// //
// // Initialize GMP integers // // Initialize GMP integers
// mpz_init(base_mpz); // mpz_inits(base_mpz, exp_mpz, mod_mpz, result_mpz, NULL);
// mpz_init(exp_mpz);
// mpz_init(mod_mpz);
// mpz_init(result_mpz);
// //
// // Import big-endian byte arrays into GMP integers // // Import big-endian byte arrays into GMP integers
// // Handle empty arrays specially - GMP treats NULL with size 0 as 0 // // 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); // mpz_import(mod_mpz, mod_len, 1, 1, 0, 0, mod);
// } // }
// //
// // Special case: modulus is zero - return empty result (EVM behavior) // // Perform modular exponentiation
// if (mpz_cmp_ui(mod_mpz, 0) == 0) { // mpz_powm(result_mpz, base_mpz, exp_mpz, mod_mpz);
// *result_len = 0;
// mpz_clear(base_mpz);
// mpz_clear(exp_mpz);
// mpz_clear(mod_mpz);
// mpz_clear(result_mpz);
// return 0;
// }
// //
// // Special case: base has bit length 1 (base == 1) // // Get exact size needed for result
// // Just return base % mod // size_t needed = 0;
// if (mpz_sizeinbase(base_mpz, 2) == 1) { // mpz_export(NULL, &needed, 1, 1, 0, 0, result_mpz);
// 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
// //
// // Check if result buffer is large enough // // Check if result buffer is large enough
// if (*result_len < needed) { // if (*result_len < needed) {
// mpz_clear(base_mpz); // mpz_clears(base_mpz, exp_mpz, mod_mpz, result_mpz, NULL);
// mpz_clear(exp_mpz);
// mpz_clear(mod_mpz);
// mpz_clear(result_mpz);
// return -2; // return -2;
// } // }
// //
// // Export result to big-endian byte array // // Export result to big-endian byte array
// size_t count; // mpz_export(result, &needed, 1, 1, 0, 0, result_mpz);
// mpz_export(result, &count, 1, 1, 0, 0, result_mpz); // *result_len = needed;
// *result_len = count;
//
// // Handle zero result specially
// if (count == 0) {
// result[0] = 0;
// *result_len = 1;
// }
// //
// // Clean up // // Clean up
// mpz_clear(base_mpz); // mpz_clears(base_mpz, exp_mpz, mod_mpz, result_mpz, NULL);
// mpz_clear(exp_mpz);
// mpz_clear(mod_mpz);
// mpz_clear(result_mpz);
// //
// return 0; // return 0;
// } // }
import "C" import "C"
import ( import (
"errors" "errors"
"runtime"
"unsafe" "unsafe"
) )
@ -105,12 +74,55 @@ func ModExp(base, exp, mod []byte) ([]byte, error) {
return []byte{}, nil 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) // Allocate result buffer (size of modulus is the max possible result)
result := make([]byte, len(mod)) result := make([]byte, len(mod))
resultLen := C.size_t(len(result)) resultLen := C.size_t(len(result))
// Handle empty slices - pass a dummy non-nil pointer with length 0 // 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) dummy := C.uint8_t(0)
var basePtr, expPtr, modPtr *C.uint8_t = &dummy, &dummy, &dummy 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, (*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 // Check for errors
switch ret { switch ret {
case 0: case 0: