1use std::cmp::Reverse;
6
7use binius_core::word::Word;
8use binius_field::{BinaryField, Field, field::FieldOps};
9use binius_ip::channel::{IPVerifierChannel, WordIPVerifierChannel};
10use binius_math::multilinear::eq::eq_ind;
11
12use crate::channel::{
13 Error, IOPVerifierChannel, OracleSchedule, OracleSpec, TransparentEvalFn, merged_log_msg_len,
14};
15
16#[derive(Debug, Clone, Copy)]
18pub struct MergeOracle {
19 index: usize,
20}
21
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24pub struct Placement {
25 pub round: usize,
27
28 pub block_index: usize,
32
33 pub combined_log_len: usize,
35}
36
37pub fn place_oracles(schedule: &OracleSchedule) -> Vec<Placement> {
44 let mut placements = Vec::with_capacity(schedule.specs().len());
45 for (round, specs) in schedule.rounds().enumerate() {
46 let combined_log_len = merged_log_msg_len(specs.iter().map(|spec| spec.log_msg_len));
47
48 let mut order: Vec<usize> = (0..specs.len()).collect();
49 order.sort_by_key(|&k| Reverse(specs[k].log_msg_len));
50
51 let mut block_indices = vec![0; specs.len()];
52 let mut offset = 0usize;
53 for k in order {
54 let n_k = specs[k].log_msg_len;
55 block_indices[k] = offset >> n_k;
56 offset += 1 << n_k;
57 }
58
59 placements.extend(block_indices.into_iter().map(|block_index| Placement {
60 round,
61 block_index,
62 combined_log_len,
63 }));
64 }
65 assert_eq!(placements.len(), schedule.specs().len(), "every oracle must be in a closed round");
66 placements
67}
68
69pub struct MergeVerifierChannel<'a, F, C>
142where
143 F: Field,
144 C: IOPVerifierChannel<F>,
145{
146 inner: C,
148
149 schedule: &'a OracleSchedule,
151
152 placements: Vec<Placement>,
154
155 outers: Vec<C::Oracle>,
157
158 round_specs: Vec<OracleSpec>,
162
163 n_received: usize,
165}
166
167impl<'a, F, C> MergeVerifierChannel<'a, F, C>
168where
169 F: Field,
170 C: IOPVerifierChannel<F>,
171{
172 pub fn new(inner: C, schedule: &'a OracleSchedule) -> Self {
184 let round_specs = schedule.merged_specs();
185 assert_eq!(
186 inner.remaining_oracle_specs(),
187 round_specs,
188 "inner channel must be configured with the schedule's merged specs"
189 );
190 Self {
191 inner,
192 schedule,
193 placements: place_oracles(schedule),
194 outers: Vec::new(),
195 round_specs,
196 n_received: 0,
197 }
198 }
199
200 pub fn into_inner(self) -> C {
206 let n_remaining = self.placements.len() - self.n_received;
207 assert!(n_remaining == 0, "into_inner called but {n_remaining} oracle specs remaining",);
208 self.inner
209 }
210}
211
212impl<F, C> IPVerifierChannel<F> for MergeVerifierChannel<'_, F, C>
213where
214 F: Field,
215 C: IOPVerifierChannel<F>,
216{
217 type Elem = C::Elem;
218
219 fn recv_one(&mut self) -> Result<Self::Elem, binius_ip::channel::Error> {
220 self.inner.recv_one()
221 }
222
223 fn recv_many(&mut self, n: usize) -> Result<Vec<Self::Elem>, binius_ip::channel::Error> {
224 self.inner.recv_many(n)
225 }
226
227 fn recv_array<const N: usize>(&mut self) -> Result<[Self::Elem; N], binius_ip::channel::Error> {
228 self.inner.recv_array()
229 }
230
231 fn recv_public_claim(&mut self) -> Result<Self::Elem, binius_ip::channel::Error> {
232 self.inner.recv_public_claim()
233 }
234
235 fn sample(&mut self) -> Self::Elem {
236 self.inner.sample()
237 }
238
239 fn observe_one(&mut self, val: F) -> Self::Elem {
240 self.inner.observe_one(val)
241 }
242
243 fn observe_many(&mut self, vals: &[F]) -> Vec<Self::Elem> {
244 self.inner.observe_many(vals)
245 }
246
247 fn assert_zero(&mut self, val: Self::Elem) -> Result<(), binius_ip::channel::Error> {
248 self.inner.assert_zero(val)
249 }
250}
251
252impl<F, C> WordIPVerifierChannel<F> for MergeVerifierChannel<'_, F, C>
253where
254 F: BinaryField,
255 C: IOPVerifierChannel<F> + WordIPVerifierChannel<F>,
256{
257 type Word = C::Word;
258
259 fn observe_words(&mut self, words: &[Word]) -> Vec<Self::Word> {
260 self.inner.observe_words(words)
261 }
262
263 fn subset_sum(&mut self, elems: &[Self::Elem], word: &Self::Word) -> Self::Elem {
264 self.inner.subset_sum(elems, word)
265 }
266
267 fn select(&mut self, elems: &[Self::Elem], word: &Self::Word) -> Self::Elem {
268 self.inner.select(elems, word)
269 }
270
271 fn sample_bits(&mut self, bits: usize) -> Self::Word {
272 self.inner.sample_bits(bits)
273 }
274
275 fn pack_words(&mut self, words: &[Self::Word]) -> Vec<Self::Elem> {
276 self.inner.pack_words(words)
277 }
278}
279
280impl<'a, F, C> IOPVerifierChannel<F> for MergeVerifierChannel<'a, F, C>
281where
282 F: Field,
283 C: IOPVerifierChannel<F>,
284{
285 type Oracle = MergeOracle;
286
287 fn remaining_oracle_specs(&self) -> &[OracleSpec] {
288 &self.schedule.specs()[self.n_received..]
289 }
290
291 fn recv_oracle(
292 &mut self,
293 log_msg_len: usize,
294 is_witness_dependent: bool,
295 ) -> Result<Self::Oracle, Error> {
296 let remaining = self.remaining_oracle_specs();
301 assert!(!remaining.is_empty(), "recv_oracle called but no remaining oracle specs");
302 let spec = remaining[0];
303 assert_eq!(log_msg_len, spec.log_msg_len, "oracle size must match its spec");
304
305 assert!(
309 !spec.is_zk || is_witness_dependent,
310 "a zero-knowledge oracle spec must be received as witness-dependent"
311 );
312
313 let index = self.n_received;
314 self.n_received += 1;
315
316 let Placement {
318 round,
319 combined_log_len,
320 ..
321 } = self.placements[index];
322 let is_last_of_round = self
323 .placements
324 .get(index + 1)
325 .is_none_or(|next| next.round != round);
326 if is_last_of_round {
327 let outer = self
333 .inner
334 .recv_oracle(combined_log_len, self.round_specs[round].is_zk)?;
335 self.outers.push(outer);
336 }
337
338 Ok(MergeOracle { index })
339 }
340
341 fn verify_oracle_relation(
342 &mut self,
343 oracle: Self::Oracle,
344 transparent: TransparentEvalFn<Self::Elem>,
345 claim: Self::Elem,
346 ) -> Result<(), Error> {
347 let n_i = self.schedule.specs()[oracle.index].log_msg_len;
348 let Placement {
349 round,
350 block_index,
351 combined_log_len,
352 } = self.placements[oracle.index];
353 let outer = self.outers[round].clone();
354
355 let padding_len = combined_log_len - n_i;
361 let block_pattern: Vec<Self::Elem> = (0..padding_len)
362 .map(|bit| {
363 if (block_index >> bit) & 1 == 1 {
364 Self::Elem::one()
365 } else {
366 Self::Elem::zero()
367 }
368 })
369 .collect();
370
371 let padded_transparent: TransparentEvalFn<Self::Elem> = Box::new(move |point| {
376 let (low, high) = point.split_at(n_i);
377 eq_ind(high, &block_pattern) * transparent(low)
378 });
379
380 self.inner
385 .verify_oracle_relation(outer, padded_transparent, claim)
386 }
387}
388
389#[cfg(test)]
390mod tests {
391 use binius_field::Ghash128b;
392 use binius_hash::StdDigest;
393 use binius_transcript::{VerifierTranscript, fiat_shamir::HasherChallenger};
394
395 use super::*;
396 use crate::channel::naive::NaiveVerifierChannel;
397
398 type F = Ghash128b;
399
400 fn schedule(rounds: &[&[usize]]) -> OracleSchedule {
402 let mut schedule = OracleSchedule::new();
403 for sizes in rounds {
404 for &n in *sizes {
405 schedule.push(OracleSpec::new(n));
406 }
407 schedule.end_round();
408 }
409 schedule
410 }
411
412 #[test]
413 fn placements_sort_each_round_largest_first() {
414 let placements = place_oracles(&schedule(&[&[2, 4, 2], &[1]]));
419 let place = |round, block_index, combined_log_len| Placement {
420 round,
421 block_index,
422 combined_log_len,
423 };
424 assert_eq!(
425 placements,
426 [
427 place(0, 4, 5),
428 place(0, 0, 5),
429 place(0, 5, 5),
430 place(1, 0, 1)
431 ]
432 );
433 }
434
435 #[test]
436 #[should_panic(expected = "recv_oracle called but no remaining oracle specs")]
437 fn recv_oracle_past_remaining_specs_panics() {
438 let schedule = schedule(&[&[2]]);
443 let merged_specs = schedule.merged_specs();
444 let mut transcript =
445 VerifierTranscript::new(HasherChallenger::<StdDigest>::default(), vec![0; 4 * 16]);
446 let mut channel = MergeVerifierChannel::new(
447 NaiveVerifierChannel::<F, _>::new(&mut transcript, &merged_specs),
448 &schedule,
449 );
450 channel.recv_oracle(2, true).unwrap();
451 channel.recv_oracle(2, true).unwrap();
452 }
453
454 #[test]
455 #[should_panic(expected = "a zero-knowledge oracle spec must be received as witness-dependent")]
456 fn zk_spec_received_as_structural_panics() {
457 let mut schedule = OracleSchedule::new();
460 schedule.push(OracleSpec::new_zk(2));
461 schedule.end_round();
462 let merged_specs = schedule.merged_specs();
463 let mut transcript =
464 VerifierTranscript::new(HasherChallenger::<StdDigest>::default(), Vec::new());
465 let mut channel = MergeVerifierChannel::new(
466 NaiveVerifierChannel::<F, _>::new(&mut transcript, &merged_specs),
467 &schedule,
468 );
469 let _ = channel.recv_oracle(2, false);
470 }
471
472 #[test]
473 #[should_panic(expected = "into_inner called but 1 oracle specs remaining")]
474 fn into_inner_before_all_specs_received_panics() {
475 let schedule = schedule(&[&[2, 2]]);
480 let merged_specs = schedule.merged_specs();
481 let mut transcript =
482 VerifierTranscript::new(HasherChallenger::<StdDigest>::default(), Vec::new());
483 let mut channel = MergeVerifierChannel::new(
484 NaiveVerifierChannel::<F, _>::new(&mut transcript, &merged_specs),
485 &schedule,
486 );
487 channel.recv_oracle(2, true).unwrap();
488 let _ = channel.into_inner();
489 }
490}