libdpf/include/dpf/json.hpp

274 lines
9.4 KiB
C++

/// @file dpf/json.hpp
/// @brief nlohmann::json serializers for DPF keys and beaver triples.
/// @details ADL `adl_serializer` specializations so `nlohmann::json` can
/// convert the library's key and triple types.
/// @author Ryan Henry <ryan.henry@ucalgary.ca>
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
/// see [LICENSE.md](@ref license) for details.
#ifndef LIBDPF_INCLUDE_DPF_JSON_HPP__
#define LIBDPF_INCLUDE_DPF_JSON_HPP__
#include <cstddef>
#include <cstdint>
#include <tuple>
#include <array>
#include <string>
#include <bitset>
#include <type_traits>
#include <utility>
#include "json/include/nlohmann/json.hpp"
#include "portable-snippets/exact-int/exact-int.h"
#include "dpf/dpf_key.hpp"
namespace nlohmann
{
template <typename NodeT,
typename OutputT>
struct adl_serializer<dpf::beaver<true, NodeT, OutputT>>
{
static void from_json(const nlohmann::json & j, dpf::beaver<true, NodeT, OutputT> & beaver) // NOLINT(runtime/references)
{
j.get_to(beaver.output_blind);
j.get_to(beaver.vector_blind);
j.get_to(beaver.blinded_vector);
}
static void to_json(nlohmann::json & j, const dpf::beaver<true, NodeT, OutputT> & beaver) // NOLINT(runtime/references)
{
j = nlohmann::json{
{"output_blind", beaver.output_blind},
{"vector_blind", beaver.vector_blind},
{"blinded_vector", beaver.blinded_vector}
};
}
};
template <>
struct adl_serializer<simde__m128i>
{
static void from_json(const nlohmann::json & j, simde__m128i & a) // NOLINT(runtime/references)
{
std::array<psnip_uint64_t, 2> A;
j.get_to(A);
a = simde_mm_set_epi64x(A[1], A[0]);
}
static void to_json(nlohmann::json & j, const simde__m128i & a) // NOLINT(runtime/references)
{
j = nlohmann::json{a[0], a[1]};
}
};
template <>
struct adl_serializer<simde__m256i>
{
static void from_json(const nlohmann::json & j, simde__m256i & a) // NOLINT(runtime/references)
{
std::array<psnip_uint64_t, 4> A;
j.get_to(A);
a = simde_mm256_set_epi64x(A[3], A[2], A[1], A[0]);
}
static void to_json(nlohmann::json & j, const simde__m256i & a) // NOLINT(runtime/references)
{
j = nlohmann::json{a[0], a[1], a[2], a[3]};
}
};
template <>
struct adl_serializer<dpf::detail::cmp_meta>
{
static void from_json(const nlohmann::json & j, dpf::detail::cmp_meta & c) // NOLINT(runtime/references)
{
j.at("nbits").get_to(c.nbits);
j.at("mask").get_to(c.mask);
c.kind = static_cast<dpf::cmp_kind>(j.at("kind").get<psnip_uint8_t>());
c.trivial =
static_cast<dpf::cmp_trivial>(j.at("trivial").get<psnip_uint8_t>());
j.at("eval_as_ge").get_to(c.eval_as_ge);
j.at("include_eq").get_to(c.include_eq);
j.at("active").get_to(c.active);
c.incremental = j.value("incremental", false);
c.block_width = j.value("block_width", 0);
c.tail_bits = j.value("tail_bits", 0);
}
static void to_json(nlohmann::json & j, const dpf::detail::cmp_meta & c) // NOLINT(runtime/references)
{
j = nlohmann::json{
{"nbits", c.nbits},
{"mask", c.mask},
{"kind", static_cast<psnip_uint8_t>(c.kind)},
{"trivial", static_cast<psnip_uint8_t>(c.trivial)},
{"eval_as_ge", c.eval_as_ge},
{"include_eq", c.include_eq},
{"active", c.active}
};
if (c.incremental)
j["incremental"] = true;
if (c.block_width != 0)
{
j["block_width"] = c.block_width;
j["tail_bits"] = c.tail_bits;
}
}
};
// Classic single-level key (no `at<>` / no comparison channel).
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs>
struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>,
std::enable_if_t<!dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT,
OutputTs...>::is_multilevel>>
{
using dpf_type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>;
using interior_node = typename dpf_type::interior_node;
using leaf_tuple = typename dpf_type::leaf_tuple;
using beaver_tuple = typename dpf_type::beaver_tuple;
static dpf_type from_json(const nlohmann::json & j)
{
interior_node root;
j.at("root").get_to(root);
std::array<interior_node, dpf_type::depth> correction_words;
j.at("correction_words").get_to(correction_words);
std::array<psnip_uint8_t, dpf_type::depth> correction_advice;
j.at("correction_advice").get_to(correction_advice);
leaf_tuple leaves;
j.at("leaves").get_to(leaves);
std::string wildcard_mask_str;
j.at("wildcards").get_to(wildcard_mask_str);
beaver_tuple beavers;
j.at("beavers").get_to(beavers);
return dpf_type{
root,
correction_words,
correction_advice,
leaves,
std::bitset<std::tuple_size_v<leaf_tuple>>(wildcard_mask_str),
beavers
};
}
static void to_json(nlohmann::json & j, const dpf_type & dpf) // NOLINT(runtime/references)
{
j = nlohmann::json{
{"root", dpf.root()},
{"correction_words", dpf.correction_words()},
{"correction_advice", dpf.correction_advice()},
{"leaves", dpf.mutable_leaf_tuple()},
{"wildcards", dpf.mutable_wildcard_mask()},
{"beavers", dpf.mutable_beaver_tuple()}
};
}
};
// Multi-level / comparison key (`at<>` and/or a `cmp` channel). Round-trips
// the public tree (root, CWs, advice) and the comparison channel (cmp meta,
// value CWs, `cw_last`, and this party's `cmp_addend` share). Leaf outputs are
// not yet serialized here, so this path currently supports comparison-only
// keys (`num_outputs == 0`, e.g. `make_dpf(x, dpf::lt(...))`).
template <typename InteriorPRG,
typename ExteriorPRG,
typename InputT,
typename OutputT,
typename ...OutputTs>
struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>,
std::enable_if_t<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT,
OutputTs...>::is_multilevel>>
{
using dpf_type = dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, OutputTs...>;
using interior_node = typename dpf_type::interior_node;
using input_type = typename dpf_type::input_type;
static dpf_type from_json(const nlohmann::json & j)
{
static_assert(dpf_type::num_outputs == 0,
"dpf::json round-trip currently supports comparison-only "
"multi-level keys (no leaf outputs)");
interior_node root;
j.at("root").get_to(root);
typename dpf_type::correction_words_array correction_words;
j.at("correction_words").get_to(correction_words);
typename dpf_type::correction_advice_array correction_advice;
j.at("correction_advice").get_to(correction_advice);
dpf::detail::cmp_meta cmp;
j.at("cmp").get_to(cmp);
typename dpf_type::value_cw_array value_cws;
j.at("value_cw").get_to(value_cws);
uint64_t cw_last = j.at("cw_last").template get<uint64_t>();
uint64_t cmp_addend = j.at("cmp_addend").template get<uint64_t>();
typename dpf_type::tail_array tail{};
if constexpr (dpf_type::cmp_block > 0)
j.at("tail_cw").get_to(tail);
typename dpf_type::prefix_cw_array prefix{};
if constexpr (dpf_type::cmp_idcf)
j.at("prefix_cw").get_to(prefix);
typename dpf_type::leaf_wrapper_tuple leaves{};
input_type offset_share{};
typename dpf_type::addend_tuple addends{};
return dpf_type{root, correction_words, correction_advice,
std::move(leaves), offset_share, cmp, value_cws, cw_last,
cmp_addend, addends, {}, 0, tail, {}, prefix};
}
static void to_json(nlohmann::json & j, const dpf_type & dpf) // NOLINT(runtime/references)
{
static_assert(dpf_type::num_outputs == 0,
"dpf::json round-trip currently supports comparison-only "
"multi-level keys (no leaf outputs)");
j = nlohmann::json{
{"root", dpf.root()},
{"correction_words", dpf.correction_words()},
{"correction_advice", dpf.correction_advice()},
{"cmp", dpf.cmp()},
{"value_cw", dpf.value_cw()},
{"cw_last", static_cast<uint64_t>(dpf.cw_last())},
{"cmp_addend", static_cast<uint64_t>(dpf.cmp_addend())}
};
if constexpr (dpf_type::cmp_block > 0)
j["tail_cw"] = dpf.tail_cw();
if constexpr (dpf_type::cmp_idcf)
j["prefix_cw"] = dpf.prefix_cws();
}
};
} // namespace nlohmann
namespace dpf
{
namespace json
{
template <typename DpfKey>
static std::string to_json(const DpfKey & dpf)
{
nlohmann::json json = dpf;
return json.dump();
}
template <typename DpfType>
static auto from_json(const std::string & json_string)
{
nlohmann::json json = nlohmann::json::parse(json_string);
return static_cast<DpfType>(json);
}
} // namespace json
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_JSON_HPP__