#include #include #include #include #include #include #include "pydpf_types.hpp" #include "pydpf_eval.hpp" namespace py = pybind11; namespace pydpf_detail { namespace { template std::vector eval_full_limbs(const Key & key) { auto [buf, iter] = dpf::eval_full(key); (void)buf; return collect_limbs(iter); } template std::vector eval_interval_limbs(const Key & key, std::uint8_t from, std::uint8_t to) { if (to < from) throw std::invalid_argument("eval_interval: to < from"); auto [buf, iter] = dpf::eval_interval(key, from, to); (void)buf; return collect_limbs(iter); } template std::vector eval_sequence_limbs(const Key & key, const std::vector & xs) { require_sorted(xs); auto [buf, iter] = dpf::eval_sequence(key, xs.begin(), xs.end()); (void)buf; return collect_limbs(iter); } template std::vector eval_recipe_limbs(const Key & key, const std::vector & xs) { require_sorted(xs); auto recipe = dpf::make_sequence_recipe(key, xs.begin(), xs.end()); auto [buf, iter] = dpf::eval_sequence(key, recipe); (void)buf; return collect_limbs(iter); } } // namespace void register_eval(py::module_ & m) { py::class_(m, "MultiKeyPair"); py::class_(m, "WildcardKeyPair"); m.def("eval_full", [](const PointKeyPair & keys, int party) { require_party(party); return party == 0 ? eval_full_limbs(keys.k0) : eval_full_limbs(keys.k1); }, py::arg("keys"), py::arg("party")); m.def("eval_interval", [](const PointKeyPair & keys, int party, std::uint8_t from, std::uint8_t to) { require_party(party); return party == 0 ? eval_interval_limbs(keys.k0, from, to) : eval_interval_limbs(keys.k1, from, to); }, py::arg("keys"), py::arg("party"), py::arg("from"), py::arg("to")); m.def("eval_sequence", [](const PointKeyPair & keys, int party, const std::vector & xs) { require_party(party); return party == 0 ? eval_sequence_limbs(keys.k0, xs) : eval_sequence_limbs(keys.k1, xs); }, py::arg("keys"), py::arg("party"), py::arg("points")); m.def("eval_sequence_recipe", [](const PointKeyPair & keys, int party, const std::vector & xs) { require_party(party); return party == 0 ? eval_recipe_limbs(keys.k0, xs) : eval_recipe_limbs(keys.k1, xs); }, py::arg("keys"), py::arg("party"), py::arg("points")); m.def("make_dpf_multi", [](std::uint8_t alpha, std::uint64_t beta0, std::uint64_t beta1) { auto [a, b] = dpf::make_dpf(alpha, beta0, beta1); return MultiKeyPair{std::move(a), std::move(b)}; }, py::arg("alpha"), py::arg("beta0"), py::arg("beta1")); m.def("eval_point_leaf", [](const MultiKeyPair & keys, int party, int leaf, std::uint8_t x) { require_party(party); if (leaf == 0) { if (party == 0) return share_limb(*dpf::eval_point<0>(keys.k0, x)); return share_limb(*dpf::eval_point<0>(keys.k1, x)); } if (leaf == 1) { if (party == 0) return share_limb(*dpf::eval_point<1>(keys.k0, x)); return share_limb(*dpf::eval_point<1>(keys.k1, x)); } throw std::invalid_argument("eval_point_leaf: leaf must be 0 or 1"); }, py::arg("keys"), py::arg("party"), py::arg("leaf"), py::arg("x")); m.def("eval_full_leaf", [](const MultiKeyPair & keys, int party, int leaf) { require_party(party); if (leaf == 0) { if (party == 0) { auto [buf, iter] = dpf::eval_full<0>(keys.k0); (void)buf; return collect_limbs(iter); } auto [buf, iter] = dpf::eval_full<0>(keys.k1); (void)buf; return collect_limbs(iter); } if (leaf == 1) { if (party == 0) { auto [buf, iter] = dpf::eval_full<1>(keys.k0); (void)buf; return collect_limbs(iter); } auto [buf, iter] = dpf::eval_full<1>(keys.k1); (void)buf; return collect_limbs(iter); } throw std::invalid_argument("eval_full_leaf: leaf must be 0 or 1"); }, py::arg("keys"), py::arg("party"), py::arg("leaf")); m.def("make_dpf_wildcard", [](std::uint8_t alpha) { auto [a, b] = dpf::make_dpf(alpha, dpf::wildcard_value{}); return WildcardKeyPair{std::move(a), std::move(b)}; }, py::arg("alpha")); m.def("assign_wildcard", [](WildcardKeyPair & keys, std::uint64_t beta) { // Beaver assign takes additive shares of β (same as wildcard_test). const std::uint64_t shr0 = dpf::uniform_sample(); const std::uint64_t shr1 = beta - shr0; assign_leaf_local(keys.k0, keys.k1, shr0, shr1); }, py::arg("keys"), py::arg("beta")); m.def("eval_point", [](const WildcardKeyPair & keys, int party, std::uint8_t x) { require_party(party); if (party == 0) return share_limb(*dpf::eval_point(keys.k0, x)); return share_limb(*dpf::eval_point(keys.k1, x)); }, py::arg("keys"), py::arg("party"), py::arg("x")); m.def("eval_full", [](const WildcardKeyPair & keys, int party) { require_party(party); return party == 0 ? eval_full_limbs(keys.k0) : eval_full_limbs(keys.k1); }, py::arg("keys"), py::arg("party")); } } // namespace pydpf_detail