use core::array;
pub use ruint::aliases::{U256, U320, U512};
pub type U576 = ruint::Uint<576, 9>;
pub fn from_limbs16(v: &[u16; 16]) -> U256 {
U256::from_limbs(array::from_fn(|w| {
(0..4).fold(0u64, |acc, i| acc | (u64::from(v[4 * w + i]) << (16 * i)))
}))
}
pub fn to_limbs16(v: U256) -> [u16; 16] {
array::from_fn(|i| (v.as_limbs()[i / 4] >> (16 * (i % 4))) as u16)
}
pub fn from_limbs32(v: &[u32; 8]) -> U256 {
U256::from_limbs(array::from_fn(|w| u64::from(v[2 * w]) | (u64::from(v[2 * w + 1]) << 32)))
}
pub fn to_limbs32(v: U256) -> [u32; 8] {
array::from_fn(|i| (v.as_limbs()[i / 2] >> (32 * (i % 2))) as u32)
}
pub fn from_hex(s: &str) -> U256 {
U256::from_str_radix(s, 16).expect("valid big-endian hex")
}
fn p320(bound: U256) -> U320 {
U320::from(bound) + U320::ONE
}
pub fn add_reduce(a: U256, b: U256, bound: U256) -> U256 {
U320::from(a).add_mod(U320::from(b), p320(bound)).to()
}
pub fn sub_reduce(a: U256, b: U256, bound: U256) -> U256 {
let p = p320(bound);
U320::from(a).add_mod(p - U320::from(b), p).to()
}
pub fn mac_reduce(kappa_a: u16, a: U256, b: U256, kappa_c: u16, c: U256, bound: U256) -> U256 {
mac_div_rem(kappa_a, a, b, kappa_c, c, bound).1
}
pub fn mac_div_rem(
kappa_a: u16,
a: U256,
b: U256,
kappa_c: u16,
c: U256,
bound: U256,
) -> (U320, U256) {
let ab: U512 = a.widening_mul(b);
let n = U576::from(ab) * U576::from(kappa_a) + U576::from(c) * U576::from(kappa_c);
let (q, r) = n.div_rem(U576::from(bound) + U576::ONE);
(q.to(), r.to())
}
pub fn mac_sub_reduce(kappa_a: u16, a: U256, b: U256, kappa_c: u16, c: U256, bound: U256) -> U256 {
mac_sub_div_rem(kappa_a, a, b, kappa_c, c, bound).1
}
pub fn mac_sub_div_rem(
kappa_a: u16,
a: U256,
b: U256,
kappa_c: u16,
c: U256,
bound: U256,
) -> (U320, U256, u8) {
let ab: U512 = a.widening_mul(b);
let p = U576::from(bound) + U576::ONE;
let prod = U576::from(ab) * U576::from(kappa_a);
let sub = U576::from(c) * U576::from(kappa_c);
if prod >= sub {
let (q, r) = (prod - sub).div_rem(p);
(q.to(), r.to(), 0)
} else {
let d = sub - prod;
let (qd, rd) = d.div_rem(p);
debug_assert!(qd < U576::from(2u32), "mac_sub underflow exceeds 2p (κ_c > 2?)");
let qd = qd.as_limbs()[0] as u8;
if rd == U576::ZERO {
(U320::ZERO, U256::ZERO, qd)
} else {
(U320::ZERO, (p - rd).to(), qd + 1)
}
}
}
pub fn mod_inv(v: U256, bound: U256) -> U256 {
U320::from(v)
.inv_mod(p320(bound))
.expect("value not invertible under the modulus")
.to()
}