From 7f0eb69bc50b65aaafe277c228cc552c7287de5d Mon Sep 17 00:00:00 2001 From: sh1k4ku Date: Tue, 30 Jun 2026 20:10:48 +0800 Subject: [PATCH] Add Inspire PSI launcher support --- MODULE.bazel.lock | 122 ++++ psi/algorithm/inspire/BUILD.bazel | 187 ++++++ psi/algorithm/inspire/client.cc | 61 ++ psi/algorithm/inspire/client.h | 47 ++ psi/algorithm/inspire/entry.cc | 222 +++++++ psi/algorithm/inspire/entry.h | 32 + psi/algorithm/inspire/inspire_flow_smoke.cc | 41 ++ psi/algorithm/inspire/inspire_flow_test.sh | 4 + psi/algorithm/inspire/internal.cc | 609 +++++++++++++++++ psi/algorithm/inspire/internal.h | 37 ++ psi/algorithm/inspire/internal_primitives.cc | 238 +++++++ psi/algorithm/inspire/internal_primitives.h | 78 +++ psi/algorithm/inspire/large_record.cc | 612 ++++++++++++++++++ psi/algorithm/inspire/large_record.h | 32 + .../inspire/large_record_flow_smoke.cc | 70 ++ .../inspire/large_record_full_test.cc | 56 ++ psi/algorithm/inspire/large_record_he.cc | 353 ++++++++++ psi/algorithm/inspire/large_record_he.h | 123 ++++ psi/algorithm/inspire/params.h | 26 + psi/algorithm/inspire/serialize.cc | 111 ++++ psi/algorithm/inspire/serialize.h | 16 + psi/algorithm/inspire/server.cc | 93 +++ psi/algorithm/inspire/server.h | 52 ++ psi/algorithm/inspire/types.h | 83 +++ psi/apps/psi_launcher/BUILD.bazel | 34 + psi/apps/psi_launcher/inspire_launch.cc | 34 + psi/apps/psi_launcher/inspire_launch.h | 17 + psi/apps/psi_launcher/inspire_launch_stub.cc | 17 + psi/apps/psi_launcher/inspire_test.cc | 90 +++ psi/apps/psi_launcher/launch.h | 6 + psi/apps/psi_launcher/main.cc | 12 + psi/proto/entry.proto | 4 + psi/proto/pir.proto | 30 + 33 files changed, 3549 insertions(+) create mode 100644 psi/algorithm/inspire/BUILD.bazel create mode 100644 psi/algorithm/inspire/client.cc create mode 100644 psi/algorithm/inspire/client.h create mode 100644 psi/algorithm/inspire/entry.cc create mode 100644 psi/algorithm/inspire/entry.h create mode 100644 psi/algorithm/inspire/inspire_flow_smoke.cc create mode 100644 psi/algorithm/inspire/inspire_flow_test.sh create mode 100644 psi/algorithm/inspire/internal.cc create mode 100644 psi/algorithm/inspire/internal.h create mode 100644 psi/algorithm/inspire/internal_primitives.cc create mode 100644 psi/algorithm/inspire/internal_primitives.h create mode 100644 psi/algorithm/inspire/large_record.cc create mode 100644 psi/algorithm/inspire/large_record.h create mode 100644 psi/algorithm/inspire/large_record_flow_smoke.cc create mode 100644 psi/algorithm/inspire/large_record_full_test.cc create mode 100644 psi/algorithm/inspire/large_record_he.cc create mode 100644 psi/algorithm/inspire/large_record_he.h create mode 100644 psi/algorithm/inspire/params.h create mode 100644 psi/algorithm/inspire/serialize.cc create mode 100644 psi/algorithm/inspire/serialize.h create mode 100644 psi/algorithm/inspire/server.cc create mode 100644 psi/algorithm/inspire/server.h create mode 100644 psi/algorithm/inspire/types.h create mode 100644 psi/apps/psi_launcher/inspire_launch.cc create mode 100644 psi/apps/psi_launcher/inspire_launch.h create mode 100644 psi/apps/psi_launcher/inspire_launch_stub.cc create mode 100644 psi/apps/psi_launcher/inspire_test.cc diff --git a/MODULE.bazel.lock b/MODULE.bazel.lock index 2b1e424d..6196ea06 100644 --- a/MODULE.bazel.lock +++ b/MODULE.bazel.lock @@ -910,6 +910,128 @@ }, "selectedYankedVersions": {}, "moduleExtensions": { + "//bazel:defs.bzl%non_module_dependencies": { + "general": { + "bzlTransitiveDigest": "NAC0+n0skvguFGiSWQSZMhbrQQVV+nNsrNIVDde/xQE=", + "usagesDigest": "89PlExpoF4wVcL+M7Z1YGe8ZyC3E1M5lGgLzpZSTyvE=", + "recordedFileInputs": {}, + "recordedDirentsInputs": {}, + "envVariables": {}, + "generatedRepoSpecs": { + "apsi": { + "bzlFile": "@@bazel_tools//tools/build_defs/repo:http.bzl", + "ruleClassName": "http_archive", + "attributes": { + "sha256": "82c0f9329c79222675109d4a3682d204acd3ea9a724bcd98fa58eabe53851333", + "strip_prefix": "APSI-0.11.0", + "urls": [ + "https://github.com/microsoft/APSI/archive/refs/tags/v0.11.0.tar.gz" + ], + "build_file": "@@//bazel:microsoft_apsi.BUILD", + "patch_args": [ + "-p1" + ], + "patches": [ + "@@//bazel/patches:apsi.patch", + "@@//bazel/patches:apsi-fourq.patch" + ], + "patch_cmds": [ + "rm -rf common/apsi/fourq" + ] + } + }, + "kuku": { + "bzlFile": "@@bazel_tools//tools/build_defs/repo:http.bzl", + "ruleClassName": "http_archive", + "attributes": { + "sha256": "96ed5fad82ea8c8a8bb82f6eaf0b5dce744c0c2566b4baa11d8f5443ad1f83b7", + "strip_prefix": "Kuku-2.1.0", + "type": "tar.gz", + "urls": [ + "https://github.com/microsoft/Kuku/archive/refs/tags/v2.1.0.tar.gz" + ], + "build_file": "@@//bazel:microsoft_kuku.BUILD" + } + }, + "com_google_flatbuffers": { + "bzlFile": "@@bazel_tools//tools/build_defs/repo:http.bzl", + "ruleClassName": "http_archive", + "attributes": { + "sha256": "4157c5cacdb59737c5d627e47ac26b140e9ee28b1102f812b36068aab728c1ed", + "strip_prefix": "flatbuffers-24.3.25", + "urls": [ + "https://github.com/google/flatbuffers/archive/refs/tags/v24.3.25.tar.gz" + ], + "patch_cmds": [ + "rm grpc/BUILD.bazel", + "rm grpc/src/compiler/BUILD.bazel", + "rm src/BUILD.bazel" + ], + "build_file": "@@//bazel:flatbuffers.BUILD" + } + }, + "curve25519-donna": { + "bzlFile": "@@bazel_tools//tools/build_defs/repo:http.bzl", + "ruleClassName": "http_archive", + "attributes": { + "strip_prefix": "curve25519-donna-2fe66b65ea1acb788024f40a3373b8b3e6f4bbb2", + "sha256": "ba57d538c241ad30ff85f49102ab2c8dd996148456ed238a8c319f263b7b149a", + "type": "tar.gz", + "build_file": "@@//bazel:curve25519-donna.BUILD", + "urls": [ + "https://github.com/floodyberry/curve25519-donna/archive/2fe66b65ea1acb788024f40a3373b8b3e6f4bbb2.tar.gz" + ] + } + }, + "com_github_zeromq_cppzmq": { + "bzlFile": "@@bazel_tools//tools/build_defs/repo:http.bzl", + "ruleClassName": "http_archive", + "attributes": { + "build_file": "@@//bazel:cppzmq.BUILD", + "strip_prefix": "cppzmq-4.10.0", + "sha256": "c81c81bba8a7644c84932225f018b5088743a22999c6d82a2b5f5cd1e6942b74", + "type": ".tar.gz", + "urls": [ + "https://github.com/zeromq/cppzmq/archive/refs/tags/v4.10.0.tar.gz" + ] + } + }, + "com_github_zeromq_libzmq": { + "bzlFile": "@@bazel_tools//tools/build_defs/repo:http.bzl", + "ruleClassName": "http_archive", + "attributes": { + "build_file": "@@//bazel:libzmq.BUILD", + "strip_prefix": "libzmq-4.3.5", + "sha256": "6c972d1e6a91a0ecd79c3236f04cf0126f2f4dfbbad407d72b4606a7ba93f9c6", + "type": ".tar.gz", + "urls": [ + "https://github.com/zeromq/libzmq/archive/refs/tags/v4.3.5.tar.gz" + ] + } + }, + "com_github_open_source_parsers_jsoncpp": { + "bzlFile": "@@bazel_tools//tools/build_defs/repo:http.bzl", + "ruleClassName": "http_archive", + "attributes": { + "build_file": "@@//bazel:jsoncpp.BUILD", + "strip_prefix": "jsoncpp-1.9.6", + "sha256": "f93b6dd7ce796b13d02c108bc9f79812245a82e577581c4c9aabe57075c90ea2", + "type": ".tar.gz", + "urls": [ + "https://github.com/open-source-parsers/jsoncpp/archive/refs/tags/1.9.6.tar.gz" + ] + } + } + }, + "recordedRepoMappingEntries": [ + [ + "", + "bazel_tools", + "bazel_tools" + ] + ] + } + }, "@@apple_support~//crosstool:setup.bzl%apple_cc_configure_extension": { "general": { "bzlTransitiveDigest": "7ii+gFxWSxHhQPrBxfMEHhtrGvHmBTvsh+KOyGunP/s=", diff --git a/psi/algorithm/inspire/BUILD.bazel b/psi/algorithm/inspire/BUILD.bazel new file mode 100644 index 00000000..de91879c --- /dev/null +++ b/psi/algorithm/inspire/BUILD.bazel @@ -0,0 +1,187 @@ +# Copyright 2026 The secretflow authors. + +load("//bazel:psi.bzl", "psi_cc_binary", "psi_cc_library", "psi_cc_test") + +package(default_visibility = ["//visibility:public"]) + +X86_64_COMPATIBLE = ["@platforms//cpu:x86_64"] + +psi_cc_library( + name = "types", + hdrs = ["types.h"], + target_compatible_with = X86_64_COMPATIBLE, +) + +psi_cc_library( + name = "params", + hdrs = ["params.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = ["//psi/algorithm/ypir:params"], +) + +psi_cc_library( + name = "serialize", + srcs = ["serialize.cc"], + hdrs = ["serialize.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":types", + "@yacl//yacl/base:buffer", + "@yacl//yacl/base:byte_container_view", + "@yacl//yacl/base:exception", + ], +) + +psi_cc_library( + name = "internal_primitives", + srcs = ["internal_primitives.cc"], + hdrs = ["internal_primitives.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + "//psi/algorithm/ypir:legacy_ypir_impl", + "//psi/algorithm/ypir:params", + "//psi/algorithm/ypir:ypir_internal_params", + "@yacl//yacl/base:exception", + ], +) + +psi_cc_library( + name = "internal", + srcs = ["internal.cc"], + hdrs = ["internal.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":types", + "//psi/algorithm/ypir:legacy_ypir_impl", + "//psi/algorithm/ypir:params", + "//psi/algorithm/ypir:util", + "//psi/algorithm/ypir:ypir_internal_params", + "@yacl//yacl/base:exception", + ], +) + +psi_cc_library( + name = "large_record_he", + srcs = ["large_record_he.cc"], + hdrs = ["large_record_he.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + "@hexl", + ], +) + +psi_cc_library( + name = "large_record", + srcs = ["large_record.cc"], + hdrs = ["large_record.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":internal_primitives", + ":large_record_he", + ":types", + "//psi/algorithm/ypir:legacy_ypir_impl", + "//psi/algorithm/ypir:ypir_internal_params", + "@hexl", + "@yacl//yacl/base:exception", + ], +) + +psi_cc_library( + name = "client", + srcs = ["client.cc"], + hdrs = ["client.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":internal", + ":params", + ":serialize", + ":types", + "//psi/algorithm/pir_interface:index_pir", + "//psi/algorithm/pir_interface:pir_db", + "//psi/algorithm/ypir:ypir_internal_params", + "@yacl//yacl/base:buffer", + "@yacl//yacl/base:byte_container_view", + "@yacl//yacl/base:exception", + ], +) + +psi_cc_library( + name = "server", + srcs = ["server.cc"], + hdrs = ["server.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":internal", + ":params", + ":serialize", + ":types", + "//psi/algorithm/pir_interface:pir_db", + "//psi/algorithm/ypir:ypir_internal_params", + "@yacl//yacl/base:buffer", + "@yacl//yacl/base:byte_container_view", + "@yacl//yacl/base:exception", + ], +) + +psi_cc_library( + name = "entry", + srcs = ["entry.cc"], + hdrs = ["entry.h"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":client", + ":params", + ":server", + "//psi/algorithm/pir_interface:pir_db", + "@abseil-cpp//absl/strings", + "@yacl//yacl/base:buffer", + "@yacl//yacl/base:byte_container_view", + "@yacl//yacl/base:exception", + "@yacl//yacl/link:context", + ], +) + +psi_cc_binary( + name = "inspire_flow_smoke", + srcs = ["inspire_flow_smoke.cc"], + linkstatic = True, + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":internal", + ":params", + "//psi/algorithm/ypir:ypir_internal_params", + ], +) + +psi_cc_test( + name = "inspire_flow_test", + srcs = ["inspire_flow_smoke.cc"], + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":internal", + ":params", + "//psi/algorithm/ypir:ypir_internal_params", + ], +) + +psi_cc_binary( + name = "large_record_flow_smoke", + srcs = ["large_record_flow_smoke.cc"], + linkstatic = True, + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":large_record", + ":types", + ], +) + +psi_cc_binary( + name = "large_record_full_test", + srcs = ["large_record_full_test.cc"], + linkstatic = True, + target_compatible_with = X86_64_COMPATIBLE, + deps = [ + ":large_record", + ":types", + ], +) diff --git a/psi/algorithm/inspire/client.cc b/psi/algorithm/inspire/client.cc new file mode 100644 index 00000000..443fe2e1 --- /dev/null +++ b/psi/algorithm/inspire/client.cc @@ -0,0 +1,61 @@ +#include "psi/algorithm/inspire/client.h" + +#include +#include + +#include "yacl/base/exception.h" + +#include "psi/algorithm/inspire/serialize.h" +#include "psi/algorithm/ypir/ypir_internal_params.h" + +namespace psi::inspire { + +InspireClient::InspireClient(InspireParameters params) + : params_(std::move(params)), + context_(std::make_unique( + psi::ypir::internal::ypir::CreateContext(params_))) { + YACL_ENFORCE(params_.mode == psi::ypir::YpirMode::kDoublepir); +} + +yacl::Buffer InspireClient::GeneratePksBuffer() const { + return yacl::Buffer(); +} + +std::string InspireClient::GeneratePksString() const { return {}; } + +InspireQuery InspireClient::GenerateQuery(uint64_t raw_idx) const { + YACL_ENFORCE_LT(raw_idx, params_.NumItems()); + return internal::GenerateQuery(raw_idx, params_, client_secrets_, *context_); +} + +yacl::Buffer InspireClient::GenerateQueryBuffer(uint64_t raw_idx) const { + return SerializeQuery(GenerateQuery(raw_idx)); +} + +yacl::Buffer InspireClient::GenerateIndexQuery(uint64_t raw_idx) const { + return GenerateQueryBuffer(raw_idx); +} + +std::string InspireClient::GenerateIndexQueryStr(uint64_t raw_idx) const { + auto buffer = GenerateQueryBuffer(raw_idx); + return std::string(static_cast(buffer)); +} + +std::vector InspireClient::DecodeResponse( + const InspireResponse& response, uint64_t raw_idx) const { + YACL_ENFORCE_LT(raw_idx, params_.NumItems()); + return internal::RecoverResponse(response, params_, client_secrets_, + *context_); +} + +std::vector InspireClient::DecodeResponseBuffer( + const yacl::ByteContainerView& response_buffer, uint64_t raw_idx) const { + return DecodeResponse(DeserializeResponse(response_buffer), raw_idx); +} + +std::vector InspireClient::DecodeIndexResponse( + const yacl::ByteContainerView& response_buffer, uint64_t raw_idx) const { + return DecodeResponseBuffer(response_buffer, raw_idx); +} + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/client.h b/psi/algorithm/inspire/client.h new file mode 100644 index 00000000..acd8ed01 --- /dev/null +++ b/psi/algorithm/inspire/client.h @@ -0,0 +1,47 @@ +#pragma once + +#include +#include +#include + +#include "yacl/base/buffer.h" +#include "yacl/base/byte_container_view.h" + +#include "psi/algorithm/inspire/internal.h" +#include "psi/algorithm/inspire/params.h" +#include "psi/algorithm/inspire/types.h" +#include "psi/algorithm/pir_interface/index_pir.h" + +namespace psi::inspire { + +class InspireClient : public psi::pir::IndexPirClient { + public: + explicit InspireClient(InspireParameters params); + + const InspireParameters& GetParameters() const { return params_; } + + pir::PirType GetPirType() const override { return pir::PirType::YPIR_PIR; } + + yacl::Buffer GeneratePksBuffer() const override; + std::string GeneratePksString() const override; + + InspireQuery GenerateQuery(uint64_t raw_idx) const; + yacl::Buffer GenerateQueryBuffer(uint64_t raw_idx) const; + yacl::Buffer GenerateIndexQuery(uint64_t raw_idx) const override; + std::string GenerateIndexQueryStr(uint64_t raw_idx) const override; + + std::vector DecodeResponse(const InspireResponse& response, + uint64_t raw_idx) const; + std::vector DecodeResponseBuffer( + const yacl::ByteContainerView& response_buffer, uint64_t raw_idx) const; + std::vector DecodeIndexResponse( + const yacl::ByteContainerView& response_buffer, + uint64_t raw_idx) const override; + + private: + InspireParameters params_; + std::unique_ptr context_; + mutable internal::ClientSecrets client_secrets_; +}; + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/entry.cc b/psi/algorithm/inspire/entry.cc new file mode 100644 index 00000000..04aecf0c --- /dev/null +++ b/psi/algorithm/inspire/entry.cc @@ -0,0 +1,222 @@ +#include "psi/algorithm/inspire/entry.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +#include "absl/strings/escaping.h" +#include "fmt/format.h" +#include "spdlog/spdlog.h" +#include "yacl/base/buffer.h" +#include "yacl/base/byte_container_view.h" +#include "yacl/base/exception.h" + +#include "psi/algorithm/inspire/client.h" +#include "psi/algorithm/inspire/params.h" +#include "psi/algorithm/inspire/server.h" +#include "psi/algorithm/pir_interface/pir_db.h" + +namespace psi::inspire { +namespace { + +constexpr char kQueryCountTag[] = "inspire/query_count"; + +std::string TrimAscii(std::string_view text) { + size_t begin = 0; + size_t end = text.size(); + while (begin < end && + std::isspace(static_cast(text[begin])) != 0) { + ++begin; + } + while (end > begin && + std::isspace(static_cast(text[end - 1])) != 0) { + --end; + } + return std::string(text.substr(begin, end - begin)); +} + +uint8_t ParseHexNibble(char ch) { + if (ch >= '0' && ch <= '9') { + return static_cast(ch - '0'); + } + if (ch >= 'a' && ch <= 'f') { + return static_cast(10 + ch - 'a'); + } + if (ch >= 'A' && ch <= 'F') { + return static_cast(10 + ch - 'A'); + } + YACL_THROW("invalid hex character: {}", ch); +} + +std::vector ParseHexBytes(std::string text, size_t value_bytes) { + if (text.size() >= 2 && text[0] == '0' && + (text[1] == 'x' || text[1] == 'X')) { + text = text.substr(2); + } + YACL_ENFORCE_EQ(text.size(), value_bytes * 2, + "hex value width mismatch, expect {} bytes", value_bytes); + std::vector out(value_bytes, 0); + for (size_t i = 0; i < value_bytes; ++i) { + out[i] = static_cast((ParseHexNibble(text[2 * i]) << 4) | + ParseHexNibble(text[2 * i + 1])); + } + return out; +} + +std::vector ReadLines(const std::string& path) { + std::ifstream input(path); + YACL_ENFORCE(input.is_open(), "failed to open file: {}", path); + + std::vector lines; + std::string line; + while (std::getline(input, line)) { + auto trimmed = TrimAscii(line); + if (!trimmed.empty()) { + lines.push_back(std::move(trimmed)); + } + } + return lines; +} + +void WriteLines(const std::string& path, + const std::vector& lines) { + std::ofstream output(path, std::ios::out | std::ios::trunc); + YACL_ENFORCE(output.is_open(), "failed to open output file: {}", path); + for (const auto& line : lines) { + output << line << '\n'; + } +} + +template +yacl::Buffer SerializeScalar(T value) { + yacl::Buffer out(sizeof(T)); + std::memcpy(out.data(), &value, sizeof(T)); + return out; +} + +template +T DeserializeScalar(const yacl::ByteContainerView& buffer) { + YACL_ENFORCE_EQ(buffer.size(), sizeof(T)); + T out = 0; + std::memcpy(&out, buffer.data(), sizeof(T)); + return out; +} + +std::string QueryTag(uint64_t idx) { + return fmt::format("inspire/query/{}", idx); +} + +std::string ResponseTag(uint64_t idx) { + return fmt::format("inspire/response/{}", idx); +} + +InspireParameters BuildParameters(uint64_t db_rows, uint64_t db_cols, + uint64_t item_size_bits) { + YACL_ENFORCE_GT(db_rows, 0U); + YACL_ENFORCE_GT(db_cols, 0U); + YACL_ENFORCE_EQ(item_size_bits, 8U, + "Inspire launcher currently supports 8-bit values"); + return CreateParamsForShape(db_rows, db_cols, item_size_bits); +} + +std::vector> LoadDatabaseRows(const std::string& db_file, + size_t value_bytes, + uint64_t num_items) { + const auto lines = ReadLines(db_file); + YACL_ENFORCE_EQ(lines.size(), num_items, + "db_file must contain exactly {} values", num_items); + + std::vector> rows; + rows.reserve(lines.size()); + for (const auto& line : lines) { + rows.push_back(ParseHexBytes(line, value_bytes)); + } + return rows; +} + +std::vector LoadQueryIndices(const std::string& query_file, + uint64_t num_items) { + const auto lines = ReadLines(query_file); + std::vector indices; + indices.reserve(lines.size()); + for (const auto& line : lines) { + uint64_t raw_idx = std::stoull(line); + YACL_ENFORCE_LT(raw_idx, num_items, "query index out of range"); + indices.push_back(raw_idx); + } + return indices; +} + +std::vector EncodeOutputLines( + const std::vector>& values) { + std::vector out; + out.reserve(values.size()); + for (const auto& value : values) { + out.push_back(absl::BytesToHexString(absl::string_view( + reinterpret_cast(value.data()), value.size()))); + } + return out; +} + +} // namespace + +int SenderOnline(const InspireSenderOptions& options, + std::shared_ptr lctx) { + YACL_ENFORCE(lctx != nullptr, "link context is required for Inspire sender"); + const auto params = + BuildParameters(options.db_rows, options.db_cols, options.item_size_bits); + + SPDLOG_INFO("Starting Inspire sender, db_rows={}, db_cols={}", params.db_rows, + params.db_cols); + + InspireServer server(params); + server.GenerateFromRawData(psi::pir::RawDatabase(LoadDatabaseRows( + options.db_file, params.value_bytes, params.NumItems()))); + + lctx->ConnectToMesh(); + const uint64_t query_count = + DeserializeScalar(lctx->Recv(lctx->NextRank(), kQueryCountTag)); + for (uint64_t idx = 0; idx < query_count; ++idx) { + auto query = lctx->Recv(lctx->NextRank(), QueryTag(idx)); + auto response = server.Response(query, yacl::Buffer()); + lctx->Send(lctx->NextRank(), response, ResponseTag(idx)); + } + return 0; +} + +int ReceiverOnline(const InspireReceiverOptions& options, + std::shared_ptr lctx) { + YACL_ENFORCE(lctx != nullptr, "link context is required for Inspire receiver"); + const auto params = + BuildParameters(options.db_rows, options.db_cols, options.item_size_bits); + const auto query_indices = + LoadQueryIndices(options.query_file, params.NumItems()); + + SPDLOG_INFO("Starting Inspire receiver, query_count={}", + query_indices.size()); + + InspireClient client(params); + lctx->ConnectToMesh(); + lctx->Send(lctx->NextRank(), SerializeScalar(query_indices.size()), + kQueryCountTag); + + std::vector> decoded_values; + decoded_values.reserve(query_indices.size()); + for (size_t idx = 0; idx < query_indices.size(); ++idx) { + const auto raw_idx = query_indices[idx]; + auto query = client.GenerateIndexQuery(raw_idx); + lctx->Send(lctx->NextRank(), query, QueryTag(idx)); + auto response = lctx->Recv(lctx->NextRank(), ResponseTag(idx)); + decoded_values.push_back(client.DecodeIndexResponse(response, raw_idx)); + } + + WriteLines(options.output_file, EncodeOutputLines(decoded_values)); + return 0; +} + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/entry.h b/psi/algorithm/inspire/entry.h new file mode 100644 index 00000000..2eb0994e --- /dev/null +++ b/psi/algorithm/inspire/entry.h @@ -0,0 +1,32 @@ +#pragma once + +#include +#include +#include + +#include "yacl/link/context.h" + +namespace psi::inspire { + +struct InspireSenderOptions { + uint64_t db_rows = 0; + uint64_t db_cols = 0; + uint64_t item_size_bits = 0; + std::string db_file; +}; + +struct InspireReceiverOptions { + uint64_t db_rows = 0; + uint64_t db_cols = 0; + uint64_t item_size_bits = 0; + std::string query_file; + std::string output_file; +}; + +int SenderOnline(const InspireSenderOptions& options, + std::shared_ptr lctx); + +int ReceiverOnline(const InspireReceiverOptions& options, + std::shared_ptr lctx); + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/inspire_flow_smoke.cc b/psi/algorithm/inspire/inspire_flow_smoke.cc new file mode 100644 index 00000000..ed5b21d4 --- /dev/null +++ b/psi/algorithm/inspire/inspire_flow_smoke.cc @@ -0,0 +1,41 @@ +#include +#include + +#include "psi/algorithm/inspire/internal.h" +#include "psi/algorithm/inspire/params.h" +#include "psi/algorithm/ypir/ypir_internal_params.h" + +int main() { + auto params = psi::inspire::CreateSmallTestParams(); + auto client_context = psi::ypir::internal::ypir::CreateContext(params); + auto server_context = psi::ypir::internal::ypir::CreateContext(params); + psi::inspire::internal::ClientSecrets secrets; + + std::vector db(params.db_rows * params.db_cols, 0); + for (uint64_t row = 0; row < params.db_rows; ++row) { + for (uint64_t col = 0; col < params.db_cols; ++col) { + db[row * params.db_cols + col] = static_cast((row + col) % 251); + } + } + + const auto state = + psi::inspire::internal::PrepareOfflineState(db, params, server_context); + const uint64_t row = 111; + const uint64_t col = 222; + const uint64_t raw_idx = row * params.db_cols + col; + + const auto query = psi::inspire::internal::GenerateQuery( + raw_idx, params, secrets, client_context); + const auto response = psi::inspire::internal::ProcessQuery( + db, query, state, params, server_context); + const auto decoded = psi::inspire::internal::RecoverResponse( + response, params, secrets, client_context); + + const uint8_t expected = static_cast((row + col) % 251); + if (decoded.size() != 1 || decoded[0] != expected) { + std::fprintf(stderr, "decoded mismatch: got %u expected %u\n", + decoded.empty() ? 0 : decoded[0], expected); + return 1; + } + return 0; +} diff --git a/psi/algorithm/inspire/inspire_flow_test.sh b/psi/algorithm/inspire/inspire_flow_test.sh new file mode 100644 index 00000000..d7dc3cf6 --- /dev/null +++ b/psi/algorithm/inspire/inspire_flow_test.sh @@ -0,0 +1,4 @@ +#!/usr/bin/env bash +set -euo pipefail + +"${TEST_SRCDIR}/_main/psi/algorithm/inspire/inspire_flow_smoke" diff --git a/psi/algorithm/inspire/internal.cc b/psi/algorithm/inspire/internal.cc new file mode 100644 index 00000000..d89eb61f --- /dev/null +++ b/psi/algorithm/inspire/internal.cc @@ -0,0 +1,609 @@ +#include "psi/algorithm/inspire/internal.h" + +#include +#include +#include +#include +#include +#include + +#include "yacl/base/exception.h" + +#include "psi/algorithm/ypir/legacy/client.h" +#include "psi/algorithm/ypir/legacy/hexl.h" +#include "psi/algorithm/ypir/legacy/server.h" +#include "psi/algorithm/ypir/legacy/ypir_params.h" +#include "psi/algorithm/ypir/legacy/ypir_util.h" +#include "psi/algorithm/ypir/util.h" + +namespace psi::inspire::internal { +namespace { + +using psi::ypir::YpirMode; +using psi::ypir::ypir_internal::AESCTR_PRNG; +using psi::ypir::ypir_internal::Cdks21Lwe2RlweInplace; +using psi::ypir::ypir_internal::EltwiseAddMod; +using psi::ypir::ypir_internal::EltwiseFMAMod; +using psi::ypir::ypir_internal::EltwiseMultMod; +using psi::ypir::ypir_internal::EltwiseSubMod; +using psi::ypir::ypir_internal::FheParams; +using psi::ypir::ypir_internal::GetLog2; +using psi::ypir::ypir_internal::IsPowerOfTwo; +using psi::ypir::ypir_internal::kSecondDimensionSeed; +using psi::ypir::ypir_internal::kThirdDimensionSeed; +using psi::ypir::ypir_internal::MatrixMultiplicationFlat; +using psi::ypir::ypir_internal::MatrixMultiplicationFlatU16; +using psi::ypir::ypir_internal::MatrixRowDecompose; +using psi::ypir::ypir_internal::MatrixTranspose; +using psi::ypir::ypir_internal::MatrixVectorFirstDimension; +using psi::ypir::ypir_internal::MatrixVectorMultiplicationU16; +using psi::ypir::ypir_internal::MatVecU8U32Mod2p32; +using psi::ypir::ypir_internal::ModInverse; +using psi::ypir::ypir_internal::PirParams; +using psi::ypir::ypir_internal::PrecomputeAutomap; +using psi::ypir::ypir_internal::PseudorandomMatrixGenerate; +using psi::ypir::ypir_internal::SampleGauss; +using psi::ypir::ypir_internal::Secret; +using psi::ypir::ypir_internal::VectorColDecompose; +using psi::ypir::ypir_internal::YpirHexlNtt; + +uint64_t MakeSeedMaterial() { + const uint64_t time_seed = static_cast( + std::chrono::high_resolution_clock::now().time_since_epoch().count()); + std::random_device rd; + return time_seed ^ (static_cast(rd()) << 32) ^ rd(); +} + +std::mt19937_64 MakePrng() { + return std::mt19937_64(MakeSeedMaterial()); +} + +uint64_t GetBase(uint64_t b, uint64_t z, uint64_t ti) { + return 1ULL << (b + ti * z); +} + +uint64_t GetInspiringGenerator(uint64_t gamma, uint64_t degree) { + if (gamma == degree) { + return 5; + } + return ((degree << 1) / gamma) + 1; +} + +void PrecomputeAutomaps(uint32_t length, const std::vector& idx, + std::vector>& automaps) { + automaps.assign(idx.size(), std::vector(length, 0)); + for (uint64_t i = 0; i < idx.size(); ++i) { + PrecomputeAutomap(length, static_cast(idx[i]), automaps[i]); + } +} + +void ApplyAutoNttFormGeneral(const std::vector& vec, + std::vector& result, + const std::vector& automap) { + const uint64_t length = vec.size(); + if (result.size() != length) { + result.resize(length); + } + if (vec.data() == result.data()) { + std::vector scratch(length, 0); + for (uint64_t i = 0; i < length; ++i) { + scratch[i] = vec[automap[i]]; + } + result.swap(scratch); + return; + } + for (uint64_t i = 0; i < length; ++i) { + result[i] = vec[automap[i]]; + } +} + +void RotateKskRows(const std::vector>& in, + const std::vector& automap, + std::vector>& out) { + const uint64_t rows = in.size(); + const uint64_t degree = rows == 0 ? 0 : in[0].size(); + out.assign(rows, std::vector(degree, 0)); + for (uint64_t i = 0; i < rows; ++i) { + ApplyAutoNttFormGeneral(in[i], out[i], automap); + } +} + +void BuildNttMonomials(uint64_t degree, uint64_t count, YpirHexlNtt& ntt, + std::vector>& monomials_ntt) { + monomials_ntt.assign(count, std::vector(degree, 0)); + for (uint64_t i = 0; i < count; ++i) { + monomials_ntt[i][i] = 1; + ntt.Forward(monomials_ntt[i].data(), degree); + } +} + +void ApplyAutoCoefForm(std::vector& result, + const std::vector& input, int32_t index, + uint64_t modulus) { + const uint64_t length = input.size(); + result.assign(length, 0); + for (uint64_t i = 0; i < length; ++i) { + uint64_t destination = (i * index) % (2 * length); + if (destination >= length) { + result[destination - length] = (modulus - input[i]) % modulus; + } else { + result[destination] = input[i]; + } + } +} + +void ApproximateGadgetDecomp(const std::vector& poly, + std::vector>& mat, + uint64_t b, uint64_t z, uint64_t t) { + const uint64_t degree = poly.size(); + const uint64_t mask = (1ULL << z) - 1; + mat.assign(t, std::vector(degree, 0)); + for (uint64_t i = 0; i < degree; ++i) { + uint64_t val = poly[i] >> b; + for (uint64_t j = 0; j < t; ++j) { + mat[j][i] = val & mask; + val >>= z; + } + } +} + +void KeyswitchPreprocess( + const std::vector>& ksk_a, + std::vector& a_in, std::vector& a_out, + std::vector>>& decomp_buf, + const FheParams& fparm) { + const uint64_t length = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + const uint64_t t = fparm.get_t_auto(); + YpirHexlNtt& ntt = fparm.get_ntt(); + + std::vector> decomp_a; + ntt.Inverse(a_in.data(), length); + ApproximateGadgetDecomp(a_in, decomp_a, fparm.get_b_auto(), + fparm.get_z_auto(), t); + for (uint64_t i = 0; i < t; ++i) { + ntt.Forward(decomp_a[i].data(), length); + } + + decomp_buf.push_back(std::move(decomp_a)); + a_out.assign(length, 0); + std::vector tmp(length, 0); + for (uint64_t i = 0; i < t; ++i) { + EltwiseMultMod(tmp.data(), ksk_a[i].data(), decomp_buf.back()[i].data(), + length, modulus); + EltwiseSubMod(a_out.data(), a_out.data(), tmp.data(), length, modulus); + } +} + +void KeyswitchOnline(const std::vector>& ksk_b, + std::vector& b_in, std::vector& b_out, + const std::vector>& decomp_buf, + const FheParams& fparm) { + const uint64_t length = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + const uint64_t t = fparm.get_t_auto(); + b_out = b_in; + std::vector tmp(length, 0); + for (uint64_t i = 0; i < t; ++i) { + EltwiseMultMod(tmp.data(), ksk_b[i].data(), decomp_buf[i].data(), length, + modulus); + EltwiseSubMod(b_out.data(), b_out.data(), tmp.data(), length, modulus); + } +} + +void LweEncrypt(Secret& sk, const std::vector& a, uint64_t message, + uint64_t& b, uint64_t pmod, double sig) { + const uint64_t lwe_dimension = sk.get_len(); + const uint64_t cmod = sk.get_mod(); + const long double delta = + static_cast(cmod) / static_cast(pmod); + + auto rng = MakePrng(); + const uint64_t e = SampleGauss(sig, cmod, rng); + std::vector tmp(lwe_dimension, 0); + EltwiseMultMod(tmp.data(), a.data(), sk.data.data(), lwe_dimension, cmod); + b = 0; + for (uint64_t i = 0; i < lwe_dimension; ++i) { + b = (b + tmp[i]) % cmod; + } + b = (b + e) % cmod; + b = (b + static_cast(message * delta)) % cmod; +} + +void RlweEncode(Secret& sk, const std::vector& a, + std::vector& message, std::vector& b, + const FheParams& fparm) { + const uint64_t poly_degree = fparm.get_poly_degree(); + const uint64_t rlwe_cmod = fparm.get_rlwe_cmod(); + YpirHexlNtt& ntt = fparm.get_ntt(); + if (!sk.get_ntt_form()) { + ntt.Forward(sk.data.data(), poly_degree); + sk.switch_ntt_format(); + } + + std::vector err(poly_degree, 0); + auto rng = MakePrng(); + SampleGauss(err, fparm.get_sig_ring(), rlwe_cmod, rng); + EltwiseFMAMod(err.data(), message.data(), 1, err.data(), poly_degree, + rlwe_cmod); + ntt.Forward(err.data(), poly_degree); + EltwiseMultMod(b.data(), sk.data.data(), a.data(), poly_degree, rlwe_cmod); + EltwiseAddMod(b.data(), err.data(), b.data(), poly_degree, rlwe_cmod); +} + +void GadgetEncrypt(Secret& sk, const std::vector>& a, + std::vector& message, + std::vector>& b, + const FheParams& fparm) { + const uint64_t t_auto = fparm.get_t_auto(); + const uint64_t poly_degree = fparm.get_poly_degree(); + const uint64_t rlwe_cmod = fparm.get_rlwe_cmod(); + std::vector message_tmp(poly_degree, 0); + for (uint64_t i = 0; i < t_auto; ++i) { + const uint64_t base = + GetBase(fparm.get_b_auto(), fparm.get_z_auto(), i) % rlwe_cmod; + EltwiseFMAMod(message_tmp.data(), message.data(), base, nullptr, + poly_degree, rlwe_cmod); + RlweEncode(sk, a[i], message_tmp, b[i], fparm); + } +} + +void GenerateInspiringPartialKey( + Secret& sk, uint64_t gamma, std::vector>& ksk_b, + AESCTR_PRNG& prng, const FheParams& fparm) { + const uint64_t degree = fparm.get_poly_degree(); + const uint64_t cmod = fparm.get_rlwe_cmod(); + const uint64_t t_auto = fparm.get_t_auto(); + const uint64_t generator = GetInspiringGenerator(gamma, degree); + YpirHexlNtt& ntt = fparm.get_ntt(); + + if (sk.get_ntt_form()) { + ntt.Inverse(sk.data.data(), degree); + sk.switch_ntt_format(); + } + + std::vector> ksk_a( + t_auto, std::vector(degree, 0)); + ksk_b.assign(t_auto, std::vector(degree, 0)); + prng.refresh(kThirdDimensionSeed); + PseudorandomMatrixGenerate(ksk_a, cmod, prng); + + std::vector newkey(degree, 0); + const std::vector copykey = sk.data; + ApplyAutoCoefForm(newkey, copykey, static_cast(generator), cmod); + GadgetEncrypt(sk, ksk_a, newkey, ksk_b, fparm); +} + +void DoublepirQuery(uint64_t c_idx, uint64_t r_idx, Secret& simple_sk, + Secret& double_sk, std::vector& qu0, + std::vector& qu1, AESCTR_PRNG& prng, + const FheParams& fparm, const PirParams& pparm) { + const uint64_t cols = pparm.get_col(); + qu0.resize(cols); + const auto& matrix0 = fparm.get_persudo_matrix_simplepir(); + for (uint64_t i = 0; i < cols; ++i) { + LweEncrypt(simple_sk, matrix0[i], i == c_idx ? 1 : 0, qu0[i], + fparm.get_lwe_pmod(), fparm.get_sig()); + } + + const uint64_t rows = pparm.get_row(); + const uint64_t degree = fparm.get_poly_degree(); + qu1.resize(rows); + std::vector> matrix1(rows, + std::vector(degree, 0)); + prng.refresh(kSecondDimensionSeed); + PseudorandomMatrixGenerate(matrix1, fparm.get_rlwe_cmod(), prng); + for (uint64_t i = 0; i < rows; ++i) { + LweEncrypt(double_sk, matrix1[i], i == r_idx ? 1 : 0, qu1[i], + fparm.get_rlwe_pmod(), fparm.get_sig_ring()); + } +} + +std::vector InspiringPartialPreprocess( + const std::vector>& a_ct_tilde, uint64_t gamma, + const std::vector>& ksk_a, + std::vector>>& decomp_buf, + const FheParams& fparm) { + const uint64_t degree = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + YpirHexlNtt& ntt = fparm.get_ntt(); + if (gamma == 1) { + return a_ct_tilde[0]; + } + YACL_ENFORCE_LE(gamma, degree >> 1); + YACL_ENFORCE(IsPowerOfTwo(gamma)); + + const uint64_t generator = GetInspiringGenerator(gamma, degree); + std::vector generator_powers(gamma, 1); + for (uint64_t i = 1; i < gamma; ++i) { + generator_powers[i] = + static_cast((static_cast( + generator_powers[i - 1]) * + generator) % + (degree << 1)); + } + + std::vector> automaps; + PrecomputeAutomaps(static_cast(degree), generator_powers, automaps); + + std::vector> monomials_ntt; + BuildNttMonomials(degree, gamma, ntt, monomials_ntt); + + const uint64_t mod_inv = ModInverse(static_cast(gamma), + static_cast(modulus)); + std::vector> scaled_a = a_ct_tilde; + for (uint64_t i = 0; i < gamma; ++i) { + EltwiseFMAMod(scaled_a[i].data(), scaled_a[i].data(), mod_inv, nullptr, + degree, modulus); + } + + std::vector> aggregated( + gamma, std::vector(degree, 0)); + std::vector rotated(degree, 0); + std::vector shifted(degree, 0); + for (uint64_t k = 0; k < gamma; ++k) { + for (uint64_t j = 0; j < gamma; ++j) { + ApplyAutoNttFormGeneral(scaled_a[k], rotated, automaps[j]); + EltwiseMultMod(shifted.data(), rotated.data(), monomials_ntt[k].data(), + degree, modulus); + EltwiseAddMod(aggregated[j].data(), aggregated[j].data(), shifted.data(), + degree, modulus); + } + } + + std::vector> current = std::move(aggregated); + std::vector> rotated_ksk_a; + for (uint64_t level = gamma - 1; level > 0; --level) { + RotateKskRows(ksk_a, automaps[level - 1], rotated_ksk_a); + std::vector source = current[level]; + std::vector switched(degree, 0); + KeyswitchPreprocess(rotated_ksk_a, source, switched, decomp_buf, fparm); + EltwiseAddMod(current[level - 1].data(), current[level - 1].data(), + switched.data(), degree, modulus); + } + return current[0]; +} + +std::vector InspiringPartialOnline( + const std::vector& b_values, uint64_t gamma, + const std::vector>& ksk_b, + const std::vector>>& decomp_buf, + const FheParams& fparm) { + const uint64_t degree = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + YpirHexlNtt& ntt = fparm.get_ntt(); + + std::vector packed_b(degree, 0); + for (uint64_t i = 0; i < gamma; ++i) { + packed_b[i] = b_values[i] % modulus; + } + ntt.Forward(packed_b.data(), degree); + if (gamma == 1) { + return packed_b; + } + YACL_ENFORCE_LE(gamma, degree >> 1); + YACL_ENFORCE(IsPowerOfTwo(gamma)); + YACL_ENFORCE_EQ(decomp_buf.size(), gamma - 1); + + const uint64_t generator = GetInspiringGenerator(gamma, degree); + std::vector generator_powers(gamma, 1); + for (uint64_t i = 1; i < gamma; ++i) { + generator_powers[i] = + static_cast((static_cast( + generator_powers[i - 1]) * + generator) % + (degree << 1)); + } + std::vector> automaps; + PrecomputeAutomaps(static_cast(degree), generator_powers, automaps); + + std::vector> rotated_ksk_b; + for (uint64_t level = gamma - 1, step = 0; level > 0; --level, ++step) { + RotateKskRows(ksk_b, automaps[level - 1], rotated_ksk_b); + std::vector switched(degree, 0); + KeyswitchOnline(rotated_ksk_b, packed_b, switched, decomp_buf[step], + fparm); + packed_b.swap(switched); + } + return packed_b; +} + +void DoublepirHintGenerate( + std::vector>& db, + std::vector>& double_server_hint, + std::vector>& double_client_hint, + const FheParams& fparm) { + const uint64_t lwe_dimension = fparm.get_lwe_dimension(); + const uint64_t poly_degree = fparm.get_poly_degree(); + const uint64_t rlwe_cmod = fparm.get_rlwe_cmod(); + + const uint64_t* matrix_flat = fparm.get_persudo_matrix_simplepir_flat(); + std::vector> simple_hint( + db.size(), std::vector(lwe_dimension, 0)); + MatrixMultiplicationFlat(simple_hint, db, matrix_flat, lwe_dimension, + fparm.get_lwe_cmod()); + + std::vector> matrix_decomp; + MatrixRowDecompose(simple_hint, matrix_decomp, fparm.get_b_decomp(), + fparm.get_z_decomp(), fparm.get_t_decomp()); + MatrixTranspose(matrix_decomp, double_server_hint); + + const uint64_t* matrix_d2_flat = fparm.get_persudo_matrix_doublepir_flat(); + MatrixMultiplicationFlat(double_client_hint, double_server_hint, + matrix_d2_flat, poly_degree, rlwe_cmod); +} + +void DoublepirAnswer(const uint8_t* db, uint32_t* qu0, + std::vector& qu1, + const std::vector>& server_hint, + std::vector& h2, + std::vector>& h3, + std::vector& h4, const FheParams& fparm, + const PirParams& pparm) { + std::vector simple_res(pparm.get_row(), 0); + MatVecU8U32Mod2p32(db, qu0, simple_res.data(), pparm.get_row(), + pparm.get_col()); + + std::vector> trans_simple_res_decomp; + VectorColDecompose(simple_res, trans_simple_res_decomp, + fparm.get_b_decomp(), fparm.get_z_decomp(), + fparm.get_t_decomp()); + + MatrixVectorFirstDimension(h2, server_hint, qu1, fparm.get_rlwe_cmod()); + MatrixMultiplicationFlatU16( + h3, trans_simple_res_decomp, + fparm.get_persudo_matrix_doublepir_flat(), fparm.get_poly_degree(), + fparm.get_rlwe_cmod()); + MatrixVectorMultiplicationU16(h4, trans_simple_res_decomp, qu1, + fparm.get_rlwe_cmod()); +} + +std::vector> ExpandDatabase(const std::vector& db, + uint64_t rows, + uint64_t cols) { + YACL_ENFORCE_EQ(db.size(), rows * cols); + std::vector> out(rows, std::vector(cols, 0)); + for (uint64_t row = 0; row < rows; ++row) { + for (uint64_t col = 0; col < cols; ++col) { + out[row][col] = db[row * cols + col]; + } + } + return out; +} + +} // namespace + +InspireQuery GenerateQuery(uint64_t raw_idx, + const psi::ypir::YpirParameters& params, + ClientSecrets& secrets, + const psi::ypir::internal::ypir::Context& context) { + YACL_ENFORCE(params.mode == YpirMode::kDoublepir); + YACL_ENFORCE_LT(raw_idx, params.NumItems()); + + const uint64_t row_idx = raw_idx / params.db_cols; + const uint64_t col_idx = raw_idx % params.db_cols; + secrets.simple_secret = + Secret(context.fhe_params->get_lwe_dimension(), + context.fhe_params->get_lwe_cmod()); + secrets.double_secret = + Secret(context.fhe_params->get_poly_degree(), + context.fhe_params->get_rlwe_cmod()); + secrets.initialized = true; + + InspireQuery query; + DoublepirQuery(col_idx, row_idx, secrets.simple_secret, secrets.double_secret, + query.qu0, query.qu1, *context.prng, *context.fhe_params, + *context.pir_params); + GenerateInspiringPartialKey( + secrets.double_secret, + context.fhe_params->get_lwe_dimension() * + context.fhe_params->get_t_decomp(), + query.ksk_b, *context.prng, *context.fhe_params); + return query; +} + +std::vector RecoverResponse( + const InspireResponse& response, const psi::ypir::YpirParameters& params, + const ClientSecrets& secrets, + const psi::ypir::internal::ypir::Context& context) { + YACL_ENFORCE(params.mode == YpirMode::kDoublepir); + YACL_ENFORCE(secrets.initialized, + "GenerateQuery must be called before decode"); + auto simple_secret = secrets.simple_secret; + auto double_secret = secrets.double_secret; + auto result = response.doublepir_response; + uint64_t message = 0; + psi::ypir::ypir_internal::YpirRecover( + simple_secret, double_secret, result, message, *context.fhe_params, + *context.pir_params); + return psi::ypir::EncodeIntegerValue(message, params.value_bytes); +} + +InspirePrecomputedState PrepareOfflineState( + const std::vector& db, const psi::ypir::YpirParameters& params, + const psi::ypir::internal::ypir::Context& context) { + YACL_ENFORCE(params.mode == YpirMode::kDoublepir); + auto db_matrix = ExpandDatabase(db, params.db_rows, params.db_cols); + + InspirePrecomputedState state; + std::vector> double_server_hint; + const uint64_t pack_num = + context.fhe_params->get_lwe_dimension() * + context.fhe_params->get_t_decomp(); + std::vector> double_client_hint( + pack_num, + std::vector(context.fhe_params->get_poly_degree(), 0)); + DoublepirHintGenerate(db_matrix, double_server_hint, double_client_hint, + *context.fhe_params); + + std::vector> ksk_a( + context.fhe_params->get_t_auto(), + std::vector(context.fhe_params->get_poly_degree(), 0)); + context.prng->refresh(kThirdDimensionSeed); + PseudorandomMatrixGenerate(ksk_a, context.fhe_params->get_rlwe_cmod(), + *context.prng); + + YpirHexlNtt& ntt = context.fhe_params->get_ntt(); + for (uint64_t i = 0; i < pack_num; ++i) { + Cdks21Lwe2RlweInplace(double_client_hint[i].data(), + context.fhe_params->get_poly_degree(), + context.fhe_params->get_rlwe_cmod(), ntt); + } + + state.decomp_buf.clear(); + state.hint_0 = InspiringPartialPreprocess( + double_client_hint, pack_num, ksk_a, state.decomp_buf, + *context.fhe_params); + MatrixTranspose(double_server_hint, state.server_hint); + return state; +} + +InspireResponse ProcessQuery( + const std::vector& db, const InspireQuery& query, + const InspirePrecomputedState& state, + const psi::ypir::YpirParameters& params, + const psi::ypir::internal::ypir::Context& context) { + YACL_ENFORCE(params.mode == YpirMode::kDoublepir); + YACL_ENFORCE_EQ(context.fhe_params->get_t_decomp(), 1ULL, + "Inspire currently requires t_decomp == 1"); + + std::vector qu0(query.qu0.size(), 0); + for (size_t i = 0; i < query.qu0.size(); ++i) { + qu0[i] = static_cast(query.qu0[i]); + } + + auto qu1 = query.qu1; + std::vector> result; + result.push_back(state.hint_0); + + const uint64_t lwe_dimension = context.fhe_params->get_lwe_dimension(); + const uint64_t poly_degree = context.fhe_params->get_poly_degree(); + const uint64_t pack_num = + lwe_dimension * context.fhe_params->get_t_decomp(); + std::vector> h3( + context.fhe_params->get_t_decomp(), + std::vector(poly_degree, 0)); + std::vector h2(pack_num, 0); + std::vector h4(context.fhe_params->get_t_decomp(), 0); + DoublepirAnswer(db.data(), qu0.data(), qu1, state.server_hint, h2, h3, h4, + *context.fhe_params, *context.pir_params); + + result.push_back(InspiringPartialOnline(h2, pack_num, query.ksk_b, + state.decomp_buf, + *context.fhe_params)); + + std::vector pack_b(poly_degree, 0); + pack_b[0] = h4[0]; + context.fhe_params->get_ntt().Forward(pack_b.data(), poly_degree); + Cdks21Lwe2RlweInplace(h3[0].data(), poly_degree, + context.fhe_params->get_rlwe_cmod(), + context.fhe_params->get_ntt()); + result.push_back(h3[0]); + result.push_back(pack_b); + + InspireResponse response; + response.doublepir_response = std::move(result); + return response; +} + +} // namespace psi::inspire::internal diff --git a/psi/algorithm/inspire/internal.h b/psi/algorithm/inspire/internal.h new file mode 100644 index 00000000..0c8f183f --- /dev/null +++ b/psi/algorithm/inspire/internal.h @@ -0,0 +1,37 @@ +#pragma once + +#include + +#include "psi/algorithm/inspire/types.h" +#include "psi/algorithm/ypir/params.h" +#include "psi/algorithm/ypir/ypir_internal_params.h" + +namespace psi::inspire::internal { + +struct ClientSecrets { + psi::ypir::ypir_internal::Secret simple_secret; + psi::ypir::ypir_internal::Secret double_secret; + bool initialized = false; +}; + +InspireQuery GenerateQuery(uint64_t raw_idx, + const psi::ypir::YpirParameters& params, + ClientSecrets& secrets, + const psi::ypir::internal::ypir::Context& context); + +std::vector RecoverResponse( + const InspireResponse& response, const psi::ypir::YpirParameters& params, + const ClientSecrets& secrets, + const psi::ypir::internal::ypir::Context& context); + +InspirePrecomputedState PrepareOfflineState( + const std::vector& db, const psi::ypir::YpirParameters& params, + const psi::ypir::internal::ypir::Context& context); + +InspireResponse ProcessQuery( + const std::vector& db, const InspireQuery& query, + const InspirePrecomputedState& state, + const psi::ypir::YpirParameters& params, + const psi::ypir::internal::ypir::Context& context); + +} // namespace psi::inspire::internal diff --git a/psi/algorithm/inspire/internal_primitives.cc b/psi/algorithm/inspire/internal_primitives.cc new file mode 100644 index 00000000..4de1d1e6 --- /dev/null +++ b/psi/algorithm/inspire/internal_primitives.cc @@ -0,0 +1,238 @@ +#include "psi/algorithm/inspire/internal_primitives.h" + +#include +#include +#include +#include +#include + +#include "yacl/base/exception.h" + +#include "psi/algorithm/ypir/legacy/ypir_util.h" + +namespace psi::inspire::internal { + +namespace { + +using psi::ypir::ypir_internal::kThirdDimensionSeed; +using psi::ypir::ypir_internal::PseudorandomMatrixGenerate; + +uint64_t MakeSeedMaterial() { + const uint64_t time_seed = static_cast( + std::chrono::high_resolution_clock::now().time_since_epoch().count()); + std::random_device rd; + return time_seed ^ (static_cast(rd()) << 32) ^ rd(); +} + +std::mt19937_64 MakePrng() { return std::mt19937_64(MakeSeedMaterial()); } + +} // namespace + +uint64_t GetBase(uint64_t b, uint64_t z, uint64_t ti) { + return 1ULL << (b + ti * z); +} + +uint64_t GetInspiringGenerator(uint64_t gamma, uint64_t degree) { + if (gamma == degree) { + return 5; + } + return ((degree << 1) / gamma) + 1; +} + +void PrecomputeAutomaps(uint32_t length, + const std::vector& idx, + std::vector>& automaps) { + automaps.assign(idx.size(), std::vector(length, 0)); + for (uint64_t i = 0; i < idx.size(); ++i) { + psi::ypir::ypir_internal::PrecomputeAutomap( + length, static_cast(idx[i]), automaps[i]); + } +} + +void ApplyAutoCoefForm(std::vector& result, + const std::vector& input, int32_t index, + uint64_t modulus) { + const uint64_t length = input.size(); + result.assign(length, 0); + for (uint64_t i = 0; i < length; ++i) { + uint64_t destination = (i * static_cast(index)) % (2 * length); + if (destination >= length) { + result[destination - length] = (modulus - input[i]) % modulus; + } else { + result[destination] = input[i]; + } + } +} + +void ApplyAutoNttFormGeneral(const std::vector& vec, + std::vector& result, + const std::vector& automap) { + const uint64_t length = vec.size(); + if (result.size() != length) { + result.resize(length); + } + if (vec.data() == result.data()) { + std::vector scratch(length, 0); + for (uint64_t i = 0; i < length; ++i) { + scratch[i] = vec[automap[i]]; + } + result.swap(scratch); + return; + } + for (uint64_t i = 0; i < length; ++i) { + result[i] = vec[automap[i]]; + } +} + +void RotateKskRows(const std::vector>& in, + const std::vector& automap, + std::vector>& out) { + const uint64_t rows = in.size(); + const uint64_t degree = rows == 0 ? 0 : in[0].size(); + out.assign(rows, std::vector(degree, 0)); + for (uint64_t i = 0; i < rows; ++i) { + ApplyAutoNttFormGeneral(in[i], out[i], automap); + } +} + +void BuildNttMonomials(uint64_t degree, uint64_t count, YpirHexlNtt& ntt, + std::vector>& monomials_ntt) { + monomials_ntt.assign(count, std::vector(degree, 0)); + for (uint64_t i = 0; i < count; ++i) { + monomials_ntt[i][i] = 1; + ntt.Forward(monomials_ntt[i].data(), degree); + } +} + +void ApproximateGadgetDecomp(const std::vector& poly, + std::vector>& mat, + uint64_t b, uint64_t z, uint64_t t) { + const uint64_t degree = poly.size(); + const uint64_t mask = (1ULL << z) - 1; + mat.assign(t, std::vector(degree, 0)); + for (uint64_t i = 0; i < degree; ++i) { + uint64_t val = poly[i] >> b; + for (uint64_t j = 0; j < t; ++j) { + mat[j][i] = val & mask; + val >>= z; + } + } +} + +void KeyswitchPreprocess( + const std::vector>& ksk_a, + std::vector& a_in, std::vector& a_out, + std::vector>>& decomp_buf, + const FheParams& fparm) { + const uint64_t length = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + const uint64_t t = fparm.get_t_auto(); + YpirHexlNtt& ntt = fparm.get_ntt(); + + std::vector> decomp_a; + ntt.Inverse(a_in.data(), length); + ApproximateGadgetDecomp(a_in, decomp_a, fparm.get_b_auto(), + fparm.get_z_auto(), t); + for (uint64_t i = 0; i < t; ++i) { + ntt.Forward(decomp_a[i].data(), length); + } + + decomp_buf.push_back(std::move(decomp_a)); + a_out.assign(length, 0); + std::vector tmp(length, 0); + for (uint64_t i = 0; i < t; ++i) { + EltwiseMultMod(tmp.data(), ksk_a[i].data(), decomp_buf.back()[i].data(), + length, modulus); + EltwiseSubMod(a_out.data(), a_out.data(), tmp.data(), length, modulus); + } +} + +void KeyswitchOnline(const std::vector>& ksk_b, + std::vector& b_in, + std::vector& b_out, + const std::vector>& decomp_buf, + const FheParams& fparm) { + const uint64_t length = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + const uint64_t t = fparm.get_t_auto(); + b_out = b_in; + std::vector tmp(length, 0); + for (uint64_t i = 0; i < t; ++i) { + EltwiseMultMod(tmp.data(), ksk_b[i].data(), decomp_buf[i].data(), length, + modulus); + EltwiseSubMod(b_out.data(), b_out.data(), tmp.data(), length, modulus); + } +} + +void LweEncrypt(Secret& sk, const std::vector& a, uint64_t message, + uint64_t& b, uint64_t pmod, double sig) { + const uint64_t lwe_dimension = sk.get_len(); + const uint64_t cmod = sk.get_mod(); + const long double delta = + static_cast(cmod) / static_cast(pmod); + + auto rng = MakePrng(); + const uint64_t e = SampleGauss(sig, cmod, rng); + std::vector tmp(lwe_dimension, 0); + EltwiseMultMod(tmp.data(), a.data(), sk.data.data(), lwe_dimension, cmod); + b = 0; + for (uint64_t i = 0; i < lwe_dimension; ++i) { + b = (b + tmp[i]) % cmod; + } + b = (b + e) % cmod; + b = (b + static_cast(message * delta)) % cmod; +} + +void RlweEncode(Secret& sk, const std::vector& a, + std::vector& message, std::vector& b, + const FheParams& fparm) { + const uint64_t poly_degree = fparm.get_poly_degree(); + const uint64_t rlwe_cmod = fparm.get_rlwe_cmod(); + YpirHexlNtt& ntt = fparm.get_ntt(); + if (!sk.get_ntt_form()) { + ntt.Forward(sk.data.data(), poly_degree); + sk.switch_ntt_format(); + } + + std::vector err(poly_degree, 0); + auto rng = MakePrng(); + SampleGauss(err, fparm.get_sig_ring(), rlwe_cmod, rng); + EltwiseFMAMod(err.data(), message.data(), 1, err.data(), poly_degree, + rlwe_cmod); + ntt.Forward(err.data(), poly_degree); + EltwiseMultMod(b.data(), sk.data.data(), a.data(), poly_degree, rlwe_cmod); + EltwiseAddMod(b.data(), err.data(), b.data(), poly_degree, rlwe_cmod); +} + +void GadgetEncrypt(Secret& sk, + const std::vector>& a, + std::vector& message, + std::vector>& b, + const FheParams& fparm) { + const uint64_t t_auto = fparm.get_t_auto(); + const uint64_t poly_degree = fparm.get_poly_degree(); + const uint64_t rlwe_cmod = fparm.get_rlwe_cmod(); + std::vector message_tmp(poly_degree, 0); + for (uint64_t i = 0; i < t_auto; ++i) { + const uint64_t base = + GetBase(fparm.get_b_auto(), fparm.get_z_auto(), i) % rlwe_cmod; + EltwiseFMAMod(message_tmp.data(), message.data(), base, nullptr, + poly_degree, rlwe_cmod); + RlweEncode(sk, a[i], message_tmp, b[i], fparm); + } +} + +void SimplepirQuery(uint64_t idx, Secret& sk, std::vector& qu, + const FheParams& fparm, const PirParams& pparm) { + const uint64_t length = pparm.get_col(); + const uint64_t lwe_pmod = fparm.get_lwe_pmod(); + const float sig = fparm.get_sig(); + qu.resize(length, 0); + + const auto& matrix = fparm.get_persudo_matrix_simplepir(); + for (uint64_t i = 0; i < length; ++i) { + LweEncrypt(sk, matrix[i], i == idx ? 1 : 0, qu[i], lwe_pmod, sig); + } +} + +} // namespace psi::inspire::internal diff --git a/psi/algorithm/inspire/internal_primitives.h b/psi/algorithm/inspire/internal_primitives.h new file mode 100644 index 00000000..c19a9594 --- /dev/null +++ b/psi/algorithm/inspire/internal_primitives.h @@ -0,0 +1,78 @@ +#pragma once + +#include +#include +#include + +#include "psi/algorithm/ypir/legacy/hexl.h" +#include "psi/algorithm/ypir/legacy/ypir_params.h" + +namespace psi::inspire::internal { + +using psi::ypir::ypir_internal::AESCTR_PRNG; +using psi::ypir::ypir_internal::EltwiseAddMod; +using psi::ypir::ypir_internal::EltwiseFMAMod; +using psi::ypir::ypir_internal::EltwiseMultMod; +using psi::ypir::ypir_internal::EltwiseSubMod; +using psi::ypir::ypir_internal::FheParams; +using psi::ypir::ypir_internal::PirParams; +using psi::ypir::ypir_internal::SampleGauss; +using psi::ypir::ypir_internal::Secret; +using psi::ypir::ypir_internal::YpirHexlNtt; + +uint64_t GetBase(uint64_t b, uint64_t z, uint64_t ti); + +uint64_t GetInspiringGenerator(uint64_t gamma, uint64_t degree); + +void PrecomputeAutomaps(uint32_t length, + const std::vector& idx, + std::vector>& automaps); + +void ApplyAutoCoefForm(std::vector& result, + const std::vector& input, int32_t index, + uint64_t modulus); + +void ApplyAutoNttFormGeneral(const std::vector& vec, + std::vector& result, + const std::vector& automap); + +void RotateKskRows(const std::vector>& in, + const std::vector& automap, + std::vector>& out); + +void BuildNttMonomials(uint64_t degree, uint64_t count, YpirHexlNtt& ntt, + std::vector>& monomials_ntt); + +void ApproximateGadgetDecomp(const std::vector& poly, + std::vector>& mat, + uint64_t b, uint64_t z, uint64_t t); + +void KeyswitchPreprocess( + const std::vector>& ksk_a, + std::vector& a_in, std::vector& a_out, + std::vector>>& decomp_buf, + const FheParams& fparm); + +void KeyswitchOnline(const std::vector>& ksk_b, + std::vector& b_in, + std::vector& b_out, + const std::vector>& decomp_buf, + const FheParams& fparm); + +void LweEncrypt(Secret& sk, const std::vector& a, uint64_t message, + uint64_t& b, uint64_t pmod, double sig); + +void RlweEncode(Secret& sk, const std::vector& a, + std::vector& message, std::vector& b, + const FheParams& fparm); + +void GadgetEncrypt(Secret& sk, + const std::vector>& a, + std::vector& message, + std::vector>& b, + const FheParams& fparm); + +void SimplepirQuery(uint64_t idx, Secret& sk, std::vector& qu, + const FheParams& fparm, const PirParams& pparm); + +} // namespace psi::inspire::internal diff --git a/psi/algorithm/inspire/large_record.cc b/psi/algorithm/inspire/large_record.cc new file mode 100644 index 00000000..ce2f16b2 --- /dev/null +++ b/psi/algorithm/inspire/large_record.cc @@ -0,0 +1,612 @@ +#include "psi/algorithm/inspire/large_record.h" + +#include +#include +#include +#include +#include +#include + +#include "yacl/base/exception.h" + +#include "psi/algorithm/inspire/internal_primitives.h" +#include "psi/algorithm/inspire/large_record_he.h" +#include "psi/algorithm/ypir/legacy/ypir_util.h" + +using namespace psi::inspire::lrhe; + +namespace psi::inspire::internal { + +namespace { + +using psi::ypir::ypir_internal::Cdks21Lwe2RlweInplace; +using psi::ypir::ypir_internal::EltwiseAddMod; +using psi::ypir::ypir_internal::EltwiseFMAMod; +using psi::ypir::ypir_internal::EltwiseMultMod; +using psi::ypir::ypir_internal::EltwiseSubMod; +using psi::ypir::ypir_internal::IsPowerOfTwo; +using psi::ypir::ypir_internal::MatrixMultiplicationFlatU16; +using psi::ypir::ypir_internal::MatrixVectorMultiplicationU16; +using psi::ypir::ypir_internal::ModInverse; +using psi::ypir::ypir_internal::PirParams; +using psi::ypir::ypir_internal::PseudorandomMatrixGenerate; +using psi::ypir::ypir_internal::Secret; +using psi::ypir::ypir_internal::YpirHexlNtt; +using psi::ypir::ypir_internal::kThirdDimensionSeed; + +using Record = std::array; + +uint64_t ModAdd(uint64_t lhs, uint64_t rhs, uint64_t mod) { + uint64_t sum = lhs + rhs; + if (sum >= mod) { + sum -= mod; + } + return sum; +} + +uint64_t ModMul(uint64_t lhs, uint64_t rhs, uint64_t mod) { + return static_cast((static_cast<__uint128_t>(lhs) * rhs) % mod); +} + +void NegacyclicMonomialMul(const std::vector& input, uint64_t shift, + uint64_t mod, std::vector& output) { + const uint64_t degree = input.size(); + std::fill(output.begin(), output.end(), 0); + for (uint64_t i = 0; i < degree; ++i) { + uint64_t exponent = i + shift; + bool neg = false; + while (exponent >= degree) { + exponent -= degree; + neg = !neg; + } + uint64_t value = input[i] % mod; + if (!neg) { + output[exponent] = ModAdd(output[exponent], value, mod); + } else if (value != 0) { + output[exponent] = (output[exponent] + mod - value) % mod; + } + } +} + +std::vector> InterpolateColumn( + const std::vector>& values, + const LargeRecordConfig& cfg) { + const uint64_t row = cfg.folding; + const uint64_t degree = cfg.degree; + const uint64_t mod = cfg.pmod; + const uint64_t step = (degree << 1) / row; + const uint64_t row_inv = static_cast( + ModInverse(static_cast(row), static_cast(mod))); + + std::vector> coeffs( + row, std::vector(degree, 0)); + std::vector rotated(degree, 0); + + for (uint64_t j = 0; j < row; ++j) { + for (uint64_t r = 0; r < row; ++r) { + uint64_t power = (j * r) % row; + uint64_t shift = ((row - power) % row) * step; + NegacyclicMonomialMul(values[r], shift, mod, rotated); + for (uint64_t k = 0; k < degree; ++k) { + coeffs[j][k] = ModAdd(coeffs[j][k], rotated[k], mod); + } + } + for (uint64_t k = 0; k < degree; ++k) { + coeffs[j][k] = ModMul(coeffs[j][k], row_inv, mod); + } + } + return coeffs; +} + +uint64_t GetBundleFactor(const LargeRecordConfig& cfg) { + return cfg.packed_coeffs / LargeRecordConfig::kRecordBytes; +} + +Record GenerateRecord(uint64_t record_idx) { + Record record{}; + for (uint64_t b = 0; b < LargeRecordConfig::kRecordBytes; ++b) { + record[b] = static_cast((17 * record_idx + b) & 0xFFULL); + } + return record; +} + +std::vector BuildBundledEntry(uint64_t bundle_idx, + const LargeRecordConfig& cfg) { + const uint64_t bundle_factor = GetBundleFactor(cfg); + std::vector entry(cfg.degree, 0); + for (uint64_t slot = 0; slot < bundle_factor; ++slot) { + Record record = GenerateRecord(bundle_idx * bundle_factor + slot); + uint64_t base = slot * LargeRecordConfig::kRecordBytes; + for (uint64_t b = 0; b < LargeRecordConfig::kRecordBytes; ++b) { + entry[base + b] = static_cast(record[b]); + } + } + return entry; +} + +void EncodeDatabaseLargeRecord(const LargeRecordConfig& cfg, + std::vector>& encoded_db) { + const uint64_t bundle_factor = GetBundleFactor(cfg); + const uint64_t bundled_entries = cfg.total_records / bundle_factor; + const uint64_t t = cfg.folding; + const uint64_t cols = bundled_entries / t; + const uint64_t rows = t * cfg.packed_coeffs; + + encoded_db.assign(rows, std::vector(cols, 0)); + +#pragma omp parallel for + for (uint64_t col = 0; col < cols; ++col) { + std::vector> values( + t, std::vector(cfg.degree, 0)); + for (uint64_t row = 0; row < t; ++row) { + values[row] = BuildBundledEntry(col * t + row, cfg); + } + std::vector> coeffs = InterpolateColumn(values, cfg); + for (uint64_t k = 0; k < t; ++k) { + for (uint64_t i = 0; i < cfg.packed_coeffs; ++i) { + encoded_db[k * cfg.packed_coeffs + i][col] = + static_cast(coeffs[k][i] % cfg.pmod); + } + } + } +} + +void BuildPaperPackKeyMaterial( + Secret& sk, AESCTR_PRNG& prng, + const psi::ypir::ypir_internal::FheParams& fparm, + PaperPackKeyMaterial& key_material) { + const uint64_t degree = fparm.get_poly_degree(); + const uint64_t cmod = fparm.get_rlwe_cmod(); + const uint64_t t_auto = fparm.get_t_auto(); + const uint64_t half = degree / 2; + const uint64_t g = 5; + const uint64_t h = (degree << 1) - 1; + + std::vector> wg( + t_auto, std::vector(degree, 0)); + std::vector> yg( + t_auto, std::vector(degree, 0)); + std::vector> wh( + t_auto, std::vector(degree, 0)); + std::vector> yh( + t_auto, std::vector(degree, 0)); + + prng.refresh(kThirdDimensionSeed); + PseudorandomMatrixGenerate(wg, cmod, prng); + PseudorandomMatrixGenerate(wh, cmod, prng); + + std::vector sk_copy = sk.data; + std::vector g_key(degree, 0); + std::vector h_key(degree, 0); + ApplyAutoCoefForm(g_key, sk_copy, static_cast(g), cmod); + ApplyAutoCoefForm(h_key, sk_copy, static_cast(h), cmod); + GadgetEncrypt(sk, wg, g_key, yg, fparm); + GadgetEncrypt(sk, wh, h_key, yh, fparm); + + key_material.g_automaps.assign(half, std::vector(degree, 0)); + key_material.gh_automaps.assign(half, std::vector(degree, 0)); + uint64_t g_power = 1; + for (uint64_t i = 0; i < half; ++i) { + psi::ypir::ypir_internal::PrecomputeAutomap( + static_cast(degree), static_cast(g_power), + key_material.g_automaps[i]); + psi::ypir::ypir_internal::PrecomputeAutomap( + static_cast(degree), + static_cast( + (static_cast<__uint128_t>(h) * g_power) % (degree << 1)), + key_material.gh_automaps[i]); + g_power = static_cast( + (static_cast<__uint128_t>(g_power) * g) % (degree << 1)); + } + + key_material.kg_a_powers.assign( + half - 1, std::vector>( + t_auto, std::vector(degree, 0))); + key_material.kg_b_powers.assign( + half - 1, std::vector>( + t_auto, std::vector(degree, 0))); + key_material.gh_kg_a_powers.assign( + half - 1, std::vector>( + t_auto, std::vector(degree, 0))); + key_material.gh_kg_b_powers.assign( + half - 1, std::vector>( + t_auto, std::vector(degree, 0))); + for (uint64_t i = 0; i + 1 < half; ++i) { + RotateKskRows(wg, key_material.g_automaps[i], + key_material.kg_a_powers[i]); + RotateKskRows(yg, key_material.g_automaps[i], + key_material.kg_b_powers[i]); + RotateKskRows(wg, key_material.gh_automaps[i], + key_material.gh_kg_a_powers[i]); + RotateKskRows(yg, key_material.gh_automaps[i], + key_material.gh_kg_b_powers[i]); + } + + key_material.kh_a = std::move(wh); + key_material.kh_b = std::move(yh); + BuildNttMonomials(degree, degree, fparm.get_ntt(), + key_material.monomials_ntt); +} + +void AggregatePaperTransformHalves( + const std::vector>& block_rows, + const PaperPackKeyMaterial& key_material, + const psi::ypir::ypir_internal::FheParams& fparm, + std::vector>& first_half, + std::vector>& second_half) { + const uint64_t degree = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + const uint64_t half = degree / 2; + const uint64_t degree_inv = static_cast( + ModInverse(static_cast(degree), static_cast(modulus))); + YpirHexlNtt& ntts = fparm.get_ntt(); + + first_half.assign(half, std::vector(degree, 0)); + second_half.assign(half, std::vector(degree, 0)); + + std::vector a_tilde(degree, 0); + std::vector rotated(degree, 0); + std::vector shifted(degree, 0); + + for (uint64_t k = 0; k < degree; ++k) { + a_tilde = block_rows[k]; + Cdks21Lwe2RlweInplace(a_tilde.data(), degree, modulus, ntts); + EltwiseFMAMod(a_tilde.data(), a_tilde.data(), degree_inv, nullptr, degree, + modulus); + + for (uint64_t j = 0; j < half; ++j) { + ApplyAutoNttFormGeneral(a_tilde, rotated, key_material.g_automaps[j]); + EltwiseMultMod(shifted.data(), rotated.data(), + key_material.monomials_ntt[k].data(), degree, modulus); + EltwiseAddMod(first_half[j].data(), first_half[j].data(), + shifted.data(), degree, modulus); + + ApplyAutoNttFormGeneral(a_tilde, rotated, key_material.gh_automaps[j]); + EltwiseMultMod(shifted.data(), rotated.data(), + key_material.monomials_ntt[k].data(), degree, modulus); + EltwiseAddMod(second_half[j].data(), second_half[j].data(), + shifted.data(), degree, modulus); + } + } +} + +void CollapseOnePreprocessPaper( + std::vector>& current, + const std::vector>& ksk_a, + std::vector>>& decomp_steps, + const psi::ypir::ypir_internal::FheParams& fparm) { + const uint64_t degree = fparm.get_poly_degree(); + const uint64_t modulus = fparm.get_rlwe_cmod(); + + std::vector switched(degree, 0); + std::vector>> local_decomp; + std::vector last = current.back(); + KeyswitchPreprocess(ksk_a, last, switched, local_decomp, fparm); + current.pop_back(); + EltwiseAddMod(current.back().data(), current.back().data(), + switched.data(), degree, modulus); + decomp_steps.push_back(std::move(local_decomp[0])); +} + +void CollapseHalfPreprocessPaper( + std::vector> current, + const std::vector>>& rotated_ksk_a, + std::vector>>& decomp_steps, + std::vector& final_a_ntt, + const psi::ypir::ypir_internal::FheParams& fparm) { + decomp_steps.clear(); + while (current.size() > 1) { + CollapseOnePreprocessPaper(current, rotated_ksk_a[current.size() - 2], + decomp_steps, fparm); + } + final_a_ntt = current[0]; +} + +void CollapseOneOnlinePaper( + std::vector& current_b_ntt, + const std::vector>& ksk_b, + std::vector>& decomp_step, + const psi::ypir::ypir_internal::FheParams& fparm) { + std::vector next_b(current_b_ntt.size(), 0); + KeyswitchOnline(ksk_b, current_b_ntt, next_b, decomp_step, fparm); + current_b_ntt.swap(next_b); +} + +void CollapseHalfOnlinePaper( + std::vector& current_b_ntt, + const std::vector>>& rotated_ksk_b, + const std::vector>>& decomp_steps, + const psi::ypir::ypir_internal::FheParams& fparm) { + const uint64_t num_steps = decomp_steps.size(); + for (uint64_t i = 0; i < num_steps; ++i) { + CollapseOneOnlinePaper( + current_b_ntt, rotated_ksk_b[num_steps - 1 - i], + const_cast>&>(decomp_steps[i]), + fparm); + } +} + +void PreprocessPackedHintsPaper( + const std::vector>& hint, + const PaperPackKeyMaterial& key_material, + std::vector& packed_hint_blocks, + const psi::ypir::ypir_internal::FheParams& fparm, + const LargeRecordConfig& cfg) { + packed_hint_blocks.resize(cfg.folding); + + for (uint64_t block = 0; block < cfg.folding; ++block) { + std::vector> block_rows( + cfg.packed_coeffs, std::vector(cfg.degree, 0)); + for (uint64_t i = 0; i < cfg.packed_coeffs; ++i) { + block_rows[i] = hint[block * cfg.packed_coeffs + i]; + } + + std::vector> first_half; + std::vector> second_half; + AggregatePaperTransformHalves(block_rows, key_material, fparm, first_half, + second_half); + + std::vector a1_ntt; + std::vector a2_ntt; + CollapseHalfPreprocessPaper( + std::move(first_half), key_material.kg_a_powers, + packed_hint_blocks[block].decomp_half1, a1_ntt, fparm); + CollapseHalfPreprocessPaper( + std::move(second_half), key_material.gh_kg_a_powers, + packed_hint_blocks[block].decomp_half2, a2_ntt, fparm); + + std::vector> final_pair; + final_pair.push_back(std::move(a1_ntt)); + final_pair.push_back(std::move(a2_ntt)); + std::vector>> final_decomp_steps; + CollapseOnePreprocessPaper(final_pair, key_material.kh_a, + final_decomp_steps, fparm); + packed_hint_blocks[block].packed_a_ntt = std::move(final_pair[0]); + packed_hint_blocks[block].decomp_final = std::move(final_decomp_steps[0]); + } +} + +RlweCiphertext PackSelectedBlockOnlinePaper( + uint64_t block_idx, const std::vector& selected_column_response, + const PaperPackKeyMaterial& key_material, + const std::vector& packed_hint_blocks, + const psi::ypir::ypir_internal::FheParams& fparm, + const LargeRecordConfig& cfg) { + std::vector b_agg(cfg.degree, 0); + for (uint64_t i = 0; i < cfg.packed_coeffs; ++i) { + b_agg[i] = + selected_column_response[block_idx * cfg.packed_coeffs + i] % cfg.qmod; + } + fparm.get_ntt().Forward(b_agg.data(), cfg.degree); + + CollapseHalfOnlinePaper(b_agg, key_material.kg_b_powers, + packed_hint_blocks[block_idx].decomp_half1, fparm); + CollapseHalfOnlinePaper(b_agg, key_material.gh_kg_b_powers, + packed_hint_blocks[block_idx].decomp_half2, fparm); + CollapseOneOnlinePaper( + b_agg, key_material.kh_b, + const_cast>&>( + packed_hint_blocks[block_idx].decomp_final), + fparm); + + RlweCiphertext ct(cfg.degree, cfg.qmod, true); + ct.a = Poly(packed_hint_blocks[block_idx].packed_a_ntt, cfg.qmod, true); + ct.b = Poly(b_agg, cfg.qmod, true); + return ct; +} + +RlweCiphertext AnswerLargeRecordOnlinePaper( + const std::vector& selected_column_response, + const PaperPackKeyMaterial& key_material, + const std::vector& packed_hint_blocks, + ApproximateRgswCiphertext& point_ct, intel::hexl::NTT& ntts, + const psi::ypir::ypir_internal::FheParams& fparm, + const LargeRecordConfig& cfg) { + std::vector coeff_cts; + coeff_cts.reserve(cfg.folding); + for (uint64_t k = 0; k < cfg.folding; ++k) { + coeff_cts.push_back(PackSelectedBlockOnlinePaper( + k, selected_column_response, key_material, packed_hint_blocks, fparm, + cfg)); + } + + RlweCiphertext acc = coeff_cts.back(); + for (int64_t j = static_cast(coeff_cts.size()) - 2; j >= 0; --j) { + RlweCiphertext mul_in = acc; + RlweCiphertext mul_out(acc.get_degree(), acc.get_modulus(), true); + ExternalProduct(point_ct, mul_in, mul_out, ntts); + RlweBfvAdd( + mul_out, const_cast(coeff_cts[static_cast(j)]), + acc); + } + return acc; +} + +ApproximateRgswCiphertext EncryptPointQuery(uint64_t row_idx, + lrhe::Secret& rlwe_sk, + const LargeRecordConfig& cfg) { + std::vector point_poly(cfg.degree, 0); + uint64_t shift = (row_idx % cfg.folding) * ((cfg.degree << 1) / cfg.folding); + uint64_t exponent = shift; + bool neg = false; + while (exponent >= cfg.degree) { + exponent -= cfg.degree; + neg = !neg; + } + point_poly[exponent] = neg ? (cfg.qmod - 1) : 1; + + Plaintext point_pt(point_poly, cfg.qmod, false); + ApproximateRgswCiphertext point_ct(cfg.degree, cfg.qmod, cfg.rgsw_b, + cfg.rgsw_z, cfg.rgsw_t); + RgswEncode(point_pt, point_ct, rlwe_sk, cfg.sigma); + return point_ct; +} + +Record ExtractRecordFromBundle(const std::vector& bundle, + uint64_t local_offset) { + uint64_t base = local_offset * LargeRecordConfig::kRecordBytes; + Record record{}; + for (uint64_t i = 0; i < LargeRecordConfig::kRecordBytes; ++i) { + record[i] = static_cast(bundle[base + i] & 0xFFULL); + } + return record; +} + +void SerializeRgswCt( + const ApproximateRgswCiphertext& ct, + std::vector>>& ct_m_out, + std::vector>>& ct_sm_out) { + ct_m_out.resize(ct.get_t()); + ct_sm_out.resize(ct.get_t()); + for (uint64_t i = 0; i < ct.get_t(); ++i) { + ct_m_out[i] = {ct.ct_m[i].a.get_data(), ct.ct_m[i].b.get_data()}; + ct_sm_out[i] = {ct.ct_sm[i].a.get_data(), ct.ct_sm[i].b.get_data()}; + } +} + +ApproximateRgswCiphertext DeserializeRgswCt( + const std::vector>>& ct_m_ser, + const std::vector>>& ct_sm_ser, + const LargeRecordConfig& cfg) { + ApproximateRgswCiphertext ct(cfg.degree, cfg.qmod, cfg.rgsw_b, cfg.rgsw_z, + cfg.rgsw_t); + for (uint64_t i = 0; i < ct.get_t(); ++i) { + ct.ct_m[i].a = Poly(ct_m_ser[i][0], cfg.qmod, true); + ct.ct_m[i].b = Poly(ct_m_ser[i][1], cfg.qmod, true); + ct.ct_sm[i].a = Poly(ct_sm_ser[i][0], cfg.qmod, true); + ct.ct_sm[i].b = Poly(ct_sm_ser[i][1], cfg.qmod, true); + } + return ct; +} + +} // namespace + +LargeRecordPrecomputedState PrepareLargeRecordState( + const LargeRecordConfig& cfg) { + const uint64_t bundle_factor = GetBundleFactor(cfg); + const uint64_t bundled_entries = cfg.total_records / bundle_factor; + const uint64_t cols = bundled_entries / cfg.folding; + const uint64_t encoded_rows = cfg.folding * cfg.packed_coeffs; + + psi::ypir::ypir_internal::FheParams fparm( + cfg.degree, cfg.qmod, cfg.pmod, cfg.lwe_dimension, cfg.lwe_cmod, + cfg.lwe_pmod, 0.0, 3.19, {20, 18, 2}, {16, 16, 1}); + [[maybe_unused]] PirParams pparm(encoded_rows, cols); + + const uint8_t key[16] = "I am xct's son"; + AESCTR_PRNG prng(key); + + srand(123); + std::vector sk_vec_signed(cfg.degree, 0); + std::vector sk_vec(cfg.degree, 0); + for (uint64_t i = 0; i < cfg.degree; ++i) { + sk_vec_signed[i] = rand() & 1; + sk_vec[i] = static_cast(sk_vec_signed[i]); + } + + Secret packing_sk(cfg.degree, cfg.qmod); + packing_sk.data = sk_vec; + + LargeRecordPrecomputedState state; + + EncodeDatabaseLargeRecord(cfg, state.encoded_db); + fparm.set_persudo_matrix_simplepir(cols); + + state.hint.assign(encoded_rows, + std::vector(cfg.lwe_dimension, 0)); + MatrixMultiplicationFlatU16( + state.hint, state.encoded_db, fparm.get_persudo_matrix_simplepir_flat(), + cfg.lwe_dimension, cfg.lwe_cmod); + + BuildPaperPackKeyMaterial(packing_sk, prng, fparm, state.key_material); + PreprocessPackedHintsPaper(state.hint, state.key_material, + state.packed_hint_blocks, fparm, cfg); + return state; +} + +LargeRecordQuery GenerateLargeRecordQuery( + uint64_t target_index, const LargeRecordConfig& cfg, + LargeRecordClientSecrets& secrets) { + const uint64_t bundle_factor = GetBundleFactor(cfg); + const uint64_t bundled_entries = cfg.total_records / bundle_factor; + const uint64_t cols = bundled_entries / cfg.folding; + const uint64_t encoded_rows = cfg.folding * cfg.packed_coeffs; + + psi::ypir::ypir_internal::FheParams fparm( + cfg.degree, cfg.qmod, cfg.pmod, cfg.lwe_dimension, cfg.lwe_cmod, + cfg.lwe_pmod, 0.0, 3.19, {20, 18, 2}, {16, 16, 1}); + [[maybe_unused]] PirParams pparm(encoded_rows, cols); + + const uint64_t bundle_idx = target_index / bundle_factor; + const uint64_t col_idx = bundle_idx / cfg.folding; + const uint64_t row_idx = bundle_idx % cfg.folding; + + srand(123); + secrets.sk_vec_signed.resize(cfg.degree, 0); + secrets.sk_vec.resize(cfg.degree, 0); + for (uint64_t i = 0; i < cfg.degree; ++i) { + secrets.sk_vec_signed[i] = rand() & 1; + secrets.sk_vec[i] = static_cast(secrets.sk_vec_signed[i]); + } + + Secret simple_sk(cfg.lwe_dimension, cfg.lwe_cmod); + simple_sk.data = secrets.sk_vec; + + fparm.set_persudo_matrix_simplepir(cols); + + LargeRecordQuery query; + SimplepirQuery(col_idx, simple_sk, query.column_query, fparm, pparm); + + lrhe::Secret rlwe_sk(secrets.sk_vec_signed, cfg.qmod, false); + auto point_ct = EncryptPointQuery(row_idx, rlwe_sk, cfg); + SerializeRgswCt(point_ct, query.rgsw_ct_m, query.rgsw_ct_sm); + + return query; +} + +LargeRecordResponse ProcessLargeRecordQuery( + const LargeRecordQuery& query, const LargeRecordPrecomputedState& state, + const LargeRecordConfig& cfg) { + psi::ypir::ypir_internal::FheParams fparm( + cfg.degree, cfg.qmod, cfg.pmod, cfg.lwe_dimension, cfg.lwe_cmod, + cfg.lwe_pmod, 0.0, 3.19, {20, 18, 2}, {16, 16, 1}); + + std::vector selected_column_response; + MatrixVectorMultiplicationU16(selected_column_response, state.encoded_db, + query.column_query, cfg.lwe_cmod); + + auto point_ct = DeserializeRgswCt(query.rgsw_ct_m, query.rgsw_ct_sm, cfg); + + auto response_ct = + AnswerLargeRecordOnlinePaper(selected_column_response, state.key_material, + state.packed_hint_blocks, point_ct, + fparm.get_ntt().Raw(), fparm, cfg); + + LargeRecordResponse response; + response.ct_a = response_ct.a.get_data(); + response.ct_b = response_ct.b.get_data(); + return response; +} + +std::vector RecoverLargeRecordResponse( + const LargeRecordResponse& response, uint64_t target_index, + const LargeRecordConfig& cfg, const LargeRecordClientSecrets& secrets) { + const uint64_t bundle_factor = GetBundleFactor(cfg); + const uint64_t local_offset = target_index % bundle_factor; + + lrhe::Secret rlwe_sk( + const_cast&>(secrets.sk_vec_signed), cfg.qmod, + false); + + RlweCiphertext ct(cfg.degree, cfg.qmod, true); + ct.a = Poly(response.ct_a, cfg.qmod, true); + ct.b = Poly(response.ct_b, cfg.qmod, true); + + Plaintext dec_pt(cfg.degree, cfg.pmod); + RlweBfvDecrypt(ct, dec_pt, rlwe_sk); + + Record recovered = ExtractRecordFromBundle(dec_pt.pt.get_data(), local_offset); + return std::vector(recovered.begin(), recovered.end()); +} + +} // namespace psi::inspire::internal diff --git a/psi/algorithm/inspire/large_record.h b/psi/algorithm/inspire/large_record.h new file mode 100644 index 00000000..381261a3 --- /dev/null +++ b/psi/algorithm/inspire/large_record.h @@ -0,0 +1,32 @@ +#pragma once + +#include +#include + +#include "psi/algorithm/inspire/types.h" + +namespace psi::inspire::internal { + +struct LargeRecordClientSecrets { + std::vector sk_vec_signed; + std::vector sk_vec; +}; + +LargeRecordPrecomputedState PrepareLargeRecordState( + const LargeRecordConfig& cfg); + +LargeRecordQuery GenerateLargeRecordQuery( + uint64_t target_index, const LargeRecordConfig& cfg, + LargeRecordClientSecrets& secrets); + +LargeRecordResponse ProcessLargeRecordQuery( + const LargeRecordQuery& query, + const LargeRecordPrecomputedState& state, + const LargeRecordConfig& cfg); + +std::vector RecoverLargeRecordResponse( + const LargeRecordResponse& response, uint64_t target_index, + const LargeRecordConfig& cfg, + const LargeRecordClientSecrets& secrets); + +} // namespace psi::inspire::internal diff --git a/psi/algorithm/inspire/large_record_flow_smoke.cc b/psi/algorithm/inspire/large_record_flow_smoke.cc new file mode 100644 index 00000000..fd2f194d --- /dev/null +++ b/psi/algorithm/inspire/large_record_flow_smoke.cc @@ -0,0 +1,70 @@ +#include +#include +#include + +#include "psi/algorithm/inspire/large_record.h" +#include "psi/algorithm/inspire/types.h" + +int main() { + psi::inspire::LargeRecordConfig cfg; + cfg.total_records = 1ULL << 16; + cfg.folding = 8; + cfg.degree = 2048; + cfg.packed_coeffs = 2048; + cfg.qmod = psi::inspire::LargeRecordConfig::kCrtMod; + cfg.pmod = 257; + cfg.sigma = 0.41f; + cfg.rgsw_b = 0; + cfg.rgsw_z = 8; + cfg.rgsw_t = 7; + cfg.lwe_dimension = 2048; + cfg.lwe_cmod = psi::inspire::LargeRecordConfig::kCrtMod; + cfg.lwe_pmod = 257; + + const uint64_t target_index = 12345; + const uint64_t kRecordBytes = psi::inspire::LargeRecordConfig::kRecordBytes; + + std::fprintf(stderr, "large_record_paper smoke test\n"); + std::fprintf(stderr, "total_records: %lu\n", cfg.total_records); + std::fprintf(stderr, "record_bytes: %lu\n", kRecordBytes); + std::fprintf(stderr, "target_index: %lu\n", target_index); + + std::vector expected(kRecordBytes, 0); + for (uint64_t b = 0; b < kRecordBytes; ++b) { + expected[b] = static_cast((17 * target_index + b) & 0xFFULL); + } + + auto state = psi::inspire::internal::PrepareLargeRecordState(cfg); + std::fprintf(stderr, "setup done\n"); + + psi::inspire::internal::LargeRecordClientSecrets secrets; + auto query = psi::inspire::internal::GenerateLargeRecordQuery( + target_index, cfg, secrets); + std::fprintf(stderr, "query done\n"); + + auto response = psi::inspire::internal::ProcessLargeRecordQuery( + query, state, cfg); + std::fprintf(stderr, "answer done\n"); + + auto recovered = psi::inspire::internal::RecoverLargeRecordResponse( + response, target_index, cfg, secrets); + std::fprintf(stderr, "recover done\n"); + + bool ok = recovered.size() == kRecordBytes; + if (ok) { + for (uint64_t i = 0; i < kRecordBytes; ++i) { + if (recovered[i] != expected[i]) { + ok = false; + std::fprintf(stderr, "mismatch at byte %lu: got %u expected %u\n", i, + recovered[i], expected[i]); + break; + } + } + } else { + std::fprintf(stderr, "size mismatch: got %zu expected %lu\n", + recovered.size(), kRecordBytes); + } + + std::fprintf(stderr, "full PIR recover check: %s\n", ok ? "ok" : "FAIL"); + return ok ? 0 : 1; +} diff --git a/psi/algorithm/inspire/large_record_full_test.cc b/psi/algorithm/inspire/large_record_full_test.cc new file mode 100644 index 00000000..d414a5c1 --- /dev/null +++ b/psi/algorithm/inspire/large_record_full_test.cc @@ -0,0 +1,56 @@ +#include +#include +#include + +#include "psi/algorithm/inspire/large_record.h" +#include "psi/algorithm/inspire/types.h" + +int main() { + psi::inspire::LargeRecordConfig cfg; + cfg.total_records = 1ULL << 20; + cfg.folding = 8; + cfg.degree = 2048; + cfg.packed_coeffs = 2048; + + const uint64_t target_index = 123456; + const uint64_t kRecordBytes = psi::inspire::LargeRecordConfig::kRecordBytes; + + std::fprintf(stderr, "=== Large Record Paper - Full 2^20 Test ===\n"); + std::fprintf(stderr, "total_records: %lu\n", cfg.total_records); + std::fprintf(stderr, "target_index: %lu\n", target_index); + + std::vector expected(kRecordBytes, 0); + for (uint64_t b = 0; b < kRecordBytes; ++b) { + expected[b] = static_cast((17 * target_index + b) & 0xFFULL); + } + + auto state = psi::inspire::internal::PrepareLargeRecordState(cfg); + std::fprintf(stderr, "setup done\n"); + + psi::inspire::internal::LargeRecordClientSecrets secrets; + auto query = psi::inspire::internal::GenerateLargeRecordQuery( + target_index, cfg, secrets); + std::fprintf(stderr, "query done\n"); + + auto response = psi::inspire::internal::ProcessLargeRecordQuery( + query, state, cfg); + std::fprintf(stderr, "answer done\n"); + + auto recovered = psi::inspire::internal::RecoverLargeRecordResponse( + response, target_index, cfg, secrets); + std::fprintf(stderr, "recover done\n"); + + bool ok = recovered.size() == kRecordBytes; + if (ok) { + for (uint64_t i = 0; i < kRecordBytes; ++i) { + if (recovered[i] != expected[i]) { + ok = false; + std::fprintf(stderr, "mismatch at byte %lu: got %u expected %u\n", i, + recovered[i], expected[i]); + break; + } + } + } + std::fprintf(stderr, "full PIR recover check: %s\n", ok ? "ok" : "FAIL"); + return ok ? 0 : 1; +} diff --git a/psi/algorithm/inspire/large_record_he.cc b/psi/algorithm/inspire/large_record_he.cc new file mode 100644 index 00000000..ce5a4e8b --- /dev/null +++ b/psi/algorithm/inspire/large_record_he.cc @@ -0,0 +1,353 @@ +#include "psi/algorithm/inspire/large_record_he.h" + +#include +#include +#include +#include +#include + +namespace psi::inspire::lrhe { + +namespace { + +void SampleRandom(Poly& a, uint64_t modulus, uint64_t length) { + for (uint64_t i = 0; i < length; ++i) { + uint64_t value = (static_cast(std::rand()) << 31) | + static_cast(std::rand()); + a[i] = value % modulus; + } +} + +void SampleGauss(std::vector& err, float st_dev, uint64_t modulus) { + std::default_random_engine engine( + std::chrono::system_clock::now().time_since_epoch().count()); + std::normal_distribution gaussian(0.0, st_dev); + for (size_t i = 0; i < err.size(); ++i) { + int64_t value = static_cast(std::llround(gaussian(engine))); + err[i] = (value < 0) ? (modulus - static_cast(-value)) + : static_cast(value); + } +} + +void ApproximateGadgetDecomp(Poly& poly, std::vector& mat, uint64_t b, + uint64_t z, uint64_t t) { + const uint64_t degree = poly.get_length(); + const uint64_t mask = (1ULL << z) - 1; + for (uint64_t i = 0; i < degree; ++i) { + for (uint64_t j = 0; j < t; ++j) { + mat[j][i] = (poly[i] >> (b + z * j)) & mask; + } + } +} + +void ApproximateGadgetDecomp(RlweCiphertext& ct_bfv, + std::vector& result_b, + std::vector& result_a, uint64_t b, + uint64_t z, uint64_t t) { + ApproximateGadgetDecomp(ct_bfv.b, result_b, b, z, t); + ApproximateGadgetDecomp(ct_bfv.a, result_a, b, z, t); +} + +} // namespace + +Poly::Poly() : length_(0), modulus_(0), nttform_(false) {} + +Poly::Poly(uint64_t len, uint64_t modulus, bool ntt) + : length_(len), modulus_(modulus), nttform_(ntt), payload(len, 0) {} + +Poly::Poly(const std::vector& vec, uint64_t modulus, bool ntt) + : length_(vec.size()), modulus_(modulus), nttform_(ntt), payload(vec) {} + +uint64_t Poly::get_length() const { return length_; } +uint64_t Poly::get_modulus() const { return modulus_; } +bool Poly::get_nttform() const { return nttform_; } +void Poly::set_nttform(bool ntt) { nttform_ = ntt; } +uint64_t* Poly::data() { return payload.data(); } +const uint64_t* Poly::data() const { return payload.data(); } +uint64_t& Poly::operator[](size_t idx) { return payload[idx]; } +const uint64_t& Poly::operator[](size_t idx) const { return payload[idx]; } +std::vector& Poly::get_data() { return payload; } +const std::vector& Poly::get_data() const { return payload; } + +void PolyMult(Poly& lhs, Poly& rhs, Poly& result) { + intel::hexl::EltwiseMultMod(result.data(), lhs.data(), rhs.data(), + result.get_length(), result.get_modulus(), 1); +} + +void PolyAdd(Poly& lhs, Poly& rhs, Poly& result) { + intel::hexl::EltwiseAddMod(result.data(), lhs.data(), rhs.data(), + result.get_length(), result.get_modulus()); +} + +void PolySub(Poly& lhs, Poly& rhs, Poly& result) { + intel::hexl::EltwiseSubMod(result.data(), lhs.data(), rhs.data(), + result.get_length(), result.get_modulus()); +} + +void PolyFMA(Poly& mul, uint64_t constant, Poly& result) { + intel::hexl::EltwiseFMAMod(result.data(), mul.data(), constant, nullptr, + result.get_length(), result.get_modulus(), 1); +} + +void PolyToNTT(Poly& poly, intel::hexl::NTT& ntts) { + if (poly.get_nttform()) { + return; + } + ntts.ComputeForward(poly.data(), poly.data(), 1, 1); + poly.set_nttform(true); +} + +void PolyToCoef(Poly& poly, intel::hexl::NTT& ntts) { + if (!poly.get_nttform()) { + return; + } + ntts.ComputeInverse(poly.data(), poly.data(), 1, 1); + poly.set_nttform(false); +} + +Secret::Secret(std::vector& vec, uint64_t modulus, bool ntt) { + std::vector tmp(vec.size(), 0); + for (size_t i = 0; i < vec.size(); ++i) { + tmp[i] = (vec[i] < 0) ? (modulus - static_cast(-vec[i])) + : static_cast(vec[i]); + } + data_ = Poly(tmp, modulus, ntt); + if (modulus == kLrCrtMod && vec.size() == 4096) { + ntts_ = intel::hexl::NTT(vec.size(), modulus, kLrRootOfUnityCrt4096); + } else if (modulus == kLrCrtMod && vec.size() == 2048) { + ntts_ = intel::hexl::NTT(vec.size(), modulus, kLrRootOfUnityCrt2048); + } else { + ntts_ = intel::hexl::NTT(vec.size(), modulus); + } + if (ntt) { + ntts_.ComputeForward(data_.data(), data_.data(), 1, 1); + data_.set_nttform(true); + } +} + +intel::hexl::NTT& Secret::get_ntt() { return ntts_; } +Poly& Secret::get_data() { return data_; } +const Poly& Secret::get_data() const { return data_; } +uint64_t Secret::get_modulus() const { return data_.get_modulus(); } +uint64_t Secret::get_length() const { return data_.get_length(); } +bool Secret::get_nttform() const { return data_.get_nttform(); } +void Secret::to_ntt_form() { PolyToNTT(data_, ntts_); } + +RlweCiphertext::RlweCiphertext() = default; + +RlweCiphertext::RlweCiphertext(uint64_t len, uint64_t modulus, bool ntt) + : a(len, modulus, ntt), b(len, modulus, ntt) {} + +uint64_t RlweCiphertext::get_degree() const { return a.get_length(); } +uint64_t RlweCiphertext::get_modulus() const { return a.get_modulus(); } +void RlweCiphertext::set_nttform(bool ntt) { + a.set_nttform(ntt); + b.set_nttform(ntt); +} + +Plaintext::Plaintext(uint64_t len, uint64_t pmod) : pt(len, pmod, false) {} + +Plaintext::Plaintext(const std::vector& vec, uint64_t pmod, bool ntt) + : pt(vec, pmod, ntt) {} + +uint64_t Plaintext::get_pmod() const { return pt.get_modulus(); } +bool Plaintext::get_isntt() const { return pt.get_nttform(); } + +ApproximateRgswCiphertext::ApproximateRgswCiphertext(uint64_t len, + uint64_t modulus, + uint64_t b, uint64_t z, + uint64_t t) + : degree_(len), + modulus_(modulus), + b_(b), + z_(z), + t_(t), + ct_m(t, RlweCiphertext(len, modulus, true)), + ct_sm(t, RlweCiphertext(len, modulus, true)) {} + +uint64_t ApproximateRgswCiphertext::get_degree() const { return degree_; } +uint64_t ApproximateRgswCiphertext::get_modulus() const { return modulus_; } +uint64_t ApproximateRgswCiphertext::get_b() const { return b_; } +uint64_t ApproximateRgswCiphertext::get_z() const { return z_; } +uint64_t ApproximateRgswCiphertext::get_t() const { return t_; } + +void RlweBfvEncrypt(Plaintext& pt, RlweCiphertext& ct, Secret& sk, float sig) { + const uint64_t modulus = sk.get_modulus(); + const uint64_t length = sk.get_length(); + const uint64_t pmod = pt.get_pmod(); + const uint64_t delta = static_cast( + std::floor(static_cast(modulus) / static_cast(pmod))); + + Poly tmp(length, modulus, false); + ct.set_nttform(true); + SampleRandom(ct.a, modulus, length); + + if (!sk.get_nttform()) { + sk.to_ntt_form(); + } + + intel::hexl::NTT& ntts = sk.get_ntt(); + if (pt.get_isntt()) { + PolyToCoef(pt.pt, ntts); + } + + PolyMult(ct.a, sk.get_data(), ct.b); + PolyFMA(pt.pt, delta, tmp); + + Poly err(length, modulus, false); + SampleGauss(err.get_data(), sig, modulus); + + PolyToCoef(ct.b, ntts); + PolyAdd(ct.b, err, ct.b); + PolyAdd(ct.b, tmp, ct.b); + PolyToNTT(ct.b, ntts); +} + +void RlweBfvEncode(Plaintext& pt, RlweCiphertext& ct, Secret& sk, float sig) { + const uint64_t modulus = sk.get_modulus(); + const uint64_t length = sk.get_length(); + + ct.set_nttform(true); + SampleRandom(ct.a, modulus, length); + if (!sk.get_nttform()) { + sk.to_ntt_form(); + } + + PolyMult(ct.a, sk.get_data(), ct.b); + intel::hexl::NTT& ntts = sk.get_ntt(); + if (pt.get_isntt()) { + PolyToCoef(pt.pt, ntts); + } + + Poly err(length, modulus, false); + SampleGauss(err.get_data(), sig, modulus); + + PolyToCoef(ct.b, ntts); + PolyAdd(ct.b, err, ct.b); + PolyAdd(ct.b, pt.pt, ct.b); + PolyToNTT(ct.b, ntts); +} + +void RlweBfvDecrypt(RlweCiphertext& ct, Plaintext& pt, Secret& sk) { + const uint64_t modulus = sk.get_modulus(); + const uint64_t pmod = pt.get_pmod(); + const uint64_t delta = static_cast( + std::floor(static_cast(modulus) / static_cast(pmod))); + + if (!sk.get_nttform()) { + sk.to_ntt_form(); + } + + Poly tmp(ct.get_degree(), modulus, true); + PolyMult(ct.a, sk.get_data(), tmp); + PolySub(ct.b, tmp, tmp); + intel::hexl::NTT& ntts = sk.get_ntt(); + PolyToCoef(tmp, ntts); + + for (uint64_t i = 0; i < ct.get_degree(); ++i) { + int64_t centered = (tmp.payload[i] > (modulus >> 1)) + ? static_cast(tmp.payload[i] - modulus) + : static_cast(tmp.payload[i]); + int64_t decoded = static_cast( + std::llround(static_cast(centered) / + static_cast(delta))); + decoded %= static_cast(pmod); + if (decoded < 0) { + decoded += static_cast(pmod); + } + pt.pt[i] = static_cast(decoded); + } +} + +void RlweBfvAdd(RlweCiphertext& add1, RlweCiphertext& add2, + RlweCiphertext& result) { + PolyAdd(add1.a, add2.a, result.a); + PolyAdd(add1.b, add2.b, result.b); +} + +void ExternalProduct(const ApproximateRgswCiphertext& ct_rgsw, + RlweCiphertext& ct_bfv, RlweCiphertext& result, + intel::hexl::NTT& ntts) { + const uint64_t length = ct_bfv.get_degree(); + const uint64_t modulus = ct_bfv.get_modulus(); + const uint64_t t = ct_rgsw.get_t(); + + std::vector ct_bfv_b(t, Poly(length, modulus)); + std::vector ct_bfv_a(t, Poly(length, modulus)); + + ct_bfv.b.set_nttform(true); + ct_bfv.a.set_nttform(true); + PolyToCoef(ct_bfv.b, ntts); + PolyToCoef(ct_bfv.a, ntts); + ApproximateGadgetDecomp(ct_bfv, ct_bfv_b, ct_bfv_a, ct_rgsw.get_b(), + ct_rgsw.get_z(), t); + + for (uint64_t i = 0; i < t; ++i) { + PolyToNTT(ct_bfv_b[i], ntts); + PolyToNTT(ct_bfv_a[i], ntts); + } + + RlweCiphertext tmp(length, modulus, true); + RlweCiphertext result1(length, modulus, true); + RlweCiphertext result2(length, modulus, true); + for (uint64_t i = 0; i < t; ++i) { + Poly ct_m_b = ct_rgsw.ct_m[i].b; + Poly ct_m_a = ct_rgsw.ct_m[i].a; + PolyMult(ct_bfv_b[i], ct_m_b, tmp.b); + PolyMult(ct_bfv_b[i], ct_m_a, tmp.a); + RlweBfvAdd(tmp, result1, result1); + } + for (uint64_t i = 0; i < t; ++i) { + Poly ct_sm_b = ct_rgsw.ct_sm[i].b; + Poly ct_sm_a = ct_rgsw.ct_sm[i].a; + PolyMult(ct_bfv_a[i], ct_sm_b, tmp.b); + PolyMult(ct_bfv_a[i], ct_sm_a, tmp.a); + RlweBfvAdd(tmp, result2, result2); + } + PolySub(result1.a, result2.a, result.a); + PolySub(result1.b, result2.b, result.b); +} + +void RgswEncode(Plaintext& pt, ApproximateRgswCiphertext& ct, Secret sk, + float sig) { + const uint64_t z_gsw = ct.get_z(); + const uint64_t t_gsw = ct.get_t(); + const uint64_t b_gsw = ct.get_b(); + const uint64_t length = ct.get_degree(); + + Poly sm(length, ct.get_modulus(), true); + + if (!sk.get_nttform()) { + sk.to_ntt_form(); + } + intel::hexl::NTT& ntts = sk.get_ntt(); + if (!pt.get_isntt()) { + PolyToNTT(pt.pt, ntts); + } + PolyMult(pt.pt, sk.get_data(), sm); + + PolyToCoef(sm, ntts); + PolyToCoef(pt.pt, ntts); + + RlweCiphertext tmp(length, ct.get_modulus(), true); + PolyFMA(pt.pt, (1ULL << b_gsw), pt.pt); + PolyFMA(sm, (1ULL << b_gsw), sm); + const uint64_t base = (1ULL << z_gsw); + for (uint64_t i = 0; i < t_gsw; ++i) { + RlweBfvEncode(pt, tmp, sk, sig); + tmp.b.set_nttform(true); + tmp.a.set_nttform(true); + ct.ct_m[i] = tmp; + + Plaintext pt_sm(sm.get_data(), ct.get_modulus(), false); + RlweBfvEncode(pt_sm, tmp, sk, sig); + tmp.b.set_nttform(true); + tmp.a.set_nttform(true); + ct.ct_sm[i] = tmp; + + PolyFMA(pt.pt, base, pt.pt); + PolyFMA(sm, base, sm); + } +} + +} // namespace psi::inspire::lrhe diff --git a/psi/algorithm/inspire/large_record_he.h b/psi/algorithm/inspire/large_record_he.h new file mode 100644 index 00000000..8a80841b --- /dev/null +++ b/psi/algorithm/inspire/large_record_he.h @@ -0,0 +1,123 @@ +#pragma once + +#include +#include + +#include "hexl/hexl.hpp" + +namespace psi::inspire::lrhe { + +constexpr uint64_t kLrCrtQ1 = 268369921ULL; +constexpr uint64_t kLrCrtQ2 = 249561089ULL; +constexpr uint64_t kLrCrtMod = kLrCrtQ1 * kLrCrtQ2; +constexpr uint64_t kLrRootOfUnityCrt4096 = 3375402822066082ULL; +constexpr uint64_t kLrRootOfUnityCrt2048 = 38878761190133527ULL; + +class Poly { + public: + Poly(); + Poly(uint64_t len, uint64_t modulus, bool ntt = false); + Poly(const std::vector& vec, uint64_t modulus, bool ntt); + + uint64_t get_length() const; + uint64_t get_modulus() const; + bool get_nttform() const; + void set_nttform(bool ntt); + uint64_t* data(); + const uint64_t* data() const; + uint64_t& operator[](size_t idx); + const uint64_t& operator[](size_t idx) const; + std::vector& get_data(); + const std::vector& get_data() const; + + private: + uint64_t length_; + uint64_t modulus_; + bool nttform_; + + public: + std::vector payload; +}; + +void PolyMult(Poly& lhs, Poly& rhs, Poly& result); +void PolyAdd(Poly& lhs, Poly& rhs, Poly& result); +void PolySub(Poly& lhs, Poly& rhs, Poly& result); +void PolyFMA(Poly& mul, uint64_t constant, Poly& result); +void PolyToNTT(Poly& poly, intel::hexl::NTT& ntts); +void PolyToCoef(Poly& poly, intel::hexl::NTT& ntts); + +class Secret { + public: + Secret(std::vector& vec, uint64_t modulus, bool ntt); + intel::hexl::NTT& get_ntt(); + Poly& get_data(); + const Poly& get_data() const; + uint64_t get_modulus() const; + uint64_t get_length() const; + bool get_nttform() const; + void to_ntt_form(); + + private: + intel::hexl::NTT ntts_; + Poly data_; +}; + +class RlweCiphertext { + public: + RlweCiphertext(); + RlweCiphertext(uint64_t len, uint64_t modulus, bool ntt); + + uint64_t get_degree() const; + uint64_t get_modulus() const; + void set_nttform(bool ntt); + + Poly a; + Poly b; +}; + +class Plaintext { + public: + Plaintext(uint64_t len, uint64_t pmod); + Plaintext(const std::vector& vec, uint64_t pmod, bool ntt); + + uint64_t get_pmod() const; + bool get_isntt() const; + + Poly pt; +}; + +class ApproximateRgswCiphertext { + public: + ApproximateRgswCiphertext(uint64_t len, uint64_t modulus, uint64_t b, + uint64_t z, uint64_t t); + + uint64_t get_degree() const; + uint64_t get_modulus() const; + uint64_t get_b() const; + uint64_t get_z() const; + uint64_t get_t() const; + + private: + uint64_t degree_; + uint64_t modulus_; + uint64_t b_; + uint64_t z_; + uint64_t t_; + + public: + std::vector ct_m; + std::vector ct_sm; +}; + +void RlweBfvEncrypt(Plaintext& pt, RlweCiphertext& ct, Secret& sk, float sig); +void RlweBfvEncode(Plaintext& pt, RlweCiphertext& ct, Secret& sk, float sig); +void RlweBfvDecrypt(RlweCiphertext& ct, Plaintext& pt, Secret& sk); +void RlweBfvAdd(RlweCiphertext& add1, RlweCiphertext& add2, + RlweCiphertext& result); +void ExternalProduct(const ApproximateRgswCiphertext& ct_rgsw, + RlweCiphertext& ct_bfv, RlweCiphertext& result, + intel::hexl::NTT& ntts); +void RgswEncode(Plaintext& pt, ApproximateRgswCiphertext& ct, Secret sk, + float sig); + +} // namespace psi::inspire::lrhe diff --git a/psi/algorithm/inspire/params.h b/psi/algorithm/inspire/params.h new file mode 100644 index 00000000..bd9f47ff --- /dev/null +++ b/psi/algorithm/inspire/params.h @@ -0,0 +1,26 @@ +#pragma once + +#include "psi/algorithm/ypir/params.h" + +namespace psi::inspire { + +using InspireParameters = psi::ypir::YpirParameters; + +inline InspireParameters CreateParamsForScenario(uint64_t num_items, + uint64_t item_size_bits) { + return psi::ypir::CreateParamsForScenarioDoublePIR(num_items, + item_size_bits); +} + +inline InspireParameters CreateParamsForShape(uint64_t db_rows, + uint64_t db_cols, + uint64_t item_size_bits) { + return psi::ypir::CreateParamsForShapeDoublePIR(db_rows, db_cols, + item_size_bits); +} + +inline InspireParameters CreateSmallTestParams() { + return psi::ypir::CreateSmallTestParamsDoublePIR(); +} + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/serialize.cc b/psi/algorithm/inspire/serialize.cc new file mode 100644 index 00000000..2ad000dc --- /dev/null +++ b/psi/algorithm/inspire/serialize.cc @@ -0,0 +1,111 @@ +#include "psi/algorithm/inspire/serialize.h" + +#include +#include +#include +#include +#include + +#include "yacl/base/exception.h" + +namespace psi::inspire { +namespace { + +template +void AppendPod(std::string& out, const T& value) { + out.append(reinterpret_cast(&value), sizeof(T)); +} + +template +T ReadPod(const char*& ptr, const char* end) { + YACL_ENFORCE_GE(end - ptr, static_cast(sizeof(T))); + T value; + std::memcpy(&value, ptr, sizeof(T)); + ptr += sizeof(T); + return value; +} + +template +void AppendVector(std::string& out, const std::vector& values) { + AppendPod(out, values.size()); + if (!values.empty()) { + out.append(reinterpret_cast(values.data()), + values.size() * sizeof(T)); + } +} + +template +std::vector ReadVector(const char*& ptr, const char* end) { + const auto size = ReadPod(ptr, end); + YACL_ENFORCE_GE(end - ptr, static_cast(size * sizeof(T))); + std::vector out(size); + if (size > 0) { + std::memcpy(out.data(), ptr, size * sizeof(T)); + ptr += size * sizeof(T); + } + return out; +} + +template +void AppendVector2D(std::string& out, + const std::vector>& values) { + AppendPod(out, values.size()); + for (const auto& inner : values) { + AppendVector(out, inner); + } +} + +template +std::vector> ReadVector2D(const char*& ptr, const char* end) { + const auto outer = ReadPod(ptr, end); + std::vector> out; + out.reserve(outer); + for (uint64_t i = 0; i < outer; ++i) { + out.push_back(ReadVector(ptr, end)); + } + return out; +} + +yacl::Buffer ToBuffer(const std::string& bytes) { + return yacl::Buffer(bytes.data(), bytes.size()); +} + +} // namespace + +yacl::Buffer SerializeQuery(const InspireQuery& query) { + std::string bytes; + AppendVector(bytes, query.qu0); + AppendVector(bytes, query.qu1); + AppendVector2D(bytes, query.ksk_b); + return ToBuffer(bytes); +} + +InspireQuery DeserializeQuery(const yacl::ByteContainerView& buffer) { + const char* ptr = reinterpret_cast(buffer.data()); + const char* end = ptr + buffer.size(); + + InspireQuery query; + query.qu0 = ReadVector(ptr, end); + query.qu1 = ReadVector(ptr, end); + query.ksk_b = ReadVector2D(ptr, end); + YACL_ENFORCE(ptr == end, "unexpected trailing bytes in InspireQuery"); + return query; +} + +yacl::Buffer SerializeResponse(const InspireResponse& response) { + std::string bytes; + AppendVector2D(bytes, response.doublepir_response); + return ToBuffer(bytes); +} + +InspireResponse DeserializeResponse(const yacl::ByteContainerView& buffer) { + const char* ptr = reinterpret_cast(buffer.data()); + const char* end = ptr + buffer.size(); + + InspireResponse response; + response.doublepir_response = ReadVector2D(ptr, end); + YACL_ENFORCE(ptr == end, "unexpected trailing bytes in InspireResponse"); + return response; +} + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/serialize.h b/psi/algorithm/inspire/serialize.h new file mode 100644 index 00000000..55967a39 --- /dev/null +++ b/psi/algorithm/inspire/serialize.h @@ -0,0 +1,16 @@ +#pragma once + +#include "yacl/base/buffer.h" +#include "yacl/base/byte_container_view.h" + +#include "psi/algorithm/inspire/types.h" + +namespace psi::inspire { + +yacl::Buffer SerializeQuery(const InspireQuery& query); +InspireQuery DeserializeQuery(const yacl::ByteContainerView& buffer); + +yacl::Buffer SerializeResponse(const InspireResponse& response); +InspireResponse DeserializeResponse(const yacl::ByteContainerView& buffer); + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/server.cc b/psi/algorithm/inspire/server.cc new file mode 100644 index 00000000..ef108dd6 --- /dev/null +++ b/psi/algorithm/inspire/server.cc @@ -0,0 +1,93 @@ +#include "psi/algorithm/inspire/server.h" + +#include +#include +#include +#include + +#include "yacl/base/exception.h" + +#include "psi/algorithm/inspire/serialize.h" +#include "psi/algorithm/ypir/ypir_internal_params.h" + +namespace psi::inspire { + +InspireServer::InspireServer(InspireParameters params) + : psi::pir::IndexPirDataBase(psi::pir::PirType::YPIR_PIR), + params_(std::move(params)), + context_(std::make_unique( + psi::ypir::internal::ypir::CreateContext(params_))) { + YACL_ENFORCE(params_.mode == psi::ypir::YpirMode::kDoublepir); + YACL_ENFORCE_EQ(params_.value_bytes, 1U, + "Inspire currently expects uint8_t database values"); +} + +void InspireServer::GenerateFromRawData( + const psi::pir::RawDatabase& raw_database) { + const bool item_layout = raw_database.Rows() <= params_.NumItems() && + raw_database.RowByteLen() == params_.value_bytes; + const bool matrix_layout = raw_database.Rows() == params_.db_rows && + raw_database.RowByteLen() == params_.db_cols; + YACL_ENFORCE(item_layout || matrix_layout, + "raw database shape does not match Inspire parameters"); + + db_row_major_.assign(params_.db_rows * params_.db_cols, 0); + if (item_layout) { + for (uint64_t raw_idx = 0; raw_idx < raw_database.Rows(); ++raw_idx) { + db_row_major_[raw_idx] = raw_database.At(raw_idx)[0]; + } + } else { + for (uint64_t row = 0; row < params_.db_rows; ++row) { + const auto& row_bytes = raw_database.At(row); + std::memcpy(db_row_major_.data() + row * params_.db_cols, + row_bytes.data(), row_bytes.size()); + } + } + db_set_ = true; +} + +void InspireServer::GenerateFromSimpleHashTable( + const psi::pir::RawDatabase& raw_database) { + GenerateFromRawData(raw_database); +} + +void InspireServer::Dump(std::ostream& out_stream) const { + out_stream << "InspireServer{db_rows=" << params_.db_rows + << ", db_cols=" << params_.db_cols + << ", value_bytes=" << params_.value_bytes + << ", db_set=" << db_set_ << "}"; +} + +InspirePrecomputedState InspireServer::PerformOfflinePrecomputation() const { + YACL_ENFORCE(db_set_, "database must be loaded before precomputation"); + return internal::PrepareOfflineState(db_row_major_, params_, *context_); +} + +InspireResponse InspireServer::ProcessQuery(const InspireQuery& query) const { + return ProcessQuery(query, PerformOfflinePrecomputation()); +} + +InspireResponse InspireServer::ProcessQuery( + const InspireQuery& query, const InspirePrecomputedState& state) const { + YACL_ENFORCE(db_set_, "database must be loaded before query processing"); + return internal::ProcessQuery(db_row_major_, query, state, params_, *context_); +} + +yacl::Buffer InspireServer::Response( + const yacl::ByteContainerView& query_buffer) const { + return SerializeResponse(ProcessQuery(DeserializeQuery(query_buffer))); +} + +yacl::Buffer InspireServer::Response( + const yacl::ByteContainerView& query_buffer, + const yacl::Buffer& /*pks_buffer*/) const { + return Response(query_buffer); +} + +std::string InspireServer::Response(const yacl::ByteContainerView& query_buffer, + const std::string& /*pks_buffer*/) const { + auto buffer = Response(query_buffer); + return std::string(static_cast(buffer)); +} + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/server.h b/psi/algorithm/inspire/server.h new file mode 100644 index 00000000..3829fa15 --- /dev/null +++ b/psi/algorithm/inspire/server.h @@ -0,0 +1,52 @@ +#pragma once + +#include +#include +#include +#include +#include + +#include "yacl/base/buffer.h" +#include "yacl/base/byte_container_view.h" + +#include "psi/algorithm/inspire/internal.h" +#include "psi/algorithm/inspire/params.h" +#include "psi/algorithm/inspire/types.h" +#include "psi/algorithm/pir_interface/pir_db.h" + +namespace psi::inspire { + +class InspireServer : public psi::pir::IndexPirDataBase { + public: + explicit InspireServer(InspireParameters params); + + void GenerateFromRawData(const psi::pir::RawDatabase& raw_database) override; + void GenerateFromSimpleHashTable( + const psi::pir::RawDatabase& raw_database) override; + void Dump(std::ostream& out_stream) const override; + std::size_t MaxElementsOfOnePt() const override { return 1; } + bool DbSeted() const override { return db_set_; } + + [[nodiscard]] bool DbSet() const { return db_set_; } + [[nodiscard]] const InspireParameters& GetParameters() const { + return params_; + } + + InspirePrecomputedState PerformOfflinePrecomputation() const; + InspireResponse ProcessQuery(const InspireQuery& query) const; + InspireResponse ProcessQuery(const InspireQuery& query, + const InspirePrecomputedState& state) const; + yacl::Buffer Response(const yacl::ByteContainerView& query_buffer) const; + yacl::Buffer Response(const yacl::ByteContainerView& query_buffer, + const yacl::Buffer& pks_buffer) const override; + std::string Response(const yacl::ByteContainerView& query_buffer, + const std::string& pks_buffer) const override; + + private: + InspireParameters params_; + bool db_set_ = false; + std::vector db_row_major_; + std::unique_ptr context_; +}; + +} // namespace psi::inspire diff --git a/psi/algorithm/inspire/types.h b/psi/algorithm/inspire/types.h new file mode 100644 index 00000000..64bece78 --- /dev/null +++ b/psi/algorithm/inspire/types.h @@ -0,0 +1,83 @@ +#pragma once + +#include +#include +#include + +namespace psi::inspire { + +struct InspirePrecomputedState { + std::vector hint_0; + std::vector> server_hint; + std::vector>> decomp_buf; +}; + +struct InspireQuery { + std::vector qu0; + std::vector qu1; + std::vector> ksk_b; +}; + +struct InspireResponse { + std::vector> doublepir_response; +}; + +struct LargeRecordConfig { + static constexpr uint64_t kCrtQ1 = 268369921ULL; + static constexpr uint64_t kCrtQ2 = 249561089ULL; + static constexpr uint64_t kCrtMod = kCrtQ1 * kCrtQ2; + static constexpr size_t kRecordBytes = 256; + + uint64_t total_records = 1ULL << 20; + uint64_t folding = 8; + uint64_t degree = 2048; + uint64_t packed_coeffs = 2048; + uint64_t qmod = kCrtMod; + uint64_t pmod = 257; + float sigma = 0.41f; + uint64_t rgsw_b = 0; + uint64_t rgsw_z = 8; + uint64_t rgsw_t = 7; + uint64_t lwe_dimension = 2048; + uint64_t lwe_cmod = kCrtMod; + uint64_t lwe_pmod = 257; +}; + +struct PaperPackKeyMaterial { + std::vector>> kg_a_powers; + std::vector>> kg_b_powers; + std::vector>> gh_kg_a_powers; + std::vector>> gh_kg_b_powers; + std::vector> kh_a; + std::vector> kh_b; + std::vector> g_automaps; + std::vector> gh_automaps; + std::vector> monomials_ntt; +}; + +struct PaperPackedHintBlock { + std::vector packed_a_ntt; + std::vector>> decomp_half1; + std::vector>> decomp_half2; + std::vector> decomp_final; +}; + +struct LargeRecordPrecomputedState { + std::vector> encoded_db; + std::vector> hint; + std::vector packed_hint_blocks; + PaperPackKeyMaterial key_material; +}; + +struct LargeRecordQuery { + std::vector column_query; + std::vector>> rgsw_ct_m; + std::vector>> rgsw_ct_sm; +}; + +struct LargeRecordResponse { + std::vector ct_a; + std::vector ct_b; +}; + +} // namespace psi::inspire diff --git a/psi/apps/psi_launcher/BUILD.bazel b/psi/apps/psi_launcher/BUILD.bazel index 05aea38e..8ee24c83 100644 --- a/psi/apps/psi_launcher/BUILD.bazel +++ b/psi/apps/psi_launcher/BUILD.bazel @@ -64,12 +64,35 @@ psi_cc_library( }), ) +psi_cc_library( + name = "launch_inspire", + srcs = select({ + "@platforms//cpu:x86_64": ["inspire_launch.cc"], + "//conditions:default": ["inspire_launch_stub.cc"], + }), + hdrs = ["inspire_launch.h"], + deps = select({ + "@platforms//cpu:x86_64": [ + "//psi/algorithm/inspire:entry", + "//psi/proto:pir_cc_proto", + "@yacl//yacl/base:exception", + "@yacl//yacl/link:context", + ], + "//conditions:default": [ + "//psi/proto:pir_cc_proto", + "@yacl//yacl/base:exception", + "@yacl//yacl/link:context", + ], + }), +) + psi_cc_library( name = "launch", srcs = ["launch.cc"], hdrs = ["launch.h"], deps = [ ":factory", + ":launch_inspire", ":launch_ypir", ":report", "//psi:trace_categories", @@ -106,6 +129,17 @@ psi_cc_test( ], ) +psi_cc_test( + name = "inspire_test", + srcs = ["inspire_test.cc"], + target_compatible_with = ["@platforms//cpu:x86_64"], + deps = [ + ":launch_inspire", + "@abseil-cpp//absl/strings", + "@yacl//yacl/link:test_util", + ], +) + psi_cc_library( name = "kuscia_adapter", srcs = [ diff --git a/psi/apps/psi_launcher/inspire_launch.cc b/psi/apps/psi_launcher/inspire_launch.cc new file mode 100644 index 00000000..34d354d4 --- /dev/null +++ b/psi/apps/psi_launcher/inspire_launch.cc @@ -0,0 +1,34 @@ +#include "psi/apps/psi_launcher/inspire_launch.h" + +#include "yacl/base/exception.h" + +#include "psi/algorithm/inspire/entry.h" + +namespace psi { + +PirResultReport RunPir(const InspireReceiverConfig& inspire_receiver_config, + const std::shared_ptr& lctx) { + psi::inspire::InspireReceiverOptions options; + options.db_rows = inspire_receiver_config.db_rows(); + options.db_cols = inspire_receiver_config.db_cols(); + options.item_size_bits = inspire_receiver_config.item_size_bits(); + options.query_file = inspire_receiver_config.query_file(); + options.output_file = inspire_receiver_config.output_file(); + + YACL_ENFORCE_EQ(psi::inspire::ReceiverOnline(options, lctx), 0); + return PirResultReport(); +} + +PirResultReport RunPir(const InspireSenderConfig& inspire_sender_config, + const std::shared_ptr& lctx) { + psi::inspire::InspireSenderOptions options; + options.db_rows = inspire_sender_config.db_rows(); + options.db_cols = inspire_sender_config.db_cols(); + options.item_size_bits = inspire_sender_config.item_size_bits(); + options.db_file = inspire_sender_config.db_file(); + + YACL_ENFORCE_EQ(psi::inspire::SenderOnline(options, lctx), 0); + return PirResultReport(); +} + +} // namespace psi diff --git a/psi/apps/psi_launcher/inspire_launch.h b/psi/apps/psi_launcher/inspire_launch.h new file mode 100644 index 00000000..68c26395 --- /dev/null +++ b/psi/apps/psi_launcher/inspire_launch.h @@ -0,0 +1,17 @@ +#pragma once + +#include + +#include "yacl/link/context.h" + +#include "psi/proto/pir.pb.h" + +namespace psi { + +PirResultReport RunPir(const InspireReceiverConfig& inspire_receiver_config, + const std::shared_ptr& lctx); + +PirResultReport RunPir(const InspireSenderConfig& inspire_sender_config, + const std::shared_ptr& lctx); + +} // namespace psi diff --git a/psi/apps/psi_launcher/inspire_launch_stub.cc b/psi/apps/psi_launcher/inspire_launch_stub.cc new file mode 100644 index 00000000..8953d2be --- /dev/null +++ b/psi/apps/psi_launcher/inspire_launch_stub.cc @@ -0,0 +1,17 @@ +#include "yacl/base/exception.h" + +#include "psi/apps/psi_launcher/inspire_launch.h" + +namespace psi { + +PirResultReport RunPir(const InspireReceiverConfig&, + const std::shared_ptr&) { + YACL_THROW("Inspire is only supported on x86_64"); +} + +PirResultReport RunPir(const InspireSenderConfig&, + const std::shared_ptr&) { + YACL_THROW("Inspire is only supported on x86_64"); +} + +} // namespace psi diff --git a/psi/apps/psi_launcher/inspire_test.cc b/psi/apps/psi_launcher/inspire_test.cc new file mode 100644 index 00000000..03ad9748 --- /dev/null +++ b/psi/apps/psi_launcher/inspire_test.cc @@ -0,0 +1,90 @@ +#include +#include +#include +#include +#include + +#include "absl/strings/escaping.h" +#include "gtest/gtest.h" +#include "yacl/link/test_util.h" + +#include "psi/apps/psi_launcher/inspire_launch.h" + +namespace psi { +namespace { + +std::string HexByte(uint64_t value) { + std::vector bytes = {static_cast(value)}; + return absl::BytesToHexString(absl::string_view( + reinterpret_cast(bytes.data()), bytes.size())); +} + +void WriteLines(const std::filesystem::path& path, + const std::vector& lines) { + std::ofstream output(path, std::ios::out | std::ios::trunc); + ASSERT_TRUE(output.is_open()); + for (const auto& line : lines) { + output << line << '\n'; + } +} + +std::vector ReadLines(const std::filesystem::path& path) { + std::ifstream input(path); + EXPECT_TRUE(input.is_open()); + std::vector lines; + std::string line; + while (std::getline(input, line)) { + if (!line.empty()) { + lines.push_back(line); + } + } + return lines; +} + +TEST(InspireLauncherTest, RunPirUsesLauncherEntry) { + const auto tmp_dir = + std::filesystem::temp_directory_path() / "inspire_launcher"; + std::filesystem::create_directories(tmp_dir); + const auto db_path = tmp_dir / "db.hex"; + const auto query_path = tmp_dir / "query.txt"; + const auto output_path = tmp_dir / "result.hex"; + + std::vector db_lines; + db_lines.reserve(1ULL << 20); + for (uint64_t raw_idx = 0; raw_idx < (1ULL << 20); ++raw_idx) { + const uint64_t row = raw_idx / (1ULL << 10); + const uint64_t col = raw_idx % (1ULL << 10); + db_lines.push_back(HexByte((row + col) % 251)); + } + WriteLines(db_path, db_lines); + WriteLines(query_path, {"0", "113886", "1048575"}); + + InspireSenderConfig sender_config; + sender_config.set_db_rows(1ULL << 10); + sender_config.set_db_cols(1ULL << 10); + sender_config.set_item_size_bits(8); + sender_config.set_db_file(db_path); + + InspireReceiverConfig receiver_config; + receiver_config.set_db_rows(1ULL << 10); + receiver_config.set_db_cols(1ULL << 10); + receiver_config.set_item_size_bits(8); + receiver_config.set_query_file(query_path); + receiver_config.set_output_file(output_path); + + auto lctxs = yacl::link::test::SetupWorld(2); + auto sender = std::async(std::launch::async, + [&] { return RunPir(sender_config, lctxs[0]); }); + auto receiver = std::async(std::launch::async, + [&] { return RunPir(receiver_config, lctxs[1]); }); + + EXPECT_EQ(sender.get().match_cnt(), 0); + EXPECT_EQ(receiver.get().match_cnt(), 0); + EXPECT_EQ(ReadLines(output_path), + (std::vector{HexByte(0), HexByte(82), HexByte(38)})); + + std::filesystem::remove_all(tmp_dir); +} + +} // namespace +} // namespace psi diff --git a/psi/apps/psi_launcher/launch.h b/psi/apps/psi_launcher/launch.h index 0b5dad32..4c210ae4 100644 --- a/psi/apps/psi_launcher/launch.h +++ b/psi/apps/psi_launcher/launch.h @@ -53,6 +53,12 @@ PirResultReport RunPir(const YpirReceiverConfig& ypir_receiver_config, PirResultReport RunPir(const YpirSenderConfig& ypir_sender_config, const std::shared_ptr& lctx); +PirResultReport RunPir(const InspireReceiverConfig& inspire_receiver_config, + const std::shared_ptr& lctx); + +PirResultReport RunPir(const InspireSenderConfig& inspire_sender_config, + const std::shared_ptr& lctx); + PirResultReport RunDkPir(const DkPirReceiverConfig& dk_pir_receiver_config, const std::shared_ptr& lctx); diff --git a/psi/apps/psi_launcher/main.cc b/psi/apps/psi_launcher/main.cc index 1d65811e..fd8f303b 100644 --- a/psi/apps/psi_launcher/main.cc +++ b/psi/apps/psi_launcher/main.cc @@ -133,6 +133,18 @@ int main(int argc, char* argv[]) { YACL_ENFORCE(google::protobuf::util::MessageToJsonString( report, &report_json, json_print_options) .ok()); + } else if (launch_config.has_inspire_sender_config()) { + psi::PirResultReport report = + psi::RunPir(launch_config.inspire_sender_config(), lctx); + YACL_ENFORCE(google::protobuf::util::MessageToJsonString( + report, &report_json, json_print_options) + .ok()); + } else if (launch_config.has_inspire_receiver_config()) { + psi::PirResultReport report = + psi::RunPir(launch_config.inspire_receiver_config(), lctx); + YACL_ENFORCE(google::protobuf::util::MessageToJsonString( + report, &report_json, json_print_options) + .ok()); } else if (launch_config.has_dk_pir_sender_config()) { psi::PirResultReport report = psi::RunDkPir(launch_config.dk_pir_sender_config(), lctx); diff --git a/psi/proto/entry.proto b/psi/proto/entry.proto index 49c65bca..3c3d2ef0 100644 --- a/psi/proto/entry.proto +++ b/psi/proto/entry.proto @@ -51,5 +51,9 @@ message LaunchConfig { YpirSenderConfig ypir_sender_config = 10; YpirReceiverConfig ypir_receiver_config = 11; + + InspireSenderConfig inspire_sender_config = 12; + + InspireReceiverConfig inspire_receiver_config = 13; } } diff --git a/psi/proto/pir.proto b/psi/proto/pir.proto index 87a426ed..69a7f2aa 100644 --- a/psi/proto/pir.proto +++ b/psi/proto/pir.proto @@ -268,6 +268,36 @@ message YpirReceiverConfig { string output_file = 6; } +message InspireSenderConfig { + // Logical database shape. Inspire currently supports 8-bit values. + uint64 db_rows = 1; + uint64 db_cols = 2; + + // Plaintext bit width of each value. Currently must be 8. + uint64 item_size_bits = 3; + + // Path to a text file that contains exactly db_rows * db_cols hex-encoded + // values, one per line, without a header. + string db_file = 4; +} + +message InspireReceiverConfig { + // Logical database shape. Must match sender config. + uint64 db_rows = 1; + uint64 db_cols = 2; + + // Plaintext bit width of each value. Must match sender config. + uint64 item_size_bits = 3; + + // Path to a text file containing raw indexes in decimal, one per line, + // without a header. + string query_file = 4; + + // Path to a text file where hex-encoded results will be written, one per + // line, without a header. + string output_file = 5; +} + // The report of pir task. message PirResultReport { int64 match_cnt = 1;