1use std::{fs::File, io::Write, iter::repeat_with, slice};
4
5use binius_field::Field;
6use binius_utils::{DeserializeBytes, SerializeBytes};
7use bytes::{Buf, BufMut, Bytes, BytesMut};
8
9use super::{
10 error::Error,
11 fiat_shamir::{Challenger, FiatShamirBuf},
12};
13use crate::fiat_shamir::{CanSample, CanSampleBits, sample_bits_reader};
14
15pub const MAX_GRINDING_BITS: usize = 32;
21
22#[derive(Debug, Clone, Copy)]
24pub struct Options {
25 pub debug_assertions: bool,
27}
28
29impl Default for Options {
30 fn default() -> Self {
31 Self {
32 debug_assertions: cfg!(debug_assertions),
33 }
34 }
35}
36
37#[derive(Debug, Clone)]
43pub struct VerifierTranscript<Challenger> {
44 combined: FiatShamirBuf<Bytes, Challenger>,
45 options: Options,
46}
47
48impl<Challenger_: Challenger> VerifierTranscript<Challenger_> {
49 pub fn new(challenger: Challenger_, vec: Vec<u8>) -> Self {
50 Self::with_opts(challenger, vec, Options::default())
51 }
52
53 pub fn with_opts(challenger: Challenger_, vec: Vec<u8>, options: Options) -> Self {
54 Self {
55 combined: FiatShamirBuf {
56 buffer: Bytes::from(vec),
57 challenger,
58 },
59 options,
60 }
61 }
62
63 pub fn finalize(self) -> Result<(), Error> {
64 if self.combined.buffer.has_remaining() {
65 return Err(Error::TranscriptNotEmpty {
66 remaining: self.combined.buffer.remaining(),
67 });
68 }
69 Ok(())
70 }
71
72 pub fn observe<'a, 'b>(&'a mut self) -> TranscriptWriter<'b, impl BufMut + 'b>
77 where
78 'a: 'b,
79 {
80 TranscriptWriter {
81 buffer: self.combined.challenger.observer(),
82 options: self.options,
83 }
84 }
85
86 pub fn decommitment(&mut self) -> TranscriptReader<'_, impl Buf + '_> {
92 TranscriptReader {
93 buffer: &mut self.combined.buffer,
94 options: self.options,
95 }
96 }
97
98 pub fn message<'a, 'b>(&'a mut self) -> TranscriptReader<'b, impl Buf>
102 where
103 'a: 'b,
104 {
105 TranscriptReader {
106 buffer: &mut self.combined,
107 options: self.options,
108 }
109 }
110
111 pub fn verify_grind(&mut self, bits: usize) -> Result<u64, Error> {
127 assert!(
128 bits <= MAX_GRINDING_BITS,
129 "precondition: bits must be at most {MAX_GRINDING_BITS}"
130 );
131
132 let nonce = self.message().read::<u64>()?;
133 let sampled = CanSampleBits::<u32>::sample_bits(self, bits);
134 if sampled != 0 {
135 return Err(Error::InsufficientWork { bits, sampled });
136 }
137 Ok(nonce)
138 }
139}
140
141impl<Challenger> Drop for VerifierTranscript<Challenger> {
143 fn drop(&mut self) {
144 if self.combined.buffer.has_remaining() {
145 tracing::warn!(
146 "Transcript reader is not fully read out: {:?} bytes left",
147 self.combined.buffer.remaining()
148 );
149 }
150 }
151}
152
153impl<F, Challenger_> CanSample<F> for VerifierTranscript<Challenger_>
154where
155 F: Field,
156 Challenger_: Challenger,
157{
158 fn sample(&mut self) -> F {
159 DeserializeBytes::deserialize(self.combined.challenger.sampler())
160 .expect("challenger has infinite buffer")
161 }
162}
163
164impl<Challenger_> CanSampleBits<u32> for VerifierTranscript<Challenger_>
165where
166 Challenger_: Challenger,
167{
168 fn sample_bits(&mut self, bits: usize) -> u32 {
169 sample_bits_reader(self.combined.challenger.sampler(), bits)
170 }
171}
172
173pub struct TranscriptReader<'a, B: Buf> {
174 buffer: &'a mut B,
175 options: Options,
176}
177
178impl<B: Buf> TranscriptReader<'_, B> {
179 pub const fn buffer(&mut self) -> &mut B {
180 self.buffer
181 }
182
183 pub fn read<T: DeserializeBytes>(&mut self) -> Result<T, Error> {
184 T::deserialize(self.buffer()).map_err(Into::into)
185 }
186
187 pub fn read_vec<T: DeserializeBytes>(&mut self, n: usize) -> Result<Vec<T>, Error> {
188 let mut buffer = self.buffer();
189 repeat_with(move || T::deserialize(&mut buffer).map_err(Into::into))
190 .take(n)
191 .collect()
192 }
193
194 pub fn read_bytes(&mut self, buf: &mut [u8]) -> Result<(), Error> {
195 let buffer = self.buffer();
196 if buffer.remaining() < buf.len() {
197 return Err(Error::NotEnoughBytes);
198 }
199 buffer.copy_to_slice(buf);
200 Ok(())
201 }
202
203 pub fn read_scalar<F: Field>(&mut self) -> Result<F, Error> {
204 let mut out = F::default();
205 self.read_scalar_slice_into(slice::from_mut(&mut out))?;
206 Ok(out)
207 }
208
209 pub fn read_scalar_slice_into<F: Field>(&mut self, buf: &mut [F]) -> Result<(), Error> {
210 let mut buffer = self.buffer();
211 for elem in buf {
212 *elem = DeserializeBytes::deserialize(&mut buffer)?;
213 }
214 Ok(())
215 }
216
217 pub fn read_scalar_slice<F: Field>(&mut self, len: usize) -> Result<Vec<F>, Error> {
218 let mut elems = vec![F::default(); len];
219 self.read_scalar_slice_into(&mut elems)?;
220 Ok(elems)
221 }
222
223 pub fn read_debug(&mut self, msg: &str) {
224 if self.options.debug_assertions {
225 let msg_bytes = msg.as_bytes();
226 let mut buffer = vec![0; msg_bytes.len()];
227 assert!(self.read_bytes(&mut buffer).is_ok());
228 assert_eq!(msg_bytes, buffer);
229 }
230 }
231}
232
233#[derive(Debug, Clone)]
239pub struct ProverTranscript<Challenger> {
240 combined: FiatShamirBuf<BytesMut, Challenger>,
241 options: Options,
242}
243
244impl<Challenger_: Challenger> ProverTranscript<Challenger_> {
245 pub fn new(challenger: Challenger_) -> Self {
249 Self::with_opts(challenger, Options::default())
250 }
251
252 pub fn with_opts(challenger: Challenger_, options: Options) -> Self {
253 Self {
254 combined: FiatShamirBuf {
255 buffer: BytesMut::default(),
256 challenger,
257 },
258 options,
259 }
260 }
261
262 pub fn finalize(self) -> Vec<u8> {
263 let transcript = self.combined.buffer.to_vec();
264
265 let proof_size_bytes = transcript.len();
267 tracing::event!(
268 name: "proof_size",
269 tracing::Level::INFO,
270 category = "metrics",
271 proof_size_bytes = proof_size_bytes,
272 );
273
274 if let Ok(path) = std::env::var("BINIUS_DUMP_PROOF") {
276 let path = if cfg!(test) {
277 let current_thread = std::thread::current();
280 let test_name = current_thread.name().unwrap_or("unknown");
281 let rebased = path.strip_prefix("./").map(|s| format!("../../{s}"));
284 let path = rebased.unwrap_or(path);
285 std::fs::create_dir_all(&path)
286 .unwrap_or_else(|_| panic!("Failed to create directories for path: {path}",));
287 format!("{path}/{test_name}.bin")
288 } else {
289 path
290 };
291
292 let mut file = File::create(&path)
293 .unwrap_or_else(|_| panic!("Failed to create proof dump file: {path}"));
294 file.write_all(&transcript)
295 .expect("Failed to write proof to dump file");
296 }
297 transcript
298 }
299
300 pub fn observe<'a, 'b>(&'a mut self) -> TranscriptWriter<'b, impl BufMut + 'b>
305 where
306 'a: 'b,
307 {
308 TranscriptWriter {
309 buffer: self.combined.challenger.observer(),
310 options: self.options,
311 }
312 }
313
314 pub fn decommitment(&mut self) -> TranscriptWriter<'_, impl BufMut> {
323 TranscriptWriter {
324 buffer: &mut self.combined.buffer,
325 options: self.options,
326 }
327 }
328
329 pub fn message<'a, 'b>(&'a mut self) -> TranscriptWriter<'b, impl BufMut>
333 where
334 'a: 'b,
335 {
336 TranscriptWriter {
337 buffer: &mut self.combined,
338 options: self.options,
339 }
340 }
341
342 pub fn grind(&mut self, bits: usize) -> u64
361 where
362 Challenger_: Clone,
363 {
364 assert!(
365 bits <= MAX_GRINDING_BITS,
366 "precondition: bits must be at most {MAX_GRINDING_BITS}"
367 );
368
369 let mut nonce = 0u64;
373 loop {
374 let mut trial = self.combined.challenger.clone();
375 trial.observer().put_slice(&nonce.to_le_bytes());
376 if sample_bits_reader(trial.sampler(), bits) == 0 {
377 break;
378 }
379 nonce += 1;
380 }
381
382 self.message().write(&nonce);
383 let sampled = CanSampleBits::<u32>::sample_bits(self, bits);
384 debug_assert_eq!(sampled, 0, "the nonce search only exits on a landing nonce");
385 nonce
386 }
387}
388
389impl<Challenger_: Default + Challenger> ProverTranscript<Challenger_> {
390 pub fn into_verifier(self) -> VerifierTranscript<Challenger_> {
391 let options = self.options;
392 let transcript = self.finalize();
393
394 VerifierTranscript::with_opts(Challenger_::default(), transcript, options)
395 }
396}
397
398impl<Challenger_: Default + Challenger> Default for ProverTranscript<Challenger_> {
399 fn default() -> Self {
400 Self::new(Challenger_::default())
401 }
402}
403
404pub struct TranscriptWriter<'a, B: BufMut> {
410 buffer: &'a mut B,
411 options: Options,
412}
413
414impl<B: BufMut> TranscriptWriter<'_, B> {
415 pub const fn buffer(&mut self) -> &mut B {
416 self.buffer
417 }
418
419 pub fn write<T: SerializeBytes>(&mut self, value: &T) {
426 self.proof_size_event_wrapper(move |buffer| {
427 value
428 .serialize(buffer)
429 .expect("serialization to a growable transcript buffer is infallible");
430 });
431 }
432
433 pub fn write_slice<T: SerializeBytes>(&mut self, values: &[T]) {
440 self.proof_size_event_wrapper(move |buffer| {
441 for value in values {
442 value
443 .serialize(&mut *buffer)
444 .expect("serialization to a growable transcript buffer is infallible");
445 }
446 });
447 }
448
449 pub fn write_bytes(&mut self, data: &[u8]) {
450 self.proof_size_event_wrapper(|buffer| {
451 buffer.put_slice(data);
452 });
453 }
454
455 pub fn write_scalar<F: Field>(&mut self, f: F) {
456 self.write_scalar_slice(slice::from_ref(&f));
457 }
458
459 pub fn write_scalar_iter<F: Field>(&mut self, it: impl IntoIterator<Item = F>) {
466 self.proof_size_event_wrapper(move |buffer| {
467 for elem in it {
468 SerializeBytes::serialize(&elem, &mut *buffer)
469 .expect("serialization to a growable transcript buffer is infallible");
470 }
471 });
472 }
473
474 pub fn write_scalar_slice<F: Field>(&mut self, elems: &[F]) {
475 self.write_scalar_iter(elems.iter().copied());
476 }
477
478 pub fn write_debug(&mut self, msg: &str) {
479 if self.options.debug_assertions {
480 self.write_bytes(msg.as_bytes());
481 }
482 }
483
484 fn proof_size_event_wrapper<F: FnOnce(&mut B)>(&mut self, f: F) {
485 let buffer = self.buffer();
486 let start_bytes = buffer.remaining_mut();
487 f(buffer);
488 let end_bytes = buffer.remaining_mut();
489 tracing::event!(
490 name: "incremental_proof_size",
491 tracing::Level::TRACE,
492 counter=true,
493 incremental=true,
494 value=start_bytes - end_bytes,
495 );
496 }
497}
498
499impl<F, Challenger_> CanSample<F> for ProverTranscript<Challenger_>
500where
501 F: Field,
502 Challenger_: Challenger,
503{
504 fn sample(&mut self) -> F {
505 DeserializeBytes::deserialize(self.combined.challenger.sampler())
506 .expect("challenger has infinite buffer")
507 }
508}
509
510impl<Challenger_> CanSampleBits<u32> for ProverTranscript<Challenger_>
511where
512 Challenger_: Challenger,
513{
514 fn sample_bits(&mut self, bits: usize) -> u32 {
515 sample_bits_reader(self.combined.challenger.sampler(), bits)
516 }
517}
518
519#[cfg(test)]
520mod tests {
521 use binius_field::Ghash128b as B128;
522 use sha2::Sha256;
523
524 use super::*;
525 use crate::fiat_shamir::{CanSample, HasherChallenger};
526
527 #[test]
528 fn test_transcript_interactions() {
529 let mut prover_transcript = ProverTranscript::new(HasherChallenger::<Sha256>::default());
530
531 prover_transcript
533 .message()
534 .write_scalar(B128::new(0x11111111222222223333333344444444));
535 prover_transcript
536 .message()
537 .write_scalar(B128::new(0xAAAAAAAABBBBBBBBCCCCCCCCDDDDDDDD));
538
539 prover_transcript
541 .decommitment()
542 .write_scalar(B128::new(0x5555555566666666777777778888888));
543
544 prover_transcript
546 .observe()
547 .write_scalar(B128::new(0xFFFFFFFFEEEEEEEEDDDDDDDDCCCCCCCC));
548
549 let sampled_challenge: B128 = prover_transcript.sample();
551
552 let mut verifier_transcript = prover_transcript.into_verifier();
554
555 let msg1: B128 = verifier_transcript.message().read_scalar().unwrap();
557 let msg2: B128 = verifier_transcript.message().read_scalar().unwrap();
558 assert_eq!(msg1, B128::new(0x11111111222222223333333344444444));
559 assert_eq!(msg2, B128::new(0xAAAAAAAABBBBBBBBCCCCCCCCDDDDDDDD));
560
561 let decommit: B128 = verifier_transcript.decommitment().read_scalar().unwrap();
563 assert_eq!(decommit, B128::new(0x5555555566666666777777778888888));
564
565 verifier_transcript
567 .observe()
568 .write_scalar(B128::new(0xFFFFFFFFEEEEEEEEDDDDDDDDCCCCCCCC));
569
570 let verifier_challenge: B128 = verifier_transcript.sample();
572 assert_eq!(verifier_challenge, sampled_challenge);
573
574 verifier_transcript.finalize().unwrap();
576 }
577
578 #[test]
579 fn test_transcript_debug() {
580 let options = Options {
581 debug_assertions: true,
582 };
583 let mut transcript =
584 ProverTranscript::with_opts(HasherChallenger::<Sha256>::default(), options);
585
586 transcript.message().write_debug("test_transcript_debug");
587 transcript
588 .into_verifier()
589 .message()
590 .read_debug("test_transcript_debug");
591 }
592
593 #[test]
594 #[should_panic]
595 fn test_transcript_debug_fail() {
596 let options = Options {
597 debug_assertions: true,
598 };
599 let mut transcript =
600 ProverTranscript::with_opts(HasherChallenger::<Sha256>::default(), options);
601
602 transcript.message().write_debug("test_transcript_debug");
603 transcript
604 .into_verifier()
605 .message()
606 .read_debug("test_transcript_debug_should_fail");
607 }
608 #[test]
609 fn grinding_round_trips_and_keeps_both_challengers_in_step() {
610 const BITS: usize = 12;
611
612 let mut prover = ProverTranscript::new(HasherChallenger::<Sha256>::default());
613 prover.message().write_scalar(B128::new(7));
614 let nonce = prover.grind(BITS);
615 assert!(nonce > 0);
617 let prover_challenge: B128 = prover.sample();
618
619 let mut verifier = prover.into_verifier();
620 let echoed: B128 = verifier.message().read_scalar().unwrap();
621 assert_eq!(echoed, B128::new(7));
622 assert_eq!(verifier.verify_grind(BITS).unwrap(), nonce);
623
624 let verifier_challenge: B128 = verifier.sample();
626 assert_eq!(verifier_challenge, prover_challenge);
627 verifier.finalize().unwrap();
628 }
629
630 #[test]
631 fn zero_difficulty_is_met_by_the_first_nonce() {
632 let mut prover = ProverTranscript::new(HasherChallenger::<Sha256>::default());
633 assert_eq!(prover.grind(0), 0);
635
636 let mut verifier = prover.into_verifier();
638 assert_eq!(verifier.verify_grind(0).unwrap(), 0);
639 verifier.finalize().unwrap();
640 }
641
642 #[test]
643 fn a_tampered_nonce_is_rejected() {
644 const BITS: usize = 12;
645
646 let mut prover = ProverTranscript::new(HasherChallenger::<Sha256>::default());
647 prover.grind(BITS);
648 let mut tape = prover.finalize();
649
650 tape[0] ^= 1;
653
654 let mut verifier =
655 VerifierTranscript::new(HasherChallenger::<Sha256>::default(), tape.clone());
656 let err = verifier.verify_grind(BITS).unwrap_err();
657 let Error::InsufficientWork { bits, sampled } = err else {
658 panic!("expected InsufficientWork, got {err:?}");
659 };
660 assert_eq!(bits, BITS);
661 assert_ne!(sampled, 0);
662 assert!(sampled < 1 << BITS);
664 verifier.finalize().unwrap();
665 }
666
667 #[test]
668 fn an_empty_tape_has_no_nonce_to_read() {
669 let mut verifier = VerifierTranscript::new(HasherChallenger::<Sha256>::default(), vec![]);
670 let err = verifier.verify_grind(8).unwrap_err();
671 let Error::Serialization(inner) = err else {
672 panic!("expected Serialization, got {err:?}");
673 };
674 assert!(matches!(inner, binius_utils::SerializationError::NotEnoughBytes));
675 verifier.finalize().unwrap();
676 }
677
678 #[test]
679 #[should_panic(expected = "bits must be at most 32")]
680 fn a_difficulty_past_the_sampler_width_is_rejected() {
681 ProverTranscript::new(HasherChallenger::<Sha256>::default()).grind(MAX_GRINDING_BITS + 1);
682 }
683}