1use std::collections::hash_map::Entry;
7
8use binius_core::word::Word;
9use cranelift_entity::{EntityRef, PrimaryMap, SecondaryMap, entity_impl};
10use rustc_hash::FxHashMap;
11use smallvec::SmallVec;
12
13use crate::{
14 gates::opcode::{Opcode, OpcodeShape},
15 ir::{
16 hints::{HintId, HintRegistry},
17 path::{PathSpec, PathSpecTree},
18 },
19};
20
21pub mod hints;
22pub mod path;
23
24#[derive(Copy, Clone, PartialEq, Eq, Hash, Debug, PartialOrd, Ord)]
29pub struct Wire(u32);
30entity_impl!(Wire);
31
32#[derive(Copy, Clone, Debug)]
33pub enum WireKind {
34 Constant(Word),
35 Inout,
36 Witness,
37 Internal,
39 Scratch,
41}
42impl WireKind {
43 pub const fn is_const(&self) -> bool {
45 matches!(self, WireKind::Constant(_))
46 }
47}
48
49#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
51pub struct Gate(u32);
52
53entity_impl!(Gate);
54
55#[derive(Copy, Clone)]
60pub struct GateParam<'a> {
61 pub constants: &'a [Wire],
62 pub inputs: &'a [Wire],
63 pub outputs: &'a [Wire],
64 pub aux: &'a [Wire],
65 pub scratch: &'a [Wire],
66 pub imm: &'a [u32],
67}
68
69impl GateParam<'_> {
70 pub fn const_wires<const N: usize>(&self) -> [Wire; N] {
72 fixed(self.constants, "constant")
73 }
74
75 pub fn in_wires<const N: usize>(&self) -> [Wire; N] {
77 fixed(self.inputs, "input")
78 }
79
80 pub fn out_wires<const N: usize>(&self) -> [Wire; N] {
82 fixed(self.outputs, "output")
83 }
84
85 pub fn aux_wires<const N: usize>(&self) -> [Wire; N] {
87 fixed(self.aux, "auxiliary")
88 }
89
90 pub fn scratch_wires<const N: usize>(&self) -> [Wire; N] {
92 fixed(self.scratch, "scratch")
93 }
94
95 pub fn imms<const N: usize>(&self) -> [u32; N] {
97 fixed(self.imm, "immediate")
98 }
99}
100
101fn fixed<T: Copy, const N: usize>(group: &[T], what: &str) -> [T; N] {
108 <[T; N]>::try_from(group).unwrap_or_else(|_| {
109 panic!("a gate carries {} {what} entries, but its shape declares {N}", group.len())
110 })
111}
112
113#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
115pub enum GateBody {
116 Op(Opcode),
118 Hint(HintId),
120}
121
122pub struct GateData {
125 pub body: GateBody,
127
128 pub wires: SmallVec<[Wire; 5]>,
144
145 pub immediates: SmallVec<[u32; 2]>,
155
156 pub dimensions: Box<[usize]>,
166}
167
168impl GateData {
169 pub fn gate_param(&self, registry: &HintRegistry) -> GateParam<'_> {
171 self.gate_param_for_shape(self.shape(registry))
172 }
173
174 fn gate_param_for_shape(&self, shape: OpcodeShape) -> GateParam<'_> {
175 let start_const = 0;
176 let end_const = shape.const_in.len();
177 let start_input = end_const;
178 let end_input = start_input + shape.n_in;
179 let start_output = end_input;
180 let end_output = start_output + shape.n_out;
181 let start_aux = end_output;
182 let end_aux = start_aux + shape.n_aux;
183 let start_scratch = end_aux;
184 let end_scratch = start_scratch + shape.n_scratch;
185 GateParam {
186 constants: &self.wires[start_const..end_const],
187 inputs: &self.wires[start_input..end_input],
188 outputs: &self.wires[start_output..end_output],
189 aux: &self.wires[start_aux..end_aux],
190 scratch: &self.wires[start_scratch..end_scratch],
191 imm: &self.immediates,
192 }
193 }
194
195 pub fn shape(&self, registry: &HintRegistry) -> OpcodeShape {
197 match self.body {
198 GateBody::Op(opcode) => opcode.shape(&self.dimensions),
199 GateBody::Hint(hint_id) => {
200 let (n_in, n_out) = registry.shape(hint_id, &self.dimensions);
201 OpcodeShape::new(n_in, n_out)
202 }
203 }
204 }
205
206 pub fn validate_shape(&self, registry: &HintRegistry) {
208 let shape = self.shape(registry);
209 let expected_wires =
210 shape.const_in.len() + shape.n_in + shape.n_out + shape.n_aux + shape.n_scratch;
211 assert_eq!(self.wires.len(), expected_wires);
212 assert_eq!(self.immediates.len(), shape.n_imm);
213 }
214}
215
216pub struct GateGraph {
218 pub gates: PrimaryMap<Gate, GateData>,
220 pub wires: PrimaryMap<Wire, WireKind>,
221
222 pub path_spec_tree: PathSpecTree,
223 pub gate_origin: SecondaryMap<Gate, PathSpec>,
224 pub assertion_names: SecondaryMap<Gate, PathSpec>,
225
226 pub const_pool: FxHashMap<Word, Wire>,
230
231 pub all_one: Wire,
237
238 pub wire_def: SecondaryMap<Wire, Option<Gate>>,
241 use_edges: Vec<Gate>,
246 use_offsets: Vec<u32>,
250}
251
252impl GateGraph {
253 pub fn new() -> Self {
254 let path_spec_tree = PathSpecTree::new();
255 let root = path_spec_tree.root();
256 let mut graph = Self {
257 gates: PrimaryMap::new(),
258 wires: PrimaryMap::new(),
259 path_spec_tree,
260 gate_origin: SecondaryMap::with_default(root),
261 assertion_names: SecondaryMap::with_default(root),
262 const_pool: FxHashMap::default(),
263 all_one: Wire::from_u32(0),
265 wire_def: SecondaryMap::new(),
266 use_edges: Vec::new(),
267 use_offsets: Vec::new(),
268 };
269 graph.all_one = graph.add_constant(Word::ALL_ONE);
272 graph
273 }
274
275 pub fn validate(&self, hint_registry: &HintRegistry) {
277 for gate in self.gates.values() {
279 gate.validate_shape(hint_registry);
280 }
281 }
282
283 pub fn add_inout(&mut self) -> Wire {
284 self.wires.push(WireKind::Inout)
285 }
286
287 pub fn add_witness(&mut self) -> Wire {
288 self.wires.push(WireKind::Witness)
289 }
290
291 pub fn add_internal(&mut self) -> Wire {
292 self.wires.push(WireKind::Internal)
293 }
294
295 pub fn add_scratch(&mut self) -> Wire {
296 self.wires.push(WireKind::Scratch)
297 }
298
299 pub fn add_constant(&mut self, word: Word) -> Wire {
301 match self.const_pool.entry(word) {
303 Entry::Occupied(entry) => *entry.get(),
304 Entry::Vacant(entry) => *entry.insert(self.wires.push(WireKind::Constant(word))),
305 }
306 }
307
308 pub fn emit_gate(
310 &mut self,
311 gate_origin: PathSpec,
312 opcode: Opcode,
313 inputs: impl IntoIterator<Item = Wire>,
314 outputs: impl IntoIterator<Item = Wire>,
315 ) -> Gate {
316 self.emit_gate_generic(gate_origin, opcode, inputs, outputs, &[], &[])
317 }
318
319 pub fn emit_gate_generic(
324 &mut self,
325 gate_origin: PathSpec,
326 opcode: Opcode,
327 inputs: impl IntoIterator<Item = Wire>,
328 outputs: impl IntoIterator<Item = Wire>,
329 dimensions: &[usize],
330 immediates: &[u32],
331 ) -> Gate {
332 let shape = opcode.shape(dimensions);
333 let mut wires: SmallVec<[Wire; 5]> = SmallVec::with_capacity(
334 shape.const_in.len() + shape.n_in + shape.n_out + shape.n_aux + shape.n_scratch,
335 );
336 for c in shape.const_in {
337 wires.push(self.add_constant(*c));
338 }
339 wires.extend(inputs);
340 wires.extend(outputs);
341 for _ in 0..shape.n_aux {
342 wires.push(self.add_internal());
344 }
345 for _ in 0..shape.n_scratch {
346 wires.push(self.add_scratch());
347 }
348 let data = GateData {
349 body: GateBody::Op(opcode),
350 wires,
351 dimensions: dimensions.into(),
352 immediates: SmallVec::from_slice(immediates),
353 };
354 let expected_wires =
356 shape.const_in.len() + shape.n_in + shape.n_out + shape.n_aux + shape.n_scratch;
357 assert_eq!(data.wires.len(), expected_wires);
358 assert_eq!(data.immediates.len(), shape.n_imm);
359
360 let gate = self.gates.push(data);
361
362 self.gate_origin[gate] = gate_origin;
363
364 gate
365 }
366
367 pub fn emit_hint_gate(
371 &mut self,
372 gate_origin: PathSpec,
373 hint_id: HintId,
374 dimensions: &[usize],
375 inputs: impl IntoIterator<Item = Wire>,
376 outputs: impl IntoIterator<Item = Wire>,
377 ) -> Gate {
378 let mut wires: SmallVec<[Wire; 5]> = SmallVec::new();
379 wires.extend(inputs);
380 wires.extend(outputs);
381 let data = GateData {
382 body: GateBody::Hint(hint_id),
383 wires,
384 dimensions: dimensions.into(),
385 immediates: SmallVec::new(),
386 };
387 let gate = self.gates.push(data);
388 self.gate_origin[gate] = gate_origin;
389 gate
390 }
391
392 pub fn rebuild_wire_defs(&mut self, hint_registry: &HintRegistry) {
398 self.wire_def.clear();
399 for (gate, data) in self.gates.iter() {
400 let param = data.gate_param(hint_registry);
401 for &wire in param.outputs.iter().chain(param.aux) {
402 self.wire_def[wire] = Some(gate);
403 }
404 }
405 }
406
407 pub fn rebuild_wire_uses(&mut self, hint_registry: &HintRegistry) {
427 let n_wires = self.wires.len();
428
429 let mut last_seen: SecondaryMap<Wire, Option<Gate>> = SecondaryMap::new();
431
432 let mut offsets = vec![0u32; n_wires + 1];
435 for (gate, data) in self.gates.iter() {
436 let param = data.gate_param(hint_registry);
437 for &wire in param.constants.iter().chain(param.inputs) {
438 if last_seen[wire] != Some(gate) {
439 last_seen[wire] = Some(gate);
440 offsets[wire.index() + 1] += 1;
441 }
442 }
443 }
444
445 for i in 0..n_wires {
447 offsets[i + 1] += offsets[i];
448 }
449
450 let mut edges = vec![Gate::from_u32(0); offsets[n_wires] as usize];
452 let mut cursor = offsets.clone();
453 last_seen.clear();
454 for (gate, data) in self.gates.iter() {
455 let param = data.gate_param(hint_registry);
456 for &wire in param.constants.iter().chain(param.inputs) {
457 if last_seen[wire] != Some(gate) {
458 last_seen[wire] = Some(gate);
459 edges[cursor[wire.index()] as usize] = gate;
460 cursor[wire.index()] += 1;
461 }
462 }
463 }
464
465 self.use_edges = edges;
466 self.use_offsets = offsets;
467 }
468
469 pub fn rebuild_use_def_chains(&mut self, hint_registry: &HintRegistry) {
473 self.rebuild_wire_defs(hint_registry);
474 self.rebuild_wire_uses(hint_registry);
475 }
476
477 pub fn get_wire_uses(&self, wire: Wire) -> impl Iterator<Item = Gate> + '_ {
481 let run = match self.use_offsets.get(wire.index() + 1) {
483 Some(&end) => self.use_offsets[wire.index()] as usize..end as usize,
484 None => 0..0,
485 };
486 self.use_edges[run].iter().copied()
487 }
488
489 pub fn iter_const_wires(&self) -> impl Iterator<Item = (Wire, &WireKind)> {
491 self.wires.iter().filter(|(_, kind)| kind.is_const())
492 }
493
494 pub fn wire_kind(&self, wire: Wire) -> WireKind {
496 self.wires[wire]
497 }
498
499 pub fn gate_data(&self, gate: Gate) -> &GateData {
501 &self.gates[gate]
502 }
503
504 pub fn replace_gate_wire(&mut self, gate: Gate, old_wire: Wire, new_wire: Wire) -> usize {
507 let gate_data = &mut self.gates[gate];
508 let mut rewritten = 0;
509 for wire in &mut gate_data.wires {
510 if *wire == old_wire {
511 *wire = new_wire;
512 rewritten += 1;
513 }
514 }
515 rewritten
516 }
517
518 pub fn replace_wire_with_constant(&mut self, old_wire: Wire, value: Word) -> WireReplacement {
520 let const_wire = self.add_constant(value);
521 self.replace_wire_with_wire(old_wire, const_wire)
522 }
523
524 pub fn replace_wire_with_wire(&mut self, old_wire: Wire, new_wire: Wire) -> WireReplacement {
533 if new_wire == old_wire {
534 return WireReplacement {
535 n_slots_rewritten: 0,
536 affected_gates: Vec::new(),
537 };
538 }
539
540 let affected_gates: Vec<Gate> = self.get_wire_uses(old_wire).collect();
542
543 let n_slots_rewritten = affected_gates
546 .iter()
547 .map(|&gate| self.replace_gate_wire(gate, old_wire, new_wire))
548 .sum();
549
550 WireReplacement {
551 n_slots_rewritten,
552 affected_gates,
553 }
554 }
555
556 pub fn gates(&self) -> impl Iterator<Item = Gate> + '_ {
558 self.gates.iter().map(|(gate, _)| gate)
559 }
560}
561
562pub struct WireReplacement {
564 pub n_slots_rewritten: usize,
566 pub affected_gates: Vec<Gate>,
568}
569
570impl Default for GateGraph {
571 fn default() -> Self {
572 Self::new()
573 }
574}
575
576#[cfg(test)]
577mod tests {
578 use super::*;
579 use crate::gates::opcode::Opcode;
580
581 fn get_wire_def(graph: &GateGraph, wire: Wire) -> Option<Gate> {
583 graph.wire_def[wire]
584 }
585
586 fn wire_use_count(graph: &GateGraph, wire: Wire) -> usize {
587 graph.get_wire_uses(wire).count()
588 }
589
590 fn is_wire_single_use(graph: &GateGraph, wire: Wire) -> bool {
591 wire_use_count(graph, wire) == 1
592 }
593
594 fn get_wire_single_use(graph: &GateGraph, wire: Wire) -> Option<Gate> {
595 let mut uses = graph.get_wire_uses(wire);
596 match (uses.next(), uses.next()) {
597 (Some(only), None) => Some(only),
598 _ => None,
599 }
600 }
601
602 fn get_gate_inputs(graph: &GateGraph, gate: Gate) -> Vec<Wire> {
603 let gate_data = &graph.gates[gate];
604 let gate_param = gate_data.gate_param(&HintRegistry::new());
605
606 let mut inputs = Vec::new();
607 inputs.extend_from_slice(gate_param.constants);
608 inputs.extend_from_slice(gate_param.inputs);
609 inputs
610 }
611
612 fn get_gate_outputs(graph: &GateGraph, gate: Gate) -> Vec<Wire> {
613 let gate_data = &graph.gates[gate];
614 let gate_param = gate_data.gate_param(&HintRegistry::new());
615
616 let mut outputs = Vec::new();
617 outputs.extend_from_slice(gate_param.outputs);
618 outputs
619 }
620
621 #[test]
622 fn test_use_def_analysis() {
623 let mut graph = GateGraph::new();
624 let root = graph.path_spec_tree.root();
625
626 let in1 = graph.add_inout();
628 let in2 = graph.add_inout();
629 let out1 = graph.add_witness();
630 let out2 = graph.add_witness();
631
632 let gate1 = graph.emit_gate(root, Opcode::Bxor, vec![in1, in2], vec![out1]);
634
635 let gate2 = graph.emit_gate(root, Opcode::Band, vec![out1, in1], vec![out2]);
637
638 graph.rebuild_use_def_chains(&HintRegistry::new());
640
641 assert_eq!(get_wire_def(&graph, out1), Some(gate1));
643
644 assert_eq!(get_wire_def(&graph, out2), Some(gate2));
646
647 assert!(graph.get_wire_uses(in1).any(|g| g == gate1));
649 assert!(graph.get_wire_uses(in2).any(|g| g == gate1));
650
651 assert!(graph.get_wire_uses(out1).any(|g| g == gate2));
653
654 assert_eq!(wire_use_count(&graph, in1), 2); assert_eq!(wire_use_count(&graph, in2), 1);
657 assert_eq!(wire_use_count(&graph, out1), 1);
658 assert_eq!(wire_use_count(&graph, out2), 0);
659
660 assert!(!is_wire_single_use(&graph, in1)); assert!(is_wire_single_use(&graph, in2));
663 assert!(is_wire_single_use(&graph, out1));
664 assert!(!is_wire_single_use(&graph, out2)); assert_eq!(get_wire_single_use(&graph, in1), None); assert_eq!(get_wire_single_use(&graph, out1), Some(gate2));
669 assert_eq!(get_wire_single_use(&graph, out2), None); }
671
672 #[test]
673 fn wire_uses_are_ordered_and_deduplicated() {
674 let mut graph = GateGraph::new();
684 let root = graph.path_spec_tree.root();
685
686 let x = graph.add_inout();
687 let y = graph.add_inout();
688
689 let o0 = graph.add_internal();
690 let g0 = graph.emit_gate(root, Opcode::Band, vec![x, y], vec![o0]);
691 let o1 = graph.add_internal();
692 let g1 = graph.emit_gate(root, Opcode::Bxor, vec![x, x], vec![o1]);
693 let o2 = graph.add_internal();
694 let g2 = graph.emit_gate(root, Opcode::Band, vec![x, y], vec![o2]);
695
696 graph.rebuild_use_def_chains(&HintRegistry::new());
697
698 let readers: Vec<Gate> = graph.get_wire_uses(x).collect();
700 assert_eq!(readers, vec![g0, g1, g2]);
701
702 assert_eq!(graph.get_wire_uses(o0).collect::<Vec<_>>(), Vec::new());
704 assert_eq!(graph.get_wire_uses(y).collect::<Vec<_>>(), vec![g0, g2]);
705 }
706
707 #[test]
708 fn wire_uses_are_empty_for_a_wire_added_after_the_rebuild() {
709 let mut graph = GateGraph::new();
712 graph.rebuild_use_def_chains(&HintRegistry::new());
713
714 let late = graph.add_inout();
715 assert_eq!(graph.get_wire_uses(late).count(), 0);
716 }
717
718 #[test]
719 fn test_constant_use_def() {
720 let mut graph = GateGraph::new();
721 let root = graph.path_spec_tree.root();
722
723 let const_wire = graph.add_constant(Word(42u64));
725 let in_wire = graph.add_inout();
726 let out = graph.add_witness();
727
728 let gate = graph.emit_gate(root, Opcode::Bxor, vec![const_wire, in_wire], vec![out]);
730
731 graph.rebuild_use_def_chains(&HintRegistry::new());
733
734 assert_eq!(get_wire_def(&graph, const_wire), None);
736
737 assert!(graph.get_wire_uses(const_wire).any(|g| g == gate));
739 assert_eq!(wire_use_count(&graph, const_wire), 1);
740 }
741
742 #[test]
743 fn test_rebuild_use_def_chains() {
744 let mut graph = GateGraph::new();
745 let root = graph.path_spec_tree.root();
746
747 let in1 = graph.add_inout();
749 let in2 = graph.add_inout();
750 let out = graph.add_witness();
751
752 graph.emit_gate(root, Opcode::Bxor, vec![in1, in2], vec![out]);
753
754 graph.wire_def.clear();
756 graph.use_edges.clear();
757 graph.use_offsets.clear();
758
759 assert_eq!(get_wire_def(&graph, out), None);
761 assert!(graph.get_wire_uses(in1).next().is_none());
762
763 graph.rebuild_use_def_chains(&HintRegistry::new());
765
766 assert!(get_wire_def(&graph, out).is_some());
768 assert!(!graph.get_wire_uses(in1).next().is_none());
769 assert!(!graph.get_wire_uses(in2).next().is_none());
770 }
771
772 #[test]
773 fn test_gate_inputs_outputs() {
774 let mut graph = GateGraph::new();
775 let root = graph.path_spec_tree.root();
776
777 let a = graph.add_inout();
778 let b = graph.add_inout();
779 let bin = graph.add_inout();
780 let diff = graph.add_witness();
781 let bout = graph.add_witness();
782
783 let gate = graph.emit_gate(root, Opcode::IsubBinBout, vec![a, b, bin], vec![diff, bout]);
785
786 let inputs = get_gate_inputs(&graph, gate);
790 assert_eq!(inputs.len(), 4);
792 assert!(inputs.contains(&a));
793 assert!(inputs.contains(&b));
794 assert!(inputs.contains(&bin));
795 let const_wire = inputs[0];
797 match graph.wires[const_wire] {
798 WireKind::Constant(word) => assert_eq!(word, Word::ALL_ONE),
799 _ => panic!("Expected constant wire"),
800 }
801
802 let outputs = get_gate_outputs(&graph, gate);
803 assert_eq!(outputs.len(), 2);
804 assert!(outputs.contains(&diff));
805 assert!(outputs.contains(&bout));
806 }
807}