use core::cmp::Ordering;
use miden_core::Felt;
use super::arithmetic::{
add_mod, barrett_mu, cmp, inv_mod_prime_barrett, mul_mod_barrett, sub_mod, sub_small,
wrapping_add, wrapping_mul, wrapping_sub,
};
pub type Limbs = [u32; 8];
pub const ZERO_LIMBS: Limbs = [0; 8];
pub const ONE_LIMBS: Limbs = [1, 0, 0, 0, 0, 0, 0, 0];
pub const TWO_LIMBS: Limbs = [2, 0, 0, 0, 0, 0, 0, 0];
pub trait UintSpec: 'static {
const ID: Felt;
const ENCODED_MODULUS: Limbs;
const BARRETT_MU: [u32; 9] = barrett_mu(Self::ENCODED_MODULUS);
const IS_PRIME_FIELD: bool = false;
fn is_canonical(value: &Limbs) -> bool {
if Self::ENCODED_MODULUS == ZERO_LIMBS {
true
} else {
cmp(value, &Self::ENCODED_MODULUS) == Ordering::Less
}
}
fn add(lhs: Limbs, rhs: Limbs) -> Limbs {
if Self::ENCODED_MODULUS == ZERO_LIMBS {
wrapping_add(lhs, rhs)
} else {
add_mod(lhs, rhs, Self::ENCODED_MODULUS)
}
}
fn sub(lhs: Limbs, rhs: Limbs) -> Limbs {
if Self::ENCODED_MODULUS == ZERO_LIMBS {
wrapping_sub(lhs, rhs)
} else {
sub_mod(lhs, rhs, Self::ENCODED_MODULUS)
}
}
fn mul(lhs: Limbs, rhs: Limbs) -> Limbs {
if Self::ENCODED_MODULUS == ZERO_LIMBS {
wrapping_mul(lhs, rhs)
} else {
mul_mod_barrett(lhs, rhs, Self::ENCODED_MODULUS, Self::BARRETT_MU)
}
}
fn inv(value: Limbs) -> Option<Limbs> {
if Self::IS_PRIME_FIELD && Self::ENCODED_MODULUS != ZERO_LIMBS {
inv_mod_prime_barrett(value, Self::ENCODED_MODULUS, Self::BARRETT_MU)
} else {
None
}
}
fn minus_one() -> Limbs {
if Self::ENCODED_MODULUS == ZERO_LIMBS {
[u32::MAX; 8]
} else {
sub_small(Self::ENCODED_MODULUS, 1)
}
}
fn half() -> Option<Limbs> {
Self::inv(TWO_LIMBS)
}
fn pow2_mod(exponent: usize) -> Option<Limbs> {
if !Self::IS_PRIME_FIELD || Self::ENCODED_MODULUS == ZERO_LIMBS {
return None;
}
let mut value = ONE_LIMBS;
for _ in 0..exponent {
value = Self::add(value, value);
}
Some(value)
}
}