From 0d9dbf64d75240a411afab5f1a455c2f7f4ccafe Mon Sep 17 00:00:00 2001 From: leejonghyeong Date: Fri, 18 Sep 2026 14:50:55 +0900 Subject: [PATCH] Release v0.5.0 Sync with the internal repository. Changes since the last sync (v0.4.1): * Seed-only ciphertext: drop the PUBLICKEY seed mode. Only the UNIFORM a-part can be reconstructed from a seed without the encryption key, so the public-key variant is removed from CipherSeedMode, Encryptor and the serialization format. * Serialization: add a 2 GiB size guard on both serialize and deserialize. FlatBuffers' internal size counter is a uint32 that only asserts in debug builds, so oversized buffers silently corrupted in release. * SeedGenerator: serialize Gen()/Reseed() behind a mutex, make Reseed() honor its documented empty-optional behavior, and stop evaluating Gen() eagerly through value_or() at the Encryptor/KeyGenerator call sites. * OmpUtils: make the saved thread count thread_local and add an OMP thread limit guard. * Install: fix find_package(deb) on a default (static) install by making the alea/flatbuffers find_dependency calls conditional, install the generated headers flat next to the headers that include them, and add the per-project include root to the install interface. Skip BLAKE3's broken subproject install rules. * Warnings: add set_deb_no_warnings() and apply it to the flatbuffers targets so a consumer's global warning flags don't surface warnings we don't own. * Restore the DEB_ARCH cache variable, which CMakePresets.json already referenced but CMakeLists.txt no longer declared. * .clang-tidy: fix the "ccpcoreguidelines-*" typo that disabled that check group. Co-Authored-By: Claude Opus 5 (1M context) --- .clang-tidy | 2 +- CMakeLists.txt | 63 +++++++- benchmark/CMakeLists.txt | 18 +++ cmake/debConfig.cmake.in | 18 ++- cmake/warnings.cmake | 17 ++ examples/SeedOnlyCiphertext.cpp | 28 ++-- external/CMakeLists.txt | 20 +++ include/deb/CKKSTypes.hpp | 19 ++- include/deb/Encryptor.hpp | 32 ++-- include/deb/KeyGenerator.hpp | 9 ++ include/deb/SeedGenerator.hpp | 11 +- include/deb/Serialize.hpp | 267 ++++++++++++++++++++++++++++++-- include/deb/utils/OmpUtils.hpp | 34 ++++ src/CKKSTypes.cpp | 4 + src/Decryptor.cpp | 14 +- src/Encryptor.cpp | 192 ++++++++++------------- src/KeyGenerator.cpp | 7 +- src/OmpUtils.cpp | 55 ++++++- src/SeedGenerator.cpp | 45 +++--- src/Serialize.cpp | 254 ++++++++++++++++++++++++++---- test/EnDecryption-test.cpp | 105 +++++-------- test/Operation-test.cpp | 16 +- test/Serialize-test.cpp | 219 +++++++++++++++++++++++++- 23 files changed, 1155 insertions(+), 294 deletions(-) diff --git a/.clang-tidy b/.clang-tidy index 5f8379e..b9bc5a4 100644 --- a/.clang-tidy +++ b/.clang-tidy @@ -6,7 +6,7 @@ Checks: '-*, -bugprone-implicit-widening-of-multiplication-result, performance-*, -performance-unnecessary-value-param, - ccpcoreguidelines-*, + cppcoreguidelines-*, misc-static-assert, misc-throw-by-value-catch-by-reference, misc-unconventional-assign-operator, diff --git a/CMakeLists.txt b/CMakeLists.txt index ab1edad..e19189e 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -17,7 +17,7 @@ cmake_minimum_required(VERSION 3.21) project( deb - VERSION 0.4.1 + VERSION 0.5.0 LANGUAGES CXX DESCRIPTION "CryptoLab's official cryptosystem library for FHE.") @@ -73,6 +73,12 @@ option(DEB_RUNTIME_RESOURCE_CHECK "Enable runtime resource check." ON) option(DEB_SERIALIZE_API "Enable serialize API." ON) option(DEB_SUPPORT_U64 "Compile u64 coefficient word type support." ON) option(DEB_SUPPORT_U32 "Compile u32 coefficient word type support." OFF) +set(DEB_ARCH + "" + CACHE + STRING + "Target CPU microarchitecture passed to -march (e.g. x86-64-v3, native). Empty keeps the portable default. x86-64-v3 (AVX2+FMA+BMI2) is recommended for modern x86 and speeds up the hot NTT/modular-arithmetic kernels without the AVX-512 downclocking that -march=native can cause." +) if(NOT DEB_SUPPORT_U64 AND NOT DEB_SUPPORT_U32) message( @@ -101,6 +107,26 @@ if(BUILD_SHARED_LIBS) FORCE) endif() +# The serialize API puts flatbuffers into the PUBLIC include graph: the +# installed Serialize.hpp includes DebFBType.h, which includes +# . A consumer cannot compile against an install that +# omits those headers, so this is not optional whenever we install at all. +if(DEB_INSTALL AND DEB_SERIALIZE_API) + set(DEB_INSTALL_FLATBUFFERS + ON + CACHE BOOL + "Enable installation of flatbuffers required by the serialize API." + FORCE) +endif() + +# When tuning for a specific architecture, also tune the C dependencies (most +# importantly the alea RNG's Keccak permutation, which is a measurable part of +# encryption). Done before add_subdirectory(external) so those targets inherit +# it; the deb C++ library itself gets -march via target_compile_options below. +if(DEB_ARCH AND NOT MSVC) + string(APPEND CMAKE_C_FLAGS " -march=${DEB_ARCH}") +endif() + add_subdirectory(external) add_subdirectory(prebuild) @@ -150,6 +176,15 @@ if(NOT MSVC) target_compile_options(${PROJECT_NAME} PRIVATE -Wno-pedantic) endif() +# Optional target-architecture tuning. Applied PUBLIC so the hot kernels in the +# deb library itself and any inlined template code in dependents (benchmark, +# examples) are all compiled for the same ISA. Empty by default to preserve a +# portable baseline build. +if(DEB_ARCH AND NOT MSVC) + target_compile_options(${PROJECT_NAME} PUBLIC -march=${DEB_ARCH}) + message(STATUS "deb: tuning for -march=${DEB_ARCH}") +endif() + string(TOUPPER "${DEB_EXT_LIB_FOR_SECURE_ZERO}" _deb_secure_zero_backend) if(_deb_secure_zero_backend STREQUAL "LIBSODIUM") target_compile_definitions(${PROJECT_NAME} PUBLIC DEB_SECURE_ZERO_LIBSODIUM) @@ -162,11 +197,18 @@ elseif(_deb_secure_zero_backend STREQUAL "NATIVE") endif() unset(_deb_secure_zero_backend) +# The install interface must mirror the build interface. The public headers +# include each other unqualified relative to include/${PROJECT_NAME} (e.g. +# utils/RandomGenerator.hpp does #include "Types.hpp", and Preset.hpp does +# #include "DebParam.hpp"), so that directory has to be an include root for +# consumers too. The plain include dir is kept as a root as well so consumers +# can write the qualified #include <${PROJECT_NAME}/Encryptor.hpp>. target_include_directories( ${PROJECT_NAME} PUBLIC $ $ - $) + $ + $) if(DEB_RUNTIME_RESOURCE_CHECK) target_compile_definitions(${PROJECT_NAME} PUBLIC DEB_RESOURCE_CHECK) @@ -266,9 +308,20 @@ if(DEB_INSTALL) INCLUDES DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}) - # Install header files - install(DIRECTORY ${CMAKE_CURRENT_SOURCE_DIR}/include/ - DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}) + # Install header files. The generated headers are excluded here and installed + # separately below, because they must not keep their extra "generated/" level + # in the installed tree. + install( + DIRECTORY ${CMAKE_CURRENT_SOURCE_DIR}/include/ + DESTINATION ${CMAKE_INSTALL_INCLUDEDIR} + PATTERN "generated" EXCLUDE) + + # ${PRE_BUILD_DIR} is a second include root at build time, so headers such as + # Preset.hpp and Serialize.hpp include the generated ones unqualified + # (#include "DebParam.hpp"). The install tree exposes a single include root, + # so put the generated headers flat next to the headers that include them. + install(DIRECTORY ${PRE_BUILD_DIR}/ + DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}/${PROJECT_NAME}) # Export target for use with find_package install( diff --git a/benchmark/CMakeLists.txt b/benchmark/CMakeLists.txt index 32baafa..9c74572 100644 --- a/benchmark/CMakeLists.txt +++ b/benchmark/CMakeLists.txt @@ -26,6 +26,15 @@ add_executable(deb-bench-ct-vs-rt bench_ct_vs_rt.cpp) target_link_libraries(deb-bench-ct-vs-rt PRIVATE deb) enable_language(C) +# BLAKE3 is used only by the benchmark below, never by the library, so none of +# it belongs in the install tree. Its install rules are also broken for a +# subproject build: upstream writes libblake3.pc into its own binary dir but +# installs it from ${CMAKE_BINARY_DIR} (the top-level one), so `cmake --install` +# fails on a file that was never there. There is no BLAKE3_INSTALL option to +# turn this off, and EXCLUDE_FROM_ALL does not suppress install() rules, so the +# rules are skipped for the duration of the subdirectory instead. +set(_deb_skip_install_backup ${CMAKE_SKIP_INSTALL_RULES}) +set(CMAKE_SKIP_INSTALL_RULES TRUE) cpmaddpackage( NAME blake3 @@ -35,6 +44,15 @@ cpmaddpackage( 1.8.1 SOURCE_SUBDIR c) +set(CMAKE_SKIP_INSTALL_RULES ${_deb_skip_install_backup}) +unset(_deb_skip_install_backup) +# CMAKE_SKIP_INSTALL_RULES suppresses the generated cmake_install.cmake for that +# directory, but this directory's own script still include()s it, so leave an +# empty one in its place. +file( + WRITE ${blake3_BINARY_DIR}/cmake_install.cmake + "# BLAKE3 is a benchmark-only dependency and is intentionally not installed.\n" +) add_executable(deb-benchmark-blake3 benchmark_blake3.cpp) target_link_libraries(deb-benchmark-blake3 diff --git a/cmake/debConfig.cmake.in b/cmake/debConfig.cmake.in index 2430866..c514a2b 100644 --- a/cmake/debConfig.cmake.in +++ b/cmake/debConfig.cmake.in @@ -18,8 +18,22 @@ include(CMakeFindDependencyMacro) -find_dependency(alea REQUIRED) -find_dependency(flatbuffers REQUIRED) +# alea and flatbuffers are only dependencies a consumer has to find when they +# were installed alongside deb, which is the shared-library build. A static deb +# links them through BUILD_INTERFACE and bakes their objects into libdeb.a, so +# requiring them here would make find_package(deb) fail on every default +# (static) install, where neither package is present. +set(DEB_INSTALL_ALEA @DEB_INSTALL_ALEA@) +if(DEB_INSTALL_ALEA) + find_dependency(alea REQUIRED) +endif() +unset(DEB_INSTALL_ALEA) + +set(DEB_INSTALL_FLATBUFFERS @DEB_INSTALL_FLATBUFFERS@) +if(DEB_INSTALL_FLATBUFFERS) + find_dependency(flatbuffers REQUIRED) +endif() +unset(DEB_INSTALL_FLATBUFFERS) set(DEB_BUILD_WITH_OMP @DEB_BUILD_WITH_OMP@) if(DEB_BUILD_WITH_OMP) diff --git a/cmake/warnings.cmake b/cmake/warnings.cmake index 88d7047..3cee5f4 100644 --- a/cmake/warnings.cmake +++ b/cmake/warnings.cmake @@ -31,3 +31,20 @@ function(set_deb_warnings target) $<$: /W4>) endfunction() + +# Silence all warnings for a third-party target we pull in via CPM. A consumer +# project that includes deb with global warning flags (e.g. +# add_compile_options(-Wall ...) or CMAKE_CXX_FLAGS) leaks those flags into +# every add_subdirectory(), including our dependencies. A target-level -w / /w +# is appended after those global flags and disables the warnings we neither own +# nor can fix, so the dependency stays quiet regardless of who builds deb. +function(set_deb_no_warnings target) + if(NOT TARGET ${target}) + return() + endif() + target_compile_options( + ${target} + PRIVATE + $<$,$,$>:-w> + $<$:/w>) +endfunction() diff --git a/examples/SeedOnlyCiphertext.cpp b/examples/SeedOnlyCiphertext.cpp index 782a66c..e0314b4 100644 --- a/examples/SeedOnlyCiphertext.cpp +++ b/examples/SeedOnlyCiphertext.cpp @@ -114,33 +114,35 @@ int main() { } // --------------------------------------------------------------------- - // 3. Public-key seed-only encryption (needs the enc key to expand 'a') + // 3. Seed compression is rejected for public-key encryption // --------------------------------------------------------------------- + // Under a public key the 'a' part is a = v*ax + e_a, derived from the same + // stream as the ephemeral encryption randomness v. Storing that seed in the + // ciphertext would let anyone holding the ciphertext and the (public) + // encryption key replay v and recover the plaintext without the secret key, + // so the library refuses the combination. Seed compression is only sound + // for secret-key encryption, where 'a' is public uniform randomness. { KeyGenerator keygen(preset); SwitchKey ek = keygen.genEncKey(sk); - Ciphertext ctxt(preset); - enc.encrypt(msg, ek, ctxt, EncryptOptions().SeedOnlyA(true)); std::cout << "\n[Public-key seed-only]" << std::endl; - std::cout << " isAxFlushed=" << ctxt.isAxFlushed() << std::endl; - - // a = v*ax + e cannot be regenerated without the encryption key, so the - // decryptor refuses a still-compressed public-key ciphertext. try { - Message tmp(preset); - dec.decrypt(ctxt, sk, tmp); - std::cout << " unexpected: decrypt succeeded while compressed" + Ciphertext rejected(preset); + enc.encrypt(msg, ek, rejected, EncryptOptions().SeedOnlyA(true)); + std::cout << " unexpected: encrypt accepted seed-only 'a'" << std::endl; } catch (const std::exception &e) { - std::cout << " decrypt refused (expected): " << e.what() + std::cout << " encrypt refused (expected): " << e.what() << std::endl; } - enc.completeCiphertext(ctxt, ek); // supply the encryption key + // Public-key encryption without seed compression works as usual. + Ciphertext ctxt(preset); + enc.encrypt(msg, ek, ctxt); Message dec_msg(preset); dec.decrypt(ctxt, sk, dec_msg); - std::cout << " log2 error (after completeCiphertext) = " + std::cout << " plain public-key encrypt log2 error = " << compareMessage(msg, dec_msg) << std::endl; } diff --git a/external/CMakeLists.txt b/external/CMakeLists.txt index a8ff811..d129c8f 100644 --- a/external/CMakeLists.txt +++ b/external/CMakeLists.txt @@ -15,6 +15,7 @@ # ~~~ include(CPM) +include(warnings) cpmaddpackage( NAME @@ -45,6 +46,25 @@ if(DEB_SERIALIZE_API) set(flatbuffers_SOURCE_DIR ${flatbuffers_SOURCE_DIR} CACHE PATH "" FORCE) + + # flatbuffers is built from source via CPM. Keep its (and flatc's) compilation + # quiet so a consumer project's global warning flags don't surface warnings we + # neither own nor can fix. Also mark its headers as SYSTEM so the warnings + # don't leak through the generated headers we include either. + foreach(_fb_target flatbuffers flatbuffers_shared flatc flatlib) + set_deb_no_warnings(${_fb_target}) + endforeach() + if(TARGET flatbuffers) + get_target_property(_fb_inc flatbuffers INTERFACE_INCLUDE_DIRECTORIES) + if(_fb_inc) + # Keep INTERFACE_INCLUDE_DIRECTORIES intact (it is what actually adds the + # headers to the compile line); INTERFACE_SYSTEM_INCLUDE_DIRECTORIES only + # flags which of those already-listed dirs are treated as -isystem. + set_target_properties( + flatbuffers PROPERTIES INTERFACE_SYSTEM_INCLUDE_DIRECTORIES + "${_fb_inc}") + endif() + endif() endif() string(TOUPPER "${DEB_EXT_LIB_FOR_SECURE_ZERO}" _deb_secure_zero_backend) diff --git a/include/deb/CKKSTypes.hpp b/include/deb/CKKSTypes.hpp index 11b70a9..f5dd2df 100644 --- a/include/deb/CKKSTypes.hpp +++ b/include/deb/CKKSTypes.hpp @@ -350,8 +350,12 @@ enum class CipherSeedMode : u8 { NONE = 0, /**< Not seed-compressed; @c a is stored in full. */ UNIFORM = 1, /**< Secret-key origin: @c a is a uniform sample of the seed. */ - PUBLICKEY = 2, /**< Public-key origin: @c a = v*ax + e; expanding requires - the encryption key. */ + // Value 2 was PUBLICKEY (public-key origin, a = v*ax + e_a). It was removed + // because the stored seed also reproduces the ephemeral encryption + // randomness v, so publishing it reveals the plaintext to anyone holding + // the ciphertext and the public encryption key. Seed compression is only + // sound when 'a' is uniform. Buffers carrying it are rejected on + // deserialization; do not reuse the value. }; /** @@ -568,6 +572,17 @@ template class SwitchKeyT { void setRotIdx(Size rot_idx) noexcept; Size rotIdx() const noexcept; Size dnum() const noexcept; + /** + * @brief Sets the decomposition count, which is also the number of @c ax + * polynomials the key holds. + * + * Every key kind keeps the invariant @c axSize()==dnum() and + * @c bxSize()==dnum()*num_secret. A self mod-pack key is sized by its + * @c pad_rank rather than the preset's gadget rank, so its generator uses + * this to record that rank; serialization relies on it to validate the + * key's shape. + */ + void setDnum(Size dnum) noexcept; void addAx(const Size num_polyunit, std::optional size = std::nullopt, diff --git a/include/deb/Encryptor.hpp b/include/deb/Encryptor.hpp index c19855c..3253fde 100644 --- a/include/deb/Encryptor.hpp +++ b/include/deb/Encryptor.hpp @@ -44,7 +44,13 @@ struct EncryptOptions { release its storage after encryption. The regenerated @c a follows the ciphertext's domain (NTT, or coefficient when - ntt_out==false). Requires rank==1. */ + ntt_out==false). Requires rank==1. + SECRET-KEY ENCRYPTION ONLY: encrypting under + a public key with this option throws, because + there @c a depends on the ephemeral + encryption randomness @c v and the stored + seed would reveal the plaintext to any holder + of the ciphertext and the public key. */ std::optional a_seed = std::nullopt; /**< Optional fixed seed for the @c a part. When unset and @ref seed_only_a is enabled, a fresh seed is drawn. @@ -129,6 +135,15 @@ struct EncryptOptions { /** * @brief Provides CKKS encoding and encryption routines. + * + * @note Not thread-safe. The encode/encrypt methods are @c const, but they + * mutate the RNG streams and the reusable scratch buffers this object owns, + * so concurrent calls on the SAME instance are a data race. Use one instance + * per thread. Separate instances share no mutable state of their own, but + * constructing without an explicit seed, and encrypting with @ref + * EncryptOptions::seed_only_a but no @ref EncryptOptions::a_seed, both draw + * from the process-wide SeedGenerator singleton, which is itself + * unsynchronized; pass explicit seeds when such calls can overlap. */ template class EncryptorT : public PresetTraits { @@ -246,16 +261,6 @@ class EncryptorT : public PresetTraits { */ void completeCiphertext(CiphertextT &ctxt) const; - /** - * @brief Regenerates the @c a part of a PUBLICKEY (public-key) seed-only - * ciphertext, where @c a = v*ax + e. The same encryption key used at - * encryption time must be supplied. - * @param ctxt Seed-only ciphertext to complete in place. - * @param enckey Encryption (switching) key used to produce @p ctxt. - */ - void completeCiphertext(CiphertextT &ctxt, - const SwitchKeyT &enckey) const; - private: /** * @brief Samples a zero-one polynomial. @@ -302,9 +307,8 @@ class EncryptorT : public PresetTraits { * * Mirrors @ref completeSecretKey: reseeds from the ciphertext's stored seed and * refills the released @c a polynomial (no key required). Throws if the - * ciphertext has no seed, or if it is a PUBLICKEY seed-only ciphertext (use - * @ref EncryptorT::completeCiphertext with the encryption key for that case). A - * no-op when @c a is already present. + * ciphertext has no seed or is not @ref CipherSeedMode::UNIFORM. A no-op when + * @c a is already present. * * @param ctxt Seed-only ciphertext to complete in place. */ diff --git a/include/deb/KeyGenerator.hpp b/include/deb/KeyGenerator.hpp index c8a4d74..67ce193 100644 --- a/include/deb/KeyGenerator.hpp +++ b/include/deb/KeyGenerator.hpp @@ -30,6 +30,15 @@ namespace deb { /** * @brief Generates an encryption key and switching keys for CKKS presets. + * + * @note Not thread-safe. The genXxxKey methods are @c const, but they all + * draw from the RandomGenerator this object owns and that draw is not + * synchronized (the default ALEA backend states that its API "is not + * guaranteed to be thread-safe"), so concurrent calls on the SAME instance + * are a data race. Use one instance per thread. Separate instances share no + * mutable state once constructed; construction without an explicit seed or + * RNG draws from the process-wide SeedGenerator singleton, which is itself + * unsynchronized, so build the per-thread instances before the threads start. */ template class KeyGeneratorT : public PresetTraits { diff --git a/include/deb/SeedGenerator.hpp b/include/deb/SeedGenerator.hpp index 52fb153..b41a033 100644 --- a/include/deb/SeedGenerator.hpp +++ b/include/deb/SeedGenerator.hpp @@ -26,6 +26,9 @@ namespace deb { /** * @brief Singleton wrapper over RandomGenerator to provide deterministic RNG * streams. + * + * All static entry points are thread-safe; the underlying RNG state is + * serialized internally. */ class SeedGenerator { public: @@ -34,8 +37,10 @@ class SeedGenerator { SeedGenerator &operator=(const SeedGenerator &) = delete; /** - * @brief Accesses the singleton, optionally reseeding it. - * @param seed Optional deterministic seed. + * @brief Accesses the singleton, creating it on first use. + * @param seed Optional deterministic seed. It is only honored by the call + * that constructs the singleton; later calls ignore it. Use Reseed() to + * change the seed of an existing instance. * @return Reference to the singleton instance. */ static SeedGenerator & @@ -46,7 +51,7 @@ class SeedGenerator { * @param seed Optional deterministic seed; when empty a random seed is * chosen. */ - static void Reseed(const std::optional &seed); + static void Reseed(const std::optional &seed = std::nullopt); /** * @brief Generates a new random seed suitable for deterministic APIs. * @return Fresh RNG seed. diff --git a/include/deb/Serialize.hpp b/include/deb/Serialize.hpp index 593f472..2d098dc 100644 --- a/include/deb/Serialize.hpp +++ b/include/deb/Serialize.hpp @@ -20,14 +20,186 @@ #include "DebFBType.h" #include +#include namespace deb { +/** + * @brief Upper bound on the length prefix accepted by @ref + * deserializeFromStream. + * + * The prefix is read from an untrusted stream and used directly to size the + * read buffer, so it is bounded to keep a malformed or hostile header from + * requesting an arbitrary allocation. The bound is exclusive: this value names + * the smallest length that is rejected. 2 GiB is far above any real serialized + * key or ciphertext, and @ref Size is 32-bit, so this is half the representable + * range. + */ +constexpr Size DEB_MAX_SERIALIZED_SIZE = Size{1} << 31; + +// FlatBuffers addresses a buffer with 32-bit signed offsets, so it cannot +// represent one at or above 2 GiB either. Keeping the two limits identical is +// what lets serializeToStream reject, up front, exactly the objects +// deserializeFromStream would refuse. +static_assert(DEB_MAX_SERIALIZED_SIZE - 1 <= + static_cast(FLATBUFFERS_MAX_BUFFER_SIZE), + "DEB_MAX_SERIALIZED_SIZE must not exceed what FlatBuffers can " + "represent"); + +namespace detail { + +// Upper bounds on FlatBuffers' structural overhead. A PolyUnit table costs a +// vtable, a body (soffset, prime, degree, ntt_info, array offset), the array's +// length word and worst-case alignment padding -- measured at ~36 bytes, so +// these constants only ever over-estimate. Exactness is not the goal and would +// be fragile across a FlatBuffers bump; never under-counting is. +constexpr u64 FB_POLYUNIT_OVERHEAD = 64; +constexpr u64 FB_POLY_OVERHEAD = 48; +constexpr u64 FB_TABLE_OVERHEAD = 64; +/// Deb table, union type/value vectors, root offset, and the 4-byte length +/// prefix serializeToStream writes ahead of the buffer. +constexpr u64 FB_ENVELOPE_OVERHEAD = 128; +/// One slot in a vector of offsets. +constexpr u64 FB_OFFSET = 4; +/// A `[uint64]` seed vector: length word, payload, worst-case padding. +constexpr u64 FB_SEED_VECTOR = 16 + 8 * u64{DEB_U64_SEED_SIZE}; +constexpr u64 FB_EMPTY_VECTOR = 8; + +inline u64 boundPolyUnit(const PolyUnit &unit) { + // degree() is 0 for a released unit, which is exactly what + // serializePolyUnit writes, so emptiness needs no special case. + return u64{8} * unit.degree() + FB_POLYUNIT_OVERHEAD; +} + +inline u64 boundPoly(const Polynomial &poly) { + u64 bytes = FB_POLY_OVERHEAD; + // Units within one Polynomial can have different degrees (a sliced or + // partially-copied polynomial does), so sum them rather than multiplying by + // any preset-derived limb count. + for (Size i = 0; i < poly.size(); ++i) { + bytes += boundPolyUnit(poly[i]) + FB_OFFSET; + } + return bytes; +} + +} // namespace detail + +/** + * @brief Upper bound, in bytes, on what @ref serializeToStream writes for + * @p data, including the length prefix. + * + * Computed from the object's own shape in 64-bit arithmetic, so it stays exact + * for sizes that a serialized buffer could never represent. A seed-only + * ciphertext needs no special case: @ref CiphertextT::flushAx leaves the + * released @c a part as a zero-size polynomial, which the sum below skips. + */ +inline u64 serializedSizeUpperBound(const Ciphertext &cipher) { + u64 bytes = detail::FB_TABLE_OVERHEAD + detail::FB_ENVELOPE_OVERHEAD; + for (Size i = 0; i < cipher.numPoly(); ++i) { + bytes += detail::boundPoly(cipher[i]) + detail::FB_OFFSET; + } + bytes += + cipher.hasSeed() ? detail::FB_SEED_VECTOR : detail::FB_EMPTY_VECTOR; + return bytes; +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const SwitchKey &swk) { + u64 bytes = detail::FB_TABLE_OVERHEAD + detail::FB_ENVELOPE_OVERHEAD; + for (Size i = 0; i < swk.axSize(); ++i) { + bytes += detail::boundPoly(swk.ax(i)) + detail::FB_OFFSET; + } + // bxSize() is not always axSize(): addBx() can append dnum*num_secret + // polynomials in a single call. + for (Size i = 0; i < swk.bxSize(); ++i) { + bytes += detail::boundPoly(swk.bx(i)) + detail::FB_OFFSET; + } + return bytes; +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const SecretKey &sk) { + u64 bytes = detail::FB_TABLE_OVERHEAD + detail::FB_ENVELOPE_OVERHEAD; + // The coefficients and the embedded polynomials are independently optional. + bytes += u64{sk.coeffsSize()} + detail::FB_EMPTY_VECTOR; + for (Size i = 0; i < sk.numPoly(); ++i) { + bytes += detail::boundPoly(sk[i]) + detail::FB_OFFSET; + } + bytes += sk.hasSeed() ? detail::FB_SEED_VECTOR : detail::FB_EMPTY_VECTOR; + return bytes; +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const Polynomial &poly) { + return detail::FB_ENVELOPE_OVERHEAD + detail::boundPoly(poly); +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const PolyUnit &unit) { + return detail::FB_ENVELOPE_OVERHEAD + detail::boundPolyUnit(unit); +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const Message &msg) { + return detail::FB_ENVELOPE_OVERHEAD + detail::FB_TABLE_OVERHEAD + + u64{2 * sizeof(Real)} * msg.size(); +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const FMessage &msg) { + return detail::FB_ENVELOPE_OVERHEAD + detail::FB_TABLE_OVERHEAD + + u64{2 * sizeof(float)} * msg.size(); +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const CoeffMessage &coeff) { + return detail::FB_ENVELOPE_OVERHEAD + detail::FB_TABLE_OVERHEAD + + u64{sizeof(Real)} * coeff.size(); +} + +/** @copydoc serializedSizeUpperBound(const Ciphertext &) */ +inline u64 serializedSizeUpperBound(const FCoeffMessage &coeff) { + return detail::FB_ENVELOPE_OVERHEAD + detail::FB_TABLE_OVERHEAD + + u64{sizeof(float)} * coeff.size(); +} + /** * @brief Convenience alias for FlatBuffers vector types. */ template using Vector = flatbuffers::Vector; +/** + * @brief The union tag a serialized @p T carries, or @c DebUnion_NONE for a + * type that is not serializable. + * + * Used to reject type confusion: @c GetAs() is an unchecked reinterpret, so + * without comparing the stored tag a buffer holding one type and read back as + * another turns scalar payload bytes into table and vector offsets. + */ +template constexpr deb_fb::DebUnion debUnionTag() { + if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Swk; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Sk; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Cipher; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Poly; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_PolyUnit; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Message; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Message32; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Coeff; + } else if constexpr (std::is_same_v) { + return deb_fb::DebUnion_Coeff32; + } else { + return deb_fb::DebUnion_NONE; + } +} + /** * @brief Converts a double-precision complex into FlatBuffers format. * @param data Pointer to double-precision complex values. @@ -149,7 +321,8 @@ serializePolyUnit(flatbuffers::FlatBufferBuilder &builder, * @param polyunit FlatBuffers poly unit object. * @return PolyUnit populated from the serialized data. */ -PolyUnit deserializePolyUnit(const deb_fb::PolyUnit *polyunit); +PolyUnit deserializePolyUnit(const deb_fb::PolyUnit *polyunit, + std::optional preset = std::nullopt); /** * @brief Serializes a polynomial object. @@ -278,10 +451,32 @@ flatbuffers::Offset toDeb(flatbuffers::FlatBufferBuilder &builder, * @tparam T Supported object type (Ciphertext, SecretKey, etc.). * @param data Object to serialize. * @param os Output stream receiving the bytes. - * @throws std::runtime_error If the object type is unsupported - * or if serialization fails (e.g., output stream errors). + * @throws std::runtime_error If the object type is unsupported, if the object + * is too large to fit in one buffer (see @ref serializedSizeUpperBound and + * @ref DEB_MAX_SERIALIZED_SIZE), or if writing to @p os fails. */ template void serializeToStream(const T &data, std::ostream &os) { + // Reject an oversized object BEFORE building it. FlatBuffers' only size + // guard is a FLATBUFFERS_ASSERT, i.e. plain assert(), which this library's + // release builds compile out; past 4 GiB its internal 32-bit size counter + // simply wraps. builder.GetSize() is that same uint32, so a check after + // Finish() would be reading a number modulo 2^32 -- and a wrapped value can + // land back inside the accepted range, turning a loud failure into a + // silently truncated buffer. The bound below is the only reliable check, + // and it also costs nothing: an object too large to represent is rejected + // without allocating it. + const u64 size_bound = serializedSizeUpperBound(data); + if (size_bound >= static_cast(DEB_MAX_SERIALIZED_SIZE)) { + throw std::runtime_error( + "[serializeToStream] Object is too large to serialize: it needs " + "about " + + std::to_string(size_bound) + + " bytes, and a single buffer cannot reach " + + std::to_string(static_cast(DEB_MAX_SERIALIZED_SIZE)) + + " bytes. Split it, or use a smaller parameter (for a self mod-pack " + "key, a smaller pad_rank)."); + } + flatbuffers::FlatBufferBuilder builder; if constexpr (std::is_same_v) { builder.Finish(toDeb(builder, serializeSwk(builder, data))); @@ -306,9 +501,23 @@ template void serializeToStream(const T &data, std::ostream &os) { "[serializeToStream] Unsupported type for serialization"); } Size size = builder.GetSize(); + // Second line of defence, using the reader's exact predicate so the two + // cannot disagree. This is only sound because the bound above already ruled + // out a wrapped size; on its own it would be meaningless. + if (size == 0 || size >= DEB_MAX_SERIALIZED_SIZE) { + throw std::runtime_error( + "[serializeToStream] Serialized buffer has an unusable size"); + } os.write(reinterpret_cast(&size), sizeof(Size)); os.write(reinterpret_cast(builder.GetBufferPointer()), builder.GetSize()); + // A failed first write makes the second a silent no-op, and failbit/badbit + // are sticky, so one check covers both. Note this reports a write error, + // not durability: a buffered stream may only fail later, at flush or close. + if (!os) { + throw std::runtime_error( + "[serializeToStream] Failed to write to the output stream"); + } } /** @@ -322,20 +531,45 @@ template void serializeToStream(const T &data, std::ostream &os) { template void deserializeFromStream(std::istream &is, T &data, std::optional preset = std::nullopt) { - Size size; - is.read(reinterpret_cast(&size), sizeof(Size)); - deb_assert(size > 0, - "[deserializeFromStream] Invalid size for deserialization"); + // Validation of an untrusted buffer is a security boundary, not a + // "resource check": these checks throw unconditionally rather than through + // deb_assert, which compiles to nothing when DEB_RUNTIME_RESOURCE_CHECK is + // off and would leave the parser reading unverified bytes. + Size size = 0; + if (!is.read(reinterpret_cast(&size), sizeof(Size))) { + throw std::runtime_error( + "[deserializeFromStream] Could not read the buffer length"); + } + if (size == 0 || size >= DEB_MAX_SERIALIZED_SIZE) { + throw std::runtime_error( + "[deserializeFromStream] Invalid size for deserialization"); + } std::vector buffer(size); - is.read(buffer.data(), size); + if (!is.read(buffer.data(), static_cast(size))) { + throw std::runtime_error( + "[deserializeFromStream] Truncated serialized buffer"); + } flatbuffers::Verifier verifier( reinterpret_cast(buffer.data()), buffer.size()); - deb_assert(deb_fb::VerifyDebBuffer(verifier), - "[deserializeFromStream] Invalid buffer for deserialization"); + if (!deb_fb::VerifyDebBuffer(verifier)) { + throw std::runtime_error( + "[deserializeFromStream] Invalid buffer for deserialization"); + } const auto *deb = deb_fb::GetDeb(buffer.data()); - deb_assert(deb->list()->size() == 1, - "[deserializeFromStream] Invalid Deb buffer: expected exactly " - "one element"); + if (deb == nullptr || deb->list() == nullptr || deb->list()->size() != 1) { + throw std::runtime_error( + "[deserializeFromStream] Invalid Deb buffer: expected exactly " + "one element"); + } + // GetAs() below is an unchecked reinterpret, so the stored union tag + // must be confirmed to match the requested type first. + constexpr deb_fb::DebUnion expected_tag = debUnionTag(); + if (deb->list_type() == nullptr || deb->list_type()->size() != 1 || + deb->list_type()->Get(0) != expected_tag) { + throw std::runtime_error( + "[deserializeFromStream] Serialized object is not of the " + "requested type"); + } if constexpr (std::is_same_v) { data = deserializeSwk(deb->list()->GetAs(0)); } else if constexpr (std::is_same_v) { @@ -350,7 +584,12 @@ void deserializeFromStream(std::istream &is, T &data, data = deserializePoly(preset.value(), deb->list()->GetAs(0)); } else if constexpr (std::is_same_v) { - data = deserializePolyUnit(deb->list()->GetAs(0)); + if (!preset.has_value()) { + throw std::runtime_error("[deserializeFromStream] Preset must be " + "provided for deserializing PolyUnit"); + } + data = deserializePolyUnit(deb->list()->GetAs(0), + preset); } else if constexpr (std::is_same_v) { data = deserializeMessage(deb->list()->GetAs(0)); } else if constexpr (std::is_same_v) { diff --git a/include/deb/utils/OmpUtils.hpp b/include/deb/utils/OmpUtils.hpp index 42a0b5a..01bcae3 100644 --- a/include/deb/utils/OmpUtils.hpp +++ b/include/deb/utils/OmpUtils.hpp @@ -14,6 +14,8 @@ * limitations under the License. */ +#pragma once + namespace deb::utils { /** * @brief Sets an OpenMP thread limit for the current process. @@ -25,4 +27,36 @@ void setOmpThreadLimit(int max_threads); */ void unsetOmpThreadLimit(); +/** + * @brief Scope guard applying an OpenMP thread limit for its lifetime. + * + * The limit is applied only when @p max_threads is below the number of + * threads in effect at construction. Each guard remembers the value that + * was in effect when it was constructed and restores exactly that value in + * its destructor, so nested guards unwind correctly and an exception thrown + * inside the guarded region cannot leave the limit applied. + */ +class OmpThreadLimitGuard { +public: + /** + * @brief Applies the thread limit if it is lower than the current one. + * @param max_threads Maximum number of threads; implementation-defined. + */ + explicit OmpThreadLimitGuard(int max_threads); + /** + * @brief Restores the thread count in effect at construction, if the + * guard applied a limit. Does not throw. + */ + ~OmpThreadLimitGuard(); + + OmpThreadLimitGuard(const OmpThreadLimitGuard &) = delete; + OmpThreadLimitGuard &operator=(const OmpThreadLimitGuard &) = delete; + +private: + /// Thread count to restore; meaningful only when applied_ is true. + int prev_; + /// Whether this guard actually changed the thread count. + bool applied_; +}; + } // namespace deb::utils diff --git a/src/CKKSTypes.cpp b/src/CKKSTypes.cpp index 1b8a8fc..f2f3be1 100644 --- a/src/CKKSTypes.cpp +++ b/src/CKKSTypes.cpp @@ -583,6 +583,10 @@ template Size SwitchKeyT::dnum() const noexcept { return dnum_; } +template void SwitchKeyT::setDnum(Size dnum) noexcept { + dnum_ = dnum; +} + template void SwitchKeyT::addAx(const Size num_polyunit, std::optional size, const utils::NTTType ntt_type, diff --git a/src/Decryptor.cpp b/src/Decryptor.cpp index 076c3ad..6115908 100644 --- a/src/Decryptor.cpp +++ b/src/Decryptor.cpp @@ -92,15 +92,14 @@ void DecryptorT::decryptInplace(CiphertextT &ctxt, "[Decryptor::decrypt] Level of secret key must be greater than " "or equal to ciphertext level"); - // Seed-only 'a': regenerate before use. The secret-key (UNIFORM) case needs - // no key; the public-key (PUBLICKEY) case cannot be reconstructed here (the - // encryption key is not available) — the caller must complete it first. + // Seed-only 'a': regenerate before use. Seed compression is only supported + // for secret-key encryption (UNIFORM), where 'a' is public uniform + // randomness reproducible from the stored seed without any key. if (ctxt.hasSeed() && ctxt.isAxFlushed()) { if (ctxt.seedMode() != CipherSeedMode::UNIFORM) { throw std::runtime_error( - "[Decryptor::decrypt] Cannot decrypt a public-key seed-only " - "ciphertext without the encryption key; call " - "Encryptor::completeCiphertext(ctxt, enckey) first."); + "[Decryptor::decrypt] Ciphertext has a released 'a' part but " + "an unsupported seed mode."); } const utils::NTTType ntt_type_a = (ctxt.encoding() == REAL) ? utils::NTTType::CYCLIC @@ -117,7 +116,7 @@ void DecryptorT::decryptInplace(CiphertextT &ctxt, const int max_num_threads = static_cast(ctxt[0].size() * (degree >> 10)); - utils::setOmpThreadLimit(max_num_threads); + const utils::OmpThreadLimitGuard omp_guard(max_num_threads); const bool is_real = ctxt.encoding() == REAL; const utils::NTTType ntt_type = @@ -139,7 +138,6 @@ void DecryptorT::decryptInplace(CiphertextT &ctxt, PolynomialT ptxt_tmp = innerDecrypt(ctxt_tmp, sk[i], ax); decode(ptxt_tmp, msg[i], scale, is_real); } - utils::unsetOmpThreadLimit(); } template diff --git a/src/Encryptor.cpp b/src/Encryptor.cpp index b93d850..60bc33f 100644 --- a/src/Encryptor.cpp +++ b/src/Encryptor.cpp @@ -44,7 +44,7 @@ template EncryptorT::EncryptorT(Preset target_preset, std::optional seed) : PresetTraits(target_preset), - rng_(createRandomGenerator(seed.value_or(SeedGenerator::Gen()))), + rng_(createRandomGenerator(seed ? *seed : SeedGenerator::Gen())), ptxt_buffer_(target_preset, num_p * num_secret), vx_buffer_(target_preset, true), ex_buffer_(target_preset, true), mask_(degree), samples_(buffer_size(degree)), i_samples_(degree), @@ -127,74 +127,89 @@ void EncryptorT::encrypt(const MSG *msg, const KEY &key, [[maybe_unused]] RngResetGuard rng_reset_guard{a_rng_, error_rng_}; // a_rng_ is non-null only for this call when seed-only mode is on; it - // produces the 'a' part (uniform for secret-key, or v and e_a for - // public-key) so 'a' can later be regenerated from the stored seed. + // produces the uniform 'a' part so 'a' can later be regenerated from the + // stored seed. Secret-key encryption only: see the throw below. RNGSeed a_seed_used{}; if (opt.seed_only_a) { - if (rank != 1) { - throw std::runtime_error("[Encryptor::encrypt] seed-only 'a' is " - "only supported for rank == 1"); + if constexpr (std::is_same_v>) { + // Under a public key, 'a' is not public randomness: it is + // a = v*ax + e_a, drawn from the same stream as the ephemeral + // encryption randomness v. Storing that seed in the ciphertext + // would let anyone holding the ciphertext and the (public) + // encryption key replay v and recover m = b - v*bx, i.e. the + // plaintext, without the secret key. The compression is only + // sound when 'a' is uniform, which is the secret-key case. + throw std::runtime_error( + "[Encryptor::encrypt] seed-only 'a' is not supported for " + "public-key encryption: the stored seed would reveal the " + "encryption randomness and hence the plaintext"); + } else { + if (rank != 1) { + throw std::runtime_error( + "[Encryptor::encrypt] seed-only 'a' is " + "only supported for rank == 1"); + } + a_seed_used = opt.a_seed ? *opt.a_seed : SeedGenerator::Gen(); + a_rng_ = createRandomGenerator(a_seed_used); } - a_seed_used = opt.a_seed.value_or(SeedGenerator::Gen()); - a_rng_ = createRandomGenerator(a_seed_used); } error_rng_ = opt.error_seed ? createRandomGenerator(*opt.error_seed) : nullptr; const int max_num_threads = static_cast(single_num_polyunit * (degree >> 10)); - utils::setOmpThreadLimit(max_num_threads); + { + const utils::OmpThreadLimitGuard omp_guard(max_num_threads); - PolynomialT ptxt(ptxt_buffer_, 0, num_polyunit); - for (Size i = 0; i < num_polyunit; ++i) { - ptxt[i].setPrime(primes[i % single_num_polyunit]); - } + PolynomialT ptxt(ptxt_buffer_, 0, num_polyunit); + for (Size i = 0; i < num_polyunit; ++i) { + ptxt[i].setPrime(primes[i % single_num_polyunit]); + } - if (num_secret > 1) { - for (Size i = 0; i < num_secret; ++i) { - PolynomialT ptxt_tmp(ptxt, single_num_polyunit * i, - single_num_polyunit); - encode(msg[i], ptxt_tmp, single_num_polyunit, opt); + if (num_secret > 1) { + for (Size i = 0; i < num_secret; ++i) { + PolynomialT ptxt_tmp(ptxt, single_num_polyunit * i, + single_num_polyunit); + encode(msg[i], ptxt_tmp, single_num_polyunit, opt); + } + } else { + encode(msg[0], ptxt, single_num_polyunit, opt); } - } else { - encode(msg[0], ptxt, single_num_polyunit, opt); - } - if constexpr (std::is_same_v || - std::is_same_v) { - ctxt.setEncoding(SLOT); - } else if constexpr (std::is_same_v || - std::is_same_v) { - ctxt.setEncoding(COEFF); - } else { - throw std::runtime_error( - "[Encryptor::encrypt] Unsupported message type"); - } - if (opt.real_encrypt) { - ctxt.setEncoding(REAL); - } - innerEncrypt(ptxt, key, single_num_polyunit, ctxt); - - if (!opt.ntt_out) { - // When seed-only, the last poly ('a') is about to be released, so only - // the retained 'b' parts need the inverse transform; 'a' is regenerated - // (and matched to the coefficient domain) on demand. - const Size n_transform = - opt.seed_only_a ? ctxt.numPoly() - 1 : ctxt.numPoly(); - for (u64 i = 0; i < n_transform; ++i) { - backwardNTT(modarith, ctxt[i], single_num_polyunit, - ctxt[i][0].getNTTType()); + if constexpr (std::is_same_v || + std::is_same_v) { + ctxt.setEncoding(SLOT); + } else if constexpr (std::is_same_v || + std::is_same_v) { + ctxt.setEncoding(COEFF); + } else { + throw std::runtime_error( + "[Encryptor::encrypt] Unsupported message type"); + } + if (opt.real_encrypt) { + ctxt.setEncoding(REAL); + } + innerEncrypt(ptxt, key, single_num_polyunit, ctxt); + + if (!opt.ntt_out) { + // When seed-only, the last poly ('a') is about to be + // released, so only the retained 'b' parts need the inverse + // transform; 'a' is regenerated (and matched to the + // coefficient domain) on demand. + const Size n_transform = + opt.seed_only_a ? ctxt.numPoly() - 1 : ctxt.numPoly(); + for (u64 i = 0; i < n_transform; ++i) { + backwardNTT(modarith, ctxt[i], single_num_polyunit, + ctxt[i][0].getNTTType()); + } } } - utils::unsetOmpThreadLimit(); // --- Seed-only 'a' finalize ------------------------------------------- if (opt.seed_only_a) { + // Only reachable for secret-key encryption -- the public-key case + // throws above -- so 'a' is always a uniform sample of the stored seed. ctxt.setSeed(a_seed_used); - if constexpr (std::is_same_v>) { - ctxt.setSeedMode(CipherSeedMode::UNIFORM); - } else { - ctxt.setSeedMode(CipherSeedMode::PUBLICKEY); - } + ctxt.setSeedMode(CipherSeedMode::UNIFORM); // 'a' may be in NTT or coefficient domain depending on opt.ntt_out. ctxt.flushAx(); } else { @@ -226,15 +241,14 @@ void EncryptorT::innerEncrypt(const PolynomialT &ptxt, const KEY &key, ctxt.setNTT(ntt_type, modarith[0].getNTT(ntt_type)->getRootType()); // Select RNG streams for this call. a_rng_ is non-null iff seed-only @c a - // mode is active; in that case the @c a part (uniform for secret-key, or - // v and e_a for public-key) must come solely from a_rng_ so it can be - // regenerated from the stored seed. The Gaussian error of the stored @c b - // part is drawn from error_rng_ (if a fixed error seed was given) else - // rng_. - const bool seed_only = (a_rng_ != nullptr); - RandomGenerator *a_src = seed_only ? a_rng_.get() : rng_.get(); + // mode is active, which encrypt() restricts to secret-key encryption; there + // @c a is public uniform randomness and must come solely from a_rng_ so it + // can be regenerated from the stored seed. The Gaussian error of the stored + // @c b part is drawn from error_rng_ (if a fixed error seed was given) else + // rng_. The public-key paths below never use a_rng_: their @c a depends on + // the ephemeral v, so a stored seed would expose the plaintext. + RandomGenerator *a_src = a_rng_ ? a_rng_.get() : rng_.get(); RandomGenerator *err_src = error_rng_ ? error_rng_.get() : rng_.get(); - RandomGenerator *ea_src = seed_only ? a_rng_.get() : err_src; if constexpr (std::is_same_v>) { deb_assert(key.numPoly() == num_secret * rank, @@ -315,8 +329,8 @@ void EncryptorT::innerEncrypt(const PolynomialT &ptxt, const KEY &key, } PRAGMA_OMP(omp parallel) { - sampleZO(num_polyunit, ntt_type, a_src); - sampleGaussian(num_polyunit, ntt_type, ea_src); + sampleZO(num_polyunit, ntt_type, rng_.get()); + sampleGaussian(num_polyunit, ntt_type, err_src); mulPolyConstP

