#[cfg(zisk_guest)]
use crate::alloc_extern::vec;
#[cfg(zisk_guest)]
use crate::alloc_extern::vec::Vec;
use crate::zisklib::fcall_bin_decomp;
use super::{
mul_and_reduce_long, mulmod_short, rem_long_init, rem_short_init, square_and_reduce_long,
LongScratch, U256,
};
pub fn modexp(
base: &[U256],
exp: &[u64],
modulus: &[U256],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> Vec<U256> {
let len_b = base.len();
let len_e = exp.len();
let len_m = modulus.len();
#[cfg(debug_assertions)]
{
assert_ne!(len_b, 0, "Base must have at least one limb");
assert_ne!(len_e, 0, "Exponent must have at least one limb");
assert_ne!(len_m, 0, "Modulus must have at least one limb");
if len_b > 1 {
assert!(!base[len_b - 1].is_zero(), "Base must not have leading zeros");
}
if len_e > 1 {
assert_ne!(exp.last().unwrap(), &0, "Exponent must not have leading zeros");
}
if len_m > 1 {
assert!(!modulus[len_m - 1].is_zero(), "Modulus must not have leading zeros");
} else {
assert!(!modulus[0].is_zero(), "Modulus must not be zero");
}
}
if len_m == 1 && modulus[0].is_zero() {
return vec![U256::ZERO];
}
if len_m == 1 && modulus[0].is_one() {
return vec![U256::ZERO];
}
if len_e == 1 && exp[0] == 0 {
return vec![U256::ONE];
}
if len_b == 1 {
if base[0].is_zero() {
return vec![U256::ZERO];
}
if base[0].is_one() {
return vec![U256::ONE];
}
}
if len_m == 1 {
modexp_short(
base,
exp,
&modulus[0],
#[cfg(feature = "hints")]
hints,
)
} else {
modexp_long(
base,
exp,
modulus,
#[cfg(feature = "hints")]
hints,
)
}
}
fn modexp_short(
base: &[U256],
exp: &[u64],
modulus: &U256,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> Vec<U256> {
let len_e = exp.len();
let base = rem_short_init(
base,
modulus,
#[cfg(feature = "hints")]
hints,
);
let (len, bits) = fcall_bin_decomp(
exp,
#[cfg(feature = "hints")]
hints,
);
assert!(len > 0 && bits[0] == 1, "Exponent must be non-zero");
assert!(len <= 64 * len_e, "Exponent bit length out of range");
assert!(bits.len() == len, "Bit decomposition length mismatch");
let mut rec_exp = vec![0u64; len_e];
for (bit_idx, &bit) in bits.iter().enumerate() {
if bit == 1 {
let bits_pos = len - 1 - bit_idx;
rec_exp[bits_pos / 64] |= 1u64 << (bits_pos % 64);
}
}
assert_eq!(rec_exp[..], *exp, "Exponent decomposition mismatch");
let mut out = base;
for &bit in bits.iter().skip(1) {
if out.is_zero() {
break;
}
out = mulmod_short(
&out,
&out,
modulus,
#[cfg(feature = "hints")]
hints,
);
if bit == 1 {
out = mulmod_short(
&out,
&base,
modulus,
#[cfg(feature = "hints")]
hints,
);
}
}
vec![out]
}
fn modexp_long(
base: &[U256],
exp: &[u64],
modulus: &[U256],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> Vec<U256> {
let len_e = exp.len();
let len_m = modulus.len();
let base = rem_long_init(
base,
modulus,
#[cfg(feature = "hints")]
hints,
);
let (len, bits) = fcall_bin_decomp(
exp,
#[cfg(feature = "hints")]
hints,
);
assert!(len > 0 && bits[0] == 1, "Exponent must be non-zero");
assert!(len <= 64 * len_e, "Exponent bit length out of range");
assert!(bits.len() == len, "Bit decomposition length mismatch");
let mut rec_exp = vec![0u64; len_e];
for (bit_idx, &bit) in bits.iter().enumerate() {
if bit == 1 {
let bits_pos = len - 1 - bit_idx;
rec_exp[bits_pos / 64] |= 1u64 << (bits_pos % 64);
}
}
assert_eq!(rec_exp[..], *exp, "Exponent decomposition mismatch");
let mut scratch = LongScratch::new(len_m);
let mut out = base.clone();
for &bit in bits.iter().skip(1) {
if out.len() == 1 && out[0].is_zero() {
break;
}
out = square_and_reduce_long(
&out,
modulus,
&mut scratch,
#[cfg(feature = "hints")]
hints,
);
if bit == 1 {
out = mul_and_reduce_long(
&out,
&base,
modulus,
&mut scratch,
#[cfg(feature = "hints")]
hints,
);
}
}
out
}
#[allow(clippy::too_many_arguments)]
#[allow(dead_code)]
#[inline]
pub(crate) unsafe fn modexp_bytes_c(
base_ptr: *const u8,
base_len: usize,
exp_ptr: *const u8,
exp_len: usize,
modulus_ptr: *const u8,
modulus_len: usize,
result_ptr: *mut u8,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> usize {
let base_bytes = core::slice::from_raw_parts(base_ptr, base_len);
let exp_bytes = core::slice::from_raw_parts(exp_ptr, exp_len);
let modulus_bytes = core::slice::from_raw_parts(modulus_ptr, modulus_len);
let base_u256 = bytes_be_to_u256_le(base_bytes);
let exp_u64 = bytes_be_to_u64_le(exp_bytes);
let modulus_u256 = bytes_be_to_u256_le(modulus_bytes);
let result_u256 = modexp(
&base_u256,
&exp_u64,
&modulus_u256,
#[cfg(feature = "hints")]
hints,
);
let result = core::slice::from_raw_parts_mut(result_ptr, modulus_len);
u256_le_to_bytes_be(&result_u256, result);
modulus_len
}
#[allow(clippy::too_many_arguments)]
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_modexp_u64_c")]
pub unsafe extern "C" fn modexp_u64_c(
base_ptr: *const u64,
base_len: usize,
exp_ptr: *const u64,
exp_len: usize,
modulus_ptr: *const u64,
modulus_len: usize,
result_ptr: *mut u64,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> usize {
let base_flat = core::slice::from_raw_parts(base_ptr, base_len);
let exp = core::slice::from_raw_parts(exp_ptr, exp_len);
let modulus_flat = core::slice::from_raw_parts(modulus_ptr, modulus_len);
let base_len = base_flat.len().next_multiple_of(4);
let modulus_len = modulus_flat.len().next_multiple_of(4);
let mut base_padded = vec![0u64; base_len];
let mut modulus_padded = vec![0u64; modulus_len];
base_padded[..base_flat.len()].copy_from_slice(base_flat);
modulus_padded[..modulus_flat.len()].copy_from_slice(modulus_flat);
let base = U256::flat_to_slice(&base_padded);
let modulus = U256::flat_to_slice(&modulus_padded);
let result_u256 = modexp(
base,
exp,
modulus,
#[cfg(feature = "hints")]
hints,
);
let result_slice = U256::slice_to_flat(&result_u256);
let result_len = result_slice.len();
let result = core::slice::from_raw_parts_mut(result_ptr, modulus_len);
result[..result_len].copy_from_slice(result_slice);
result_len
}
#[allow(dead_code)]
fn bytes_be_to_u64_le(bytes: &[u8]) -> Vec<u64> {
if bytes.is_empty() {
return vec![0];
}
let first_nonzero = bytes.iter().position(|&b| b != 0).unwrap_or(bytes.len() - 1);
let bytes = &bytes[first_nonzero..];
if bytes.is_empty() {
return vec![0];
}
let num_limbs = bytes.len().div_ceil(8);
let mut result = vec![0u64; num_limbs];
for (i, &byte) in bytes.iter().rev().enumerate() {
let limb_idx = i / 8;
let byte_idx = i % 8;
result[limb_idx] |= (byte as u64) << (byte_idx * 8);
}
result
}
#[allow(dead_code)]
fn bytes_be_to_u256_le(bytes: &[u8]) -> Vec<U256> {
let u64_le = bytes_be_to_u64_le(bytes);
let padded_len = u64_le.len().next_multiple_of(4);
let mut padded = vec![0u64; padded_len];
padded[..u64_le.len()].copy_from_slice(&u64_le);
U256::flat_to_slice(&padded).to_vec()
}
#[allow(dead_code)]
fn u256_le_to_bytes_be(limbs: &[U256], output: &mut [u8]) {
let flat = U256::slice_to_flat(limbs);
let out_len = output.len();
output.fill(0);
for (i, &limb) in flat.iter().enumerate() {
for j in 0..8 {
let byte_val = ((limb >> (j * 8)) & 0xFF) as u8;
let pos_from_end = i * 8 + j;
if pos_from_end < out_len {
output[out_len - 1 - pos_from_end] = byte_val;
}
}
}
}