use crate::*;
use rand::{Rng, distributions::Distribution};
use std::marker::PhantomData;
#[derive(Clone)]
pub struct ModularCtx<T, C> {
phantom: PhantomData<C>,
modulo: T, use_barrett: bool,
mu: T, }
impl<
C,
T: CyclicOrdZeroDyn<C> + ClosedAddDyn<C> + ClosedSubDyn<C> + ZeroDyn<C> + OneDyn<C> + EuclidDyn<C>,
> ModularCtx<T, C>
{
pub fn new(modulo: T, ctx: &C) -> Self {
let max = T::zero_d(ctx).sub_d(&T::one_d(ctx), ctx);
let (mut mu, mut mu_rem) = max.euclid_div_rem_d(&modulo, ctx);
mu_rem.add_assign_d(&T::one_d(ctx), ctx);
if !mu_rem.cyclic_lt0_d(&modulo, ctx) {
mu.add_assign_d(&T::one_d(ctx), ctx);
}
let use_barrett = mu.cyclic_lt0_d(&modulo, ctx);
Self {
phantom: PhantomData,
modulo,
use_barrett,
mu,
}
}
}
#[repr(transparent)]
#[derive(Clone, Eq, PartialEq)]
pub struct ModularDyn<T> {
pub inner: T,
}
impl<T> ModularDyn<T> {
pub fn new_d<C>(value: T, ctx: &(ModularCtx<T, C>, &C)) -> Self
where
T: CyclicOrdZeroDyn<C>,
{
let m = &ctx.0.modulo;
let c = ctx.1;
debug_assert!(value.cyclic_lt0_d(m, c));
Self { inner: value }
}
}
impl<C, T: CyclicOrdDyn<C>> CyclicOrdDyn<(ModularCtx<T, C>, &C)> for ModularDyn<T> {
fn cyclic_lt_d(&self, low: &Self, high: &Self, ctx: &(ModularCtx<T, C>, &C)) -> bool {
let c = ctx.1;
self.inner.cyclic_lt_d(&low.inner, &high.inner, c)
}
}
impl<C, T: CyclicOrdZeroDyn<C>> CyclicOrdZeroDyn<(ModularCtx<T, C>, &C)> for ModularDyn<T> {
fn cyclic_lt0_d(&self, high: &Self, ctx: &(ModularCtx<T, C>, &C)) -> bool {
let c = ctx.1;
self.inner.cyclic_lt0_d(&high.inner, c)
}
}
impl<C, T: CyclicOrdZeroDyn<C> + ClosedAddDyn<C> + ClosedSubDyn<C>>
ClosedAddDyn<(ModularCtx<T, C>, &C)> for ModularDyn<T>
{
fn add_d(&self, rhs: &Self, ctx: &(ModularCtx<T, C>, &C)) -> Self {
let m = &ctx.0.modulo;
let c = ctx.1;
let sum = self.inner.add_d(&rhs.inner, c);
let r = if !sum.cyclic_lt0_d(m, c)
|| sum.cyclic_lt0_d(&self.inner, c)
|| sum.cyclic_lt0_d(&rhs.inner, c)
{
sum.sub_d(m, c)
} else {
sum
};
Self::new_d(r, ctx)
}
}
#[test]
fn add_test() {
use rand::rngs::mock::StepRng;
let mut rng = StepRng::new(0, 0x54825a7f54825a7f);
let base_dist = StandardDyn::new(&());
for _ in 0..100 {
let m: Z2_8 = rng.sample(&base_dist);
if m.is_zero_d(&()) {
continue;
}
let c = ModularCtx::new(m.clone(), &());
let ctx = (c, &());
let modular_dist = StandardDyn::new(&ctx);
let a: ModularDyn<Z2_8> = rng.sample(&modular_dist);
let b: ModularDyn<Z2_8> = rng.sample(&modular_dist);
let r = a.add_d(&b, &ctx);
let expected = ((a.inner.inner as u16 + b.inner.inner as u16) % m.inner as u16) as u8;
println!(
"{} + {} mod {} = {}",
a.inner.inner, b.inner.inner, m.inner, expected
);
assert_eq!(r.inner.inner, expected);
}
}
impl<C, T: CyclicOrdZeroDyn<C> + ClosedAddDyn<C> + ClosedSubDyn<C>>
ClosedSubDyn<(ModularCtx<T, C>, &C)> for ModularDyn<T>
{
fn sub_d(&self, rhs: &Self, ctx: &(ModularCtx<T, C>, &C)) -> Self {
let m = &ctx.0.modulo;
let c = ctx.1;
let diff = self.inner.sub_d(&rhs.inner, c);
let r = if self.inner.cyclic_lt0_d(&diff, c) {
diff.add_d(m, c)
} else {
diff
};
Self::new_d(r, ctx)
}
}
#[test]
fn sub_test() {
use rand::rngs::mock::StepRng;
let mut rng = StepRng::new(0, 0x54825a7f54825a7f);
let base_dist = StandardDyn::new(&());
for _ in 0..100 {
let m: Z2_8 = rng.sample(&base_dist);
if m.is_zero_d(&()) {
continue;
}
let c = ModularCtx::new(m.clone(), &());
let ctx = (c, &());
let modular_dist = StandardDyn::new(&ctx);
let a: ModularDyn<Z2_8> = rng.sample(&modular_dist);
let b: ModularDyn<Z2_8> = rng.sample(&modular_dist);
let r = a.sub_d(&b, &ctx);
let expected =
((a.inner.inner as u16 + m.inner as u16 - b.inner.inner as u16) % m.inner as u16) as u8;
println!(
"{} - {} mod {} = {}",
a.inner.inner, b.inner.inner, m.inner, expected
);
assert_eq!(r.inner.inner, expected);
}
}
impl<C, T: CyclicOrdZeroDyn<C> + ZeroDyn<C>> ZeroDyn<(ModularCtx<T, C>, &C)> for ModularDyn<T> {
fn zero_d(ctx: &(ModularCtx<T, C>, &C)) -> Self {
let c = ctx.1;
Self::new_d(T::zero_d(c), ctx)
}
}
impl<
C,
T: CyclicOrdZeroDyn<C>
+ ClosedSubDyn<C>
+ ZeroDyn<C>
+ CenteredMulDyn<C>
+ OneDyn<C>
+ EuclidDyn<C>,
> ClosedMulDyn<(ModularCtx<T, C>, &C)> for ModularDyn<T>
{
fn mul_d(&self, rhs: &Self, ctx: &(ModularCtx<T, C>, &C)) -> Self {
let m = &ctx.0.modulo;
let c = ctx.1;
Self::new_d(
if !&ctx.0.use_barrett {
self.inner.mul_d(&rhs.inner, c).euclid_rem_d(m, c)
} else {
let mu = &ctx.0.mu;
let (mut low, mut high) = self.inner.widening_mul_d(&rhs.inner, c);
while !high.is_zero_d(c) {
let mut q_high = high.mul_d(mu, c);
q_high.add_assign_d(&low.centered_mul_d(mu, c), c);
let (q_m_low, q_m_high) = q_high.widening_mul_d(m, c);
let r_low = low.sub_d(&q_m_low, c);
let mut r_high = high.sub_d(&q_m_high, c);
if low.cyclic_lt0_d(&q_m_low, c) {
r_high.sub_assign_d(&T::one_d(c), c);
}
low = r_low;
high = r_high;
}
low.euclid_rem_d(m, c)
},
ctx,
)
}
}
#[test]
fn mul_test() {
use rand::rngs::mock::StepRng;
let mut rng = StepRng::new(0, 0x54825a7f54825a7f);
let base_dist = StandardDyn::new(&());
for _ in 0..100 {
let m: Z2_8 = rng.sample(&base_dist);
if m.is_zero_d(&()) {
continue;
}
let c = ModularCtx::new(m.clone(), &());
let ctx = (c, &());
let modular_dist = StandardDyn::new(&ctx);
let a: ModularDyn<Z2_8> = rng.sample(&modular_dist);
let b: ModularDyn<Z2_8> = rng.sample(&modular_dist);
let r = a.mul_d(&b, &ctx);
let expected = ((a.inner.inner as u16 * b.inner.inner as u16) % m.inner as u16) as u8;
println!(
"{} * {} mod {} = {}",
a.inner.inner, b.inner.inner, m.inner, expected,
);
assert_eq!(r.inner.inner, expected);
}
}
impl<C, T: CyclicOrdZeroDyn<C> + OneDyn<C>> OneDyn<(ModularCtx<T, C>, &C)> for ModularDyn<T> {
fn one_d(ctx: &(ModularCtx<T, C>, &C)) -> Self {
let c = ctx.1;
Self::new_d(T::one_d(c), ctx)
}
}
impl<C, D, T, Rhs: ClosedAddDyn<D> + ClosedSubDyn<D> + ZeroDyn<D> + OneDyn<D> + EuclidDyn<D>>
PowDyn<C, D, Rhs> for ModularDyn<T>
where
Self: Clone + ClosedMulDyn<C> + OneDyn<C>,
{
}
impl<C, T: CyclicOrdZeroDyn<C> + ClosedAddDyn<C> + ClosedMulDyn<C> + OneDyn<C> + EuclidDyn<C>>
Distribution<ModularDyn<T>> for StandardDyn<'_, (ModularCtx<T, C>, &C)>
where
for<'a> StandardDyn<'a, C>: Distribution<T>,
{
fn sample<R: ?Sized + Rng>(&self, rng: &mut R) -> ModularDyn<T> {
let m = &self.ctx.0.modulo;
let c = self.ctx.1;
let dist = StandardDyn::new(c);
loop {
let x: T = rng.sample(&dist);
let (q, r) = x.euclid_div_rem_d(m, c);
let prod = q.mul_d(&m.add_d(&T::one_d(c), c), c);
if !prod.cyclic_lt0_d(m, c) {
return ModularDyn { inner: r };
}
}
}
}