(modarith, vx_buffer_, key.ax(0), ctxt[num_secret], num_polyunit); addPoly(modarith, ctxt[num_secret], ex_buffer_, @@ -351,8 +365,8 @@ void EncryptorT::innerEncrypt(const PolynomialT &ptxt, const KEY &key, } PRAGMA_OMP(omp parallel) { - sampleZO(num_polyunit, ntt_type, a_src); - sampleGaussian(num_polyunit, ntt_type, ea_src); + sampleZO(num_polyunit, ntt_type, rng_.get()); + sampleGaussian(num_polyunit, ntt_type, err_src); mulPolyConstP

(modarith, vx_buffer_, key.ax(0), ctxt[rank], num_polyunit); addPoly(modarith, ctxt[rank], ex_buffer_, ctxt[rank]); @@ -607,10 +621,10 @@ template void completeCiphertext(CiphertextT &ctxt) { throw std::runtime_error( "[completeCiphertext] Ciphertext has no stored seed."); } - if (ctxt.seedMode() == CipherSeedMode::PUBLICKEY) { + if (ctxt.seedMode() != CipherSeedMode::UNIFORM) { throw std::runtime_error( - "[completeCiphertext] Public-key seed-only ciphertext requires the " - "encryption key; use Encryptor::completeCiphertext(ctxt, enckey)."); + "[completeCiphertext] Ciphertext is not a uniform ('a' from seed) " + "seed-only ciphertext."); } if (!ctxt.isAxFlushed()) { return; // 'a' already present @@ -632,52 +646,6 @@ void EncryptorT::completeCiphertext(CiphertextT &ctxt) const { deb::completeCiphertext(ctxt); } -template -void EncryptorT::completeCiphertext(CiphertextT &ctxt, - const SwitchKeyT &enckey) const { - if (!ctxt.hasSeed()) { - throw std::runtime_error( - "[Encryptor::completeCiphertext] Ciphertext has no stored seed."); - } - if (ctxt.seedMode() == CipherSeedMode::UNIFORM) { - deb::completeCiphertext(ctxt); // secret-key case: no key needed - return; - } - if (ctxt.seedMode() != CipherSeedMode::PUBLICKEY) { - throw std::runtime_error( - "[Encryptor::completeCiphertext] Ciphertext is not seed-only."); - } - if (!ctxt.isAxFlushed()) { - return; // 'a' already present - } - const Size num_polyunit = ctxt[0].size(); - const Size a_idx = ctxt.numPoly() - 1; - const utils::NTTType ntt_type = (ctxt.encoding() == REAL) - ? utils::NTTType::CYCLIC - : utils::NTTType::NEGACYCLIC; - // Reallocate 'a' and recompute a = v*ax(0) + e_a, replaying v and e_a from - // the stored seed in the exact order encryption drew them. - ctxt[a_idx] = PolynomialT(preset, num_polyunit); - ctxt[a_idx].setNTT(ntt_type, modarith[0].getNTT(ntt_type)->getRootType()); - - auto rng = createRandomGenerator(ctxt.getSeed()); - const int max_num_threads = static_cast(num_polyunit * (degree >> 10)); - utils::setOmpThreadLimit(max_num_threads); - PRAGMA_OMP(omp parallel) { - sampleZO(num_polyunit, ntt_type, rng.get()); - sampleGaussian(num_polyunit, ntt_type, rng.get()); - mulPolyConst(modarith, vx_buffer_, enckey.ax(0), ctxt[a_idx], - num_polyunit); - addPoly(modarith, ctxt[a_idx], ex_buffer_, ctxt[a_idx]); - } - // Match the stored 'b' domain: a ntt_out==false ciphertext keeps 'b' (and - // hence 'a') in the coefficient domain. - if (!ctxt[0][0].isNTT()) { - backwardNTT(modarith, ctxt[a_idx], num_polyunit, ntt_type); - } - utils::unsetOmpThreadLimit(); -} - #ifdef DEB_U64 #define X(preset) DECL_ENCRYPT_TEMPLATE(PRESET_##preset, u64, ) PRESET_LIST_WITH_EMPTY diff --git a/src/KeyGenerator.cpp b/src/KeyGenerator.cpp index 1d7be52..75dea1b 100644 --- a/src/KeyGenerator.cpp +++ b/src/KeyGenerator.cpp @@ -102,7 +102,7 @@ template KeyGeneratorT::KeyGeneratorT(const Preset target_preset, std::optional seed) : PresetTraits(target_preset), - rng_(createRandomGenerator(seed.value_or(SeedGenerator::Gen()))), + rng_(createRandomGenerator(seed ? *seed : SeedGenerator::Gen())), fft_(degree) { for (u64 i = 0; i < num_p; ++i) { modarith.emplace_back(degree, primes[i]); @@ -664,6 +664,10 @@ KeyGeneratorT::genModPackKeyBundle(const Size pad_rank, const SecretKeyT &sk) const { SwitchKeyT modkey(preset, SWK_MODPACK_SELF); const auto max_length = num_p; + // A self mod-pack key is sized by pad_rank, not by the preset's gadget + // rank, so record it: every key kind keeps axSize()==dnum(), and + // serialization uses dnum to validate the shape of an incoming key. + modkey.setDnum(pad_rank); modkey.addAx(max_length, pad_rank, sk[0][0].getNTTType(), sk[0][0].getNTTRootType()); modkey.addBx(max_length, pad_rank * num_secret, sk[0][0].getNTTType(), @@ -684,6 +688,7 @@ void KeyGeneratorT::genModPackKeyBundleInplace( modkey.axSize() == pad_rank, "[KeyGenerator::genModPackKeyBundle] The provided switching key " "has invalid size."); + modkey.setDnum(pad_rank); const auto ntt_type = sk[0][0].getNTTType(); for (Size i = 0; i < pad_rank; ++i) { diff --git a/src/OmpUtils.cpp b/src/OmpUtils.cpp index a6338de..c338b18 100644 --- a/src/OmpUtils.cpp +++ b/src/OmpUtils.cpp @@ -22,13 +22,17 @@ #endif namespace deb::utils { -static int g_omp_threads = -1; +#ifdef DEB_OPENMP +// Only the OpenMP paths below touch this; declaring it unconditionally makes it +// an unused variable in a build without OpenMP. +static thread_local int tl_omp_threads = -1; +#endif -void setOmpThreadLimit([[__maybe_unused__]] int max_threads) { +void setOmpThreadLimit([[maybe_unused]] int max_threads) { #ifdef DEB_OPENMP int current = omp_get_max_threads(); - if (g_omp_threads == -1) { - g_omp_threads = current; + if (tl_omp_threads == -1) { + tl_omp_threads = current; } if (max_threads < current) { omp_set_num_threads(max_threads); @@ -38,9 +42,9 @@ void setOmpThreadLimit([[__maybe_unused__]] int max_threads) { void unsetOmpThreadLimit() { #ifdef DEB_OPENMP - if (g_omp_threads != -1) { - omp_set_num_threads(g_omp_threads); - g_omp_threads = -1; + if (tl_omp_threads != -1) { + omp_set_num_threads(tl_omp_threads); + tl_omp_threads = -1; } else { const char *env_p = std::getenv("OMP_NUM_THREADS"); if (env_p != nullptr) { @@ -51,4 +55,41 @@ void unsetOmpThreadLimit() { #endif } +namespace { +// Thin wrappers that keep the #ifdef out of OmpThreadLimitGuard, so the guard +// reads both of its members in every build configuration. Guarding the member +// accesses instead would leave them untouched without OpenMP, and silencing +// that needs [[maybe_unused]] on a non-static data member -- which GCC ignores +// with a warning. +int currentThreadCount() { +#ifdef DEB_OPENMP + return omp_get_max_threads(); +#else + return 0; +#endif +} + +void applyThreadCount([[maybe_unused]] int threads) { +#ifdef DEB_OPENMP + omp_set_num_threads(threads); +#endif +} +} // namespace + +OmpThreadLimitGuard::OmpThreadLimitGuard(int max_threads) + : prev_(currentThreadCount()), applied_(false) { + // Without OpenMP the current count reads as 0, so no limit is ever applied + // and the guard is inert. + if (max_threads < prev_) { + applied_ = true; + applyThreadCount(max_threads); + } +} + +OmpThreadLimitGuard::~OmpThreadLimitGuard() { + if (applied_) { + applyThreadCount(prev_); + } +} + } // namespace deb::utils diff --git a/src/SeedGenerator.cpp b/src/SeedGenerator.cpp index 873da57..4536ccf 100644 --- a/src/SeedGenerator.cpp +++ b/src/SeedGenerator.cpp @@ -18,40 +18,49 @@ #include #include +#include #include namespace deb { +namespace { +// Guards the singleton's RNG state: both Gen() and Reseed() mutate it, and +// either may be called concurrently from user threads. +std::mutex g_rng_mutex; + +RNGSeed makeEntropySeed() { + std::random_device rd; + RNGSeed seed = {}; + for (size_t i = 0; i < seed.size(); ++i) { + auto ptr = reinterpret_cast(&seed[i]); + for (size_t j = 0; j < sizeof(u64) / sizeof(unsigned int); ++j) { + ptr[j] = rd(); + } + } + return seed; +} +} // namespace + SeedGenerator &SeedGenerator::GetInstance(const std::optional &seed) { static SeedGenerator instance(seed); return instance; } void SeedGenerator::Reseed(const std::optional &seed) { - const auto &s = seed.value(); - GetInstance().rng_->reseed(reinterpret_cast(s.data()), - DEB_RNG_SEED_BYTE_SIZE); + const RNGSeed s = seed ? *seed : makeEntropySeed(); + SeedGenerator &instance = GetInstance(); + std::lock_guard lock(g_rng_mutex); + instance.rng_->reseed(reinterpret_cast(s.data()), + DEB_RNG_SEED_BYTE_SIZE); } RNGSeed SeedGenerator::Gen() { return GetInstance().genSeed(); } -SeedGenerator::SeedGenerator(const std::optional &seed) { - if (!seed) { - std::random_device rd; - RNGSeed nseed = {}; - for (size_t i = 0; i < nseed.size(); ++i) { - auto ptr = reinterpret_cast(&nseed[i]); - for (size_t j = 0; j < sizeof(u64) / sizeof(unsigned int); ++j) { - ptr[j] = rd(); - } - } - rng_ = createRandomGenerator(nseed); - } else { - rng_ = createRandomGenerator(seed.value()); - } -} +SeedGenerator::SeedGenerator(const std::optional &seed) + : rng_(createRandomGenerator(seed ? *seed : makeEntropySeed())) {} RNGSeed SeedGenerator::genSeed() { RNGSeed seed = {}; + std::lock_guard lock(g_rng_mutex); rng_->getRandomUint64Array(seed.data(), DEB_U64_SEED_SIZE); return seed; } diff --git a/src/Serialize.cpp b/src/Serialize.cpp index 505897d..b7da3cf 100644 --- a/src/Serialize.cpp +++ b/src/Serialize.cpp @@ -15,9 +15,77 @@ */ #include "Serialize.hpp" +#include "utils/Basic.hpp" + +#include +#include namespace deb { +namespace { + +// --- Untrusted-input guards ------------------------------------------------ +// A buffer reaching deserialize*() is attacker-controlled. The flatbuffers +// Verifier run by deserializeFromStream proves only that each vector lies +// inside the supplied bytes; it does not make an optional field present, and it +// does not force a vector's length to agree with a length the payload declares +// separately (Poly.size, Cipher.size, PolyUnit.degree, Coeff.size) or with the +// size implied by the preset. Every such field is therefore checked here before +// it is used to size a copy or to bound a loop. + +[[noreturn]] void rejectBuffer(const char *what) { + throw std::runtime_error(std::string("[deserialize] Malformed buffer: ") + + what); +} + +// Rejects an absent (nullptr) flatbuffers vector field. +template const T *requireField(const T *field, const char *what) { + if (field == nullptr) { + rejectBuffer(what); + } + return field; +} + +// Rejects a vector whose real length disagrees with the length the payload +// declares, or with the length the preset implies. +void requireLength(size_t actual, size_t expected, const char *what) { + if (actual != expected) { + rejectBuffer(what); + } +} + +// Rejects a declared count above what the preset can possibly need. Checking a +// declared length only against another declared length is not enough: both come +// from the buffer, so they can agree with each other and still be absurd. Every +// count that ends up sizing an allocation or indexing a preset table is bounded +// here before it is used. +void requireMaxLength(size_t actual, size_t limit, const char *what) { + if (actual > limit) { + rejectBuffer(what); + } +} + +// Rejects a scalar that is not a valid enumerator. Unchecked casts of these +// bytes reach code that switches on them or uses them to pick a code path. +void requireEnumRange(int value, int min, int max, const char *what) { + if (value < min || value > max) { + rejectBuffer(what); + } +} + +// Rejects a preset byte that names no known preset. Without this the raw value +// flows into the preset accessors, which look it up in a global map and would +// otherwise silently substitute another preset's parameters. +Preset requirePreset(u8 raw) { + const auto preset = static_cast(raw); + if (preset_map.find(preset) == preset_map.end()) { + rejectBuffer("unknown preset"); + } + return preset; +} + +} // namespace + std::vector toComplexVector(const Complex *data, const Size size) { std::vector complex_vec(size); @@ -65,7 +133,9 @@ serializeMessage(flatbuffers::FlatBufferBuilder &builder, } Message deserializeMessage(const deb_fb::Message *message) { - return Message(toDebComplexVector(message->data())); + const auto *data = requireField(message->data(), "Message.data"); + requireLength(data->size(), message->size(), "Message.data"); + return Message(toDebComplexVector(data)); } flatbuffers::Offset @@ -77,7 +147,9 @@ serializeFMessage(flatbuffers::FlatBufferBuilder &builder, } FMessage deserializeFMessage(const deb_fb::Message32 *message) { - return FMessage(toDebComplex32Vector(message->data())); + const auto *data = requireField(message->data(), "Message32.data"); + requireLength(data->size(), message->size(), "Message32.data"); + return FMessage(toDebComplex32Vector(data)); } flatbuffers::Offset @@ -89,9 +161,10 @@ serializeCoeff(flatbuffers::FlatBufferBuilder &builder, } CoeffMessage deserializeCoeff(const deb_fb::Coeff *coeff) { + const auto *data = requireField(coeff->data(), "Coeff.data"); + requireLength(data->size(), coeff->size(), "Coeff.data"); CoeffMessage coeff_t(coeff->size()); - std::memcpy(coeff_t.data(), coeff->data()->data(), - coeff_t.size() * sizeof(Real)); + std::memcpy(coeff_t.data(), data->data(), coeff_t.size() * sizeof(Real)); return coeff_t; } @@ -104,9 +177,10 @@ serializeFCoeff(flatbuffers::FlatBufferBuilder &builder, } FCoeffMessage deserializeFCoeff(const deb_fb::Coeff32 *coeff) { + const auto *data = requireField(coeff->data(), "Coeff32.data"); + requireLength(data->size(), coeff->size(), "Coeff32.data"); FCoeffMessage coeff_t(coeff->size()); - std::memcpy(coeff_t.data(), coeff->data()->data(), - coeff_t.size() * sizeof(float)); + std::memcpy(coeff_t.data(), data->data(), coeff_t.size() * sizeof(float)); return coeff_t; } @@ -122,14 +196,37 @@ serializePolyUnit(flatbuffers::FlatBufferBuilder &builder, builder.CreateVector(polyunit.data(), polyunit.degree())); } -PolyUnit deserializePolyUnit(const deb_fb::PolyUnit *polyunit) { +PolyUnit deserializePolyUnit(const deb_fb::PolyUnit *polyunit, + std::optional preset) { + // Check the array before allocating: `degree` is a declared scalar, so an + // unchecked one would both size a huge allocation and over-read the array. + const auto *array = requireField(polyunit->array(), "PolyUnit.array"); + requireLength(array->size(), polyunit->degree(), "PolyUnit.array"); + // Agreeing with its own array is not enough. Everything downstream (the NTT + // objects, the modular arithmetic, the encrypt/decrypt loops) is sized from + // the PRESET degree, never from this field, so a self-consistent but short + // unit becomes an undersized buffer that the first transform writes past. + // The prime is bound to the preset for the same reason: it selects the + // modulus the coefficients are reduced by. + if (preset.has_value()) { + requireLength(polyunit->degree(), get_degree(*preset), + "PolyUnit.degree does not match the preset"); + const u64 *primes = get_primes(*preset); + const Size num_p = get_num_p(*preset); + bool known_prime = false; + for (Size i = 0; i < num_p && !known_prime; ++i) { + known_prime = (primes[i] == polyunit->prime()); + } + if (!known_prime) { + rejectBuffer("PolyUnit.prime is not a prime of the preset"); + } + } PolyUnit poly_t(polyunit->prime(), polyunit->degree()); int ntt_info = polyunit->ntt_info(); // encoding ntt type and root type into a single int poly_t.setNTT(static_cast(ntt_info / 10), static_cast(ntt_info % 10)); - std::memcpy(poly_t.data(), polyunit->array()->data(), - poly_t.degree() * sizeof(u64)); + std::memcpy(poly_t.data(), array->data(), poly_t.degree() * sizeof(u64)); return poly_t; } @@ -145,9 +242,17 @@ serializePoly(flatbuffers::FlatBufferBuilder &builder, const Polynomial &poly) { } Polynomial deserializePoly(const Preset preset, const deb_fb::Poly *poly) { + const auto *rnspolys = requireField(poly->rnspolys(), "Poly.rnspolys"); + requireLength(rnspolys->size(), poly->size(), "Poly.rnspolys"); + // Bound the limb count before constructing: PolynomialT walks the preset's + // prime table once per limb, so an unbounded count reads past that table + // (and grows the unit vector without limit) before any later guard runs. + requireMaxLength(poly->size(), get_num_p(preset), "Poly.size"); Polynomial poly_t(preset, poly->size()); + requireLength(rnspolys->size(), poly_t.size(), "Poly.rnspolys"); for (Size i = 0; i < poly_t.size(); ++i) { - poly_t[i] = deserializePolyUnit(poly->rnspolys()->Get(i)); + poly_t[i] = deserializePolyUnit( + requireField(rnspolys->Get(i), "Poly.rnspolys[]"), preset); } return poly_t; } @@ -178,19 +283,47 @@ serializeCipher(flatbuffers::FlatBufferBuilder &builder, } Ciphertext deserializeCipher(const deb_fb::Cipher *cipher) { - auto preset = static_cast(cipher->preset()); - Ciphertext cipher_t(preset, cipher->bigpolys()->Get(0)->size(), - cipher->size()); + const auto preset = requirePreset(cipher->preset()); + const auto *bigpolys = requireField(cipher->bigpolys(), "Cipher.bigpolys"); + requireLength(bigpolys->size(), cipher->size(), "Cipher.bigpolys"); + if (bigpolys->size() == 0) { + rejectBuffer("Cipher.bigpolys is empty"); + } + const auto *first = requireField(bigpolys->Get(0), "Cipher.bigpolys[0]"); + // Bound BOTH declared counts against the preset before either sizes an + // allocation. Checking them only against each other is not enough: + // flatbuffers lets many vector entries alias one small table, so a tiny + // buffer can declare an enormous component count. + requireMaxLength(cipher->size(), + get_rank(preset) * get_num_secret(preset) + 1, + "Cipher.size"); + if (first->size() == 0) { + rejectBuffer("Cipher.bigpolys[0] has no limbs"); + } + requireMaxLength(first->size(), get_num_p(preset), "Cipher level"); + requireEnumRange(cipher->encoding(), UNKNOWN, REAL, "Cipher.encoding"); + // The ctor takes a level INDEX and allocates level+1 limbs, indexing the + // preset prime table by limb; passing the limb count would read one past + // it. + Ciphertext cipher_t(preset, first->size() - 1, cipher->size()); cipher_t.setEncoding(static_cast(cipher->encoding())); + requireLength(bigpolys->size(), cipher_t.numPoly(), "Cipher.bigpolys"); for (Size i = 0; i < cipher_t.numPoly(); ++i) { - cipher_t[i] = deserializePoly(preset, cipher->bigpolys()->Get(i)); + const auto *poly = requireField(bigpolys->Get(i), "Cipher.bigpolys[]"); + // Components must agree on their limb count. The sole legitimate + // exception is the released 'a' part of a seed-only ciphertext, which + // serializes as an empty Poly and is regenerated from the seed. + const bool released_ax = + (i + 1 == cipher_t.numPoly()) && poly->size() == 0; + if (!released_ax) { + requireLength(poly->size(), first->size(), "Cipher.bigpolys[]"); + } + cipher_t[i] = deserializePoly(preset, poly); } // Restore seed-only state (field absent for legacy buffers -> full cipher). if (cipher->seed() != nullptr && cipher->seed()->size() != 0) { RNGSeed seed{}; - if (cipher->seed()->size() != seed.size()) { - throw std::runtime_error("[deserializeCipher] Invalid seed size."); - } + requireLength(cipher->seed()->size(), seed.size(), "Cipher.seed"); std::memcpy(seed.data(), cipher->seed()->data(), sizeof(RNGSeed)); cipher_t.setSeed(seed); cipher_t.setSeedMode(static_cast(cipher->seed_mode())); @@ -216,20 +349,37 @@ serializeSk(flatbuffers::FlatBufferBuilder &builder, const SecretKey &sk) { SecretKey deserializeSk(const deb_fb::Sk *sk) { RNGSeed seed = {}; - SecretKey sk_t(static_cast(sk->preset()), seed); + SecretKey sk_t(requirePreset(sk->preset()), seed); sk_t.flushSeed(); - if (sk->seed()->size() != 0) { - std::memcpy(seed.data(), sk->seed()->data(), sizeof(RNGSeed)); + const auto *sk_seed = requireField(sk->seed(), "Sk.seed"); + if (sk_seed->size() != 0) { + requireLength(sk_seed->size(), seed.size(), "Sk.seed"); + std::memcpy(seed.data(), sk_seed->data(), sizeof(RNGSeed)); sk_t.setSeed(seed); } - if (sk->coeffs()->size() != 0) { + const auto *coeffs = requireField(sk->coeffs(), "Sk.coeffs"); + if (coeffs->size() != 0) { + // The destination is sized from the preset alone, so the declared + // length has to be checked against it: copying an attacker-chosen + // number of coefficients into it is a heap overflow. sk_t.allocCoeffs(); - std::copy(sk->coeffs()->begin(), sk->coeffs()->end(), sk_t.coeffs()); + requireLength(coeffs->size(), sk_t.coeffsSize(), "Sk.coeffs"); + std::copy(coeffs->begin(), coeffs->end(), sk_t.coeffs()); } - if (sk->bigpolys()->size() != 0) { - sk_t.allocPolys(sk->bigpolys()->Get(0)->rnspolys()->size()); + const auto *bigpolys = requireField(sk->bigpolys(), "Sk.bigpolys"); + if (bigpolys->size() != 0) { + const auto *first = requireField( + requireField(bigpolys->Get(0), "Sk.bigpolys[0]")->rnspolys(), + "Sk.bigpolys[0].rnspolys"); + if (first->size() == 0) { + rejectBuffer("Sk.bigpolys[0] has no limbs"); + } + requireMaxLength(first->size(), get_num_p(sk_t.preset()), + "Sk.bigpolys[0] level"); + sk_t.allocPolys(first->size()); + requireLength(bigpolys->size(), sk_t.numPoly(), "Sk.bigpolys"); for (Size i = 0; i < sk_t.numPoly(); ++i) { - sk_t[i] = deserializePoly(sk_t.preset(), sk->bigpolys()->Get(i)); + sk_t[i] = deserializePoly(sk_t.preset(), bigpolys->Get(i)); } } return sk_t; @@ -254,18 +404,62 @@ serializeSwk(flatbuffers::FlatBufferBuilder &builder, const SwitchKey &swk) { } SwitchKey deserializeSwk(const deb_fb::Swk *swk) { - const auto preset = static_cast(swk->preset()); + const auto preset = requirePreset(swk->preset()); + const auto *ax = requireField(swk->ax(), "Swk.ax"); + const auto *bx = requireField(swk->bx(), "Swk.bx"); + requireEnumRange(swk->type(), SWK_GENERIC, SWK_MODPACK_SELF, "Swk.type"); + // A key's shape is fully determined by its kind and its dnum, so pin it + // exactly rather than merely bounding it. Without this, a few KiB of + // aliased flatbuffer offsets can declare an unbounded number of full-size + // polynomials: each entry costs 4 bytes of input but one whole polynomial + // of allocation. + const auto kind = static_cast(swk->type()); + const Size dnum = swk->dnum(); + // dnum is the number of ax polynomials. For a self mod-pack key it is the + // pad_rank, a power of two that divides the degree; every other keyed kind + // uses the preset's gadget rank. + // A self mod-pack key's pad_rank is only bounded by the degree, which for a + // large preset would still permit gigabytes of polynomials. Bound it by + // what could actually have been WRITTEN instead: every ax and bx polynomial + // must carry real coefficients, so a key needs at least ax_count * (1 + + // num_secret) * num_p * degree * 8 bytes on the wire. A key that could + // never fit in a serialized buffer cannot have come from serializeToStream, + // so there is no reason to accept it -- and this is what stops a few KiB of + // aliased offsets from declaring an unbounded key. + const u64 bytes_per_ax = u64{1} + get_num_secret(preset); + const u64 wire_bytes_per_ax = + bytes_per_ax * get_num_p(preset) * get_degree(preset) * sizeof(u64); + const Size max_ax_by_wire_size = static_cast( + std::max(1, u64{DEB_MAX_SERIALIZED_SIZE} / wire_bytes_per_ax)); + + if (kind == SWK_MODPACK_SELF) { + if (!utils::isPowerOfTwo(dnum) || dnum > max_ax_by_wire_size) { + rejectBuffer("Swk.dnum is not a valid pad_rank"); + } + } else { + requireMaxLength(dnum, get_gadget_rank(preset), "Swk.dnum"); + } + if (kind == SWK_GENERIC) { + // Built empty and filled by the caller, so only bound it. + requireMaxLength(ax->size(), max_ax_by_wire_size, "Swk.ax"); + } else { + // genEncKeyInplace asserts axSize()==1; every other kind asserts + // axSize()==dnum(). + requireLength(ax->size(), (kind == SWK_ENC) ? Size{1} : dnum, "Swk.ax"); + } + requireLength(bx->size(), ax->size() * get_num_secret(preset), "Swk.bx"); SwitchKey swk_t(preset, static_cast(swk->type())); swk_t.getAx().clear(); - for (Size i = 0; i < swk->ax()->size(); ++i) { - Polynomial tmp = deserializePoly(preset, swk->ax()->Get(i)); + for (Size i = 0; i < ax->size(); ++i) { + Polynomial tmp = deserializePoly(preset, ax->Get(i)); swk_t.addAx(tmp); } swk_t.getBx().clear(); - for (Size i = 0; i < swk->bx()->size(); ++i) { - Polynomial tmp = deserializePoly(preset, swk->bx()->Get(i)); + for (Size i = 0; i < bx->size(); ++i) { + Polynomial tmp = deserializePoly(preset, bx->Get(i)); swk_t.addBx(tmp); } + swk_t.setDnum(dnum); if (swk->rot_idx() != static_cast(-1)) { swk_t.setRotIdx(swk->rot_idx()); } diff --git a/test/EnDecryption-test.cpp b/test/EnDecryption-test.cpp index 33b5f74..92e127f 100644 --- a/test/EnDecryption-test.cpp +++ b/test/EnDecryption-test.cpp @@ -321,53 +321,59 @@ TEST_P(EnDecrypt, SeedOnlyDeterministicRegeneration) { } } -// Public-key seed-only encryption: 'a' = v*ax + e cannot be regenerated without -// the encryption key, so the decryptor must refuse and -// completeCiphertext(enckey) is required first. -TEST_P(EnDecrypt, SeedOnlyEncryptAndDecryptWithEncKey) { +// Seed compression is unsound under a public key: 'a' = v*ax + e_a is drawn +// from the same stream as the ephemeral encryption randomness v, so storing +// that seed would let anyone holding the ciphertext and the (public) encryption +// key replay v and recover the plaintext without the secret key. encrypt() must +// refuse the combination outright. +TEST_P(EnDecrypt, SeedOnlyRejectedForPublicKey) { MSGS msg = gen_random_message(); SecretKey sk = SecretKeyGenerator::GenSecretKey(preset); SwitchKey enckey = KeyGenerator(preset).genEncKey(sk); - MSGS decrypted_msg = gen_empty_message(); - for (Size l = 0; l < get_num_p(preset); ++l) { + { + Ciphertext ctxt(preset); + EXPECT_THROW(encryptor.encrypt(msg, enckey, ctxt, + EncryptOptions().SeedOnlyA(true)), + std::runtime_error); + } + { // A caller-supplied a_seed must not provide a way around the check. + Ciphertext ctxt(preset); + EXPECT_THROW( + encryptor.encrypt( + msg, enckey, ctxt, + EncryptOptions().ASeed(SeedGenerator::Gen()).SeedOnlyA(true)), + std::runtime_error); + } + { // Nor does the coefficient-domain output path. + Ciphertext ctxt(preset); + EXPECT_THROW( + encryptor.encrypt(msg, enckey, ctxt, + EncryptOptions().NTTOut(false).SeedOnlyA(true)), + std::runtime_error); + } + + // Public-key encryption without seed compression is unaffected, and the + // same option is still accepted for secret-key encryption. + { + const Size l = 0; Ciphertext ctxt(preset, l); MSGS scaled_msg = scale_message(msg, l); - encryptor.encrypt(scaled_msg, enckey, ctxt, - EncryptOptions().Level(l).SeedOnlyA(true)); - EXPECT_TRUE(ctxt.hasSeed()); - EXPECT_TRUE(ctxt.isAxFlushed()); + MSGS decrypted_msg = gen_empty_message(); + EXPECT_NO_THROW(encryptor.encrypt(scaled_msg, enckey, ctxt, + EncryptOptions().Level(l))); + EXPECT_FALSE(ctxt.hasSeed()); EXPECT_EQ(static_cast(ctxt.seedMode()), - static_cast(CipherSeedMode::PUBLICKEY)); - - MSGS tmp = gen_empty_message(); - EXPECT_THROW(decryptor.decrypt(ctxt, sk, tmp), std::runtime_error); - - encryptor.completeCiphertext(ctxt, enckey); - EXPECT_FALSE(ctxt.isAxFlushed()); + static_cast(CipherSeedMode::NONE)); decryptor.decrypt(ctxt, sk, decrypted_msg); compare_msg(scaled_msg, decrypted_msg, scale_error(enc_err, l)); } -} - -// Public-key path is deterministic from (a_seed, error_seed, enckey). -TEST_P(EnDecrypt, SeedOnlyEncKeyDeterministicRegeneration) { - MSGS msg = gen_random_message(); - SecretKey sk = SecretKeyGenerator::GenSecretKey(preset); - SwitchKey enckey = KeyGenerator(preset).genEncKey(sk); - RNGSeed a_seed = SeedGenerator::Gen(); - RNGSeed e_seed = SeedGenerator::Gen(); - auto opt = EncryptOptions().ASeed(a_seed).ErrorSeed(e_seed).SeedOnlyA(true); - - Ciphertext c1(preset), c2(preset); - encryptor.encrypt(msg, enckey, c1, opt); - encryptor.encrypt(msg, enckey, c2, opt); - encryptor.completeCiphertext(c1, enckey); - encryptor.completeCiphertext(c2, enckey); - - ASSERT_EQ(c1.numPoly(), c2.numPoly()); - for (Size p = 0; p < c1.numPoly(); ++p) { - comparePoly(c1[p], c2[p]); + { + Ciphertext ctxt(preset); + EXPECT_NO_THROW( + encryptor.encrypt(msg, sk, ctxt, EncryptOptions().SeedOnlyA(true))); + EXPECT_EQ(static_cast(ctxt.seedMode()), + static_cast(CipherSeedMode::UNIFORM)); } } @@ -398,31 +404,6 @@ TEST_P(EnDecrypt, SeedOnlyNttOutFalseWithSecretKey) { } } -TEST_P(EnDecrypt, SeedOnlyNttOutFalseWithEncKey) { - MSGS msg = gen_random_message(); - SecretKey sk = SecretKeyGenerator::GenSecretKey(preset); - SwitchKey enckey = KeyGenerator(preset).genEncKey(sk); - MSGS decrypted_msg = gen_empty_message(); - - for (Size l = 0; l < get_num_p(preset); ++l) { - Ciphertext ctxt(preset, l); - MSGS scaled_msg = scale_message(msg, l); - encryptor.encrypt( - scaled_msg, enckey, ctxt, - EncryptOptions().Level(l).NTTOut(false).SeedOnlyA(true)); - EXPECT_FALSE(ctxt[0][0].isNTT()); - - MSGS tmp = gen_empty_message(); - EXPECT_THROW(decryptor.decrypt(ctxt, sk, tmp), std::runtime_error); - - encryptor.completeCiphertext(ctxt, enckey); - ASSERT_FALSE(ctxt.isAxFlushed()); - EXPECT_EQ(ctxt[ctxt.numPoly() - 1][0].isNTT(), ctxt[0][0].isNTT()); - decryptor.decrypt(ctxt, sk, decrypted_msg); - compare_msg(scaled_msg, decrypted_msg, scale_error(enc_err, l)); - } -} - /*--------------------------------------------------- Real Encryption Tests ---------------------------------------------------*/ diff --git a/test/Operation-test.cpp b/test/Operation-test.cpp index 9450b29..339fe21 100644 --- a/test/Operation-test.cpp +++ b/test/Operation-test.cpp @@ -279,8 +279,22 @@ TEST_F(I128ArithTest, Edge_MaxPlusMinIsNegOne) { TEST_F(I128ArithTest, Edge_MaxMinusMinIsAllOnes) { // I128_MAX - I128_MIN wraps: (2^127-1) - (-2^127) = 2^128-1 ≡ -1 mod 2^128 - // cast back to i128 → -1 (wraps) + // cast back to i128 → -1 (wraps). The overflow is intentional, so silence + // -Woverflow on just this constant-folded expression instead of disabling + // the warning for the whole target. +#if defined(__clang__) +#pragma clang diagnostic push +#pragma clang diagnostic ignored "-Woverflow" +#elif defined(__GNUC__) +#pragma GCC diagnostic push +#pragma GCC diagnostic ignored "-Woverflow" +#endif EXPECT_EQ(I128_MAX - I128_MIN, static_cast(-1)); +#if defined(__clang__) +#pragma clang diagnostic pop +#elif defined(__GNUC__) +#pragma GCC diagnostic pop +#endif } // Random tests diff --git a/test/Serialize-test.cpp b/test/Serialize-test.cpp index 965c154..2a65ff1 100644 --- a/test/Serialize-test.cpp +++ b/test/Serialize-test.cpp @@ -96,9 +96,17 @@ TEST_P(Serialize, PolyUnitSerializationTest) { serializeToStream(poly, os); std::istringstream is(os.str()); PolyUnit deserialized_poly(prime, 0); - deserializeFromStream(is, deserialized_poly); + // A bare PolyUnit carries its own degree and prime, so a preset is required + // to bind them -- as it already is for Polynomial. + deserializeFromStream(is, deserialized_poly, preset); comparePolyUnit(poly, deserialized_poly); + + // Without the preset there is nothing to validate the declared degree + // against, so the call is refused rather than trusting the buffer. + std::istringstream is_nopreset(os.str()); + PolyUnit out(prime, 0); + EXPECT_THROW(deserializeFromStream(is_nopreset, out), std::runtime_error); } TEST_P(Serialize, PolySerializationTest) { @@ -313,6 +321,215 @@ TEST_P(Serialize, SeedOnlyCipherSerializationTest) { compare_msg(msg, dec, scale_error(sk_err, 0)); } +// Deserialization consumes untrusted bytes, so every malformed shape must be +// rejected with an exception rather than parsed. These checks are +// unconditional: they must hold regardless of DEB_RUNTIME_RESOURCE_CHECK, +// which previously gated the flatbuffers verifier call and compiled it out +// entirely when off. +TEST_P(Serialize, MalformedBufferIsRejected) { + // A valid buffer to corrupt, and a sanity check that it round-trips. + MSGS msg = gen_random_message(); + msg = scale_message(msg, 0); + SecretKey sk = SecretKeyGenerator::GenSecretKey(preset); + Ciphertext ctxt(preset); + encryptor.encrypt(msg, sk, ctxt); + std::ostringstream os; + serializeToStream(ctxt, os); + const std::string good = os.str(); + { + std::istringstream is(good); + Ciphertext out(preset); + EXPECT_NO_THROW(deserializeFromStream(is, out)); + } + ASSERT_GT(good.size(), sizeof(Size)); + + const auto with_prefix = [](Size n, const std::string &payload) { + std::string s(reinterpret_cast(&n), sizeof(Size)); + s += payload; + return s; + }; + const std::string body = good.substr(sizeof(Size)); + + // Empty stream: the length prefix itself cannot be read. + { + std::istringstream is(std::string{}); + Ciphertext out(preset); + EXPECT_THROW(deserializeFromStream(is, out), std::runtime_error); + } + // Zero-length payload. + { + std::istringstream is(with_prefix(0, std::string{})); + Ciphertext out(preset); + EXPECT_THROW(deserializeFromStream(is, out), std::runtime_error); + } + // Length prefix at or beyond the bound: must be refused by the bound + // itself, before any allocation is attempted. The payload is left tiny on + // purpose -- if the bound were removed, this would try to allocate 2 GiB + // rather than fail the check, so the case genuinely covers the bound. + { + std::istringstream is(with_prefix(DEB_MAX_SERIALIZED_SIZE, body)); + Ciphertext out(preset); + EXPECT_THROW(deserializeFromStream(is, out), std::runtime_error); + } + { + std::istringstream is( + with_prefix(DEB_MAX_SERIALIZED_SIZE + 1, std::string{})); + Ciphertext out(preset); + EXPECT_THROW(deserializeFromStream(is, out), std::runtime_error); + } + // Length prefix longer than the bytes actually present (truncated stream). + { + std::istringstream is( + with_prefix(static_cast(body.size()), body.substr(0, 8))); + Ciphertext out(preset); + EXPECT_THROW(deserializeFromStream(is, out), std::runtime_error); + } + // Structurally invalid payload of a plausible length: caught by the + // flatbuffers verifier. + { + std::string garbage(body.size(), '\xA5'); + std::istringstream is( + with_prefix(static_cast(garbage.size()), garbage)); + Ciphertext out(preset); + EXPECT_THROW(deserializeFromStream(is, out), std::runtime_error); + } + // Truncated-but-well-prefixed payload: the declared length matches the + // bytes supplied, so only the verifier can reject it. + { + const std::string half = body.substr(0, body.size() / 2); + std::istringstream is( + with_prefix(static_cast(half.size()), half)); + Ciphertext out(preset); + EXPECT_THROW(deserializeFromStream(is, out), std::runtime_error); + } +} + +// A secret key blob is parsed with the destination sized from the preset alone, +// so the declared coefficient count must be checked against it. Round-tripping +// the library's own output must keep working. +TEST_P(Serialize, SecretKeyRoundTripKeepsCoeffs) { + SecretKey sk = SecretKeyGenerator::GenSecretKey(preset); + sk.allocCoeffs(); + SecretKeyGenerator::GenCoeffInplace(preset, sk.coeffs(), sk.getSeed()); + ASSERT_GT(sk.coeffsSize(), 0u); + + std::ostringstream os; + serializeToStream(sk, os); + std::istringstream is(os.str()); + SecretKey out(preset, false); + ASSERT_NO_THROW(deserializeFromStream(is, out)); + EXPECT_EQ(out.coeffsSize(), sk.coeffsSize()); + compareArray(sk.coeffs(), out.coeffs(), sk.coeffsSize()); +} + +// A self mod-pack key is sized by its pad_rank rather than the preset's gadget +// rank. That rank is carried as the key's dnum, which is what lets +// deserialization pin the key's shape exactly instead of merely bounding it. +TEST_P(Serialize, ModPackSelfKeySerializationTest) { + if (num_secret != 1) { + GTEST_SKIP() + << "MODPACK_SELF key generation is only for single secret."; + } + KeyGenerator keygen(preset); + SecretKey sk = SecretKeyGenerator::GenSecretKey(preset); + // Kept small on purpose: a self mod-pack key holds pad_rank*(1+num_secret) + // full polynomials, so a large pad_rank runs into the serialized-size + // ceiling rather than testing the shape validation. It must also differ + // from the gadget rank, which varies with the parameter set -- otherwise a + // dnum that silently fell back to the gadget rank would still look right. + Size pad_rank = 2; + while (pad_rank == get_gadget_rank(preset)) { + pad_rank *= 2; + } + if (pad_rank > degree) { + GTEST_SKIP() << "pad_rank must not exceed the degree."; + } + + SwitchKey modkey = keygen.genModPackKeyBundle(pad_rank, sk); + ASSERT_EQ(modkey.axSize(), pad_rank); + ASSERT_EQ(modkey.bxSize(), pad_rank * num_secret); + // The generator records pad_rank as dnum, keeping axSize()==dnum() true for + // this kind as it already is for every other one. + ASSERT_EQ(modkey.dnum(), pad_rank); + ASSERT_NE(pad_rank, get_gadget_rank(preset)) + << "pad_rank must differ from the gadget rank for this test to prove " + "that dnum really carries pad_rank"; + + std::ostringstream os; + serializeToStream(modkey, os); + std::istringstream is(os.str()); + SwitchKey back(preset, SwitchKeyKind::SWK_MODPACK_SELF); + ASSERT_NO_THROW(deserializeFromStream(is, back)); + + EXPECT_EQ(back.type(), modkey.type()); + EXPECT_EQ(back.dnum(), pad_rank); + EXPECT_EQ(back.axSize(), pad_rank); + EXPECT_EQ(back.bxSize(), pad_rank * num_secret); + for (Size i = 0; i < modkey.axSize(); ++i) { + comparePoly(modkey.ax(i), back.ax(i)); + } + for (Size i = 0; i < modkey.bxSize(); ++i) { + comparePoly(modkey.bx(i), back.bx(i)); + } +} + +// serializeToStream refuses an object too large for one buffer, using this +// bound. FlatBuffers' internal size counter is a uint32 whose only guard is an +// assert that release builds compile out, so the bound has to be computed from +// the object BEFORE building -- and it must never under-count, or the guard +// lets through exactly the buffers it exists to stop. +TEST_P(Serialize, SerializedSizeUpperBoundNeverUnderCounts) { + const auto check = [](const char *what, const auto &obj) { + std::ostringstream os; + serializeToStream(obj, os); + const u64 actual = os.str().size(); + const u64 bound = serializedSizeUpperBound(obj); + EXPECT_GE(bound, actual) << what << ": bound under-counts"; + // Loose enough to survive a FlatBuffers bump, tight enough that the + // bound still means something. + EXPECT_LT(bound, actual * 2) << what << ": bound is uselessly loose"; + }; + + Message msg = gen_random_message()[0]; + check("Message", msg); + check("CoeffMessage", gen_random_coeff()[0]); + check("FMessage", gen_random_message()[0]); + check("FCoeffMessage", gen_random_coeff()[0]); + + Polynomial poly(preset); + check("Polynomial", poly); + check("PolyUnit", poly[0]); + + SecretKey sk = SecretKeyGenerator::GenSecretKey(preset); + check("SecretKey", sk); + + MSGS msgs = gen_random_message(); + msgs = scale_message(msgs, 0); + Ciphertext ctxt(preset); + encryptor.encrypt(msgs, sk, ctxt); + check("Ciphertext", ctxt); + + // A seed-only ciphertext releases its 'a' part, so the bound must follow + // the real per-polynomial shapes rather than any preset-derived limb count. + Ciphertext seed_only(preset); + encryptor.encrypt(msgs, sk, seed_only, EncryptOptions().SeedOnlyA(true)); + check("Ciphertext seed-only", seed_only); + + KeyGenerator keygen(preset); + check("SwitchKey enc", keygen.genEncKey(sk)); + check("SwitchKey mult", keygen.genMultKey(sk)); +} + +// A write failure must not pass silently: a failed first write makes the second +// a no-op, so without a check a truncated record is indistinguishable from a +// complete one. +TEST_P(Serialize, SerializeReportsStreamFailure) { + Message msg = gen_random_message()[0]; + std::ostringstream os; + os.setstate(std::ios::badbit); + EXPECT_THROW(serializeToStream(msg, os), std::runtime_error); +} + #define X(PRESET) Preset::PRESET_##PRESET, const std::vector all_presets = {PRESET_LIST #undef X