use std::ffi::c_void;
use std::fmt::Debug;
use std::ptr::null_mut;
use std::sync::atomic::{AtomicPtr, Ordering};
use crate::error::Result;
use crate::{bindgen, try_seal, Context, MemoryPool, Plaintext};
pub struct CKKSEncoder {
handle: AtomicPtr<c_void>,
parms_id: Vec<u64>,
scale: f64,
}
impl CKKSEncoder {
pub fn new(
ctx: &Context,
scale: f64,
) -> Result<Self> {
let mut handle: *mut c_void = null_mut();
let parms_id = ctx.get_first_parms_id()?;
try_seal!(unsafe { bindgen::CKKSEncoder_Create(ctx.get_handle(), &mut handle) })?;
Ok(Self {
handle: AtomicPtr::new(handle),
parms_id,
scale,
})
}
pub(crate) unsafe fn get_handle(&self) -> *mut c_void {
self.handle.load(Ordering::SeqCst)
}
pub fn get_slot_count(&self) -> usize {
let mut count: u64 = 0;
try_seal!(unsafe { bindgen::CKKSEncoder_SlotCount(self.get_handle(), &mut count) })
.expect("Internal error in BVTEncoder::get_slot_count().");
count as usize
}
pub fn encode_f64(
&self,
data: &[f64],
) -> Result<Plaintext> {
let mem = MemoryPool::new()?;
let plaintext = Plaintext::new()?;
try_seal!(unsafe {
let mut parms_id = self.parms_id.clone();
let parms_id_ptr = parms_id.as_mut_ptr();
bindgen::CKKSEncoder_Encode1(
self.get_handle(),
data.len() as u64,
data.as_ptr() as *mut f64,
parms_id_ptr,
self.scale,
plaintext.get_handle(),
mem.get_handle(),
)
})?;
Ok(plaintext)
}
pub fn decode_f64(
&self,
plaintext: &Plaintext,
) -> Result<Vec<f64>> {
let mut data = Vec::with_capacity(self.get_slot_count());
let data_ptr = data.as_mut_ptr();
let mut size: u64 = 0;
try_seal!(unsafe {
bindgen::CKKSEncoder_Decode1(
self.get_handle(),
plaintext.get_handle(),
&mut size,
data_ptr,
null_mut(),
)
})?;
if data.capacity() < size as usize {
panic!("Allocation overflow BVTEncoder::decode_unsigned");
}
unsafe {
data.set_len(size as usize);
}
Ok(data)
}
}
impl Debug for CKKSEncoder {
fn fmt(
&self,
f: &mut std::fmt::Formatter<'_>,
) -> std::fmt::Result {
f.debug_struct("CKKSEncoder")
.field("handle", &self.handle)
.field("parms_id", &self.parms_id)
.field("scale", &self.scale)
.finish()
}
}
impl Drop for CKKSEncoder {
fn drop(&mut self) {
unsafe {
bindgen::CKKSEncoder_Destroy(self.get_handle());
}
}
}
#[cfg(test)]
mod tests {
use crate::{
CKKSEncoder, CKKSEncryptionParametersBuilder, CoefficientModulusFactory, Context,
DegreeType, EncryptionParameters, Error, SecurityLevel,
};
fn float_assert_eq(
a: f64,
b: f64,
) {
assert!((a - b).abs() < 0.0001);
}
fn float_iter_assert_eq(
a: impl IntoIterator<Item = f64>,
b: impl IntoIterator<Item = f64>,
) {
for (a, b) in a.into_iter().zip(b.into_iter()) {
float_assert_eq(a, b);
}
}
fn create_ckks_context(
degree: DegreeType,
bit_sizes: &[i32],
) -> Result<Context, Error> {
let security_level = SecurityLevel::TC128;
let expand_mod_chain = false;
let modulus_chain = CoefficientModulusFactory::build(degree, bit_sizes)?;
let encryption_parameters: EncryptionParameters = CKKSEncryptionParametersBuilder::new()
.set_poly_modulus_degree(degree)
.set_coefficient_modulus(modulus_chain)
.build()?;
let ctx = Context::new(&encryption_parameters, expand_mod_chain, security_level)?;
Ok(ctx)
}
#[test]
fn can_create_and_drop_ckks_encoder() {
let ctx = create_ckks_context(DegreeType::D8192, &[60, 40, 40, 60]).unwrap();
let encoder = CKKSEncoder::new(&ctx, 2.0f64.powi(40)).unwrap();
std::mem::drop(encoder);
}
#[test]
fn can_get_slots_ckks_encoder() {
let ctx = create_ckks_context(DegreeType::D8192, &[60, 40, 40, 60]).unwrap();
let encoder = CKKSEncoder::new(&ctx, 2.0f64.powi(40)).unwrap();
assert_eq!(encoder.get_slot_count(), 8192 / 2);
}
#[test]
fn can_get_encode_and_decode_unsigned() {
let ctx = create_ckks_context(DegreeType::D8192, &[60, 40, 40, 60]).unwrap();
let encoder = CKKSEncoder::new(&ctx, 2.0f64.powi(40)).unwrap();
let mut data = Vec::with_capacity(4096);
for i in 0..encoder.get_slot_count() {
data.push(i as f64);
}
let plaintext = encoder.encode_f64(&data).unwrap();
let data_decoded: Vec<f64> = encoder.decode_f64(&plaintext).unwrap();
float_iter_assert_eq(data, data_decoded);
}
#[test]
fn can_get_encode_and_decode_signed() {
let ctx = create_ckks_context(DegreeType::D8192, &[60, 40, 40, 60]).unwrap();
let encoder = CKKSEncoder::new(&ctx, 2.0f64.powi(40)).unwrap();
let mut data = Vec::with_capacity(4096);
for i in 0..encoder.get_slot_count() {
data.push(i as f64 - 2048.0);
}
let plaintext = encoder.encode_f64(&data).unwrap();
let data_decoded: Vec<f64> = encoder.decode_f64(&plaintext).unwrap();
float_iter_assert_eq(data, data_decoded);
}
#[test]
fn scalar_encoder_can_encode_decode_signed() {
let ctx = create_ckks_context(DegreeType::D8192, &[60, 40, 40, 60]).unwrap();
let encoder = CKKSEncoder::new(&ctx, 2.0f64.powi(40)).unwrap();
let encoded = encoder.encode_f64(&[-15.5f64]).unwrap();
let decoded: Vec<f64> = encoder.decode_f64(&encoded).unwrap();
float_assert_eq(decoded[0], -15.5);
}
#[test]
fn scalar_encoder_can_encode_decode_unsigned() {
let ctx = create_ckks_context(DegreeType::D8192, &[60, 40, 40, 60]).unwrap();
let encoder = CKKSEncoder::new(&ctx, 2.0f64.powi(40)).unwrap();
let encoded = encoder.encode_f64(&[42.0f64]).unwrap();
let decoded: Vec<f64> = encoder.decode_f64(&encoded).unwrap();
float_assert_eq(decoded[0], 42.0);
}
#[test]
fn can_get_encode_and_decode_float() {
let ctx = create_ckks_context(DegreeType::D8192, &[60, 40, 40, 60]).unwrap();
let encoder = CKKSEncoder::new(&ctx, 2.0f64.powi(40)).unwrap();
let data: Vec<f64> = (0..encoder.get_slot_count())
.map(|i| (i as f64) / 2.0)
.collect();
let plaintext = encoder.encode_f64(&data).unwrap();
let decoded_data: Vec<f64> = encoder.decode_f64(&plaintext).unwrap();
float_iter_assert_eq(data, decoded_data);
}
}