1use std::array;
23
24use binius_core::Word;
25use binius_utils::strided_array::StridedArray2DViewMut;
26
27use super::{
28 assertion::{MAX_ASSERTION_FAILURES, render_path},
29 exec::EvalContext,
30};
31use crate::{
32 artifact::witness::{AssertionFailure, PopulateError},
33 ir::path::{PathSpec, PathSpecTree},
34};
35
36struct InstanceAssertionFailure {
38 path_spec: PathSpec,
39 message: String,
40}
41
42#[derive(Debug, thiserror::Error)]
49#[error("instance {instance} is not satisfied: {source}")]
50#[non_exhaustive]
51pub struct BatchPopulateError {
52 pub instance: usize,
54 #[source]
56 pub source: PopulateError,
57}
58
59pub struct BatchExecutionContext<'a, 'v> {
61 values: &'a mut StridedArray2DViewMut<'v, Word>,
63 failures: Vec<InstanceAssertionFailure>,
68 min_failure_count: usize,
72 min_failing_instance: Option<usize>,
74}
75
76impl<'a, 'v> BatchExecutionContext<'a, 'v> {
77 pub const fn new(values: &'a mut StridedArray2DViewMut<'v, Word>) -> Self {
78 Self {
79 values,
80 failures: Vec::new(),
81 min_failure_count: 0,
82 min_failing_instance: None,
83 }
84 }
85
86 pub fn check_assertions(
88 self,
89 path_spec_tree: Option<&PathSpecTree>,
90 ) -> Result<(), BatchPopulateError> {
91 let Some(instance) = self.min_failing_instance else {
92 return Ok(());
93 };
94
95 let failures: Vec<AssertionFailure> = self
98 .failures
99 .into_iter()
100 .map(|f| AssertionFailure {
101 path: render_path(path_spec_tree, f.path_spec),
102 detail: f.message,
103 })
104 .collect();
105
106 Err(BatchPopulateError {
107 instance,
108 source: PopulateError {
109 total: self.min_failure_count,
110 failures,
111 },
112 })
113 }
114}
115
116impl EvalContext for BatchExecutionContext<'_, '_> {
117 fn n_instances(&self) -> usize {
118 self.values.width()
119 }
120
121 #[inline]
122 fn load(&self, reg: u32, instance: usize) -> Word {
123 self.values[(reg as usize, instance)]
124 }
125
126 #[inline]
127 fn store(&mut self, reg: u32, instance: usize, value: Word) {
128 self.values[(reg as usize, instance)] = value;
129 }
130
131 fn map<const D: usize, const S: usize, F>(&mut self, dsts: [u32; D], srcs: [u32; S], op: F)
133 where
134 F: Fn([Word; S]) -> [Word; D],
135 {
136 let n_instances = self.values.width();
137 let (mut dst_rows, src_rows) = self
138 .values
139 .rows_mut(dsts.map(|reg| reg as usize), srcs.map(|reg| reg as usize));
140 for i in 0..n_instances {
141 let out = op(array::from_fn(|k| src_rows[k][i]));
142 for (row, value) in dst_rows.iter_mut().zip(out) {
143 row[i] = value;
144 }
145 }
146 }
147
148 fn update<F>(&mut self, dst: u32, src: u32, op: F)
150 where
151 F: Fn(Word, Word) -> Word,
152 {
153 let ([dst_row], [src_row]) = self.values.rows_mut([dst as usize], [src as usize]);
154 for (acc, &word) in dst_row.iter_mut().zip(src_row) {
155 *acc = op(*acc, word);
156 }
157 }
158
159 fn check<const S: usize, P, M>(
163 &mut self,
164 srcs: [u32; S],
165 path_spec: PathSpec,
166 fails: P,
167 message: M,
168 ) where
169 P: Fn([Word; S]) -> bool,
170 M: Fn([Word; S]) -> String,
171 {
172 let n_instances = self.values.width();
173 let any = {
174 let rows = srcs.map(|reg| self.values.row(reg as usize));
175 (0..n_instances).any(|i| fails(array::from_fn(|k| rows[k][i])))
176 };
177 if !any {
178 return;
179 }
180 for i in 0..n_instances {
181 let words = srcs.map(|reg| self.load(reg, i));
182 if fails(words) {
183 self.note_assertion_failure(i, path_spec, message(words));
184 }
185 }
186 }
187
188 #[cold]
195 fn note_assertion_failure(&mut self, instance: usize, path_spec: PathSpec, message: String) {
196 match self.min_failing_instance {
197 Some(min) if instance > min => return,
198 Some(min) if instance < min => {
199 self.min_failing_instance = Some(instance);
200 self.failures.clear();
201 self.min_failure_count = 0;
202 }
203 _ => self.min_failing_instance = Some(instance),
205 }
206
207 self.min_failure_count += 1;
208 if self.failures.len() < MAX_ASSERTION_FAILURES {
209 self.failures
210 .push(InstanceAssertionFailure { path_spec, message });
211 }
212 }
213}
214
215#[cfg(test)]
216mod tests {
217 use binius_core::Word;
218 use binius_utils::strided_array::StridedArray2DViewMut;
219
220 use crate::{
221 MAX_ASSERTION_FAILURES, Wire, artifact::witness::AssertionFailure, builder::CircuitBuilder,
222 };
223
224 #[test]
227 fn batched_matches_scalar_per_instance() {
228 let builder = CircuitBuilder::new();
231 let a = builder.add_witness();
232 let b = builder.add_witness();
233 let zero = builder.add_witness();
235 let msb = builder.add_witness();
236 let k = builder.add_constant_64(0x0123_4567_89ab_cdef);
237 let c = builder.band(a, b);
238 let d = builder.bxor(a, k);
239 let (sum, cout) = builder.iadd(a, b);
240 let e = builder.rotr(b, 7);
241 let f = builder.bor(c, e);
242 let g = builder.fax(a, b, k);
243 let xs = builder.bxor_multi(&[a, b, k, e, f]);
244 let lt = builder.icmp_ult(a, b);
246 let sel = builder.select(lt, d, e);
247 let (diff, bout) = builder.isub_bin_bout(a, b, lt);
248 let (mul_hi, mul_lo) = builder.imul(a, b);
249 let (gmul_lo, gmul_hi) = builder.bmul(a, b, k, e);
250 let lanes = builder.iadd_32(a, b);
251 let (lanes_sum, lanes_cout) = builder.iadd32_cin_cout(a, b, lt);
252 let shifts = [
254 builder.shl(a, 13),
255 builder.shr(a, 13),
256 builder.sar(a, 13),
257 builder.rotr(a, 13),
258 builder.sll32(a, 13),
259 builder.srl32(a, 13),
260 builder.sra32(a, 13),
261 builder.rotr32(a, 13),
262 ];
263 builder.assert_eq("xor_roundtrip", builder.bxor(d, k), a);
265 builder.assert_eq_cond("xor_roundtrip_cond", builder.bxor(d, k), a, lt);
266 builder.assert_zero("zero_is_zero", zero);
267 builder.assert_false("zero_msb_false", zero);
268 builder.assert_non_zero("msb_is_non_zero", msb);
269 builder.assert_true("msb_is_true", msb);
270 for wire in [
271 c, d, sum, cout, e, f, g, xs, lt, sel, diff, bout, mul_hi, mul_lo, gmul_lo, gmul_hi,
272 lanes, lanes_sum, lanes_cout,
273 ]
274 .into_iter()
275 .chain(shifts)
276 {
277 builder.force_commit(wire);
278 }
279 let circuit = builder.build();
280
281 let layout = circuit.value_vec_layout().clone();
282 assert_eq!(layout.n_inout, 0, "fixture should have no inout wires");
283 let combined = layout.combined_len();
284 let full_len = combined + layout.n_scratch;
285 let n = 1024usize;
288
289 let inputs: Vec<(u64, u64)> = (0..n)
291 .map(|i| {
292 let i = i as u64;
293 (i.wrapping_mul(0x9e37_79b9_7f4a_7c15), i ^ 0x0000_0000_dead_beef)
294 })
295 .collect();
296
297 let scalar: Vec<Vec<Word>> = inputs
299 .iter()
300 .map(|&(x, y)| {
301 let mut filler = circuit.new_witness_filler();
302 filler[a] = Word(x);
303 filler[b] = Word(y);
304 filler[zero] = Word::ZERO;
305 filler[msb] = Word::MSB_ONE;
306 circuit.populate_wire_witness(&mut filler).unwrap();
307 filler.value_vec().combined_witness().to_vec()
308 })
309 .collect();
310
311 let a_row = circuit.witness_row(a);
313 let b_row = circuit.witness_row(b);
314 let zero_row = circuit.witness_row(zero);
315 let msb_row = circuit.witness_row(msb);
316 let mut data = vec![Word::ZERO; full_len * n];
317 let mut view = StridedArray2DViewMut::without_stride(&mut data, full_len, n).unwrap();
318 for (instance, &(x, y)) in inputs.iter().enumerate() {
319 view[(a_row, instance)] = Word(x);
320 view[(b_row, instance)] = Word(y);
321 view[(zero_row, instance)] = Word::ZERO;
322 view[(msb_row, instance)] = Word::MSB_ONE;
323 }
324 circuit.populate_wire_witness_batched(&mut view).unwrap();
325
326 for instance in 0..n {
328 for row in 0..combined {
329 assert_eq!(
330 view[(row, instance)],
331 scalar[instance][row],
332 "mismatch at row {row}, instance {instance}"
333 );
334 }
335 }
336 }
337
338 #[test]
340 fn batched_reports_lowest_failing_instance() {
341 let builder = CircuitBuilder::new();
343 let a = builder.add_witness();
344 let b = builder.add_witness();
345 builder.assert_eq("a_eq_b", a, b);
346 let circuit = builder.build();
347
348 let layout = circuit.value_vec_layout().clone();
349 let full_len = layout.combined_len() + layout.n_scratch;
350 let n = 4usize;
351
352 let inputs = [(1u64, 1u64), (7, 7), (4, 5), (9, 8)];
354 let a_row = circuit.witness_row(a);
355 let b_row = circuit.witness_row(b);
356 let mut data = vec![Word::ZERO; full_len * n];
357 let mut view = StridedArray2DViewMut::without_stride(&mut data, full_len, n).unwrap();
358 for (instance, &(x, y)) in inputs.iter().enumerate() {
359 view[(a_row, instance)] = Word(x);
360 view[(b_row, instance)] = Word(y);
361 }
362
363 let err = circuit
364 .populate_wire_witness_batched(&mut view)
365 .expect_err("instances 2 and 3 violate a == b");
366 assert_eq!(err.instance, 2);
367 assert_eq!(err.source.total, 1);
368 assert_eq!(
369 err.source.failures,
370 vec![AssertionFailure {
371 path: ".a_eq_b".to_string(),
372 detail: "Word(0x0000000000000004) != Word(0x0000000000000005)".to_string(),
373 }]
374 );
375 let rendered = err.to_string();
377 assert!(rendered.starts_with("instance 2 is not satisfied: "), "{rendered}");
378 assert!(rendered.contains(".a_eq_b: Word(0x0000000000000004)"), "{rendered}");
379 assert!(!rendered.ends_with('\n'), "{rendered}");
381 assert!(std::error::Error::source(&err).is_some(), "the inner error must be the source");
382 }
383
384 #[test]
395 fn batch_min_failing_instance_survives_a_full_cap_of_higher_instances() {
396 let builder = CircuitBuilder::new();
397 let a = builder.add_witness();
398 let b = builder.add_witness();
399 let c = builder.add_witness();
400 let d = builder.add_witness();
401 builder.assert_eq("assertion_one", a, b);
403 builder.assert_eq("assertion_two", c, d);
405 let circuit = builder.build();
406
407 let layout = circuit.value_vec_layout().clone();
408 let full_len = layout.combined_len() + layout.n_scratch;
409 let n = 200usize;
412
413 let a_row = circuit.witness_row(a);
414 let b_row = circuit.witness_row(b);
415 let c_row = circuit.witness_row(c);
416 let d_row = circuit.witness_row(d);
417 let mut data = vec![Word::ZERO; full_len * n];
418 let mut view = StridedArray2DViewMut::without_stride(&mut data, full_len, n).unwrap();
419 for instance in 0..n {
420 let a_val = Word(1);
422 let b_val = Word(if instance >= 50 { 2 } else { 1 });
423 let c_val = Word(3);
425 let d_val = Word(if instance == 10 { 4 } else { 3 });
426 view[(a_row, instance)] = a_val;
427 view[(b_row, instance)] = b_val;
428 view[(c_row, instance)] = c_val;
429 view[(d_row, instance)] = d_val;
430 }
431
432 let err = circuit
433 .populate_wire_witness_batched(&mut view)
434 .expect_err("instance 10 and instances 50..199 all violate an assertion");
435
436 assert_eq!(err.instance, 10);
438 assert_eq!(err.source.total, 1);
440 assert_eq!(
441 err.source.failures,
442 vec![AssertionFailure {
443 path: ".assertion_two".to_string(),
444 detail: "Word(0x0000000000000003) != Word(0x0000000000000004)".to_string(),
445 }]
446 );
447 }
448
449 #[test]
452 fn batch_total_counts_every_violation_past_the_cap_for_the_reported_instance() {
453 const N: usize = MAX_ASSERTION_FAILURES + 50;
455
456 let builder = CircuitBuilder::new();
457 let x: [Wire; N] = core::array::from_fn(|_| builder.add_witness());
458 let y: [Wire; N] = core::array::from_fn(|_| builder.add_witness());
459 builder.assert_eq_v("pairs", x, y);
460 let circuit = builder.build();
461
462 let layout = circuit.value_vec_layout().clone();
463 let full_len = layout.combined_len() + layout.n_scratch;
464 let n_instances = 4usize;
465
466 let x_rows: Vec<usize> = x.iter().map(|&w| circuit.witness_row(w)).collect();
467 let y_rows: Vec<usize> = y.iter().map(|&w| circuit.witness_row(w)).collect();
468 let mut data = vec![Word::ZERO; full_len * n_instances];
469 let mut view =
470 StridedArray2DViewMut::without_stride(&mut data, full_len, n_instances).unwrap();
471 for instance in 0..n_instances {
472 for i in 0..N {
473 let x_val = Word(i as u64);
475 let y_val = if instance == 0 {
476 Word(i as u64 + 1)
477 } else {
478 x_val
479 };
480 view[(x_rows[i], instance)] = x_val;
481 view[(y_rows[i], instance)] = y_val;
482 }
483 }
484
485 let err = circuit
486 .populate_wire_witness_batched(&mut view)
487 .expect_err("instance 0 violates every one of the N pairwise assertions");
488
489 assert_eq!(err.instance, 0);
490 assert_eq!(err.source.total, N);
492 assert_eq!(err.source.failures.len(), MAX_ASSERTION_FAILURES);
494 assert!(err.source.total > err.source.failures.len());
495 }
496}