#pragma once #include #include #include #include #include #include "cloudkey.hpp" #include "detwfa.hpp" #include "keyswitch.hpp" #include "params.hpp" #include "trlwe.hpp" #include "utils.hpp" #ifdef USE_KEY_BUNDLE #include "keybundle.hpp" #endif namespace TFHEpp { // Most TFHEpp parameter families use all 2N automorphisms of the target // ring. Shallow bootstrap parameters may instead use a smaller power-of-two // LWE modulus q, provided q divides 2N; this is the setting of ePrint // 2026/1730 (q = 512, N = 1024). template inline constexpr uint32_t BlindRotationModulus = [] { if constexpr (requires { P::blind_rotation_modulus; }) return P::blind_rotation_modulus; else return 2 * P::targetP::n; }(); template void EvalIdentityKeySwitch(TLWE& res, const TLWE& tlwe, const EvalKey& ek) { if constexpr (std::is_same_v || std::is_same_v) { static_assert(P::domainP::k * P::domainP::n >= P::targetP::k * P::targetP::n); SubsetIdentityKeySwitch

(res, tlwe, ek.getsubiksk

()); } else { IdentityKeySwitch

(res, tlwe, ek.getiksk

()); } } // https://eprint.iacr.org/2025/809 template void BRModSwitch(ModswitchTLWE& moded, const TLWE& tlwe) { using domainT = typename P::domainP::T; static_assert(std::numeric_limits::digits <= 64, "BRModSwitch correction requires at most a 64-bit torus"); using correctionT = std::conditional_t<(std::numeric_limits::digits <= 32), int64_t, __int128_t>; constexpr uint32_t modulus = BlindRotationModulus

; static_assert(std::has_single_bit(modulus), "blind rotation modulus must be a power of two"); static_assert(modulus <= 2 * P::targetP::n, "blind rotation modulus cannot exceed the ring automorphisms"); static_assert((2 * P::targetP::n) % modulus == 0, "blind rotation modulus must divide the ring automorphisms"); constexpr uint32_t modulusbit = std::countr_zero(modulus); constexpr uint32_t exponent_scale = 2 * P::targetP::n / modulus; constexpr uint32_t bitwidth = bits_needed(); correctionT c = 0; constexpr domainT roundoffset = 1ULL << (std::numeric_limits::digits - 1 - modulusbit + bitwidth); for (int i = 0; i < P::domainP::k * P::domainP::n; i++) { const domainT reduced = (tlwe[i] + roundoffset) >> (std::numeric_limits::digits - modulusbit + bitwidth) << bitwidth; moded[i] = reduced * exponent_scale; const domainT rounded = reduced << (std::numeric_limits::digits - modulusbit); c += static_cast>(tlwe[i] - rounded); } const domainT reduced_b = modulus - (static_cast( static_cast( tlwe[P::domainP::k * P::domainP::n]) - c / 2 + roundoffset) >> (std::numeric_limits::digits - modulusbit + bitwidth) << bitwidth); moded[P::domainP::k * P::domainP::n] = reduced_b * exponent_scale; } template void BlindRotate(TRLWE& res, const TLWE& tlwe, const BootstrappingKeyFFT

& bkfft, const Polynomial& testvector) { ModswitchTLWE moded; BRModSwitch(moded, tlwe); res = {}; PolynomialMulByXai( res[P::targetP::k], testvector, moded[P::domainP::k * P::domainP::n]); #ifdef USE_KEY_BUNDLE for (int i = 0; i < P::domainP::k * P::domainP::n / P::Addends; i++) { alignas(64) TRGSWFFT BKadded; KeyBundleFFT

(BKadded, bkfft[i], std::span(moded) .subspan(P::Addends * i, P::Addends) .template first()); ExternalProduct(res, res, BKadded); } #else if constexpr (requires { P::domainP::ell; }) { static_assert(P::domainP::k == 1, "block-binary blind rotation expects domain k = 1"); static_assert(P::domainP::n % P::domainP::ell == 0, "block-binary dimension must be divisible by ell"); static_assert(P::Addends == 1, "block-binary blind rotation does not use key bundling"); constexpr uint32_t ell = P::domainP::ell; constexpr uint32_t blocks = P::domainP::n / ell; for (uint32_t block = 0; block < blocks; block++) { const uint32_t base = block * ell; if (block + 1 < blocks) { const char* next_bk = reinterpret_cast(&bkfft[base + ell]); for (int p = 0; p < 8; p++) __builtin_prefetch(next_bk + p * 4096, 0, 1); } std::span, ell> bkfft_block( bkfft.begin() + base, ell); std::span bara( moded.begin() + base, ell); CMUXFFTwithBlockBinaryPolynomialMulByXaiMinusOne

