Skip to main content

binius_circuits/
concat.rs

1// Copyright 2025 Irreducible Inc.
2use binius_core::word::Word;
3use binius_frontend::{CircuitBuilder, Wire, hints::Hint};
4
5use crate::{
6	fixed_byte_vec::{ByteVec, extract_const_range},
7	slice::{assert_slice_eq, slice},
8};
9
10/// Hint computing the concatenation of a list of fixed-capacity byte vectors.
11///
12/// Each input is described by a wire-count dimension `d_i` and contributes `d_i` data wires
13/// (8 bytes each, little-endian) plus one `len_bytes` wire. The hint reads the actual byte
14/// prefix of each input (per its `len_bytes`) and packs the concatenated bytes back into
15/// `sum(d_i)` little-endian output words, zero-padding any trailing space.
16struct ByteVecConcatHint;
17
18impl ByteVecConcatHint {
19	const fn new() -> Self {
20		Self
21	}
22}
23
24impl Hint for ByteVecConcatHint {
25	const NAME: &'static str = "binius.byte_vec_concat";
26
27	fn shape(&self, dimensions: &[usize]) -> (usize, usize) {
28		let total_data: usize = dimensions.iter().sum();
29		(total_data + dimensions.len(), total_data)
30	}
31
32	fn execute(&self, dimensions: &[usize], inputs: &[Word], outputs: &mut [Word]) {
33		let total_data: usize = dimensions.iter().sum();
34		let (data_wires, len_wires) = inputs.split_at(total_data);
35
36		let mut bytes = Vec::with_capacity(total_data * 8);
37		let mut cursor = 0;
38		for (&d, &len_word) in dimensions.iter().zip(len_wires) {
39			let words = &data_wires[cursor..cursor + d];
40			cursor += d;
41			let len = (len_word.as_u64() as usize).min(d * 8);
42			bytes.extend(
43				words
44					.iter()
45					.flat_map(|w| w.as_u64().to_le_bytes())
46					.take(len),
47			);
48		}
49
50		for (out, chunk) in outputs.iter_mut().zip(bytes.chunks(8)) {
51			let mut buf = [0u8; 8];
52			buf[..chunk.len()].copy_from_slice(chunk);
53			*out = Word(u64::from_le_bytes(buf));
54		}
55		for out in &mut outputs[bytes.len().div_ceil(8)..] {
56			*out = Word::ZERO;
57		}
58	}
59}
60
61/// Computes the concatenation of a list of [`ByteVec`]s as a new [`ByteVec`].
62///
63/// The returned vector has:
64/// - capacity (number of data wires) equal to the sum of the inputs' capacities,
65/// - runtime length equal to the sum of the inputs' `len_bytes` values, and
66/// - `len_range` equal to the sum of the inputs' ranges (`sum(start)..sum(end)`).
67///
68/// The output data wires are populated by a prover-side concatenation hint; soundness is enforced
69/// by constraining each input's region of the output to equal that input's data. How that
70/// constraint is emitted depends on what is known at circuit-build time:
71///
72/// - When a term's byte offset into the output is a compile-time constant (i.e. every preceding
73///   term has a constant length), its region is extracted with constant shifts and masks (no
74///   multiplexer or dynamic-shift machinery).
75/// - When that term *also* has a constant length, the comparison degenerates further to a plain
76///   per-word `assert_eq` (with a constant mask on the final partial word), dropping the
77///   saturating-diff / variable-shift / select machinery of [`assert_slice_eq`] entirely.
78/// - Otherwise (a dynamic offset, once some preceding term has a dynamic length) the term falls
79///   back to the fully-dynamic [`slice()`] + [`assert_slice_eq`] path.
80///
81/// Bytes of `output.data` beyond `output.len_bytes` are unconstrained.
82pub fn concat(b: &CircuitBuilder, inputs: &[ByteVec]) -> ByteVec {
83	let dimensions: Vec<usize> = inputs.iter().map(|v| v.data.len()).collect();
84	let mut hint_inputs: Vec<Wire> = inputs.iter().flat_map(|v| v.data.iter().copied()).collect();
85	hint_inputs.extend(inputs.iter().map(|v| v.len_bytes));
86
87	let output_data = b.call_hint(ByteVecConcatHint::new(), &dimensions, &hint_inputs);
88
89	// Running byte offset of the current term into the output, tracked both as a wire (`offset`)
90	// and as a compile-time range (`offset_range`). The offset is a compile-time constant exactly
91	// when `offset_range` is a point (all preceding terms have constant lengths), in which case the
92	// constant value is `offset_range.start`.
93	let mut offset = b.add_constant(Word::ZERO);
94	let mut offset_range = 0usize..0usize;
95	// Cumulative number of output words spanned up to and including the current term; bounds the
96	// `input` slice handed to the dynamic `slice` extraction.
97	let mut words_upper_bound = 0usize;
98
99	for (i, input) in inputs.iter().enumerate() {
100		words_upper_bound += input.data.len();
101		let name = format!("subslice eq[{i}]");
102
103		let offset_is_const = offset_range.start == offset_range.end;
104		let len_is_const = input.len_range.start() == input.len_range.end();
105
106		// Advance the running offset wire, emitting as few adder gates as possible:
107		// - while both offset and length stay constant, keep it a folded constant (no `iadd`);
108		// - when the offset is a constant zero (the leading term), `0 + len_bytes = len_bytes`;
109		// - otherwise add the runtime length.
110		let next_offset = if offset_is_const && len_is_const {
111			b.add_constant_64((offset_range.start + input.len_range.start()) as u64)
112		} else if offset_is_const && offset_range.start == 0 {
113			input.len_bytes
114		} else {
115			b.iadd(offset, input.len_bytes).0
116		};
117
118		if offset_is_const {
119			let off = offset_range.start;
120			if len_is_const {
121				// Constant offset and constant length: a fully static per-word comparison.
122				assert_const_slice_eq(
123					b,
124					&name,
125					&output_data,
126					off,
127					&input.data,
128					*input.len_range.start(),
129				);
130			} else {
131				// Constant offset, dynamic length: extract the term's (capacity-sized) region with
132				// constant shifts, then mask the comparison to the runtime length.
133				let region =
134					extract_const_range(b, &output_data, off..off + input.data.len() * Word::BYTES);
135				assert_slice_eq(b, &name, input.len_bytes, &region, &input.data);
136			}
137		} else {
138			// Dynamic offset: fall back to the fully-dynamic slice extraction.
139			let sb = b.subcircuit(format!("concat_term[{i}]"));
140			let extracted = slice(
141				&sb,
142				next_offset,
143				input.len_bytes,
144				&output_data[..words_upper_bound],
145				offset,
146				input.data.len(),
147			);
148			assert_slice_eq(b, &name, input.len_bytes, &extracted, &input.data);
149		}
150
151		offset = next_offset;
152		offset_range = (offset_range.start + input.len_range.start())
153			..(offset_range.end + input.len_range.end());
154	}
155
156	ByteVec::new_with_len_range(output_data, offset, offset_range.start..=offset_range.end)
157}
158
159/// Asserts that `output[offset..offset + len]` equals the first `len` bytes of `input`, where
160/// `offset` and `len` are compile-time constants. Lowers to constant-shift extraction plus a plain
161/// per-word `assert_eq`, masking the final partial word to `len % 8` bytes (bytes of `input` beyond
162/// `len` are ignored, matching [`assert_slice_eq`]'s semantics).
163fn assert_const_slice_eq(
164	b: &CircuitBuilder,
165	name: &str,
166	output: &[Wire],
167	offset: usize,
168	input: &[Wire],
169	len: usize,
170) {
171	// Neither side's final word is zeroed past `len`: `extract_const_range` leaves the extracted
172	// word's high bytes as the next term's data, and `input`'s bytes past its length are
173	// unconstrained. So on the partial boundary word we mask both sides to the valid `len % 8`
174	// bytes before comparing.
175	let extracted = extract_const_range(b, output, offset..offset + len);
176	let n_words = extracted.len();
177	let final_bytes = len % Word::BYTES;
178	for (k, &a) in extracted.iter().enumerate() {
179		let e = input[k];
180		if k + 1 == n_words && final_bytes != 0 {
181			let mask = b.add_constant_64((1u64 << (final_bytes * 8)) - 1);
182			// (a & mask) == (e & mask) is ((a ^ e) & mask) == 0: the XOR is free and
183			// drops one of the two masking ANDs.
184			let diff = b.band(b.bxor(a, e), mask);
185			b.assert_eq(format!("{name}[{k}]"), diff, b.add_constant(Word::ZERO));
186		} else {
187			b.assert_eq(format!("{name}[{k}]"), a, e);
188		}
189	}
190}
191
192#[cfg(test)]
193mod tests {
194	use anyhow::{Result, anyhow};
195
196	use super::*;
197
198	/// Build a circuit calling [`concat`] with `inputs` of the given capacities (in wires),
199	/// populate it from the provided bytes, and return the actual concatenated bytes plus the
200	/// circuit-verification result. Inputs use inout wires so the test driver can populate them.
201	fn run_concat(input_max_lens: &[usize], input_data: &[&[u8]]) -> Result<Vec<u8>> {
202		assert_eq!(input_max_lens.len(), input_data.len());
203
204		let b = CircuitBuilder::new();
205		let inputs: Vec<ByteVec> = input_max_lens
206			.iter()
207			.map(|&n| ByteVec::new_inout(&b, n))
208			.collect();
209		let output = concat(&b, &inputs);
210
211		let circuit = b.build();
212		let mut filler = circuit.new_witness_filler();
213		for (input, &data) in inputs.iter().zip(input_data) {
214			input.populate_len_bytes(&mut filler, data.len());
215			input.populate_data(&mut filler, data);
216		}
217
218		circuit
219			.populate_wire_witness(&mut filler)
220			.map_err(|e| anyhow!("populate_wire_witness: {e}"))?;
221
222		let total_len = input_data.iter().map(|d| d.len()).sum::<usize>();
223		let mut bytes = Vec::with_capacity(total_len);
224		for &w in &output.data {
225			let word = filler[w].as_u64();
226			for j in 0..8 {
227				bytes.push(((word >> (j * 8)) & 0xff) as u8);
228			}
229		}
230		bytes.truncate(total_len);
231
232		let cs = circuit.constraint_system();
233		cs.verify(&filler.into_value_vec())
234			.map_err(|err| anyhow!("constraint verification failed: {err}"))?;
235
236		Ok(bytes)
237	}
238
239	fn assert_concat_eq(input_max_lens: &[usize], input_data: &[&[u8]], expected: &[u8]) {
240		let bytes = run_concat(input_max_lens, input_data).unwrap();
241		assert_eq!(bytes, expected);
242	}
243
244	#[test]
245	fn two_terms() {
246		assert_concat_eq(&[1, 1], &[b"hello", b"world"], b"helloworld");
247	}
248
249	#[test]
250	fn three_terms() {
251		assert_concat_eq(&[1, 1, 1], &[b"foo", b"bar", b"baz"], b"foobarbaz");
252	}
253
254	#[test]
255	fn single_term() {
256		assert_concat_eq(&[1], &[b"hello"], b"hello");
257	}
258
259	#[test]
260	fn empty_middle_term() {
261		assert_concat_eq(&[1, 1, 1], &[b"hello", b"", b"world"], b"helloworld");
262	}
263
264	#[test]
265	fn all_terms_empty() {
266		assert_concat_eq(&[1, 1], &[b"", b""], b"");
267	}
268
269	#[test]
270	fn no_inputs() {
271		assert_concat_eq(&[], &[], b"");
272	}
273
274	#[test]
275	fn unaligned_terms() {
276		assert_concat_eq(&[1, 2], &[b"hello12", b"world456"], b"hello12world456");
277	}
278
279	#[test]
280	fn single_byte_terms() {
281		assert_concat_eq(&[1, 1, 1, 1, 1], &[b"a", b"b", b"c", b"d", b"e"], b"abcde");
282	}
283
284	#[test]
285	fn domain_concat() {
286		assert_concat_eq(
287			&[1, 1, 1, 1, 1],
288			&[b"api", b".", b"example", b".", b"com"],
289			b"api.example.com",
290		);
291	}
292
293	#[test]
294	fn different_term_max_lens() {
295		assert_concat_eq(&[1, 3], &[b"short", b"a very long string"], b"shorta very long string");
296	}
297
298	#[test]
299	fn mixed_term_sizes() {
300		assert_concat_eq(
301			&[1, 1, 4, 1, 2],
302			&[b"hi", b".", b"this is a much longer term", b".", b"bye"],
303			b"hi.this is a much longer term.bye",
304		);
305	}
306
307	#[test]
308	fn many_terms() {
309		// 50 two-byte terms.
310		let input_max_lens = vec![1usize; 50];
311		let data: Vec<Vec<u8>> = (0..50u8).map(|i| vec![i, i]).collect();
312		let data_refs: Vec<&[u8]> = data.iter().map(|v| v.as_slice()).collect();
313		let expected: Vec<u8> = data.iter().flatten().copied().collect();
314		assert_concat_eq(&input_max_lens, &data_refs, &expected);
315	}
316
317	#[test]
318	fn full_word_terms() {
319		// Terms with lengths that are exact multiples of 8.
320		assert_concat_eq(&[1, 2], &[b"01234567", b"abcdefgh01234567"], b"01234567abcdefgh01234567");
321	}
322
323	#[test]
324	fn mutated_output_fails_constraints() {
325		// Build and populate a valid concatenation, then mutate one of the hint-produced output
326		// data wires. Constraint verification should reject because the slice extraction no longer
327		// matches the inputs.
328		let b = CircuitBuilder::new();
329		let inputs = vec![ByteVec::new_inout(&b, 1), ByteVec::new_inout(&b, 1)];
330		let output = concat(&b, &inputs);
331
332		let circuit = b.build();
333		let mut filler = circuit.new_witness_filler();
334		inputs[0].populate_len_bytes(&mut filler, 5);
335		inputs[0].populate_data(&mut filler, b"hello");
336		inputs[1].populate_len_bytes(&mut filler, 5);
337		inputs[1].populate_data(&mut filler, b"world");
338
339		circuit.populate_wire_witness(&mut filler).unwrap();
340		// Corrupt one byte of the hint output.
341		filler[output.data[0]] = Word(filler[output.data[0]].as_u64() ^ 1);
342
343		let cs = circuit.constraint_system();
344		assert!(cs.verify(&filler.into_value_vec()).is_err());
345	}
346
347	/// A concat input for [`run_concat_mixed`]: either a constant-length term (built with
348	/// `new_const_len`, exercising the static comparison paths) or a dynamic-length term (built
349	/// with `new_inout`, exercising the dynamic `slice` path).
350	enum Term<'a> {
351		Const(&'a [u8]),
352		Dyn { bytes: &'a [u8], cap_words: usize },
353	}
354
355	impl Term<'_> {
356		fn bytes(&self) -> &[u8] {
357			match self {
358				Term::Const(d) => d,
359				Term::Dyn { bytes, .. } => bytes,
360			}
361		}
362	}
363
364	/// Build a concat over a mix of constant- and dynamic-length terms, populate, verify
365	/// constraints, and return the concatenated bytes. Lets tests drive the constant-offset /
366	/// constant-length branches that the `new_inout`-based `run_concat` never reaches.
367	fn run_concat_mixed(terms: &[Term<'_>]) -> Result<Vec<u8>> {
368		let b = CircuitBuilder::new();
369		let inputs: Vec<ByteVec> = terms
370			.iter()
371			.map(|t| match t {
372				Term::Const(d) => {
373					let n_words = d.len().div_ceil(8);
374					let data: Vec<Wire> = (0..n_words).map(|_| b.add_witness()).collect();
375					ByteVec::new_const_len(&b, data, d.len())
376				}
377				Term::Dyn { cap_words, .. } => ByteVec::new_inout(&b, *cap_words),
378			})
379			.collect();
380
381		let output = concat(&b, &inputs);
382		// A hint emits no constraint of its own, so pinning alone leaves these uncommitted.
383		// Promoting them to public outputs is what the test needs to read them back.
384		for &wire in &output.data {
385			b.mark_inout(wire);
386		}
387
388		let circuit = b.build();
389		let mut filler = circuit.new_witness_filler();
390		for (input, t) in inputs.iter().zip(terms) {
391			match t {
392				Term::Const(d) => input.populate_data(&mut filler, d),
393				Term::Dyn { bytes, .. } => {
394					input.populate_len_bytes(&mut filler, bytes.len());
395					input.populate_data(&mut filler, bytes);
396				}
397			}
398		}
399
400		circuit
401			.populate_wire_witness(&mut filler)
402			.map_err(|e| anyhow!("populate_wire_witness: {e}"))?;
403
404		let total_len: usize = terms.iter().map(|t| t.bytes().len()).sum();
405		let mut bytes = Vec::with_capacity(total_len);
406		for &w in &output.data {
407			let word = filler[w].as_u64();
408			for j in 0..8 {
409				bytes.push(((word >> (j * 8)) & 0xff) as u8);
410			}
411		}
412		bytes.truncate(total_len);
413
414		let cs = circuit.constraint_system();
415		cs.verify(&filler.into_value_vec())
416			.map_err(|err| anyhow!("constraint verification failed: {err}"))?;
417
418		Ok(bytes)
419	}
420
421	#[test]
422	fn const_terms_aligned() {
423		// All lengths are multiples of 8: constant offsets, no boundary masking.
424		let bytes = run_concat_mixed(&[Term::Const(b"01234567"), Term::Const(b"abcdefgh01234567")])
425			.unwrap();
426		assert_eq!(bytes, b"01234567abcdefgh01234567");
427	}
428
429	#[test]
430	fn const_terms_unaligned() {
431		// Non-8-multiple lengths exercise unaligned constant-offset extraction and the masked
432		// final-word comparison.
433		let bytes = run_concat_mixed(&[
434			Term::Const(b"hello"),
435			Term::Const(b"world!"),
436			Term::Const(b"abc"),
437		])
438		.unwrap();
439		assert_eq!(bytes, b"helloworld!abc");
440	}
441
442	#[test]
443	fn const_single_term() {
444		let bytes = run_concat_mixed(&[Term::Const(b"solo")]).unwrap();
445		assert_eq!(bytes, b"solo");
446	}
447
448	#[test]
449	fn const_empty_middle_term() {
450		let bytes =
451			run_concat_mixed(&[Term::Const(b"ab"), Term::Const(b""), Term::Const(b"cd")]).unwrap();
452		assert_eq!(bytes, b"abcd");
453	}
454
455	#[test]
456	fn const_then_dynamic() {
457		// First term constant ⇒ second term has a constant offset but dynamic length: the
458		// "constant offset, dynamic length" branch.
459		let bytes = run_concat_mixed(&[
460			Term::Const(b"hdr-"),
461			Term::Dyn {
462				bytes: b"payload",
463				cap_words: 2,
464			},
465		])
466		.unwrap();
467		assert_eq!(bytes, b"hdr-payload");
468	}
469
470	#[test]
471	fn dynamic_then_const() {
472		// First term dynamic ⇒ offset is dynamic for the second (constant-length) term, which then
473		// takes the fully-dynamic slice path.
474		let bytes = run_concat_mixed(&[
475			Term::Dyn {
476				bytes: b"payload",
477				cap_words: 2,
478			},
479			Term::Const(b"-end"),
480		])
481		.unwrap();
482		assert_eq!(bytes, b"payload-end");
483	}
484
485	#[test]
486	fn const_mutated_output_fails_constraints() {
487		// The constant-length comparison path must still bind the hint output to the inputs:
488		// corrupting a constrained output byte should be rejected.
489		let b = CircuitBuilder::new();
490		let inputs = vec![
491			ByteVec::new_const_len(&b, vec![b.add_witness()], 5),
492			ByteVec::new_const_len(&b, vec![b.add_witness()], 5),
493		];
494		let output = concat(&b, &inputs);
495
496		let circuit = b.build();
497		let mut filler = circuit.new_witness_filler();
498		inputs[0].populate_data(&mut filler, b"hello");
499		inputs[1].populate_data(&mut filler, b"world");
500
501		circuit.populate_wire_witness(&mut filler).unwrap();
502		// Corrupt a byte within the constrained region of the hint output.
503		filler[output.data[0]] = Word(filler[output.data[0]].as_u64() ^ 1);
504
505		let cs = circuit.constraint_system();
506		assert!(cs.verify(&filler.into_value_vec()).is_err());
507	}
508
509	#[test]
510	fn const_unaligned_then_dynamic() {
511		// Unaligned constant offset (5) feeding a dynamic-length term.
512		let bytes = run_concat_mixed(&[
513			Term::Const(b"hello"),
514			Term::Dyn {
515				bytes: b" world",
516				cap_words: 1,
517			},
518		])
519		.unwrap();
520		assert_eq!(bytes, b"hello world");
521	}
522
523	#[cfg(test)]
524	mod proptests {
525		use proptest::prelude::*;
526		use rand::{Rng, SeedableRng, rngs::StdRng};
527
528		use super::*;
529
530		fn random_bytes(len: usize, seed: u64) -> Vec<u8> {
531			let mut rng = StdRng::seed_from_u64(seed);
532			let mut data = vec![0u8; len];
533			rng.fill_bytes(&mut data);
534			data
535		}
536
537		fn term_strategy() -> impl Strategy<Value = (Vec<u8>, usize)> {
538			(0..=24usize, any::<u64>()).prop_map(|(len, seed)| {
539				let max_len = (len.div_ceil(8)).max(1);
540				(random_bytes(len, seed), max_len)
541			})
542		}
543
544		fn terms_strategy() -> impl Strategy<Value = Vec<(Vec<u8>, usize)>> {
545			prop::collection::vec(term_strategy(), 1..=4)
546		}
547
548		proptest! {
549			#[test]
550			fn correct_concatenation(terms in terms_strategy()) {
551				let input_max_lens: Vec<usize> = terms.iter().map(|(_, n)| *n).collect();
552				let data: Vec<&[u8]> = terms.iter().map(|(d, _)| d.as_slice()).collect();
553				let expected: Vec<u8> = data.iter().flat_map(|d| d.iter().copied()).collect();
554				let bytes = run_concat(&input_max_lens, &data).unwrap();
555				prop_assert_eq!(bytes, expected);
556			}
557
558			/// Exercise the constant-length path (`new_const_len` terms) across many length
559			/// combinations, especially non-8-multiples that hit the masked boundary-word logic and
560			/// unaligned constant-offset extraction.
561			#[test]
562			fn correct_const_concatenation(
563				lens in prop::collection::vec(0..=20usize, 1..=5),
564				seed in any::<u64>(),
565			) {
566				let term_bytes: Vec<Vec<u8>> = lens
567					.iter()
568					.enumerate()
569					.map(|(i, &n)| random_bytes(n, seed.wrapping_add(i as u64)))
570					.collect();
571				let terms: Vec<Term<'_>> = term_bytes.iter().map(|d| Term::Const(d.as_slice())).collect();
572				let expected: Vec<u8> = term_bytes.iter().flatten().copied().collect();
573				let bytes = run_concat_mixed(&terms).unwrap();
574				prop_assert_eq!(bytes, expected);
575			}
576		}
577	}
578}