#pragma once
#include <cstddef>
#include <cstdint>
#include "FeistelCipher.hpp"
#include "ProofParams.hpp"
using ProofFragment = uint64_t;
class ProofFragmentCodec {
public:
ProofFragmentCodec(ProofParams const& proof_params)
: params_(proof_params)
, cipher_(proof_params.get_plot_id_bytes(), proof_params.get_k())
{
}
std::array<ProofFragment, NUM_CHAIN_LINKS> fullProofXValuesToQualityString(
std::span<uint32_t const, TOTAL_XS_IN_PROOF> const full_proof) const
{
assert(full_proof.size() == 8 * NUM_CHAIN_LINKS);
std::array<ProofFragment, NUM_CHAIN_LINKS> quality_string;
size_t num_proof_fragments = full_proof.size() / 8;
for (size_t i = 0; i < num_proof_fragments; ++i) {
uint32_t x_values[8];
for (size_t j = 0; j < 8; ++j) {
x_values[j] = full_proof[i * 8 + j];
}
ProofFragment proof_fragment = encode(x_values);
quality_string[i] = proof_fragment;
}
return quality_string;
}
uint64_t encode(uint64_t all_x_bits) const { return cipher_.encrypt(all_x_bits); }
ProofFragment encode(uint32_t const x_values[8]) const
{
uint32_t x1 = x_values[0] >> (params_.get_k() / 2);
uint32_t x3 = x_values[2] >> (params_.get_k() / 2);
uint32_t x5 = x_values[4] >> (params_.get_k() / 2);
uint32_t x7 = x_values[6] >> (params_.get_k() / 2);
uint64_t all_x_bits = 0;
all_x_bits |= (static_cast<uint64_t>(x1) << (params_.get_k() * 3 / 2));
all_x_bits |= (static_cast<uint64_t>(x3) << (params_.get_k() * 2 / 2));
all_x_bits |= (static_cast<uint64_t>(x5) << (params_.get_k() * 1 / 2));
all_x_bits |= (static_cast<uint64_t>(x7) << (params_.get_k() * 0 / 2));
return cipher_.encrypt(all_x_bits);
}
uint64_t decode(uint64_t ciphertext) const { return cipher_.decrypt(ciphertext); }
bool validate_proof_fragment(ProofFragment proof_fragment, uint32_t const x_values[8]) const
{
size_t half_k
= params_.get_k() / 2; uint32_t x1 = x_values[0] >> half_k;
uint32_t x3 = x_values[2] >> half_k;
uint32_t x5 = x_values[4] >> half_k;
uint32_t x7 = x_values[6] >> half_k;
uint64_t decrypted_xs = cipher_.decrypt(proof_fragment);
uint32_t decrypted_x1
= static_cast<uint32_t>((decrypted_xs >> (half_k * 3)) & ((uint64_t(1) << half_k) - 1));
uint32_t decrypted_x3
= static_cast<uint32_t>((decrypted_xs >> (half_k * 2)) & ((uint64_t(1) << half_k) - 1));
uint32_t decrypted_x5
= static_cast<uint32_t>((decrypted_xs >> (half_k * 1)) & ((uint64_t(1) << half_k) - 1));
uint32_t decrypted_x7 = static_cast<uint32_t>(decrypted_xs & ((uint64_t(1) << half_k) - 1));
if (x1 != decrypted_x1 || x3 != decrypted_x3 || x5 != decrypted_x5 || x7 != decrypted_x7) {
return false;
}
return true;
}
std::array<uint32_t, 4> get_x_bits_from_proof_fragment(ProofFragment proof_fragment) const
{
uint64_t decrypted_xs = cipher_.decrypt(proof_fragment);
size_t half_k = params_.get_k() / 2;
uint32_t x1
= static_cast<uint32_t>((decrypted_xs >> (half_k * 3)) & ((uint64_t(1) << half_k) - 1));
uint32_t x3
= static_cast<uint32_t>((decrypted_xs >> (half_k * 2)) & ((uint64_t(1) << half_k) - 1));
uint32_t x5
= static_cast<uint32_t>((decrypted_xs >> (half_k * 1)) & ((uint64_t(1) << half_k) - 1));
uint32_t x7 = static_cast<uint32_t>(decrypted_xs & ((uint64_t(1) << half_k) - 1));
return { x1, x3, x5, x7 };
}
private:
ProofParams params_;
FeistelCipher cipher_;
};