( res, bkfft_block, bara); } } else { for (int i = 0; i < P::domainP::k * P::domainP::n; i++) { // Prefetch the next BK element (128KB ahead) to overlap memory // latency with computation. Spread prefetch hints across multiple // cache lines at the start of the next TRGSW element. if (i + 1 < P::domainP::k * P::domainP::n) { const char* next_bk = reinterpret_cast(&bkfft[i + 1]); for (int p = 0; p < 8; p++) __builtin_prefetch(next_bk + p * 4096, 0, 1); } if (moded[i] == 0) continue; CMUXwithPolynomialMulByXaiMinusOne

(res, bkfft[i], moded[i]); } } #endif } template void BlindRotate(TRLWE& res, const TLWE& tlwe, const BootstrappingKeyFFT

& bkfft, const TRLWE& testvector) { ModswitchTLWE moded; BRModSwitch(moded, tlwe); for (int k = 0; k < P::targetP::k + 1; k++) PolynomialMulByXai( res[k], testvector[k], moded[P::domainP::k * P::domainP::n]); #ifdef USE_KEY_BUNDLE alignas(64) std::array, P::domainP::k * P::domainP::n / P::Addends> BKadded; #pragma omp parallel for num_threads(4) for (int i = 0; i < P::domainP::k * P::domainP::n / P::Addends; i++) { std::array bara; bara[0] = moded[2 * i]; bara[1] = moded[2 * i + 1]; KeyBundleFFT

(BKadded[i], bkfft[i], bara); } for (int i = 0; i < P::domainP::k * P::domainP::n / P::Addends; i++) { ExternalProduct(res, res, BKadded[i]); } #else for (int i = 0; i < P::domainP::k * P::domainP::n; i++) { const uint32_t ā = moded[i]; if (ā == 0) continue; if (i + 1 < P::domainP::k * P::domainP::n) { const char* next_bk = reinterpret_cast(&bkfft[i + 1]); for (int p = 0; p < 8; p++) __builtin_prefetch(next_bk + p * 4096, 0, 1); } CMUXwithPolynomialMulByXaiMinusOne

(res, bkfft[i], ā); } #endif } template void BlindRotate(TRLWE& res, const TLWE& tlwe, const BootstrappingKeyNTT

