use crate::params::{self, Torus};
use super::encoder::Encoder;
use super::lookup_table::LookupTable;
#[derive(Debug, Clone)]
pub struct Generator {
encoder: Encoder,
poly_degree: usize,
lookup_table_size: usize,
}
impl Generator {
pub fn new(message_modulus: usize) -> Self {
let poly_degree = params::trgsw_lv1::N;
let lookup_table_size = poly_degree;
Self {
encoder: Encoder::new(message_modulus),
poly_degree,
lookup_table_size,
}
}
pub fn with_scale(message_modulus: usize, scale: f64) -> Self {
let poly_degree = params::trgsw_lv1::N;
let lookup_table_size = poly_degree;
Self {
encoder: Encoder::with_scale(message_modulus, scale),
poly_degree,
lookup_table_size,
}
}
pub fn generate_lookup_table<F>(&self, f: F) -> LookupTable
where
F: Fn(usize) -> usize,
{
let mut lut = LookupTable::new();
self.generate_lookup_table_assign(f, &mut lut);
lut
}
pub fn generate_lookup_table_assign<F>(&self, f: F, lut_out: &mut LookupTable)
where
F: Fn(usize) -> usize,
{
let message_modulus = self.encoder.message_modulus;
let mut lut_raw = vec![0 as Torus; self.lookup_table_size];
for x in 0..message_modulus {
let start = div_round(x * self.lookup_table_size, message_modulus);
let end = div_round((x + 1) * self.lookup_table_size, message_modulus);
let y = f(x);
let encoded_y = self.encoder.encode(y);
for xx in start..end {
lut_raw[xx] = encoded_y;
}
}
let offset = div_round(self.lookup_table_size, 2 * message_modulus);
let mut rotated = vec![0 as Torus; self.lookup_table_size];
for i in 0..self.lookup_table_size {
let src_idx = (i + offset) % self.lookup_table_size;
rotated[i] = lut_raw[src_idx];
}
for i in (self.lookup_table_size - offset)..self.lookup_table_size {
rotated[i] = rotated[i].wrapping_neg();
}
for i in 0..self.lookup_table_size {
lut_out.poly.b[i] = rotated[i];
lut_out.poly.a[i] = 0;
}
}
pub fn generate_lookup_table_full<F>(&self, f: F) -> LookupTable
where
F: Fn(usize) -> Torus,
{
let mut lut = LookupTable::new();
self.generate_lookup_table_full_assign(f, &mut lut);
lut
}
pub fn generate_lookup_table_full_assign<F>(&self, f: F, lut_out: &mut LookupTable)
where
F: Fn(usize) -> Torus,
{
let message_modulus = self.encoder.message_modulus;
let mut lut_raw = vec![0 as Torus; self.lookup_table_size];
for x in 0..message_modulus {
let start = div_round(x * self.lookup_table_size, message_modulus);
let end = div_round((x + 1) * self.lookup_table_size, message_modulus);
let y = f(x);
for i in start..end {
lut_raw[i] = y;
}
}
let offset = div_round(self.lookup_table_size, 2 * message_modulus);
let mut rotated = vec![0 as Torus; self.lookup_table_size];
for i in 0..self.lookup_table_size {
let src_idx = (i + offset) % self.lookup_table_size;
rotated[i] = lut_raw[src_idx];
}
for i in (self.lookup_table_size - offset)..self.lookup_table_size {
rotated[i] = rotated[i].wrapping_neg();
}
for i in 0..self.lookup_table_size {
lut_out.poly.b[i] = rotated[i];
lut_out.poly.a[i] = 0;
}
}
pub fn generate_lookup_table_custom<F>(
&self,
f: F,
message_modulus: usize,
scale: f64,
) -> LookupTable
where
F: Fn(usize) -> usize,
{
let mut lut = LookupTable::new();
let _old_encoder = self.encoder.clone();
let mut temp_generator = self.clone();
temp_generator.encoder = Encoder::with_scale(message_modulus, scale);
temp_generator.generate_lookup_table_assign(f, &mut lut);
lut
}
pub fn mod_switch(&self, x: Torus) -> usize {
let scaled = (x as f64) / (u32::MAX as f64) * (self.lookup_table_size as f64);
let result = scaled.round() as usize % self.lookup_table_size;
result
}
pub fn message_modulus(&self) -> usize {
self.encoder.message_modulus
}
pub fn poly_degree(&self) -> usize {
self.poly_degree
}
pub fn lookup_table_size(&self) -> usize {
self.lookup_table_size
}
}
fn div_round(a: usize, b: usize) -> usize {
(a + b / 2) / b
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generator_creation() {
let generator = Generator::new(2);
assert_eq!(generator.message_modulus(), 2);
assert_eq!(generator.poly_degree(), params::trgsw_lv1::N);
assert_eq!(generator.lookup_table_size(), params::trgsw_lv1::N);
}
#[test]
fn test_identity_function() {
let generator = Generator::new(2);
let identity = |x: usize| x;
let lut = generator.generate_lookup_table(identity);
assert!(!lut.is_empty());
}
#[test]
fn test_not_function() {
let generator = Generator::new(2);
let not_func = |x: usize| 1 - x;
let lut = generator.generate_lookup_table(not_func);
assert!(!lut.is_empty());
}
#[test]
fn test_constant_function() {
let generator = Generator::new(2);
let constant_one = |_x: usize| 1;
let lut = generator.generate_lookup_table(constant_one);
assert!(!lut.is_empty());
}
#[test]
fn test_4bit_function() {
let generator = Generator::new(4);
let increment = |x: usize| (x + 1) % 4;
let lut = generator.generate_lookup_table(increment);
assert!(!lut.is_empty());
}
#[test]
fn test_custom_scale() {
let generator = Generator::with_scale(2, 0.5);
let identity = |x: usize| x;
let lut = generator.generate_lookup_table(identity);
assert!(!lut.is_empty());
}
#[test]
fn test_mod_switch() {
let generator = Generator::new(2);
let result1 = generator.mod_switch(0);
let result2 = generator.mod_switch(u32::MAX / 2);
let result3 = generator.mod_switch(u32::MAX);
assert!(result1 < generator.lookup_table_size());
assert!(result2 < generator.lookup_table_size());
assert!(result3 < generator.lookup_table_size());
}
#[test]
fn test_div_round() {
assert_eq!(div_round(5, 2), 3); assert_eq!(div_round(4, 2), 2); assert_eq!(div_round(3, 2), 2); assert_eq!(div_round(1, 2), 1); assert_eq!(div_round(0, 2), 0); }
}