use crate::params::Torus;
use crate::utils;
#[derive(Debug, Clone)]
pub struct Encoder {
pub message_modulus: usize,
pub scale: f64,
}
impl Encoder {
pub fn new(message_modulus: usize) -> Self {
let scale = 1.0 / (2.0 * message_modulus as f64);
Self {
message_modulus,
scale,
}
}
pub fn with_scale(message_modulus: usize, scale: f64) -> Self {
Self {
message_modulus,
scale,
}
}
pub fn encode(&self, message: usize) -> Torus {
let message = message % self.message_modulus;
let value = message as f64 * self.scale;
utils::f64_to_torus(value)
}
pub fn encode_with_scale(&self, message: usize, scale: f64) -> Torus {
let message = message % self.message_modulus;
let value = message as f64 * scale;
utils::f64_to_torus(value)
}
pub fn decode(&self, value: Torus) -> usize {
let f = utils::torus_to_f64(value);
let message = (f / self.scale + 0.5) as usize;
message % self.message_modulus
}
pub fn decode_bool(&self, value: Torus) -> bool {
self.decode(value) != 0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_binary_encoder() {
let encoder = Encoder::new(2);
let encoded_0 = encoder.encode(0);
let encoded_1 = encoder.encode(1);
assert_eq!(encoder.decode(encoded_0), 0);
assert_eq!(encoder.decode(encoded_1), 1);
assert_eq!(encoder.decode_bool(encoded_0), false);
assert_eq!(encoder.decode_bool(encoded_1), true);
}
#[test]
fn test_4bit_encoder() {
let encoder = Encoder::new(4);
for i in 0..4 {
let encoded = encoder.encode(i);
let decoded = encoder.decode(encoded);
assert_eq!(decoded, i);
}
}
#[test]
fn test_custom_scale() {
let encoder = Encoder::with_scale(2, 0.5);
let encoded_0 = encoder.encode(0);
let encoded_1 = encoder.encode(1);
assert_eq!(encoder.decode(encoded_0), 0);
assert_eq!(encoder.decode(encoded_1), 1);
}
}