use std::io::{Error, Read, Result, Write};
use derivative::Derivative;
use p3_field::PrimeCharacteristicRing;
use serde::{Deserialize, Serialize};
use crate::{
codec::{vec_with_capped_capacity, DecodableConfig, Decode, EncodableConfig, Encode},
StarkProtocolConfig,
};
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(
Clone(bound = ""),
Debug(bound = ""),
PartialEq(bound = ""),
Eq(bound = "")
)]
#[serde(bound = "")]
pub struct Proof<SC: StarkProtocolConfig> {
pub common_main_commit: SC::Digest,
pub trace_vdata: Vec<Option<TraceVData<SC>>>,
pub public_values: Vec<Vec<SC::F>>,
pub gkr_proof: GkrProof<SC>,
pub batch_constraint_proof: BatchConstraintProof<SC>,
pub stacking_proof: StackingProof<SC>,
pub whir_proof: WhirProof<SC>,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(
Clone(bound = ""),
Debug(bound = ""),
PartialEq(bound = ""),
Eq(bound = ""),
Default(bound = "")
)]
#[serde(bound = "")]
pub struct TraceVData<SC: StarkProtocolConfig> {
pub log_height: usize,
pub cached_commitments: Vec<SC::Digest>,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(
Clone(bound = ""),
Debug(bound = ""),
PartialEq(bound = ""),
Eq(bound = "")
)]
#[serde(bound = "")]
pub struct GkrProof<SC: StarkProtocolConfig> {
pub logup_pow_witness: SC::F,
pub q0_claim: SC::EF,
pub claims_per_layer: Vec<GkrLayerClaims<SC>>,
pub sumcheck_polys: Vec<Vec<[SC::EF; 3]>>,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(
Clone(bound = ""),
Debug(bound = ""),
PartialEq(bound = ""),
Eq(bound = "")
)]
#[serde(bound = "")]
pub struct GkrLayerClaims<SC: StarkProtocolConfig> {
pub p_xi_0: SC::EF,
pub p_xi_1: SC::EF,
pub q_xi_0: SC::EF,
pub q_xi_1: SC::EF,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(
Clone(bound = ""),
Debug(bound = ""),
PartialEq(bound = ""),
Eq(bound = "")
)]
#[serde(bound = "")]
pub struct BatchConstraintProof<SC: StarkProtocolConfig> {
pub numerator_term_per_air: Vec<SC::EF>,
pub denominator_term_per_air: Vec<SC::EF>,
pub univariate_round_coeffs: Vec<SC::EF>,
pub sumcheck_round_polys: Vec<Vec<SC::EF>>,
pub column_openings: Vec<Vec<Vec<SC::EF>>>,
}
pub fn column_openings_by_rot<'a, EF: PrimeCharacteristicRing + Copy + 'a>(
openings: &'a [EF],
need_rot: bool,
) -> Box<dyn Iterator<Item = (EF, EF)> + 'a> {
if need_rot {
Box::new(openings.chunks_exact(2).map(|chunk| (chunk[0], chunk[1])))
} else {
Box::new(openings.iter().map(|&claim| (claim, EF::ZERO)))
}
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(
Clone(bound = ""),
Debug(bound = ""),
PartialEq(bound = ""),
Eq(bound = "")
)]
#[serde(bound = "")]
pub struct StackingProof<SC: StarkProtocolConfig> {
pub univariate_round_coeffs: Vec<SC::EF>,
pub sumcheck_round_polys: Vec<[SC::EF; 2]>,
pub stacking_openings: Vec<Vec<SC::EF>>,
}
pub type MerkleProof<Digest> = Vec<Digest>;
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(
Clone(bound = ""),
Debug(bound = ""),
PartialEq(bound = ""),
Eq(bound = "")
)]
#[serde(bound = "")]
pub struct WhirProof<SC: StarkProtocolConfig> {
pub mu_pow_witness: SC::F,
pub whir_sumcheck_polys: Vec<[SC::EF; 2]>,
pub codeword_commits: Vec<SC::Digest>,
pub ood_values: Vec<SC::EF>,
pub folding_pow_witnesses: Vec<SC::F>,
pub query_phase_pow_witnesses: Vec<SC::F>,
pub initial_round_opened_rows: Vec<Vec<Vec<Vec<SC::F>>>>,
pub initial_round_merkle_proofs: Vec<Vec<MerkleProof<SC::Digest>>>,
pub codeword_opened_values: Vec<Vec<Vec<SC::EF>>>,
pub codeword_merkle_proofs: Vec<Vec<MerkleProof<SC::Digest>>>,
pub final_poly: Vec<SC::EF>,
}
impl<SC: EncodableConfig> Encode for TraceVData<SC> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
self.log_height.encode(writer)?;
SC::encode_digest_slice(&self.cached_commitments, writer)
}
}
impl<SC: EncodableConfig> Encode for GkrLayerClaims<SC> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
SC::encode_extension_field(&self.p_xi_0, writer)?;
SC::encode_extension_field(&self.p_xi_1, writer)?;
SC::encode_extension_field(&self.q_xi_0, writer)?;
SC::encode_extension_field(&self.q_xi_1, writer)?;
Ok(())
}
}
pub(crate) const CODEC_VERSION: u32 = 3;
impl<SC: EncodableConfig> Encode for Proof<SC> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
CODEC_VERSION.encode(writer)?;
SC::encode_digest(&self.common_main_commit, writer)?;
let num_airs: usize = self.trace_vdata.len();
num_airs.encode(writer)?;
for chunk in self.trace_vdata.chunks(8) {
let mut ret = 0u8;
for (i, vdata) in chunk.iter().enumerate() {
ret |= (vdata.is_some() as u8) << (i as u8);
}
ret.encode(writer)?;
}
for vdata in self.trace_vdata.iter().flatten() {
vdata.encode(writer)?;
}
self.public_values.len().encode(writer)?;
for pv in &self.public_values {
SC::encode_base_field_slice(pv, writer)?;
}
self.gkr_proof.encode(writer)?;
self.batch_constraint_proof.encode(writer)?;
self.stacking_proof.encode(writer)?;
self.whir_proof.encode(writer)
}
}
impl<SC: EncodableConfig> Encode for GkrProof<SC> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
SC::encode_base_field(&self.logup_pow_witness, writer)?;
SC::encode_extension_field(&self.q0_claim, writer)?;
self.claims_per_layer.encode(writer)?;
for round in &self.sumcheck_polys {
for arr in round {
SC::encode_extension_field_iter(arr.iter(), writer)?;
}
}
Ok(())
}
}
impl<SC: EncodableConfig> Encode for BatchConstraintProof<SC> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
SC::encode_extension_field_slice(&self.numerator_term_per_air, writer)?;
SC::encode_extension_field_iter(self.denominator_term_per_air.iter(), writer)?;
SC::encode_extension_field_slice(&self.univariate_round_coeffs, writer)?;
let n_max = self.sumcheck_round_polys.len();
n_max.encode(writer)?;
if n_max > 0 {
self.sumcheck_round_polys[0].len().encode(writer)?;
for round_polys in &self.sumcheck_round_polys {
SC::encode_extension_field_iter(round_polys.iter(), writer)?;
}
}
for part_col_openings in &self.column_openings {
part_col_openings.len().encode(writer)?;
for col_opening in part_col_openings {
SC::encode_extension_field_slice(col_opening, writer)?;
}
}
Ok(())
}
}
impl<SC: EncodableConfig> Encode for StackingProof<SC> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
SC::encode_extension_field_slice(&self.univariate_round_coeffs, writer)?;
self.sumcheck_round_polys.len().encode(writer)?;
for arr in &self.sumcheck_round_polys {
SC::encode_extension_field_iter(arr.iter(), writer)?;
}
self.stacking_openings.len().encode(writer)?;
for opening in &self.stacking_openings {
SC::encode_extension_field_slice(opening, writer)?;
}
Ok(())
}
}
impl<SC: EncodableConfig> Encode for WhirProof<SC> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
SC::encode_base_field(&self.mu_pow_witness, writer)?;
self.whir_sumcheck_polys.len().encode(writer)?;
for arr in &self.whir_sumcheck_polys {
SC::encode_extension_field_iter(arr.iter(), writer)?;
}
let num_whir_sumcheck_rounds = self.whir_sumcheck_polys.len();
SC::encode_digest_slice(&self.codeword_commits, writer)?;
SC::encode_extension_field_iter(self.ood_values.iter(), writer)?;
let num_whir_rounds = self.codeword_commits.len() + 1;
if !num_whir_sumcheck_rounds.is_multiple_of(num_whir_rounds) {
return Err(Error::new(
std::io::ErrorKind::InvalidData,
"num_whir_sumcheck_rounds must be a multiple of num_whir_rounds",
));
}
assert_eq!(num_whir_rounds, self.query_phase_pow_witnesses.len());
SC::encode_base_field_iter(self.folding_pow_witnesses.iter(), writer)?;
SC::encode_base_field_iter(self.query_phase_pow_witnesses.iter(), writer)?;
let num_commits = self.initial_round_opened_rows.len();
assert!(num_commits > 0);
num_commits.encode(writer)?;
let initial_num_whir_queries = self.initial_round_opened_rows[0].len();
initial_num_whir_queries.encode(writer)?;
if initial_num_whir_queries > 0 {
let merkle_depth = self.initial_round_merkle_proofs[0][0].len();
merkle_depth.encode(writer)?;
let widths: Vec<usize> = self
.initial_round_opened_rows
.iter()
.map(|commit_rows| {
commit_rows
.first()
.and_then(|q| q.first())
.map(|row| row.len())
.unwrap_or(0)
})
.collect();
for w in &widths {
w.encode(writer)?;
}
for (commit_rows, &width) in self.initial_round_opened_rows.iter().zip(&widths) {
debug_assert_eq!(commit_rows.len(), initial_num_whir_queries);
for query_rows in commit_rows {
for row in query_rows {
debug_assert_eq!(row.len(), width);
SC::encode_base_field_iter(row.iter(), writer)?;
}
}
}
for merkle_proofs in &self.initial_round_merkle_proofs {
for proof in merkle_proofs {
SC::encode_digest_iter(proof.iter(), writer)?;
}
}
}
for non_init_round in &self.codeword_opened_values {
let num_queries = non_init_round.len();
num_queries.encode(writer)?;
for query_vals in non_init_round {
SC::encode_extension_field_iter(query_vals.iter(), writer)?;
}
}
let mut first_merkle_depth = 0;
if num_whir_rounds > 1 && initial_num_whir_queries > 0 {
first_merkle_depth = self.codeword_merkle_proofs[0][0].len();
}
first_merkle_depth.encode(writer)?;
for round_proofs in &self.codeword_merkle_proofs {
for proof in round_proofs {
SC::encode_digest_iter(proof.iter(), writer)?;
}
}
SC::encode_extension_field_slice(&self.final_poly, writer)
}
}
impl<SC: DecodableConfig> Decode for TraceVData<SC> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
Ok(Self {
log_height: usize::decode(reader)?,
cached_commitments: SC::decode_digest_vec(reader)?,
})
}
}
impl<SC: DecodableConfig> Decode for GkrLayerClaims<SC> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
Ok(Self {
p_xi_0: SC::decode_extension_field(reader)?,
p_xi_1: SC::decode_extension_field(reader)?,
q_xi_0: SC::decode_extension_field(reader)?,
q_xi_1: SC::decode_extension_field(reader)?,
})
}
}
impl<SC: DecodableConfig> Decode for Proof<SC> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let codec_version = u32::decode(reader)?;
if codec_version != CODEC_VERSION {
return Err(Error::other(format!(
"CODEC_VERSION mismatch, expected: {}, actual: {}",
CODEC_VERSION, codec_version
)));
}
let common_main_commit = SC::decode_digest(reader)?;
let num_airs = usize::decode(reader)?;
let bitmap_len = num_airs.div_ceil(8);
let mut bitmap: Vec<u8> = vec_with_capped_capacity(bitmap_len);
for _ in 0..bitmap_len {
bitmap.push(u8::decode(reader)?);
}
let mut trace_vdata = vec_with_capped_capacity(num_airs);
for byte in bitmap {
for i in 0u8..8 {
if trace_vdata.len() >= num_airs {
if byte >> i != 0 {
return Err(Error::other("trace_vdata bitmap padding bits set"));
}
break;
}
if byte & (1u8 << i) != 0 {
trace_vdata.push(Some(TraceVData::decode(reader)?));
} else {
trace_vdata.push(None);
}
}
}
let num_pvs = usize::decode(reader)?;
let mut public_values = vec_with_capped_capacity(num_pvs);
for _ in 0..num_pvs {
public_values.push(SC::decode_base_field_vec(reader)?);
}
Ok(Self {
common_main_commit,
trace_vdata,
public_values,
gkr_proof: GkrProof::decode(reader)?,
batch_constraint_proof: BatchConstraintProof::decode(reader)?,
stacking_proof: StackingProof::decode(reader)?,
whir_proof: WhirProof::decode(reader)?,
})
}
}
impl<SC: DecodableConfig> Decode for GkrProof<SC> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let logup_pow_witness = SC::decode_base_field(reader)?;
let q0_claim = SC::decode_extension_field(reader)?;
let claims_per_layer = Vec::<GkrLayerClaims<SC>>::decode(reader)?;
let num_sumcheck_polys = claims_per_layer.len().saturating_sub(1);
let mut sumcheck_polys = vec_with_capped_capacity(num_sumcheck_polys);
for round_idx_minus_one in 0..num_sumcheck_polys {
let n = round_idx_minus_one + 1;
let mut round = vec_with_capped_capacity(n);
for _ in 0..n {
round.push([
SC::decode_extension_field(reader)?,
SC::decode_extension_field(reader)?,
SC::decode_extension_field(reader)?,
]);
}
sumcheck_polys.push(round);
}
Ok(Self {
logup_pow_witness,
q0_claim,
claims_per_layer,
sumcheck_polys,
})
}
}
impl<SC: DecodableConfig> Decode for BatchConstraintProof<SC> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let numerator_term_per_air = SC::decode_extension_field_vec(reader)?;
let num_present_airs = numerator_term_per_air.len();
let denominator_term_per_air = SC::decode_extension_field_n(reader, num_present_airs)?;
let univariate_round_coeffs = SC::decode_extension_field_vec(reader)?;
let n_max = usize::decode(reader)?;
let mut sumcheck_round_polys = vec_with_capped_capacity(n_max);
if n_max > 0 {
let max_degree_plus_one = usize::decode(reader)?;
for _ in 0..n_max {
sumcheck_round_polys
.push(SC::decode_extension_field_n(reader, max_degree_plus_one)?);
}
}
let mut column_openings = vec_with_capped_capacity(num_present_airs);
for _ in 0..num_present_airs {
let num_parts = usize::decode(reader)?;
let mut parts = vec_with_capped_capacity(num_parts);
for _ in 0..num_parts {
parts.push(SC::decode_extension_field_vec(reader)?);
}
column_openings.push(parts);
}
Ok(Self {
numerator_term_per_air,
denominator_term_per_air,
univariate_round_coeffs,
sumcheck_round_polys,
column_openings,
})
}
}
impl<SC: DecodableConfig> Decode for StackingProof<SC> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let univariate_round_coeffs = SC::decode_extension_field_vec(reader)?;
let num_rounds = usize::decode(reader)?;
let mut sumcheck_round_polys = vec_with_capped_capacity(num_rounds);
for _ in 0..num_rounds {
sumcheck_round_polys.push([
SC::decode_extension_field(reader)?,
SC::decode_extension_field(reader)?,
]);
}
let num_openings = usize::decode(reader)?;
let mut stacking_openings = vec_with_capped_capacity(num_openings);
for _ in 0..num_openings {
stacking_openings.push(SC::decode_extension_field_vec(reader)?);
}
Ok(Self {
univariate_round_coeffs,
sumcheck_round_polys,
stacking_openings,
})
}
}
impl<SC: DecodableConfig> Decode for WhirProof<SC> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let mu_pow_witness = SC::decode_base_field(reader)?;
let num_whir_sumcheck_rounds = usize::decode(reader)?;
let mut whir_sumcheck_polys = vec_with_capped_capacity(num_whir_sumcheck_rounds);
for _ in 0..num_whir_sumcheck_rounds {
whir_sumcheck_polys.push([
SC::decode_extension_field(reader)?,
SC::decode_extension_field(reader)?,
]);
}
let codeword_commits = SC::decode_digest_vec(reader)?;
let num_whir_rounds = codeword_commits.len() + 1;
if num_whir_sumcheck_rounds % num_whir_rounds != 0 {
return Err(Error::new(
std::io::ErrorKind::InvalidData,
"num_whir_sumcheck_rounds must be a multiple of num_whir_rounds",
));
}
let k_whir = num_whir_sumcheck_rounds / num_whir_rounds;
let ood_values = SC::decode_extension_field_n(reader, num_whir_rounds - 1)?;
let folding_pow_witnesses = SC::decode_base_field_n(reader, num_whir_sumcheck_rounds)?;
let query_phase_pow_witnesses = SC::decode_base_field_n(reader, num_whir_rounds)?;
let num_commits = usize::decode(reader)?;
assert!(num_commits > 0);
let initial_num_whir_queries = usize::decode(reader)?;
let k_whir_exp = 1 << k_whir;
let mut merkle_depth = 0;
if initial_num_whir_queries > 0 {
merkle_depth = usize::decode(reader)?;
}
let mut widths = vec_with_capped_capacity(num_commits);
if initial_num_whir_queries > 0 {
for _ in 0..num_commits {
widths.push(usize::decode(reader)?);
}
}
let decoded_widths = widths.len();
let mut initial_round_opened_rows = vec_with_capped_capacity(num_commits);
for width in widths.into_iter().chain(std::iter::repeat_n(
0,
num_commits.saturating_sub(decoded_widths),
)) {
let mut opened_rows = vec_with_capped_capacity(initial_num_whir_queries);
for _ in 0..initial_num_whir_queries {
let mut rows = vec_with_capped_capacity(k_whir_exp);
for _ in 0..k_whir_exp {
rows.push(SC::decode_base_field_n(reader, width)?);
}
opened_rows.push(rows);
}
initial_round_opened_rows.push(opened_rows);
}
let mut initial_round_merkle_proofs = vec_with_capped_capacity(num_commits);
for _ in 0..num_commits {
let mut merkle_proofs = vec_with_capped_capacity(initial_num_whir_queries);
for _ in 0..initial_num_whir_queries {
merkle_proofs.push(SC::decode_digest_n(reader, merkle_depth)?);
}
initial_round_merkle_proofs.push(merkle_proofs);
}
let mut codeword_opened_values = vec_with_capped_capacity(num_whir_rounds - 1);
for _ in 0..num_whir_rounds - 1 {
let num_queries = usize::decode(reader)?;
let mut opened_values = vec_with_capped_capacity(num_queries);
for _ in 0..num_queries {
opened_values.push(SC::decode_extension_field_n(reader, k_whir_exp)?);
}
codeword_opened_values.push(opened_values);
}
merkle_depth = usize::decode(reader)?;
let mut codeword_merkle_proofs = vec_with_capped_capacity(num_whir_rounds - 1);
for opened_values in codeword_opened_values.iter() {
let num_queries = opened_values.len();
let mut merkle_proof: Vec<_> = vec_with_capped_capacity(num_queries);
for _ in 0..num_queries {
merkle_proof.push(SC::decode_digest_n(reader, merkle_depth)?);
}
codeword_merkle_proofs.push(merkle_proof);
merkle_depth -= 1;
}
let final_poly = SC::decode_extension_field_vec(reader)?;
Ok(Self {
mu_pow_witness,
whir_sumcheck_polys,
codeword_commits,
ood_values,
folding_pow_witnesses,
query_phase_pow_witnesses,
initial_round_opened_rows,
initial_round_merkle_proofs,
codeword_opened_values,
codeword_merkle_proofs,
final_poly,
})
}
}