#pragma once
#include <array>
#include <cstdint>
#include <iostream>
#include <limits>
#include <optional>
#include <stdexcept>
#include <tuple>
#include <vector>
#include "ProofConstants.hpp"
#include "ProofFragment.hpp"
#include "ProofHashing.hpp"
#include "ProofParams.hpp"
using QualityChainLinks = std::array<ProofFragment, NUM_CHAIN_LINKS>;
struct QualityChain {
QualityChainLinks chain_links;
};
struct Chain {
std::array<ProofFragment, NUM_CHAIN_LINKS> fragments; };
struct T1Pairing {
uint32_t meta_lo;
uint32_t meta_hi;
uint32_t match_info;
uint64_t meta() const noexcept { return uint64_t(meta_lo) | (uint64_t(meta_hi) << 32); }
static T1Pairing make(uint64_t meta, uint32_t match) noexcept
{
T1Pairing p {};
p.meta_lo = uint32_t(meta);
p.meta_hi = uint32_t(meta >> 32);
p.match_info = match;
return p;
}
};
static_assert(sizeof(T1Pairing) == 12);
struct T2Pairing {
uint64_t meta; uint32_t match_info; uint32_t x_bits; #ifdef RETAIN_X_VALUES_TO_T3
uint32_t xs[4];
#endif
};
struct T3Pairing {
ProofFragment proof_fragment; #ifdef RETAIN_X_VALUES_TO_T3
std::array<uint32_t, 8> xs;
#endif
};
class ProofCore {
public:
ProofHashing hashing;
ProofFragmentCodec fragment_codec;
ProofCore(ProofParams const& proof_params)
: hashing(proof_params)
, fragment_codec(proof_params)
, params_(proof_params)
{
}
uint32_t matching_target(size_t table_id, uint64_t meta, uint32_t match_key)
{
size_t num_match_target_bits = params_.get_num_match_target_bits(table_id);
return hashing.matching_target(numeric_cast<uint32_t>(table_id),
match_key,
meta,
static_cast<int>(num_match_target_bits));
}
std::optional<T1Pairing> pairing_t1(uint32_t x_l, uint32_t x_r)
{
int const num_test_bits = params_.get_num_match_key_bits(1);
PairingResult pair = hashing.pairing_t1(x_l,
x_r,
static_cast<int>(params_.get_k()),
static_cast<int>(params_.get_num_pairing_meta_bits()),
num_test_bits);
if (pair.test_result != 0) {
return std::nullopt;
}
uint64_t const meta = (static_cast<uint64_t>(x_l) << params_.get_k()) | x_r;
return T1Pairing::make(meta, pair.match_info_result);
}
std::optional<T2Pairing> pairing_t2(uint64_t const meta_l, uint64_t meta_r)
{
int const num_test_bits = params_.get_num_match_key_bits(2);
PairingResult pair = hashing.pairing_t2(meta_l,
meta_r,
static_cast<int>(params_.get_k()),
static_cast<int>(params_.get_num_pairing_meta_bits()),
num_test_bits);
if (pair.test_result != 0) {
return std::nullopt;
}
T2Pairing result;
result.match_info = pair.match_info_result;
result.meta = pair.meta_result;
uint32_t half_k = params_.get_k() / 2;
uint32_t x_bits_l = numeric_cast<uint32_t>((meta_l >> params_.get_k()) >> half_k);
uint32_t x_bits_r = numeric_cast<uint32_t>((meta_r >> params_.get_k()) >> half_k);
result.x_bits = (x_bits_l << half_k) | x_bits_r;
return result;
}
std::optional<T3Pairing> pairing_t3(
uint64_t meta_l, uint64_t meta_r, uint32_t x_bits_l, uint32_t x_bits_r)
{
int const num_test_bits = params_.get_num_match_key_bits(3);
PairingResult pair = hashing.pairing_t3(meta_l, meta_r, num_test_bits);
if (pair.test_result != 0)
return std::nullopt;
uint64_t all_x_bits = (static_cast<uint64_t>(x_bits_l) << params_.get_k()) | x_bits_r;
ProofFragment proof_fragment = fragment_codec.encode(all_x_bits);
T3Pairing result;
result.proof_fragment = proof_fragment;
return result;
}
bool validate_match_info_pairing(
int table_id, uint64_t meta_l, uint32_t match_info_l, uint32_t match_info_r)
{
uint32_t section_l = params_.extract_section_from_match_info(table_id, match_info_l);
uint32_t section_r = params_.extract_section_from_match_info(table_id, match_info_r);
uint32_t match_section = matching_section(section_l);
if (section_r != match_section) {
return false;
}
uint32_t match_key_r = params_.extract_match_key_from_match_info(table_id, match_info_r);
uint32_t match_target_r
= params_.extract_match_target_from_match_info(table_id, match_info_r);
if (match_target_r != matching_target(table_id, meta_l, match_key_r)) {
return false;
}
return true;
}
uint32_t matching_section(uint32_t section)
{
uint32_t num_section_bits = params_.get_num_section_bits();
uint32_t num_sections = params_.get_num_sections();
uint32_t rotated_left = (section << 1) | (section >> (num_section_bits - 1));
uint32_t rotated_left_plus_1 = (rotated_left + 1) & (num_sections - 1);
uint32_t section_new
= (rotated_left_plus_1 >> 1) | (rotated_left_plus_1 << (num_section_bits - 1));
return section_new & (num_sections - 1);
}
uint32_t inverse_matching_section(uint32_t section)
{
uint32_t num_section_bits = params_.get_num_section_bits();
uint32_t num_sections = params_.get_num_sections();
uint32_t rotated_left
= ((section << 1) | (section >> (num_section_bits - 1))) & (num_sections - 1);
uint32_t rotated_left_minus_1 = (rotated_left - 1) & (num_sections - 1);
uint32_t section_l
= ((rotated_left_minus_1 >> 1) | (rotated_left_minus_1 << (num_section_bits - 1)))
& (num_sections - 1);
return section_l;
}
void get_matching_sections(uint32_t section, uint32_t& section1, uint32_t& section2)
{
section1 = matching_section(section);
section2 = inverse_matching_section(section);
}
struct SelectedChallengeSets {
std::array<uint32_t, NUM_CHALLENGE_SETS> fragment_set_indexes;
std::array<Range, NUM_CHALLENGE_SETS> fragment_set_ranges;
};
SelectedChallengeSets selectChallengeSets(std::span<uint8_t const, 32> const challenge)
{
BlakeHash::Result256 grouped_challenge_hash = hashing.challengeWithPlotIdHash(challenge);
uint32_t num_chaining_sets_bits = params_.get_num_chaining_sets_bits();
assert(num_chaining_sets_bits >= 2);
uint32_t const sets_mask = (1U << num_chaining_sets_bits) - 1U;
static_assert(NUM_CHALLENGE_SETS == 4,
"selectChallengeSets currently assumes NUM_CHALLENGE_SETS == 4");
uint32_t const high_bits_mask = sets_mask & ~uint32_t(NUM_CHALLENGE_SETS - 1);
SelectedChallengeSets out {};
for (uint32_t i = 0; i < NUM_CHALLENGE_SETS; ++i) {
uint32_t const set_index = (grouped_challenge_hash.r[i] & high_bits_mask) | i;
out.fragment_set_indexes[i] = set_index;
out.fragment_set_ranges[i] = params_.get_chaining_set_range(set_index);
}
return out;
}
ProofParams getProofParams() const { return params_; }
uint32_t quality_chain_pass_threshold_ = 0;
private:
ProofParams params_;
};