use crate::syscalls::{
syscall_arith256, syscall_arith256_mod, SyscallArith256ModParams, SyscallArith256Params,
};
use crate::zisklib::fcall_bin_decomp;
use crate::zisklib::lib::{
constants::{ONE_256 as ONE, ZERO_256 as ZERO},
utils::{be_bytes_to_u64_4, gt, is_one, is_zero, lt, u64_4_to_be_bytes},
};
use crate::zisklib::{fcall_uint256_inv_mod, ModInvResult};
pub fn reduce_mod256(
a: &[u64; 4],
modulus: &[u64; 4],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> [u64; 4] {
if is_zero(modulus) {
return ZERO;
}
if lt(a, modulus) {
*a
} else {
let mut d = ZERO;
let mut params =
SyscallArith256ModParams { a, b: &ONE, c: &ZERO, module: modulus, d: &mut d };
syscall_arith256_mod(
&mut params,
#[cfg(feature = "hints")]
hints,
);
d
}
}
pub fn add_mod256(
a: &[u64; 4],
b: &[u64; 4],
modulus: &[u64; 4],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> [u64; 4] {
if is_zero(modulus) {
return ZERO;
}
let mut d = ZERO;
let mut params = SyscallArith256ModParams { a, b: &ONE, c: b, module: modulus, d: &mut d };
syscall_arith256_mod(
&mut params,
#[cfg(feature = "hints")]
hints,
);
d
}
pub fn mul_mod256(
a: &[u64; 4],
b: &[u64; 4],
modulus: &[u64; 4],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> [u64; 4] {
if is_zero(modulus) {
return ZERO;
}
let mut d = ZERO;
let mut params = SyscallArith256ModParams { a, b, c: &ZERO, module: modulus, d: &mut d };
syscall_arith256_mod(
&mut params,
#[cfg(feature = "hints")]
hints,
);
d
}
pub fn square_mod256(
a: &[u64; 4],
modulus: &[u64; 4],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> [u64; 4] {
mul_mod256(
a,
a,
modulus,
#[cfg(feature = "hints")]
hints,
)
}
pub fn pow_mod256(
base: &[u64; 4],
exp: &[u64; 4],
modulus: &[u64; 4],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> [u64; 4] {
if is_zero(modulus) {
return ZERO;
}
if is_one(modulus) {
return ZERO;
}
if is_zero(exp) {
return ONE;
} else if is_one(exp) {
return reduce_mod256(
base,
modulus,
#[cfg(feature = "hints")]
hints,
);
}
if is_zero(base) {
return ZERO;
} else if is_one(base) {
return ONE;
}
let (len, bits) = fcall_bin_decomp(
exp,
#[cfg(feature = "hints")]
hints,
);
assert!(len > 0 && bits[0] == 1, "Exponent must be non-zero");
assert!(len <= 256, "Exponent bit length out of range");
assert!(bits.len() == len, "Bit decomposition length mismatch");
let mut rec_exp = [0u64; 4];
for (bit_idx, &bit) in bits.iter().enumerate() {
if bit == 1 {
let bit_pos = len - 1 - bit_idx;
rec_exp[bit_pos / 64] |= 1u64 << (bit_pos % 64);
}
}
assert_eq!(rec_exp, *exp, "Exponent decomposition mismatch");
let mut result = reduce_mod256(
base,
modulus,
#[cfg(feature = "hints")]
hints,
);
for &bit in bits.iter().skip(1) {
if is_zero(&result) {
break;
}
result = square_mod256(
&result,
modulus,
#[cfg(feature = "hints")]
hints,
);
if bit == 1 {
result = mul_mod256(
&result,
base,
modulus,
#[cfg(feature = "hints")]
hints,
);
}
}
result
}
pub fn inv_mod256(
a: &[u64; 4],
modulus: &[u64; 4],
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> Option<[u64; 4]> {
match fcall_uint256_inv_mod(
a,
modulus,
#[cfg(feature = "hints")]
hints,
) {
ModInvResult::Inverse(inv) => {
assert!(lt(&inv, modulus), "Inverse must be less than modulus");
let result = mul_mod256(
a,
&inv,
modulus,
#[cfg(feature = "hints")]
hints,
);
assert!(is_one(&result), "a * inv must equal 1 mod modulus");
Some(inv)
}
ModInvResult::NoInverse { gcd: d, qa, qm } => {
assert!(gt(&d, &ONE), "gcd witness must be greater than 1");
let mut a_lo = ZERO;
let mut a_hi = ZERO;
let mut a_params =
SyscallArith256Params { a: &qa, b: &d, c: &ZERO, dl: &mut a_lo, dh: &mut a_hi };
syscall_arith256(
&mut a_params,
#[cfg(feature = "hints")]
hints,
);
assert!(is_zero(&a_hi) && a_lo == *a, "gcd must divide a");
let mut m_lo = ZERO;
let mut m_hi = ZERO;
let mut m_params =
SyscallArith256Params { a: &qm, b: &d, c: &ZERO, dl: &mut m_lo, dh: &mut m_hi };
syscall_arith256(
&mut m_params,
#[cfg(feature = "hints")]
hints,
);
assert!(is_zero(&m_hi) && m_lo == *modulus, "gcd must divide modulus");
None
}
}
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_reduce_mod256_c")]
pub unsafe extern "C" fn reduce_mod256_c(
a_ptr: *const u64,
modulus_ptr: *const u64,
result_ptr: *mut u64,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a = &*(a_ptr as *const [u64; 4]);
let modulus = &*(modulus_ptr as *const [u64; 4]);
let res = reduce_mod256(
a,
modulus,
#[cfg(feature = "hints")]
hints,
);
let result = &mut *(result_ptr as *mut [u64; 4]);
*result = res;
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_add_mod256_c")]
pub unsafe extern "C" fn add_mod256_c(
a_ptr: *const u64,
b_ptr: *const u64,
modulus_ptr: *const u64,
result_ptr: *mut u64,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a = &*(a_ptr as *const [u64; 4]);
let b = &*(b_ptr as *const [u64; 4]);
let modulus = &*(modulus_ptr as *const [u64; 4]);
let res = add_mod256(
a,
b,
modulus,
#[cfg(feature = "hints")]
hints,
);
let result = &mut *(result_ptr as *mut [u64; 4]);
*result = res;
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_mul_mod256_c")]
pub unsafe extern "C" fn mul_mod256_c(
a_ptr: *const u64,
b_ptr: *const u64,
modulus_ptr: *const u64,
result_ptr: *mut u64,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a = &*(a_ptr as *const [u64; 4]);
let b = &*(b_ptr as *const [u64; 4]);
let modulus = &*(modulus_ptr as *const [u64; 4]);
let res = mul_mod256(
a,
b,
modulus,
#[cfg(feature = "hints")]
hints,
);
let result = &mut *(result_ptr as *mut [u64; 4]);
*result = res;
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_mul_mod_bytes256_c")]
pub unsafe extern "C" fn mul_mod_bytes256_c(
a_ptr: *const u8,
b_ptr: *const u8,
m_ptr: *const u8,
result_ptr: *mut u8,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a_bytes = &*(a_ptr as *const [u8; 32]);
let b_bytes = &*(b_ptr as *const [u8; 32]);
let m_bytes = &*(m_ptr as *const [u8; 32]);
let a = be_bytes_to_u64_4(a_bytes);
let b = be_bytes_to_u64_4(b_bytes);
let m = be_bytes_to_u64_4(m_bytes);
let result = mul_mod256(
&a,
&b,
&m,
#[cfg(feature = "hints")]
hints,
);
let result_bytes = &mut *(result_ptr as *mut [u8; 32]);
*result_bytes = u64_4_to_be_bytes(&result);
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_reduce_mod_bytes256_c")]
pub unsafe extern "C" fn reduce_mod_bytes256_c(
a_ptr: *const u8,
m_ptr: *const u8,
result_ptr: *mut u8,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a = be_bytes_to_u64_4(&*(a_ptr as *const [u8; 32]));
let m = be_bytes_to_u64_4(&*(m_ptr as *const [u8; 32]));
let result = reduce_mod256(
&a,
&m,
#[cfg(feature = "hints")]
hints,
);
*(result_ptr as *mut [u8; 32]) = u64_4_to_be_bytes(&result);
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_add_mod_bytes256_c")]
pub unsafe extern "C" fn add_mod_bytes256_c(
a_ptr: *const u8,
b_ptr: *const u8,
m_ptr: *const u8,
result_ptr: *mut u8,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a = be_bytes_to_u64_4(&*(a_ptr as *const [u8; 32]));
let b = be_bytes_to_u64_4(&*(b_ptr as *const [u8; 32]));
let m = be_bytes_to_u64_4(&*(m_ptr as *const [u8; 32]));
let result = add_mod256(
&a,
&b,
&m,
#[cfg(feature = "hints")]
hints,
);
*(result_ptr as *mut [u8; 32]) = u64_4_to_be_bytes(&result);
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_square_mod_bytes256_c")]
pub unsafe extern "C" fn square_mod_bytes256_c(
a_ptr: *const u8,
m_ptr: *const u8,
result_ptr: *mut u8,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a = be_bytes_to_u64_4(&*(a_ptr as *const [u8; 32]));
let m = be_bytes_to_u64_4(&*(m_ptr as *const [u8; 32]));
let result = square_mod256(
&a,
&m,
#[cfg(feature = "hints")]
hints,
);
*(result_ptr as *mut [u8; 32]) = u64_4_to_be_bytes(&result);
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_pow_mod_bytes256_c")]
pub unsafe extern "C" fn pow_mod_bytes256_c(
base_ptr: *const u8,
exp_ptr: *const u8,
m_ptr: *const u8,
result_ptr: *mut u8,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let base = be_bytes_to_u64_4(&*(base_ptr as *const [u8; 32]));
let exp = be_bytes_to_u64_4(&*(exp_ptr as *const [u8; 32]));
let m = be_bytes_to_u64_4(&*(m_ptr as *const [u8; 32]));
let result = pow_mod256(
&base,
&exp,
&m,
#[cfg(feature = "hints")]
hints,
);
*(result_ptr as *mut [u8; 32]) = u64_4_to_be_bytes(&result);
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_inv_mod_bytes256_c")]
pub unsafe extern "C" fn inv_mod_bytes256_c(
a_ptr: *const u8,
m_ptr: *const u8,
result_ptr: *mut u8,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> u8 {
let a = be_bytes_to_u64_4(&*(a_ptr as *const [u8; 32]));
let m = be_bytes_to_u64_4(&*(m_ptr as *const [u8; 32]));
match inv_mod256(
&a,
&m,
#[cfg(feature = "hints")]
hints,
) {
Some(res) => {
*(result_ptr as *mut [u8; 32]) = u64_4_to_be_bytes(&res);
1
}
None => 0,
}
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_square_mod256_c")]
pub unsafe extern "C" fn square_mod256_c(
a_ptr: *const u64,
modulus_ptr: *const u64,
result_ptr: *mut u64,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let a = &*(a_ptr as *const [u64; 4]);
let modulus = &*(modulus_ptr as *const [u64; 4]);
let res = square_mod256(
a,
modulus,
#[cfg(feature = "hints")]
hints,
);
let result = &mut *(result_ptr as *mut [u64; 4]);
*result = res;
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_pow_mod256_c")]
pub unsafe extern "C" fn pow_mod256_c(
base_ptr: *const u64,
exp_ptr: *const u64,
modulus_ptr: *const u64,
result_ptr: *mut u64,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) {
let base = &*(base_ptr as *const [u64; 4]);
let exp = &*(exp_ptr as *const [u64; 4]);
let modulus = &*(modulus_ptr as *const [u64; 4]);
let res = pow_mod256(
base,
exp,
modulus,
#[cfg(feature = "hints")]
hints,
);
let result = &mut *(result_ptr as *mut [u64; 4]);
*result = res;
}
#[cfg_attr(not(feature = "hints"), no_mangle)]
#[cfg_attr(feature = "hints", export_name = "hints_inv_mod256_c")]
pub unsafe extern "C" fn inv_mod256_c(
a_ptr: *const u64,
modulus_ptr: *const u64,
result_ptr: *mut u64,
#[cfg(feature = "hints")] hints: &mut Vec<u64>,
) -> u8 {
let a = &*(a_ptr as *const [u64; 4]);
let modulus = &*(modulus_ptr as *const [u64; 4]);
match inv_mod256(
a,
modulus,
#[cfg(feature = "hints")]
hints,
) {
Some(res) => {
let result = &mut *(result_ptr as *mut [u64; 4]);
*result = res;
1
}
None => 0,
}
}