binius_iop_prover/channel/
naive.rs1use binius_compute::GlobalAllocator;
6use binius_field::{Field, PackedField};
7use binius_iop::channel::OracleSpec;
8use binius_ip_prover::channel::IPProverChannel;
9use binius_math::{FieldBuffer, FieldSlice, StructuredBuffer};
10use binius_transcript::{
11 ProverTranscript,
12 fiat_shamir::{CanSample, Challenger},
13};
14
15use crate::channel::IOPProverChannel;
16
17#[derive(Debug, Clone, Copy)]
19pub struct NaiveOracle {
20 index: usize,
21}
22
23pub struct NaiveProverChannel<'a, F, Challenger_>
34where
35 F: Field,
36 Challenger_: Challenger,
37{
38 transcript: &'a mut ProverTranscript<Challenger_>,
40 oracle_specs: Vec<OracleSpec>,
42 n_committed: usize,
44 next_oracle_index: usize,
46 _f: std::marker::PhantomData<F>,
47}
48
49impl<'a, F, Challenger_> NaiveProverChannel<'a, F, Challenger_>
50where
51 F: Field,
52 Challenger_: Challenger,
53{
54 pub const fn new(
61 transcript: &'a mut ProverTranscript<Challenger_>,
62 oracle_specs: Vec<OracleSpec>,
63 ) -> Self {
64 Self {
65 transcript,
66 oracle_specs,
67 n_committed: 0,
68 next_oracle_index: 0,
69 _f: std::marker::PhantomData,
70 }
71 }
72
73 pub const fn transcript(&self) -> &ProverTranscript<Challenger_> {
75 self.transcript
76 }
77
78 pub fn finish(self) {
80 let n_remaining = self.oracle_specs.len() - self.next_oracle_index;
81 assert!(n_remaining == 0, "finish called but {n_remaining} oracle specs remaining",);
82 }
83}
84
85impl<F, Challenger_> IPProverChannel<F> for NaiveProverChannel<'_, F, Challenger_>
86where
87 F: Field,
88 Challenger_: Challenger,
89{
90 fn send_one(&mut self, elem: F) {
91 self.transcript.message().write_scalar(elem);
92 }
93
94 fn send_many(&mut self, elems: &[F]) {
95 self.transcript.message().write_scalar_slice(elems);
96 }
97
98 fn observe_one(&mut self, val: F) {
99 self.transcript.observe().write_scalar(val);
100 }
101
102 fn observe_many(&mut self, vals: &[F]) {
103 self.transcript.observe().write_scalar_slice(vals);
104 }
105
106 fn sample(&mut self) -> F {
107 CanSample::sample(&mut self.transcript)
108 }
109}
110
111impl<F, P, Challenger_> IOPProverChannel<P, GlobalAllocator>
115 for NaiveProverChannel<'_, F, Challenger_>
116where
117 F: Field,
118 P: PackedField<Scalar = F>,
119 Challenger_: Challenger,
120{
121 type Oracle = NaiveOracle;
122
123 fn remaining_oracle_specs(&self) -> &[OracleSpec] {
124 &self.oracle_specs[self.next_oracle_index..]
125 }
126
127 fn send_oracle(&mut self, buffer: FieldSlice<'_, P>) -> Self::Oracle {
128 let index = self.next_oracle_index;
129 assert!(
130 index < self.oracle_specs.len(),
131 "send_oracle called but no remaining oracle specs"
132 );
133
134 let spec_log_msg_len = self.oracle_specs[index].log_msg_len;
136 assert_eq!(
137 buffer.log_len(),
138 spec_log_msg_len,
139 "oracle buffer log_len mismatch: expected {spec_log_msg_len}, got {}",
140 buffer.log_len()
141 );
142
143 self.transcript
146 .message()
147 .write_scalar_iter(buffer.iter_scalars());
148
149 self.n_committed += 1;
150 self.next_oracle_index += 1;
151
152 NaiveOracle { index }
153 }
154
155 fn prove_oracle_relation(
156 &mut self,
157 oracle: Self::Oracle,
158 transparent: StructuredBuffer<P, Vec<P>>,
159 _claim: P::Scalar,
160 ) {
161 let index = oracle.index;
164 assert!(index < self.n_committed, "oracle index {index} out of bounds");
165
166 let log_msg_len = self.oracle_specs[index].log_msg_len;
167 assert_eq!(
168 transparent.log_len(),
169 log_msg_len,
170 "transparent log_len mismatch: expected {log_msg_len}, got {}",
171 transparent.log_len()
172 );
173
174 let transparent = transparent.materialize(&GlobalAllocator);
176 self.transcript
177 .message()
178 .write_scalar_iter(transparent.iter_scalars());
179
180 let _point: Vec<F> = CanSample::sample_vec(&mut self.transcript, log_msg_len);
182 }
183
184 fn finalize_oracle(&mut self, _oracle: Self::Oracle, _buffer: FieldBuffer<P>) {}
189}