use super::{GenericMultiBitEncryption, PKEncryptionScheme};
use qfall_math::{
error::MathError,
integer::Z,
integer_mod_q::{MatZq, Modulus, Zq},
rational::Q,
traits::{Distance, Pow},
};
use serde::{Deserialize, Serialize};
#[derive(Debug, Serialize, Deserialize)]
pub struct RegevWithDiscreteGaussianRegularity {
n: Z, m: Z, q: Modulus, r: Q, alpha: Q, }
impl RegevWithDiscreteGaussianRegularity {
pub fn new(
n: impl Into<Z>,
m: impl Into<Z>,
q: impl Into<Modulus>,
r: impl Into<Q>,
alpha: impl Into<Q>,
) -> Self {
let n: Z = n.into();
let m: Z = m.into();
let q: Modulus = q.into();
let r: Q = r.into();
let alpha: Q = alpha.into();
Self { n, m, q, r, alpha }
}
pub fn new_from_n(n: impl Into<Z>) -> Self {
let n = n.into();
if n <= Z::ONE {
panic!("n must be chosen bigger than 1.");
}
let mut m: Z;
let mut q: Modulus;
let mut r: Q;
let mut alpha: Q;
(m, q, r, alpha) = Self::gen_new_public_parameters(&n);
let mut out = Self {
n: n.clone(),
m,
q,
r,
alpha,
};
while out.check_correctness().is_err() || out.check_security().is_err() {
(m, q, r, alpha) = Self::gen_new_public_parameters(&n);
out = Self {
n: n.clone(),
m,
q,
r,
alpha,
};
}
out
}
fn gen_new_public_parameters(n: &Z) -> (Z, Modulus, Q, Q) {
let n_i64 = i64::try_from(n).unwrap();
let power = match n_i64 {
2 => 9,
3 => 8,
4..=5 => 7,
6..=8 => 6,
9..=12 => 5,
13..=30 => 4,
_ => 3,
};
let upper_bound: Z = n.pow(power).unwrap();
let lower_bound = upper_bound.div_ceil(2);
let q = Z::sample_prime_uniform(&lower_bound, &upper_bound).unwrap();
let m = (Z::from(2) * (n + Z::ONE) * q.log(10).unwrap()).ceil();
let r = m.log(2).unwrap();
let alpha = 1 / (m.sqrt() * m.log(2).unwrap().pow(2).unwrap());
let q = Modulus::from(&q);
(m, q, r, alpha)
}
pub fn check_correctness(&self) -> Result<(), MathError> {
let q: Z = Z::from(&self.q);
if self.n <= Z::ONE {
return Err(MathError::InvalidIntegerInput(String::from(
"n must be chosen bigger than 1.",
)));
}
if q < 5 * &self.r * &self.m {
return Err(MathError::InvalidIntegerInput(String::from(
"Correctness is not guaranteed as q < 5rm, but q >= 5rm is required.",
)));
}
if self.alpha > 1 / (&self.r * self.m.sqrt() * self.n.log(2).unwrap().sqrt()) {
return Err(MathError::InvalidIntegerInput(String::from(
"Correctness is not guaranteed as α > 1/(r*sqrt(m)*ω(sqrt(log n)), but α <= 1/(r*sqrt(m)*ω(sqrt(log n)) is required.",
)));
}
Ok(())
}
pub fn check_security(&self) -> Result<(), MathError> {
let q: Z = Z::from(&self.q);
if &q * &self.alpha < self.n {
return Err(MathError::InvalidIntegerInput(String::from(
"Security is not guaranteed as q * α < n, but q * α >= n is required.",
)));
}
if self.m < 2 * (&self.n + 1) * q.log(10).unwrap() {
return Err(MathError::InvalidIntegerInput(String::from(
"Security is not guaranteed as m < 2(n + 1) lg (q), but m >= 2(n + 1) lg (q) is required.",
)));
}
if self.r < self.m.log(2).unwrap().sqrt() {
return Err(MathError::InvalidIntegerInput(String::from(
"Security is not guaranteed as r < sqrt( log m ) and r >= ω(sqrt(log m)) is required.",
)));
}
Ok(())
}
}
impl Default for RegevWithDiscreteGaussianRegularity {
fn default() -> Self {
let n = Z::from(2);
let m = Z::from(16);
let q = Modulus::from(443);
let r = Q::from(4);
let alpha = Q::from((1, 64));
Self { n, m, q, r, alpha }
}
}
impl PKEncryptionScheme for RegevWithDiscreteGaussianRegularity {
type Cipher = (MatZq, Zq);
type PublicKey = (MatZq, MatZq);
type SecretKey = MatZq;
fn key_gen(&self) -> (Self::PublicKey, Self::SecretKey) {
let vec_s = MatZq::sample_uniform(&self.n, 1, &self.q);
let mat_a = MatZq::sample_uniform(&self.n, &self.m, &self.q);
let vec_x =
MatZq::sample_discrete_gauss(&self.m, 1, &self.q, 0, &(&self.alpha * Z::from(&self.q)))
.unwrap();
let vec_p = mat_a.transpose() * &vec_s + vec_x;
((mat_a, vec_p), vec_s)
}
fn enc(&self, pk: &Self::PublicKey, message: impl Into<Z>) -> Self::Cipher {
let message: Z = message.into() % 2;
let vec_e = MatZq::sample_discrete_gauss(&self.m, 1, &self.q, 0, &self.r).unwrap();
let vec_u = &pk.0 * &vec_e;
let q_half = Z::from(&self.q).div_floor(2);
let c = pk.1.dot_product(&vec_e).unwrap() + message * q_half;
(vec_u, c)
}
fn dec(&self, sk: &Self::SecretKey, cipher: &Self::Cipher) -> Z {
let result = &cipher.1 - sk.dot_product(&cipher.0).unwrap();
let result: Z = result.get_representative_least_absolute_residue().abs();
let q_half = Z::from(&self.q).div_floor(2);
if result.distance(Z::ZERO) > result.distance(q_half) {
Z::ONE
} else {
Z::ZERO
}
}
}
impl GenericMultiBitEncryption for RegevWithDiscreteGaussianRegularity {}
#[cfg(test)]
mod test_pp_generation {
use super::RegevWithDiscreteGaussianRegularity;
use super::Z;
#[test]
fn new_availability() {
let _ = RegevWithDiscreteGaussianRegularity::new(2u8, 2u16, 2u32, 2u64, 2i8);
let _ = RegevWithDiscreteGaussianRegularity::new(2u16, 2u64, 2i32, 2i64, 2i16);
let _ = RegevWithDiscreteGaussianRegularity::new(2i16, 2i64, 2u32, 2u8, 2u16);
let _ = RegevWithDiscreteGaussianRegularity::new(Z::from(2), Z::from(2), 2u8, 2i8, 2u32);
}
#[test]
fn suitable_security_params() {
let n_choices = [
2, 3, 4, 5, 6, 7, 8, 9, 10, 25, 50, 75, 100, 250, 500, 750, 1000, 2500, 5000, 5001,
10000,
];
for n in n_choices {
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(n);
}
}
#[test]
fn default_suitable() {
let dr = RegevWithDiscreteGaussianRegularity::default();
assert!(dr.check_correctness().is_ok());
assert!(dr.check_security().is_ok());
}
#[test]
fn choice_valid() {
let n_choices = [
2, 3, 4, 5, 6, 7, 8, 9, 10, 25, 50, 75, 100, 250, 500, 750, 1000, 2500, 5000, 5001,
10000,
];
for n in n_choices {
let dr = RegevWithDiscreteGaussianRegularity::new_from_n(n);
assert!(dr.check_correctness().is_ok());
assert!(dr.check_security().is_ok());
}
}
#[test]
#[allow(clippy::needless_borrows_for_generic_args)]
fn new_from_n_availability() {
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2u8);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2u16);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2u32);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2u64);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2i8);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2i16);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2i32);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(2i64);
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(Z::from(2));
let _ = RegevWithDiscreteGaussianRegularity::new_from_n(&Z::from(2));
}
#[test]
#[should_panic]
fn invalid_n() {
RegevWithDiscreteGaussianRegularity::new_from_n(1);
}
}
#[cfg(test)]
mod test_regev {
use super::RegevWithDiscreteGaussianRegularity;
use crate::pk_encryption::PKEncryptionScheme;
use qfall_math::integer::Z;
#[test]
fn cycle_zero_small_n() {
let msg = Z::ZERO;
let dr = RegevWithDiscreteGaussianRegularity::default();
let (pk, sk) = dr.key_gen();
let cipher = dr.enc(&pk, &msg);
let m = dr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_one_small_n() {
let msg = Z::ONE;
let dr = RegevWithDiscreteGaussianRegularity::default();
let (pk, sk) = dr.key_gen();
let cipher = dr.enc(&pk, &msg);
let m = dr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_zero_large_n() {
let msg = Z::ZERO;
let dr = RegevWithDiscreteGaussianRegularity::new_from_n(30);
let (pk, sk) = dr.key_gen();
let cipher = dr.enc(&pk, &msg);
let m = dr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_one_large_n() {
let msg = Z::ONE;
let dr = RegevWithDiscreteGaussianRegularity::new_from_n(30);
let (pk, sk) = dr.key_gen();
let cipher = dr.enc(&pk, &msg);
let m = dr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn modulus_application() {
let messages = [2, 3, i64::MAX, i64::MIN];
let regev = RegevWithDiscreteGaussianRegularity::default();
let (pk, sk) = regev.key_gen();
for msg in messages {
let msg_mod = Z::from(msg.rem_euclid(2));
let cipher = regev.enc(&pk, msg);
let m = regev.dec(&sk, &cipher);
assert_eq!(msg_mod, m);
}
}
}
#[cfg(test)]
mod test_multi_bits {
use super::{
GenericMultiBitEncryption, PKEncryptionScheme, RegevWithDiscreteGaussianRegularity,
};
use qfall_math::integer::Z;
#[test]
fn positive() {
let values = [3, 13, 23, 230, 501, 1024, i64::MAX];
for value in values {
let msg = Z::from(value);
let scheme = RegevWithDiscreteGaussianRegularity::default();
let (pk, sk) = scheme.key_gen();
let cipher = scheme.enc_multiple_bits(&pk, &msg);
let m = scheme.dec_multiple_bits(&sk, &cipher);
assert_eq!(msg, m);
}
}
#[test]
fn zero() {
let msg = Z::ZERO;
let scheme = RegevWithDiscreteGaussianRegularity::default();
let (pk, sk) = scheme.key_gen();
let cipher = scheme.enc_multiple_bits(&pk, &msg);
let m = scheme.dec_multiple_bits(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn negative() {
let values = [-3, -13, -23, -230, -501, -1024, i64::MIN];
for value in values {
let msg = Z::from(value);
let scheme = RegevWithDiscreteGaussianRegularity::default();
let (pk, sk) = scheme.key_gen();
let cipher = scheme.enc_multiple_bits(&pk, &msg);
let m = scheme.dec_multiple_bits(&sk, &cipher);
assert_eq!(msg.abs(), m);
}
}
}