use super::{GenericMultiBitEncryption, PKEncryptionScheme};
use qfall_math::{
error::MathError,
integer::{MatZ, Z},
integer_mod_q::{MatZq, Modulus, Zq},
rational::Q,
traits::{Concatenate, Distance, MatrixGetEntry, MatrixSetEntry, Pow},
};
use serde::{Deserialize, Serialize};
#[derive(Debug, Serialize, Deserialize)]
pub struct Regev {
n: Z, m: Z, q: Modulus, alpha: Q, }
impl Regev {
pub fn new(
n: impl Into<Z>,
m: impl Into<Z>,
q: impl Into<Modulus>,
alpha: impl Into<Q>,
) -> Self {
let n: Z = n.into();
let m: Z = m.into();
let q: Modulus = q.into();
let alpha: Q = alpha.into();
Self { n, m, q, alpha }
}
pub fn new_from_n(n: impl Into<Z>) -> Self {
let n = n.into();
if n < 10 {
panic!(
"Choose n >= 10 as this function does not return parameters ensuring proper correctness of the scheme otherwise."
);
}
let mut m: Z;
let mut q: Modulus;
let mut alpha: Q;
(m, q, alpha) = Self::gen_new_public_parameters(&n);
let mut out = Self {
n: n.clone(),
m,
q,
alpha,
};
while out.check_correctness().is_err() || out.check_security().is_err() {
(m, q, alpha) = Self::gen_new_public_parameters(&n);
out = Self {
n: n.clone(),
m,
q,
alpha,
};
}
out
}
fn gen_new_public_parameters(n: &Z) -> (Z, Modulus, Q) {
let n_i64 = i64::try_from(n).unwrap();
let power = match n_i64 {
2..=4 => 5,
5 => 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 = (n + Z::ONE) * q.log(2).unwrap().ceil();
let alpha = 1 / (2 * n.sqrt() * n.log(2).unwrap().pow(2).unwrap());
let q = Modulus::from(q);
(m, q, alpha)
}
pub fn check_correctness(&self) -> Result<(), MathError> {
let q = Z::from(&self.q);
if self.n <= Z::ONE {
return Err(MathError::InvalidIntegerInput(String::from(
"n must be chosen bigger than 1.",
)));
}
if self.alpha > 1 / (self.n.sqrt() * self.n.log(2).unwrap()) {
return Err(MathError::InvalidIntegerInput(String::from(
"Correctness is not guaranteed as α >= 1 / (sqrt(n) * log n), but α < 1 / (sqrt(n) * log n) is required.",
)));
}
if 20 * self.m.sqrt() * &self.alpha > q {
return Err(MathError::InvalidIntegerInput(String::from(
"Correctness is not guaranteed as 5 * sqrt(m) * α > q/4, but 5 * sqrt(m) * α <= q/4 is required.",
)));
}
Ok(())
}
pub fn check_security(&self) -> Result<(), MathError> {
let q = Z::from(&self.q);
if &q * &self.alpha < 2 * self.n.sqrt() {
return Err(MathError::InvalidIntegerInput(String::from(
"Security is not guaranteed as q * α < 2 * sqrt(n), but q * α >= 2 * sqrt(n) is required.",
)));
}
if self.m <= ((&self.n + Z::ONE) * q.log(2).unwrap()).ceil() {
return Err(MathError::InvalidIntegerInput(String::from(
"Security is not guaranteed as m <= (n + 1) log q, but m > (n + 1) log q is required.",
)));
}
Ok(())
}
}
impl Default for Regev {
fn default() -> Self {
let n = Z::from(13);
let m = Z::from(154);
let q = Modulus::from(1427);
let alpha = Q::from(0.01);
Self { n, m, q, alpha }
}
}
impl PKEncryptionScheme for Regev {
type Cipher = MatZq;
type PublicKey = MatZq;
type SecretKey = MatZq;
fn key_gen(&self) -> (Self::PublicKey, Self::SecretKey) {
let mat_a = MatZq::sample_uniform(&self.n, &self.m, &self.q);
let vec_s = MatZq::sample_uniform(&self.n, 1, &self.q);
let vec_e_t =
MatZq::sample_discrete_gauss(1, &self.m, &self.q, 0, &self.alpha * Z::from(&self.q))
.unwrap();
let vec_b_t = vec_s.transpose() * &mat_a + vec_e_t;
let mat_a = mat_a.concat_vertical(&vec_b_t).unwrap();
(mat_a, vec_s)
}
fn enc(&self, pk: &Self::PublicKey, message: impl Into<Z>) -> Self::Cipher {
let message: Z = message.into() % 2;
let vec_x = MatZ::sample_uniform(&self.m, 1, 0, 2).unwrap();
let mut c = pk * vec_x;
let msg_q_half = message * Z::from(&self.q).div_floor(2);
let last_entry: Zq = c.get_entry(-1, 0).unwrap();
c.set_entry(-1, 0, last_entry + msg_q_half).unwrap();
c
}
fn dec(&self, sk: &Self::SecretKey, cipher: &Self::Cipher) -> Z {
let result = (Z::MINUS_ONE * sk)
.concat_vertical(&MatZq::identity(1, 1, &self.q))
.unwrap()
.dot_product(cipher)
.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 Regev {}
#[cfg(test)]
mod test_pp_generation {
use super::Regev;
use super::Z;
#[test]
fn new_availability() {
let _ = Regev::new(2u8, 2u16, 2u32, 2u64);
let _ = Regev::new(2u16, 2u64, 2i32, 2i64);
let _ = Regev::new(2i16, 2i64, 2u32, 2u8);
let _ = Regev::new(Z::from(2), Z::from(2), 2u8, 2i8);
}
#[test]
fn suitable_security_params() {
let n_choices = [
10, 11, 12, 13, 14, 25, 50, 100, 250, 500, 1000, 2500, 5000, 5001, 10000,
];
for n in n_choices {
let _ = Regev::new_from_n(n);
}
}
#[test]
fn default_suitable() {
let regev = Regev::default();
assert!(regev.check_correctness().is_ok());
assert!(regev.check_security().is_ok());
}
#[test]
fn choice_valid() {
let n_choices = [10, 14, 25, 50, 125, 300, 600, 1200, 4000, 6000];
for n in n_choices {
let regev = Regev::new_from_n(n);
assert!(regev.check_correctness().is_ok());
assert!(regev.check_security().is_ok());
}
}
#[test]
#[allow(clippy::needless_borrows_for_generic_args)]
fn availability() {
let _ = Regev::new_from_n(10u8);
let _ = Regev::new_from_n(10u16);
let _ = Regev::new_from_n(10u32);
let _ = Regev::new_from_n(10u64);
let _ = Regev::new_from_n(10i8);
let _ = Regev::new_from_n(10i16);
let _ = Regev::new_from_n(10i32);
let _ = Regev::new_from_n(10i64);
let _ = Regev::new_from_n(Z::from(10));
let _ = Regev::new_from_n(&Z::from(10));
}
#[test]
#[should_panic]
fn invalid_n() {
Regev::new_from_n(9);
}
}
#[cfg(test)]
mod test_regev {
use super::Regev;
use crate::pk_encryption::PKEncryptionScheme;
use qfall_math::integer::Z;
#[test]
fn cycle_zero_small_n() {
let msg = Z::ZERO;
let regev = Regev::default();
let (pk, sk) = regev.key_gen();
let cipher = regev.enc(&pk, &msg);
let m = regev.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_one_small_n() {
let msg = Z::ONE;
let regev = Regev::default();
let (pk, sk) = regev.key_gen();
let cipher = regev.enc(&pk, &msg);
let m = regev.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_zero_large_n() {
let msg = Z::ZERO;
let regev = Regev::new_from_n(50);
let (pk, sk) = regev.key_gen();
let cipher = regev.enc(&pk, &msg);
let m = regev.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_one_large_n() {
let msg = Z::ONE;
let regev = Regev::new_from_n(50);
let (pk, sk) = regev.key_gen();
let cipher = regev.enc(&pk, &msg);
let m = regev.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn modulus_application() {
let messages = [2, 3, i64::MAX, i64::MIN];
let regev = Regev::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, Regev};
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 = Regev::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 = Regev::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 = Regev::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);
}
}
}