1use std::{
5 cmp::Ordering,
6 collections::HashMap,
7 mem,
8 ops::{Index, IndexMut},
9};
10
11use binius_field::Field;
12use binius_utils::checked_arithmetics::log2_ceil_usize;
13use smallvec::{SmallVec, smallvec};
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
16pub enum WireKind {
17 Constant,
18 InOut,
19 Derived,
23 Precommit,
24 Private,
25}
26
27impl WireKind {
28 pub const fn segment(self) -> WitnessSegment {
31 match self {
32 WireKind::Constant | WireKind::InOut | WireKind::Derived => WitnessSegment::Public,
33 WireKind::Precommit => WitnessSegment::Precommit,
34 WireKind::Private => WitnessSegment::Private,
35 }
36 }
37
38 pub fn is_public(self) -> bool {
46 self.segment() == WitnessSegment::Public
47 }
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
51pub struct ConstraintWire {
52 pub(crate) kind: WireKind,
53 pub(crate) id: u32,
54}
55
56impl ConstraintWire {
57 pub const fn inout(id: u32) -> Self {
61 Self {
62 kind: WireKind::InOut,
63 id,
64 }
65 }
66
67 pub const fn precommit(id: u32) -> Self {
69 Self {
70 kind: WireKind::Precommit,
71 id,
72 }
73 }
74}
75
76#[derive(Debug, Clone)]
77pub struct Operand<W>(SmallVec<[W; 4]>);
78
79impl<W> Default for Operand<W> {
80 fn default() -> Self {
81 Operand(SmallVec::new())
82 }
83}
84
85impl<W: Copy + Ord> Operand<W> {
86 pub fn new(mut term: SmallVec<[W; 4]>) -> Self {
87 term.sort_unstable();
88
89 let has_duplicate_wire = term.windows(2).any(|w| w[0] == w[1]);
90 let term = if has_duplicate_wire {
91 term.chunk_by(|a, b| a == b)
92 .flat_map(|group| {
93 let last_even_idx = group.len() / 2 * 2;
96 group[last_even_idx..].iter().copied()
97 })
98 .collect()
99 } else {
100 term
101 };
102
103 Self(term)
104 }
105
106 pub fn len(&self) -> usize {
107 self.0.len()
108 }
109
110 pub fn is_empty(&self) -> bool {
111 self.0.is_empty()
112 }
113
114 pub fn wires(&self) -> &[W] {
115 &self.0
116 }
117
118 pub fn merge(&mut self, rhs: &Self) -> (Operand<W>, Operand<W>) {
119 let lhs = mem::take(&mut self.0);
121 let dst = &mut self.0;
122
123 let mut lhs_iter = lhs.into_iter().peekable();
124 let mut rhs_iter = rhs.0.iter().copied().peekable();
125
126 let mut additions = Operand::default();
127 let mut removals = Operand::default();
128
129 loop {
130 match (lhs_iter.peek(), rhs_iter.peek()) {
131 (Some(next_lhs), Some(next_rhs)) => {
132 match next_lhs.cmp(next_rhs) {
133 Ordering::Equal => {
134 let wire = lhs_iter.next().expect("peek returned Some");
136 let _ = rhs_iter.next().expect("peek returned Some");
137
138 removals.0.push(wire);
139 }
140 Ordering::Less => dst.push(lhs_iter.next().expect("peek returned Some")),
141 Ordering::Greater => {
142 let wire = rhs_iter.next().expect("peek returned Some");
143 additions.0.push(wire);
144 dst.push(wire);
145 }
146 }
147 }
148 (Some(_), None) => dst.push(lhs_iter.next().expect("peek returned Some")),
149 (None, Some(_)) => {
150 let wire = rhs_iter.next().expect("peek returned Some");
151 additions.0.push(wire);
152 dst.push(wire);
153 }
154 (None, None) => break,
155 }
156 }
157
158 (additions, removals)
159 }
160}
161
162impl<W> From<W> for Operand<W> {
163 fn from(value: W) -> Self {
164 Operand(smallvec![value])
165 }
166}
167
168#[derive(Debug, Clone)]
169pub struct MulConstraint<W> {
170 pub a: Operand<W>,
171 pub b: Operand<W>,
172 pub c: Operand<W>,
173}
174
175#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
176pub enum WitnessSegment {
177 Public,
179 Precommit,
182 Private,
185}
186
187#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
188pub struct WitnessIndex {
189 pub segment: WitnessSegment,
190 pub index: u32,
191}
192
193impl WitnessIndex {
194 pub const fn public(index: u32) -> Self {
195 Self {
196 segment: WitnessSegment::Public,
197 index,
198 }
199 }
200
201 pub const fn precommit(index: u32) -> Self {
202 Self {
203 segment: WitnessSegment::Precommit,
204 index,
205 }
206 }
207
208 pub const fn private(index: u32) -> Self {
209 Self {
210 segment: WitnessSegment::Private,
211 index,
212 }
213 }
214}
215
216pub struct Witness<F> {
217 public: Vec<F>,
218 precommit: Vec<F>,
219 private: Vec<F>,
220}
221
222impl<F> Witness<F> {
223 pub const fn new(public: Vec<F>, precommit: Vec<F>, private: Vec<F>) -> Self {
224 Self {
225 public,
226 precommit,
227 private,
228 }
229 }
230
231 pub fn public(&self) -> &[F] {
232 &self.public
233 }
234
235 pub fn precommit(&self) -> &[F] {
236 &self.precommit
237 }
238
239 pub fn private(&self) -> &[F] {
240 &self.private
241 }
242}
243
244impl<F> Index<WitnessIndex> for Witness<F> {
245 type Output = F;
246
247 fn index(&self, index: WitnessIndex) -> &Self::Output {
248 match index.segment {
249 WitnessSegment::Public => &self.public[index.index as usize],
250 WitnessSegment::Precommit => &self.precommit[index.index as usize],
251 WitnessSegment::Private => &self.private[index.index as usize],
252 }
253 }
254}
255
256impl<F> IndexMut<WitnessIndex> for Witness<F> {
257 fn index_mut(&mut self, index: WitnessIndex) -> &mut Self::Output {
258 match index.segment {
259 WitnessSegment::Public => &mut self.public[index.index as usize],
260 WitnessSegment::Precommit => &mut self.precommit[index.index as usize],
261 WitnessSegment::Private => &mut self.private[index.index as usize],
262 }
263 }
264}
265
266#[derive(Debug, Clone)]
274pub struct ConstraintSystem<F: Field> {
275 constants: Vec<F>,
276 n_inout: u32,
277 n_precommit: u32,
278 n_private: u32,
279 log_public: u32,
280 mul_constraints: Vec<MulConstraint<WitnessIndex>>,
281 one_wire_index: u32,
282}
283
284impl<F: Field> ConstraintSystem<F> {
285 pub const fn new(
287 constants: Vec<F>,
288 n_inout: u32,
289 n_precommit: u32,
290 n_private: u32,
291 log_public: u32,
292 mul_constraints: Vec<MulConstraint<WitnessIndex>>,
293 one_wire_index: u32,
294 ) -> Self {
295 Self {
296 constants,
297 n_inout,
298 n_precommit,
299 n_private,
300 log_public,
301 mul_constraints,
302 one_wire_index,
303 }
304 }
305
306 pub fn constants(&self) -> &[F] {
307 &self.constants
308 }
309
310 pub const fn n_inout(&self) -> u32 {
311 self.n_inout
312 }
313
314 pub const fn n_precommit(&self) -> u32 {
315 self.n_precommit
316 }
317
318 pub const fn n_private(&self) -> u32 {
319 self.n_private
320 }
321
322 pub const fn log_public(&self) -> u32 {
323 self.log_public
324 }
325
326 pub const fn n_public(&self) -> u32 {
327 1 << self.log_public
328 }
329
330 pub fn mul_constraints(&self) -> &[MulConstraint<WitnessIndex>] {
331 &self.mul_constraints
332 }
333
334 pub const fn one_wire(&self) -> WitnessIndex {
335 WitnessIndex {
336 segment: WitnessSegment::Public,
337 index: self.one_wire_index,
338 }
339 }
340
341 pub fn validate(&self, witness: &Witness<F>) {
343 let operand_val = |operand: &Operand<WitnessIndex>| {
344 operand.wires().iter().map(|&idx| witness[idx]).sum::<F>()
345 };
346
347 for MulConstraint { a, b, c } in &self.mul_constraints {
348 assert_eq!(operand_val(a) * operand_val(b), operand_val(c));
349 }
350 }
351}
352
353#[derive(Debug, Clone, Copy)]
358pub struct BlindingInfo {
359 pub n_dummy_wires: usize,
361 pub n_dummy_constraints: usize,
363}
364
365const N_DUMMY_CONSTRAINTS: usize = 2;
387
388impl BlindingInfo {
389 pub const fn for_fri_queries(n_test_queries: usize) -> Self {
399 Self {
400 n_dummy_wires: n_test_queries + 1,
401 n_dummy_constraints: N_DUMMY_CONSTRAINTS,
402 }
403 }
404}
405
406#[derive(Debug, Clone)]
407pub struct WitnessLayout<F: Field> {
408 pub(crate) constants: Vec<F>,
409 n_inout: u32,
410 n_derived: u32,
411 n_precommit: u32,
412 n_private: u32,
413 log_public: u32,
414 log_precommit: u32,
415 log_private: u32,
416 derived_index_map: HashMap<u32, u32>,
417 private_index_map: HashMap<u32, u32>,
418}
419
420impl<F: Field> WitnessLayout<F> {
421 pub fn sparse(
422 constants: Vec<F>,
423 n_inout: u32,
424 n_precommit: u32,
425 derived_alive: &[bool],
426 private_alive: &[bool],
427 ) -> Self {
428 let n_constants = constants.len() as u32;
429 let log_precommit = log2_ceil_usize(n_precommit as usize) as u32;
430
431 let derived_index_map = derived_alive
435 .iter()
436 .enumerate()
437 .filter(|(_, alive)| **alive)
438 .enumerate()
439 .map(|(new_idx, (id, _))| (id as u32, new_idx as u32))
440 .collect::<HashMap<_, _>>();
441
442 let n_derived = derived_index_map.len() as u32;
443 let n_public = n_constants + n_inout + n_derived;
444 let log_public = log2_ceil_usize(n_public as usize) as u32;
445
446 let private_index_map = private_alive
447 .iter()
448 .enumerate()
449 .filter(|(_, alive)| **alive)
450 .enumerate()
451 .map(|(new_idx, (id, _))| (id as u32, new_idx as u32))
452 .collect::<HashMap<_, _>>();
453
454 let n_private = private_index_map.len() as u32;
455 let log_private = log2_ceil_usize(n_private as usize) as u32;
456
457 Self {
458 constants,
459 n_inout,
460 n_derived,
461 n_precommit,
462 n_private,
463 log_public,
464 log_precommit,
465 log_private,
466 derived_index_map,
467 private_index_map,
468 }
469 }
470
471 pub fn with_blinding(self, info: BlindingInfo) -> Self {
472 let blinding_size = info.n_dummy_wires + 3 * info.n_dummy_constraints;
479
480 let total_precommit = self.n_precommit as usize + blinding_size;
481 let log_precommit = log2_ceil_usize(total_precommit) as u32;
482
483 let total_private = self.n_private as usize + blinding_size;
484 let log_private = log2_ceil_usize(total_private) as u32;
485
486 Self {
487 log_precommit,
488 log_private,
489 ..self
490 }
491 }
492
493 pub const fn public_size(&self) -> usize {
494 1 << self.log_public as usize
495 }
496
497 pub const fn precommit_size(&self) -> usize {
498 1 << self.log_precommit as usize
499 }
500
501 pub const fn private_size(&self) -> usize {
502 1 << self.log_private as usize
503 }
504
505 pub const fn n_constants(&self) -> usize {
506 self.constants.len()
507 }
508
509 pub const fn n_inout(&self) -> usize {
510 self.n_inout as usize
511 }
512
513 pub const fn n_derived(&self) -> usize {
514 self.n_derived as usize
515 }
516
517 pub const fn n_precommit(&self) -> usize {
518 self.n_precommit as usize
519 }
520
521 pub const fn n_private(&self) -> usize {
522 self.n_private as usize
523 }
524
525 pub const fn log_public(&self) -> u32 {
526 self.log_public
527 }
528
529 pub const fn log_precommit(&self) -> u32 {
530 self.log_precommit
531 }
532
533 pub const fn log_private(&self) -> u32 {
534 self.log_private
535 }
536
537 pub fn get(&self, wire: &ConstraintWire) -> Option<WitnessIndex> {
538 match wire.kind {
539 WireKind::Constant => {
540 assert!((wire.id as usize) < self.constants.len());
541 Some(WitnessIndex::public(wire.id))
542 }
543 WireKind::InOut => {
544 assert!(wire.id < self.n_inout);
545 Some(WitnessIndex::public(self.constants.len() as u32 + wire.id))
546 }
547 WireKind::Derived => self.derived_index_map.get(&wire.id).map(|&derived_idx| {
548 WitnessIndex::public(self.constants.len() as u32 + self.n_inout + derived_idx)
549 }),
550 WireKind::Precommit => {
551 assert!(wire.id < self.n_precommit);
552 Some(WitnessIndex::precommit(wire.id))
553 }
554 WireKind::Private => self
555 .private_index_map
556 .get(&wire.id)
557 .map(|&id| WitnessIndex::private(id)),
558 }
559 }
560}
561
562#[cfg(test)]
563mod tests {
564 use smallvec::smallvec;
565
566 use super::*;
567
568 #[test]
569 fn test_wires_added_mod2() {
570 let w = [
572 ConstraintWire {
573 kind: WireKind::Constant,
574 id: 0,
575 },
576 ConstraintWire {
577 kind: WireKind::InOut,
578 id: 0,
579 },
580 ConstraintWire {
581 kind: WireKind::Private,
582 id: 0,
583 },
584 ConstraintWire {
585 kind: WireKind::Private,
586 id: 1,
587 },
588 ];
589
590 let input = smallvec![w[0], w[2], w[2], w[3], w[3], w[1], w[2], w[1], w[3], w[3]];
594 let operand = Operand::new(input);
595
596 assert_eq!(operand.wires(), &[w[0], w[2]]);
598 }
599
600 #[test]
601 fn test_sorting_when_no_duplicates() {
602 let w = [
604 ConstraintWire {
605 kind: WireKind::Constant,
606 id: 0,
607 },
608 ConstraintWire {
609 kind: WireKind::InOut,
610 id: 0,
611 },
612 ConstraintWire {
613 kind: WireKind::Private,
614 id: 0,
615 },
616 ConstraintWire {
617 kind: WireKind::Private,
618 id: 1,
619 },
620 ];
621
622 let input = smallvec![w[2], w[3], w[0], w[1]];
624 let operand = Operand::new(input);
625
626 assert_eq!(operand.wires(), &[w[0], w[1], w[2], w[3]]);
628 }
629
630 #[test]
631 fn test_merge() {
632 let w = [
634 ConstraintWire {
635 kind: WireKind::Constant,
636 id: 0,
637 },
638 ConstraintWire {
639 kind: WireKind::InOut,
640 id: 0,
641 },
642 ConstraintWire {
643 kind: WireKind::Private,
644 id: 0,
645 },
646 ConstraintWire {
647 kind: WireKind::Private,
648 id: 1,
649 },
650 ];
651
652 let mut lhs = Operand(smallvec![w[0]]);
654 let rhs = Operand(smallvec![]);
655 let (additions, removals) = lhs.merge(&rhs);
656 assert_eq!(lhs.wires(), &[w[0]]);
657 assert_eq!(additions.wires(), &[]);
658 assert_eq!(removals.wires(), &[]);
659
660 let mut lhs = Operand(smallvec![]);
662 let rhs = Operand(smallvec![w[0]]);
663 let (additions, removals) = lhs.merge(&rhs);
664 assert_eq!(lhs.wires(), &[w[0]]);
665 assert_eq!(additions.wires(), &[w[0]]);
666 assert_eq!(removals.wires(), &[]);
667
668 let mut lhs = Operand(smallvec![w[0]]);
670 let rhs = Operand(smallvec![w[0]]);
671 let (additions, removals) = lhs.merge(&rhs);
672 assert_eq!(lhs.wires(), &[]);
673 assert_eq!(additions.wires(), &[]);
674 assert_eq!(removals.wires(), &[w[0]]);
675
676 let mut lhs = Operand(smallvec![w[0]]);
678 let rhs = Operand(smallvec![w[1]]);
679 let (additions, removals) = lhs.merge(&rhs);
680 assert_eq!(lhs.wires(), &[w[0], w[1]]);
681 assert_eq!(additions.wires(), &[w[1]]);
682 assert_eq!(removals.wires(), &[]);
683
684 let mut lhs = Operand(smallvec![w[0]]);
686 let rhs = Operand(smallvec![w[0], w[1]]);
687 let (additions, removals) = lhs.merge(&rhs);
688 assert_eq!(lhs.wires(), &[w[1]]);
689 assert_eq!(additions.wires(), &[w[1]]);
690 assert_eq!(removals.wires(), &[w[0]]);
691
692 let mut lhs = Operand(smallvec![w[0], w[2]]);
694 let rhs = Operand(smallvec![w[1], w[3]]);
695 let (additions, removals) = lhs.merge(&rhs);
696 assert_eq!(lhs.wires(), &[w[0], w[1], w[2], w[3]]);
697 assert_eq!(additions.wires(), &[w[1], w[3]]);
698 assert_eq!(removals.wires(), &[]);
699
700 let mut lhs = Operand(smallvec![w[0], w[2]]);
702 let rhs = Operand(smallvec![w[0], w[1], w[2], w[3]]);
703 let (additions, removals) = lhs.merge(&rhs);
704 assert_eq!(lhs.wires(), &[w[1], w[3]]);
705 assert_eq!(additions.wires(), &[w[1], w[3]]);
706 assert_eq!(removals.wires(), &[w[0], w[2]]);
707 }
708}