use crate::core_crypto::gpu::CudaStreams;
use crate::integer::gpu::ciphertext::{CudaIntegerRadixCiphertext, CudaUnsignedRadixCiphertext};
use crate::integer::gpu::server_key::{
CudaBootstrappingKey, CudaDynamicKeyswitchingKey, CudaServerKey,
};
use crate::integer::gpu::{cuda_backend_kreyvium_init, cuda_backend_kreyvium_step};
const KREYVIUM_KEY_BITS: usize = 128;
const KREYVIUM_IV_BITS: usize = 128;
const REGISTER_A_BITS: usize = 93;
const REGISTER_B_BITS: usize = 84;
const REGISTER_C_BITS: usize = 111;
const BATCH_SIZE: usize = 64;
pub struct CudaKreyviumState {
pub a: CudaUnsignedRadixCiphertext, pub b: CudaUnsignedRadixCiphertext, pub c: CudaUnsignedRadixCiphertext, pub k: CudaUnsignedRadixCiphertext, pub iv: CudaUnsignedRadixCiphertext, pub k_offset: u32,
pub iv_offset: u32,
}
impl CudaServerKey {
pub fn kreyvium_generate_keystream(
&self,
key: &CudaUnsignedRadixCiphertext,
iv: &CudaUnsignedRadixCiphertext,
num_steps: usize,
streams: &CudaStreams,
) -> crate::Result<CudaUnsignedRadixCiphertext> {
let mut state = self.kreyvium_init(key, iv, streams)?;
self.kreyvium_next(&mut state, num_steps, streams)
}
pub fn kreyvium_init(
&self,
key: &CudaUnsignedRadixCiphertext,
iv: &CudaUnsignedRadixCiphertext,
streams: &CudaStreams,
) -> crate::Result<CudaKreyviumState> {
if key.as_ref().d_blocks.lwe_ciphertext_count().0 != KREYVIUM_KEY_BITS
|| iv.as_ref().d_blocks.lwe_ciphertext_count().0 != KREYVIUM_IV_BITS
{
return Err(format!(
"Input key must contain {KREYVIUM_KEY_BITS} and IV must contain {KREYVIUM_IV_BITS} encrypted bits."
).into());
}
let mut state = CudaKreyviumState {
a: self.create_trivial_zero_radix(REGISTER_A_BITS, streams),
b: self.create_trivial_zero_radix(REGISTER_B_BITS, streams),
c: self.create_trivial_zero_radix(REGISTER_C_BITS, streams),
k: self.create_trivial_zero_radix(KREYVIUM_KEY_BITS, streams),
iv: self.create_trivial_zero_radix(KREYVIUM_IV_BITS, streams),
k_offset: 0,
iv_offset: 0,
};
let CudaDynamicKeyswitchingKey::Standard(computing_ks_key) = &self.key_switching_key else {
panic!("Only the standard atomic pattern is supported on GPU")
};
unsafe {
match &self.bootstrapping_key {
CudaBootstrappingKey::Classic(d_bsk) => {
cuda_backend_kreyvium_init(
streams,
state.a.as_mut(),
state.b.as_mut(),
state.c.as_mut(),
state.k.as_mut(),
state.iv.as_mut(),
&mut state.k_offset,
&mut state.iv_offset,
key.as_ref(),
iv.as_ref(),
&d_bsk.d_vec,
&computing_ks_key.d_vec,
self.message_modulus,
self.carry_modulus,
d_bsk,
computing_ks_key.params_ffi(),
d_bsk.ms_noise_reduction_configuration.as_ref(),
);
}
CudaBootstrappingKey::MultiBit(d_multibit_bsk) => {
cuda_backend_kreyvium_init(
streams,
state.a.as_mut(),
state.b.as_mut(),
state.c.as_mut(),
state.k.as_mut(),
state.iv.as_mut(),
&mut state.k_offset,
&mut state.iv_offset,
key.as_ref(),
iv.as_ref(),
&d_multibit_bsk.d_vec,
&computing_ks_key.d_vec,
self.message_modulus,
self.carry_modulus,
d_multibit_bsk,
computing_ks_key.params_ffi(),
None,
);
}
}
}
Ok(state)
}
pub fn kreyvium_next(
&self,
state: &mut CudaKreyviumState,
num_steps: usize,
streams: &CudaStreams,
) -> crate::Result<CudaUnsignedRadixCiphertext> {
if !num_steps.is_multiple_of(BATCH_SIZE) {
return Err(format!("The number of steps must be a multiple of {BATCH_SIZE}.").into());
}
let mut keystream: CudaUnsignedRadixCiphertext =
self.create_trivial_zero_radix(num_steps, streams);
let CudaDynamicKeyswitchingKey::Standard(computing_ks_key) = &self.key_switching_key else {
panic!("Only the standard atomic pattern is supported on GPU")
};
unsafe {
match &self.bootstrapping_key {
CudaBootstrappingKey::Classic(d_bsk) => {
cuda_backend_kreyvium_step(
streams,
keystream.as_mut(),
state.a.as_mut(),
state.b.as_mut(),
state.c.as_mut(),
state.k.as_mut(),
state.iv.as_mut(),
&mut state.k_offset,
&mut state.iv_offset,
num_steps as u32,
&d_bsk.d_vec,
&computing_ks_key.d_vec,
self.message_modulus,
self.carry_modulus,
d_bsk,
computing_ks_key.params_ffi(),
d_bsk.ms_noise_reduction_configuration.as_ref(),
);
}
CudaBootstrappingKey::MultiBit(d_multibit_bsk) => {
cuda_backend_kreyvium_step(
streams,
keystream.as_mut(),
state.a.as_mut(),
state.b.as_mut(),
state.c.as_mut(),
state.k.as_mut(),
state.iv.as_mut(),
&mut state.k_offset,
&mut state.iv_offset,
num_steps as u32,
&d_multibit_bsk.d_vec,
&computing_ks_key.d_vec,
self.message_modulus,
self.carry_modulus,
d_multibit_bsk,
computing_ks_key.params_ffi(),
None,
);
}
}
}
Ok(keystream)
}
}