binius_spartan_prover/wrapper/
replay_channel.rs1use std::{
7 cell::RefCell,
8 rc::{Rc, Weak},
9 sync::Arc,
10 vec::IntoIter as VecIntoIter,
11};
12
13use binius_core::word::Word;
14use binius_field::{BinaryField, Field};
15use binius_iop::channel::{IOPVerifierChannel, OracleSpec, TransparentEvalFn};
16use binius_ip::channel::{
17 IPVerifierChannel, WordIPVerifierChannel, pack_words_concrete, select_word, subset_sum_word,
18};
19use binius_spartan_frontend::{
20 circuit_builder::{CircuitBuilder, WireAllocator, WitnessError, WitnessGenerator},
21 constraint_system::{WireKind, Witness, WitnessLayout},
22};
23use binius_spartan_verifier::wrapper::circuit_elem::CircuitElem;
24
25pub struct ReplayChannel<F: Field> {
35 witness_gen: Rc<RefCell<WitnessGenerator<F>>>,
36 inout_alloc: WireAllocator,
42 precommit_alloc: WireAllocator,
43 keys: VecIntoIter<F>,
44 events: VecIntoIter<F>,
45}
46
47impl<F: Field> ReplayChannel<F> {
48 pub fn new(layout: Arc<WitnessLayout<F>>, keys: Vec<F>, events: Vec<F>) -> Self {
56 Self {
57 witness_gen: Rc::new(RefCell::new(WitnessGenerator::new(layout))),
58 inout_alloc: WireAllocator::new(WireKind::InOut),
59 precommit_alloc: WireAllocator::new(WireKind::Precommit),
60 keys: keys.into_iter(),
61 events: events.into_iter(),
62 }
63 }
64
65 fn next_inout_elem(&mut self) -> CircuitElem<F, WitnessGenerator<F>> {
66 let value = self
67 .events
68 .next()
69 .unwrap_or_else(|| panic!("replay exhausted: no more events"));
70
71 self.alloc_inout_elem(value)
72 }
73
74 fn alloc_inout_elem(&mut self, value: F) -> CircuitElem<F, WitnessGenerator<F>> {
80 let wire = self.inout_alloc.alloc();
81 let witness_wire = self.witness_gen.borrow_mut().write_inout(wire, value);
82 CircuitElem::wire(&self.witness_gen, witness_wire)
83 }
84
85 fn next_precommit_elem(&mut self) -> CircuitElem<F, WitnessGenerator<F>> {
86 let value = self
87 .keys
88 .next()
89 .expect("precommit segment is sized incorrectly");
90
91 let wire = self.precommit_alloc.alloc();
92 let witness_wire = self.witness_gen.borrow_mut().write_precommit(wire, value);
93 CircuitElem::wire(&self.witness_gen, witness_wire)
94 }
95
96 pub fn finish(self) -> Result<Witness<F>, WitnessError> {
98 Rc::try_unwrap(self.witness_gen)
99 .expect("CircuitElem values should only hold Weak references")
100 .into_inner()
101 .build()
102 }
103}
104
105impl<F: Field> IPVerifierChannel<F> for ReplayChannel<F> {
106 type Elem = CircuitElem<F, WitnessGenerator<F>>;
107
108 fn recv_one(&mut self) -> Result<Self::Elem, binius_ip::channel::Error> {
109 let encrypted_elem = self.next_inout_elem();
110 let key = self.next_precommit_elem();
111 Ok(encrypted_elem + key)
112 }
113
114 fn recv_public_claim(&mut self) -> Result<Self::Elem, binius_ip::channel::Error> {
115 Ok(self.next_inout_elem())
118 }
119
120 fn sample(&mut self) -> Self::Elem {
121 self.next_inout_elem()
122 }
123
124 fn observe_one(&mut self, _val: F) -> Self::Elem {
125 self.next_inout_elem()
126 }
127
128 fn assert_zero(&mut self, val: Self::Elem) -> Result<(), binius_ip::channel::Error> {
129 match val {
130 CircuitElem::Constant(c) => {
133 if c == F::ZERO {
134 Ok(())
135 } else {
136 Err(binius_ip::channel::Error::InvalidAssert)
137 }
138 }
139 CircuitElem::Wire { builder, wire } => {
140 assert!(Weak::ptr_eq(&Rc::downgrade(&self.witness_gen), &builder));
141 self.witness_gen.borrow_mut().assert_zero(wire);
142 Ok(())
143 }
144 }
145 }
146}
147
148impl<F: BinaryField> WordIPVerifierChannel<F> for ReplayChannel<F> {
149 type Word = Word;
150
151 fn observe_words(&mut self, words: &[Word]) -> Vec<Word> {
154 words.to_vec()
155 }
156
157 fn subset_sum(&mut self, elems: &[Self::Elem], word: &Word) -> Self::Elem {
158 subset_sum_word(elems, *word)
159 }
160
161 fn select(&mut self, elems: &[Self::Elem], word: &Word) -> Self::Elem {
162 select_word(elems, *word)
163 }
164
165 fn sample_bits(&mut self, _bits: usize) -> Word {
166 Word::ZERO
167 }
168
169 fn pack_words(&mut self, words: &[Word]) -> Vec<Self::Elem> {
170 pack_words_concrete::<F, F>(words)
173 .into_iter()
174 .map(|value| self.alloc_inout_elem(value))
175 .collect()
176 }
177}
178
179impl<F: Field> IOPVerifierChannel<F> for ReplayChannel<F> {
180 type Oracle = ();
181
182 fn remaining_oracle_specs(&self) -> &[OracleSpec] {
183 &[]
184 }
185
186 fn recv_oracle(
187 &mut self,
188 _log_msg_len: usize,
189 _is_witness_dependent: bool,
190 ) -> Result<Self::Oracle, binius_iop::channel::Error> {
191 Ok(())
192 }
193
194 fn verify_oracle_relation(
195 &mut self,
196 _oracle: Self::Oracle,
197 _transparent: TransparentEvalFn<Self::Elem>,
198 claim: Self::Elem,
199 ) -> Result<(), binius_iop::channel::Error> {
200 let decrypted_claim = self.next_inout_elem();
204 self.assert_zero(claim - decrypted_claim)?;
205 Ok(())
206 }
207}
208
209#[cfg(test)]
210mod tests {
211 use std::{cell::RefCell, rc::Rc, sync::Arc};
212
213 use binius_field::{BinaryField1b as B1, ExtensionField, Ghash128b as B128, field::FieldOps};
214 use binius_spartan_frontend::circuit_builder::{ConstraintBuilder, WitnessGenerator};
215 use binius_spartan_verifier::wrapper::circuit_elem::CircuitElem;
216
217 type BuildElem = CircuitElem<B128, ConstraintBuilder<B128>>;
218 type WitnessElem = CircuitElem<B128, WitnessGenerator<B128>>;
219
220 #[test]
221 fn test_square_transpose_wires() {
222 type FSub = B1;
225 let degree = <B128 as ExtensionField<FSub>>::DEGREE;
226
227 let mut constraint_builder = ConstraintBuilder::<B128>::new();
229 let inout_wires: Vec<_> = (0..degree)
230 .map(|_| constraint_builder.alloc_inout())
231 .collect();
232
233 let rc = Rc::new(RefCell::new(constraint_builder));
235 let mut elems: Vec<BuildElem> = inout_wires
236 .iter()
237 .map(|&w| BuildElem::wire(&rc, w))
238 .collect();
239
240 <BuildElem as FieldOps>::square_transpose::<FSub>(&mut elems);
241
242 drop(elems);
244 let constraint_builder = Rc::try_unwrap(rc).unwrap().into_inner();
245 let (cs, layout) = constraint_builder.build().finalize();
246
247 assert!(!cs.mul_constraints().is_empty());
250
251 let test_values: Vec<B128> = (0..degree)
253 .map(<B128 as ExtensionField<FSub>>::basis)
254 .collect();
255
256 let layout = Arc::new(layout);
257 let mut witness_gen = WitnessGenerator::new(Arc::clone(&layout));
258 let witness_wires: Vec<_> = inout_wires
259 .iter()
260 .zip(&test_values)
261 .map(|(&wire, &val)| witness_gen.write_inout(wire, val))
262 .collect();
263
264 let witness_rc = Rc::new(RefCell::new(witness_gen));
265 let mut witness_elems: Vec<WitnessElem> = witness_wires
266 .iter()
267 .map(|&w| WitnessElem::wire(&witness_rc, w))
268 .collect();
269
270 <WitnessElem as FieldOps>::square_transpose::<FSub>(&mut witness_elems);
271
272 drop(witness_elems);
273 let witness_gen = Rc::try_unwrap(witness_rc).unwrap().into_inner();
274 let witness = witness_gen
275 .build()
276 .expect("witness generation should succeed (all constraints satisfied)");
277
278 cs.validate(&witness);
279 }
280}