& bkntt, const Polynomial& testvector) { ModswitchTLWE moded; BRModSwitch(moded, tlwe); res = {}; PolynomialMulByXai( res[P::targetP::k], testvector, moded[P::domainP::k * P::domainP::n]); for (int i = 0; i < P::domainP::k * P::domainP::n; i++) { const uint32_t ā = moded[i]; if (ā == 0) continue; // Do not use CMUXNTT to avoid unnecessary copy. CMUXwithPolynomialMulByXaiMinusOne(res, bkntt[i], ā); } } template void BlindRotate(TRLWE& res, const TLWE& tlwe, const BootstrappingKeyRAINTT

& bkraintt, const Polynomial& testvector) { ModswitchTLWE moded; BRModSwitch(moded, tlwe); res = {}; PolynomialMulByXai( res[P::targetP::k], testvector, moded[P::domainP::k * P::domainP::n]); for (int i = 0; i < P::domainP::k * P::domainP::n; i++) { const uint32_t ā = moded[i]; if (ā == 0) continue; // Do not use CMUXNTT to avoid unnecessary copy. CMUXwithPolynomialMulByXaiMinusOne(res, bkraintt[i], ā); } } template void BlindRotate(TRLWE& res, const TLWE& tlwe, const BootstrappingKeyFNT

& bkfnt, const Polynomial& testvector) { ModswitchTLWE moded; BRModSwitch(moded, tlwe); res = {}; PolynomialMulByXai( res[P::targetP::k], testvector, moded[P::domainP::k * P::domainP::n]); for (int i = 0; i < P::domainP::k * P::domainP::n; i++) { if (moded[i] == 0) continue; CMUXwithPolynomialMulByXaiMinusOne(res, bkfnt[i], moded[i]); } } template void BlindRotate(TRLWE& res, const TLWE& tlwe, const BootstrappingKeyFNT

& bkfnt, const TRLWE& testvector) { ModswitchTLWE moded; BRModSwitch(moded, tlwe); for (int k = 0; k < P::targetP::k + 1; k++) PolynomialMulByXai( res[k], testvector[k], moded[P::domainP::k * P::domainP::n]); for (int i = 0; i < P::domainP::k * P::domainP::n; i++) { const uint32_t ā = moded[i]; if (ā == 0) continue; CMUXwithPolynomialMulByXaiMinusOne(res, bkfnt[i], ā); } } template void GateBootstrappingTLWE2TLWE( TLWE& res, const TLWE& tlwe, const BootstrappingKeyFFT

& bkfft, const Polynomial& testvector) { alignas(64) TRLWE acc; BlindRotate

(acc, tlwe, bkfft, testvector); SampleExtractIndex(res, acc, 0); } template void GateBootstrappingTLWE2TLWE(TLWE& res, const TLWE& tlwe, const BootstrappingKeyFFT

& bkfft, const TRLWE& testvector) { alignas(64) TRLWE acc; BlindRotate

(acc, tlwe, bkfft, testvector); SampleExtractIndex(res, acc, 0); } template void GateBootstrappingTLWE2TLWE( TLWE& res, const TLWE& tlwe, const BootstrappingKeyNTT

& bkntt, const Polynomial& testvector) { alignas(64) TRLWE acc; BlindRotate

(acc, tlwe, bkntt, testvector); SampleExtractIndex(res, acc, 0); } template void GateBootstrappingTLWE2TLWE( TLWE& res, const TLWE& tlwe, const BootstrappingKeyRAINTT

& bkraintt, const Polynomial& testvector) { alignas(64) TRLWE acc; BlindRotate

(acc, tlwe, bkraintt, testvector); SampleExtractIndex(res, acc, 0); } template void GateBootstrappingTLWE2TLWE( TLWE& res, const TLWE& tlwe, const BootstrappingKeyFNT

& bkfnt, const Polynomial& testvector) { alignas(64) TRLWE acc; BlindRotate

(acc, tlwe, bkfnt, testvector); SampleExtractIndex(res, acc, 0); } template void GateBootstrappingTLWE2TLWE(TLWE& res, const TLWE& tlwe, const BootstrappingKeyFNT

& bkfnt, const TRLWE& testvector) { alignas(64) TRLWE acc; BlindRotate

(acc, tlwe, bkfnt, testvector); SampleExtractIndex(res, acc, 0); } template void GateBootstrappingManyLUT( std::array, num_out>& res, const TLWE& tlwe, const BootstrappingKeyFFT

& bkfft, const Polynomial& testvector) { alignas(64) TRLWE acc; BlindRotate(acc, tlwe, bkfft, testvector); for (int i = 0; i < num_out; i++) SampleExtractIndex(res[i], acc, i); } template void GateBootstrappingManyLUT( std::array, num_out>& res, const TLWE& tlwe, const BootstrappingKeyFFT

& bkfft, const TRLWE& testvector) { alignas(64) TRLWE acc; BlindRotate(acc, tlwe, bkfft, testvector); for (int i = 0; i < num_out; i++) SampleExtractIndex(res[i], acc, i); } template constexpr Polynomial

μpolygen() { Polynomial

poly; for (typename P::T& p : poly) p = μ; return poly; } template void GateBootstrapping(TLWE& res, const TLWE& tlwe, const EvalKey& ek) { alignas(64) TLWE tlwelvl1; GateBootstrappingTLWE2TLWE(tlwelvl1, tlwe, ek.getbkfft(), μpolygen()); EvalIdentityKeySwitch(res, tlwelvl1, ek); } template void GateBootstrapping(TLWE& res, const TLWE& tlwe, const EvalKey& ek) { alignas(64) TLWE tlwelvl0; EvalIdentityKeySwitch(tlwelvl0, tlwe, ek); GateBootstrappingTLWE2TLWE(res, tlwelvl0, ek.getbkfft(), μpolygen()); } template void GateBootstrappingNTT(TLWE& res, const TLWE& tlwe, const EvalKey& ek) { alignas(64) TLWE tlwelvl1; GateBootstrappingTLWE2TLWE(tlwelvl1, tlwe, ek.getbkntt(), μpolygen()); EvalIdentityKeySwitch(res, tlwelvl1, ek); } template void GateBootstrappingNTT(TLWE& res, const TLWE& tlwe, const EvalKey& ek) { alignas(64) TLWE tlwelvl0; EvalIdentityKeySwitch(tlwelvl0, tlwe, ek); GateBootstrappingTLWE2TLWE(res, tlwelvl0, ek.getbkntt(), μpolygen()); } } // namespace TFHEpp