|
1 | 1 | //! Malicious variant of the Sub prover for soundness testing. |
2 | 2 | //! |
3 | | -//! This module contains a Sub prover that forges virtual operand claims |
4 | | -//! while keeping the underlying sumcheck honest. It is used exclusively |
5 | | -//! by [`MaliciousONNXProof`] to test that the verifier correctly handles |
6 | | -//! (and rejects) such attacks. |
| 3 | +//! The honest `Sub` proof is a no-sumcheck op (see [`crate::onnx_proof::ops::sub`]): |
| 4 | +//! it opens both operands at the node's reduced output point `r` and the verifier |
| 5 | +//! checks `left(r) - right(r) == output(r)` directly. This module forges the left |
| 6 | +//! operand opening (off by one) so that `left(r) - right(r) != output(r)`, exercising |
| 7 | +//! that the verifier's direct difference check rejects the attack. Used exclusively by |
| 8 | +//! [`MaliciousONNXProof`](crate::onnx_proof::malicious_prover). |
7 | 9 |
|
8 | 10 | use crate::{ |
9 | | - onnx_proof::{malicious_prover::malicious_sumcheck_prove, ProofId, ProofType, Prover}, |
| 11 | + onnx_proof::{ProofId, Prover}, |
10 | 12 | utils::opening_access::{AccOpeningAccessor, Target}, |
11 | 13 | }; |
12 | 14 | use atlas_onnx_tracer::{ |
13 | 15 | model::trace::{LayerData, Trace}, |
14 | 16 | node::ComputationNode, |
15 | 17 | }; |
16 | 18 | use joltworks::{ |
17 | | - field::{IntoOpening, JoltField}, |
18 | | - poly::{ |
19 | | - eq_poly::EqPolynomial, |
20 | | - multilinear_polynomial::{BindingOrder, MultilinearPolynomial, PolynomialBinding}, |
21 | | - opening_proof::ProverOpeningAccumulator, |
22 | | - split_eq_poly::GruenSplitEqPolynomial, |
23 | | - unipoly::UniPoly, |
24 | | - }, |
25 | | - subprotocols::{ |
26 | | - sumcheck::SumcheckInstanceProof, sumcheck_prover::SumcheckInstanceProver, |
27 | | - sumcheck_verifier::SumcheckInstanceParams, |
28 | | - }, |
| 19 | + field::JoltField, |
| 20 | + poly::multilinear_polynomial::{MultilinearPolynomial, PolynomialEvaluation}, |
| 21 | + subprotocols::sumcheck::SumcheckInstanceProof, |
29 | 22 | transcripts::Transcript, |
30 | 23 | }; |
31 | 24 |
|
32 | | -use crate::onnx_proof::ops::sub::SubParams; |
33 | | - |
34 | 25 | /// Run the malicious Sub prover for a single node. |
35 | 26 | /// |
36 | | -/// Returns the proof entry suitable for insertion into the proof map. |
| 27 | +/// Mirrors the honest no-sumcheck `Sub::prove` but forges the left operand |
| 28 | +/// opening as `left(r) + 1`, leaving the right operand honest. The verifier's |
| 29 | +/// `left - right == output` check then necessarily fails. Returns `vec![]` (no |
| 30 | +/// execution proof), exactly like the honest op. |
37 | 31 | pub fn malicious_sub_prove<F: JoltField, T: Transcript>( |
38 | 32 | node: &ComputationNode, |
39 | 33 | prover: &mut Prover<F, T>, |
40 | 34 | ) -> Vec<(ProofId, SumcheckInstanceProof<F, T>)> { |
41 | | - let params = SubParams::new(node.clone(), &prover.accumulator); |
42 | | - let mut prover_sumcheck = MaliciousSubProver::initialize(&prover.trace, params); |
43 | | - |
44 | | - let (proof, r_sumcheck, final_claim) = malicious_sumcheck_prove( |
45 | | - &mut prover_sumcheck, |
46 | | - &mut prover.accumulator, |
47 | | - &mut prover.transcript, |
48 | | - ); |
49 | | - prover_sumcheck.final_claim = Some(final_claim); |
50 | | - prover_sumcheck.cache_openings(&mut prover.accumulator, &mut prover.transcript, &r_sumcheck); |
51 | | - vec![(ProofId(node.idx, ProofType::Execution), proof)] |
52 | | -} |
53 | | - |
54 | | -/// Malicious prover state for element-wise subtraction sumcheck protocol. |
55 | | -/// |
56 | | -/// Identical to the honest SubProver in sumcheck computation, but forges |
57 | | -/// operand claims in `cache_openings` to demonstrate the attack vector. |
58 | | -struct MaliciousSubProver<F: JoltField> { |
59 | | - params: SubParams<F>, |
60 | | - eq_r_node_output: GruenSplitEqPolynomial<F>, |
61 | | - left_operand: MultilinearPolynomial<F>, |
62 | | - right_operand: MultilinearPolynomial<F>, |
63 | | - final_claim: Option<F>, |
64 | | -} |
65 | | - |
66 | | -impl<F: JoltField> MaliciousSubProver<F> { |
67 | | - /// Initialize the prover with trace data and parameters. |
68 | | - fn initialize(trace: &Trace, params: SubParams<F>) -> Self { |
69 | | - let eq_r_node_output = |
70 | | - GruenSplitEqPolynomial::new(¶ms.r_node_output.r, BindingOrder::LowToHigh); |
71 | | - let LayerData { |
72 | | - operands, |
73 | | - output: _, |
74 | | - } = Trace::layer_data(trace, ¶ms.computation_node); |
75 | | - let [left_operand, right_operand] = operands[..] else { |
76 | | - panic!("Expected two operands for Sub operation") |
77 | | - }; |
78 | | - let left_operand = MultilinearPolynomial::from(left_operand.clone()); |
79 | | - let right_operand = MultilinearPolynomial::from(right_operand.clone()); |
80 | | - Self { |
81 | | - params, |
82 | | - eq_r_node_output, |
83 | | - left_operand, |
84 | | - right_operand, |
85 | | - final_claim: None, |
86 | | - } |
87 | | - } |
88 | | -} |
89 | | - |
90 | | -impl<F: JoltField, T: Transcript> SumcheckInstanceProver<F, T> for MaliciousSubProver<F> { |
91 | | - fn get_params(&self) -> &dyn SumcheckInstanceParams<F> { |
92 | | - &self.params |
93 | | - } |
94 | | - |
95 | | - fn compute_message(&mut self, _round: usize, previous_claim: F) -> UniPoly<F> { |
96 | | - let Self { |
97 | | - eq_r_node_output, |
98 | | - left_operand, |
99 | | - right_operand, |
100 | | - .. |
101 | | - } = self; |
102 | | - let [q_constant] = eq_r_node_output.par_fold_out_in_unreduced::<9, 1>(&|g| { |
103 | | - let lo0 = left_operand.get_bound_coeff(2 * g); |
104 | | - let ro0 = right_operand.get_bound_coeff(2 * g); |
105 | | - [lo0 - ro0] |
106 | | - }); |
107 | | - eq_r_node_output.gruen_poly_deg_2(q_constant, previous_claim) |
108 | | - } |
109 | | - |
110 | | - fn ingest_challenge(&mut self, r_j: F::Challenge, _round: usize) { |
111 | | - self.eq_r_node_output.bind(r_j); |
112 | | - self.left_operand |
113 | | - .bind_parallel(r_j, BindingOrder::LowToHigh); |
114 | | - self.right_operand |
115 | | - .bind_parallel(r_j, BindingOrder::LowToHigh); |
116 | | - } |
117 | | - |
118 | | - fn cache_openings( |
119 | | - &self, |
120 | | - accumulator: &mut ProverOpeningAccumulator<F>, |
121 | | - transcript: &mut T, |
122 | | - sumcheck_challenges: &[F::Challenge], |
123 | | - ) { |
124 | | - let opening_point = self |
125 | | - .params |
126 | | - .normalize_opening_point(&sumcheck_challenges.into_opening()); |
127 | | - |
128 | | - // Malicious behavior: forge virtual operand claims while preserving the |
129 | | - // same subtraction difference, so expected_output_claim remains unchanged. |
130 | | - let left_claim = self.left_operand.final_claim(); |
131 | | - let right_claim = self.right_operand.final_claim(); |
132 | | - let final_claim = self |
133 | | - .final_claim |
134 | | - .expect("final_claim must be set before cache_openings"); |
135 | | - let r_node_output_prime = self |
136 | | - .params |
137 | | - .normalize_opening_point(&sumcheck_challenges.into_opening()) |
138 | | - .r; |
139 | | - let eq_eval = EqPolynomial::mle(&self.params.r_node_output.r, &r_node_output_prime); |
140 | | - |
141 | | - // Choose forged claims so that: |
142 | | - // final_claim == eq_eval * (forged_left - forged_right) |
143 | | - let forged_left = left_claim + F::one(); |
144 | | - let forged_right = if eq_eval.is_zero() { |
145 | | - // If eq_eval == 0, the only valid final claim is 0. |
146 | | - debug_assert!(final_claim.is_zero()); |
147 | | - right_claim |
148 | | - } else { |
149 | | - let inv = eq_eval |
150 | | - .inverse() |
151 | | - .expect("non-zero eq_eval must be invertible"); |
152 | | - forged_left - final_claim * inv |
153 | | - }; |
154 | | - debug_assert_eq!(final_claim, eq_eval * (forged_left - forged_right)); |
155 | | - let mut provider = AccOpeningAccessor::new(accumulator, &self.params.computation_node) |
156 | | - .into_provider(transcript, opening_point); |
157 | | - |
158 | | - // Insert forged claims while keeping transcript/opener state consistent. |
159 | | - provider.append_nodeio(Target::Input(0), forged_left); |
160 | | - provider.append_nodeio(Target::Input(1), forged_right); |
161 | | - } |
| 35 | + let (opening_point, _claim) = |
| 36 | + AccOpeningAccessor::new(&prover.accumulator, node).get_reduced_opening(); |
| 37 | + |
| 38 | + let LayerData { operands, .. } = Trace::layer_data(&prover.trace, node); |
| 39 | + let [left, right] = operands[..] else { |
| 40 | + panic!("Expected two operands for Sub operation") |
| 41 | + }; |
| 42 | + let forged_left = MultilinearPolynomial::from(left.padded_next_power_of_two()) |
| 43 | + .evaluate(&opening_point.r) |
| 44 | + + F::one(); |
| 45 | + let right_claim = |
| 46 | + MultilinearPolynomial::from(right.padded_next_power_of_two()).evaluate(&opening_point.r); |
| 47 | + |
| 48 | + let mut provider = AccOpeningAccessor::new(&mut prover.accumulator, node) |
| 49 | + .into_provider(&mut prover.transcript, opening_point); |
| 50 | + provider.append_nodeio(Target::Input(0), forged_left); |
| 51 | + provider.append_nodeio(Target::Input(1), right_claim); |
| 52 | + vec![] |
162 | 53 | } |
0 commit comments