use super::{GenericMultiBitEncryption, PKEncryptionScheme};
use qfall_math::{
error::MathError,
integer::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 LPR {
n: Z, q: Modulus, alpha: Q, }
impl LPR {
pub fn new(n: impl Into<Z>, q: impl Into<Modulus>, alpha: impl Into<Q>) -> Self {
let n: Z = n.into();
let q: Modulus = q.into();
let alpha: Q = alpha.into();
Self { n, q, alpha }
}
pub fn new_from_n(n: impl Into<Z>) -> Self {
let n = n.into();
assert!(
n >= 10,
"Choose n >= 10 as this function does not return parameters ensuring proper correctness of the scheme otherwise."
);
let mut q: Modulus;
let mut alpha: Q;
(q, alpha) = Self::gen_new_public_parameters(&n);
let mut out = Self {
n: n.clone(),
q,
alpha,
};
while out.check_correctness().is_err() || out.check_security().is_err() {
(q, alpha) = Self::gen_new_public_parameters(&n);
out = Self {
n: n.clone(),
q,
alpha,
};
}
out
}
fn gen_new_public_parameters(n: &Z) -> (Modulus, Q) {
let n_i64 = i64::try_from(n).unwrap();
let upper_bound: Z = n.pow(3).unwrap();
let lower_bound = upper_bound.div_ceil(2);
let q = Z::sample_prime_uniform(&lower_bound, &upper_bound).unwrap();
let factor = match n_i64 {
1..=20 => 1,
21..=40 => 2,
41..=80 => 3,
81..=160 => 4,
_ => 5,
};
let alpha = 1 / (factor * n.sqrt() * n.log(2).unwrap().pow(3).unwrap());
let q = Modulus::from(q);
(q, alpha)
}
pub fn check_correctness(&self) -> Result<(), MathError> {
let n_i64 = i64::try_from(&self.n)?;
if self.n <= Z::ONE {
return Err(MathError::InvalidIntegerInput(String::from(
"n must be chosen bigger than 1.",
)));
}
let factor = match n_i64 {
1..=20 => 1,
21..=40 => 2,
41..=80 => 3,
81..=160 => 4,
_ => 5,
};
if self.alpha > 1 / (factor * self.n.sqrt() * self.n.log(2).unwrap().pow(3).unwrap()) {
return Err(MathError::InvalidIntegerInput(String::from(
"Correctness is not guaranteed as α >= 1 / (sqrt(n) * log^3 n), but α < 1 / (sqrt(n) * log^3 n) is required. Please check the documentation!",
)));
}
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.",
)));
}
Ok(())
}
}
impl Default for LPR {
fn default() -> Self {
let n = Z::from(10);
let q = Modulus::from(983);
let alpha = Q::from(0.0072);
Self { n, q, alpha }
}
}
impl PKEncryptionScheme for LPR {
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.n, &self.q);
let vec_s =
MatZq::sample_discrete_gauss(&self.n, 1, &self.q, 0, &self.alpha * Z::from(&self.q))
.unwrap();
let vec_e_t =
MatZq::sample_discrete_gauss(1, &self.n, &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_r =
MatZq::sample_discrete_gauss(&self.n, 1, &self.q, 0, &self.alpha * Z::from(&self.q))
.unwrap();
let vec_e = MatZq::sample_discrete_gauss(
&(&self.n + 1),
1,
&self.q,
0,
&self.alpha * Z::from(&self.q),
)
.unwrap();
let mut c = pk * vec_r + vec_e;
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.transpose())
.concat_horizontal(&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 LPR {}
#[cfg(test)]
mod test_pp_generation {
use super::LPR;
use super::Z;
#[test]
fn new_availability() {
let _ = LPR::new(2u8, 2u32, 2u64);
let _ = LPR::new(2u16, 2i32, 2i64);
let _ = LPR::new(2i16, 2u32, 2u8);
let _ = LPR::new(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 _ = LPR::new_from_n(n);
}
}
#[test]
fn default_suitable() {
let lpr = LPR::default();
assert!(lpr.check_correctness().is_ok());
assert!(lpr.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 lpr = LPR::new_from_n(n);
assert!(lpr.check_correctness().is_ok());
assert!(lpr.check_security().is_ok());
}
}
#[test]
#[allow(clippy::needless_borrows_for_generic_args)]
fn availability() {
let _ = LPR::new_from_n(10u8);
let _ = LPR::new_from_n(10u16);
let _ = LPR::new_from_n(10u32);
let _ = LPR::new_from_n(10u64);
let _ = LPR::new_from_n(10i8);
let _ = LPR::new_from_n(10i16);
let _ = LPR::new_from_n(10i32);
let _ = LPR::new_from_n(10i64);
let _ = LPR::new_from_n(Z::from(10));
let _ = LPR::new_from_n(&Z::from(10));
}
#[test]
#[should_panic]
fn invalid_n() {
LPR::new_from_n(9);
}
}
#[cfg(test)]
mod test_lpr {
use super::LPR;
use crate::pk_encryption::PKEncryptionScheme;
use qfall_math::integer::Z;
#[test]
fn cycle_zero_small_n() {
let msg = Z::ZERO;
let lpr = LPR::default();
let (pk, sk) = lpr.key_gen();
let cipher = lpr.enc(&pk, &msg);
let m = lpr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_one_small_n() {
let msg = Z::ONE;
let lpr = LPR::default();
let (pk, sk) = lpr.key_gen();
let cipher = lpr.enc(&pk, &msg);
let m = lpr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_zero_large_n() {
let msg = Z::ZERO;
let lpr = LPR::new_from_n(50);
let (pk, sk) = lpr.key_gen();
let cipher = lpr.enc(&pk, &msg);
let m = lpr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn cycle_one_large_n() {
let msg = Z::ONE;
let lpr = LPR::new_from_n(50);
let (pk, sk) = lpr.key_gen();
let cipher = lpr.enc(&pk, &msg);
let m = lpr.dec(&sk, &cipher);
assert_eq!(msg, m);
}
#[test]
fn modulus_application() {
let messages = [2, 3, i64::MAX, i64::MIN];
let dr = LPR::default();
let (pk, sk) = dr.key_gen();
for msg in messages {
let msg_mod = Z::from(msg.rem_euclid(2));
let cipher = dr.enc(&pk, msg);
let m = dr.dec(&sk, &cipher);
assert_eq!(msg_mod, m);
}
}
}
#[cfg(test)]
mod test_multi_bits {
use super::{GenericMultiBitEncryption, LPR, PKEncryptionScheme};
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 = LPR::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 = LPR::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 = LPR::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);
}
}
}