use p3_field::PrimeCharacteristicRing;
use thiserror::Error;
use tracing::debug;
use crate::{
poly_common::{eval_eq_mle, interpolate_cubic_at_0123, interpolate_linear_at_01},
proof::{GkrLayerClaims, GkrProof},
FiatShamirTranscript, StarkProtocolConfig,
};
#[derive(Debug, Error, PartialEq, Eq)]
pub enum GkrVerificationError<EF: core::fmt::Debug + core::fmt::Display + PartialEq + Eq> {
#[error("Zero-round proof: q0_claim should be 1, got {actual}")]
InvalidZeroRoundValue { actual: EF },
#[error("Zero-check failed: numerator at root should be zero, got {actual}")]
ZeroCheckFailed { actual: EF },
#[error("Denominator consistency check failed at root: expected {expected}, got {actual}")]
RootConsistencyCheckFailed { expected: EF, actual: EF },
#[error("Layer consistency check failed at round {round}: expected {expected}, got {actual}")]
LayerConsistencyCheckFailed {
round: usize,
expected: EF,
actual: EF,
},
#[error("Expected {expected} layers, got {actual}")]
IncorrectLayerCount { expected: usize, actual: usize },
#[error("Expected {expected} sumcheck polynomial entries, got {actual}")]
IncorrectSumcheckPolyCount { expected: usize, actual: usize },
#[error("Round {round} expected {expected} sumcheck sub-rounds, got {actual}")]
IncorrectSubroundCount {
round: usize,
expected: usize,
actual: usize,
},
}
#[allow(clippy::type_complexity)]
pub fn verify_gkr<SC: StarkProtocolConfig, TS: FiatShamirTranscript<SC>>(
proof: &GkrProof<SC>,
transcript: &mut TS,
total_rounds: usize,
) -> Result<(SC::EF, SC::EF, Vec<SC::EF>), GkrVerificationError<SC::EF>> {
assert!(total_rounds > 0);
if proof.claims_per_layer.len() != total_rounds {
return Err(GkrVerificationError::IncorrectLayerCount {
expected: total_rounds,
actual: proof.claims_per_layer.len(),
});
}
let expected_sumcheck_entries = total_rounds.saturating_sub(1);
if proof.sumcheck_polys.len() != expected_sumcheck_entries {
return Err(GkrVerificationError::IncorrectSumcheckPolyCount {
expected: expected_sumcheck_entries,
actual: proof.sumcheck_polys.len(),
});
}
transcript.observe_ext(proof.q0_claim);
let layer_claims = &proof.claims_per_layer[0];
observe_layer_claims::<SC, TS>(transcript, layer_claims);
let (p_cross_term, q_cross_term) = compute_recursive_relations::<SC>(layer_claims);
if p_cross_term != SC::EF::ZERO {
return Err(GkrVerificationError::ZeroCheckFailed {
actual: p_cross_term,
});
}
if q_cross_term != proof.q0_claim {
return Err(GkrVerificationError::RootConsistencyCheckFailed {
expected: proof.q0_claim,
actual: q_cross_term,
});
}
let mu = transcript.sample_ext();
debug!(gkr_round = 0, %mu);
let (mut numer_claim, mut denom_claim) = reduce_to_single_evaluation::<SC>(layer_claims, mu);
debug!(%numer_claim, %denom_claim);
let mut gkr_r = vec![mu];
for round in 1..total_rounds {
let lambda = transcript.sample_ext();
debug!(gkr_round = round, %lambda);
let claim = numer_claim + lambda * denom_claim;
let (new_claim, round_r, eq_at_r_prime) =
verify_gkr_sumcheck::<SC, TS>(proof, transcript, round, claim, &gkr_r)?;
debug_assert_eq!(eq_at_r_prime, eval_eq_mle(&gkr_r, &round_r));
let layer_claims = &proof.claims_per_layer[round];
observe_layer_claims::<SC, TS>(transcript, layer_claims);
let (p_cross_term, q_cross_term) = compute_recursive_relations::<SC>(layer_claims);
let expected_claim = (p_cross_term + lambda * q_cross_term) * eq_at_r_prime;
if expected_claim != new_claim {
return Err(GkrVerificationError::LayerConsistencyCheckFailed {
round,
expected: expected_claim,
actual: new_claim,
});
}
let mu = transcript.sample_ext();
debug!(gkr_round = round, %mu);
(numer_claim, denom_claim) = reduce_to_single_evaluation::<SC>(layer_claims, mu);
gkr_r = std::iter::once(mu).chain(round_r).collect();
}
Ok((numer_claim, denom_claim, gkr_r))
}
#[allow(clippy::type_complexity)]
fn verify_gkr_sumcheck<SC: StarkProtocolConfig, TS: FiatShamirTranscript<SC>>(
proof: &GkrProof<SC>,
transcript: &mut TS,
round: usize,
mut claim: SC::EF,
gkr_r: &[SC::EF],
) -> Result<(SC::EF, Vec<SC::EF>, SC::EF), GkrVerificationError<SC::EF>> {
debug_assert!(
round > 0,
"verify_gkr_sumcheck should not be called for round 0"
);
debug_assert_eq!(
gkr_r.len(),
round,
"gkr_r should have exactly round elements"
);
let expected_subrounds = round;
let polys = &proof.sumcheck_polys[round - 1];
if polys.len() != expected_subrounds {
return Err(GkrVerificationError::IncorrectSubroundCount {
round,
expected: expected_subrounds,
actual: polys.len(),
});
}
let mut gkr_r_prime = Vec::with_capacity(round);
let mut eq = SC::EF::ONE;
for (sumcheck_round, poly_evals) in polys.iter().enumerate() {
debug!(gkr_round = round, %sumcheck_round, sum_claim = %claim);
for &eval in poly_evals {
transcript.observe_ext(eval);
}
let ri = transcript.sample_ext();
gkr_r_prime.push(ri);
debug!(gkr_round = round, %sumcheck_round, r_round = %ri);
let ev0 = claim - poly_evals[0]; let evals = [ev0, poly_evals[0], poly_evals[1], poly_evals[2]];
claim = interpolate_cubic_at_0123(&evals, ri);
let xi = gkr_r[sumcheck_round];
eq *= xi * ri + (SC::EF::ONE - xi) * (SC::EF::ONE - ri);
}
Ok((claim, gkr_r_prime, eq))
}
fn observe_layer_claims<SC: StarkProtocolConfig, TS: FiatShamirTranscript<SC>>(
transcript: &mut TS,
claims: &GkrLayerClaims<SC>,
) {
transcript.observe_ext(claims.p_xi_0);
transcript.observe_ext(claims.q_xi_0);
transcript.observe_ext(claims.p_xi_1);
transcript.observe_ext(claims.q_xi_1);
}
fn compute_recursive_relations<SC: StarkProtocolConfig>(
claims: &GkrLayerClaims<SC>,
) -> (SC::EF, SC::EF) {
let p_cross_term = claims.p_xi_0 * claims.q_xi_1 + claims.p_xi_1 * claims.q_xi_0;
let q_cross_term = claims.q_xi_0 * claims.q_xi_1;
(p_cross_term, q_cross_term)
}
fn reduce_to_single_evaluation<SC: StarkProtocolConfig>(
claims: &GkrLayerClaims<SC>,
mu: SC::EF,
) -> (SC::EF, SC::EF) {
let numer = interpolate_linear_at_01(&[claims.p_xi_0, claims.p_xi_1], mu);
let denom = interpolate_linear_at_01(&[claims.q_xi_0, claims.q_xi_1], mu);
(numer, denom)
}