tfhe 1.7.0

TFHE-rs is a fully homomorphic encryption (FHE) library that implements Zama's variant of TFHE.
Documentation
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_trivium_init, cuda_backend_trivium_step};

const TRIVIUM_KEY_BITS: usize = 80;
const TRIVIUM_IV_BITS: usize = 80;
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 CudaTriviumState {
    pub a: CudaUnsignedRadixCiphertext, // REGISTER_A_BITS
    pub b: CudaUnsignedRadixCiphertext, // REGISTER_B_BITS
    pub c: CudaUnsignedRadixCiphertext, // REGISTER_C_BITS
}

impl CudaServerKey {
    /// Generates a Trivium keystream homomorphically on the GPU in a single
    /// call. Initializes a transient state, runs the warmup, generates the
    /// keystream, then drops the state.
    pub fn trivium_generate_keystream(
        &self,
        key: &CudaUnsignedRadixCiphertext,
        iv: &CudaUnsignedRadixCiphertext,
        num_steps: usize,
        streams: &CudaStreams,
    ) -> crate::Result<CudaUnsignedRadixCiphertext> {
        let mut state = self.trivium_init(key, iv, streams)?;
        self.trivium_next(&mut state, num_steps, streams)
    }

    pub fn trivium_init(
        &self,
        key: &CudaUnsignedRadixCiphertext,
        iv: &CudaUnsignedRadixCiphertext,
        streams: &CudaStreams,
    ) -> crate::Result<CudaTriviumState> {
        if key.as_ref().d_blocks.lwe_ciphertext_count().0 != TRIVIUM_KEY_BITS
            || iv.as_ref().d_blocks.lwe_ciphertext_count().0 != TRIVIUM_IV_BITS
        {
            return Err(format!(
                "Input key must contain {TRIVIUM_KEY_BITS} and IV must contain {TRIVIUM_IV_BITS} encrypted bits."
            ).into());
        }

        let mut state = CudaTriviumState {
            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),
        };

        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_trivium_init(
                        streams,
                        state.a.as_mut(),
                        state.b.as_mut(),
                        state.c.as_mut(),
                        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_trivium_init(
                        streams,
                        state.a.as_mut(),
                        state.b.as_mut(),
                        state.c.as_mut(),
                        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 trivium_next(
        &self,
        state: &mut CudaTriviumState,
        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_trivium_step(
                        streams,
                        keystream.as_mut(),
                        state.a.as_mut(),
                        state.b.as_mut(),
                        state.c.as_mut(),
                        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_trivium_step(
                        streams,
                        keystream.as_mut(),
                        state.a.as_mut(),
                        state.b.as_mut(),
                        state.c.as_mut(),
                        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)
    }
}