1use 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
10struct 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
61pub 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 let mut offset = b.add_constant(Word::ZERO);
94 let mut offset_range = 0usize..0usize;
95 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 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 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 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, ®ion, &input.data);
136 }
137 } else {
138 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
159fn assert_const_slice_eq(
164 b: &CircuitBuilder,
165 name: &str,
166 output: &[Wire],
167 offset: usize,
168 input: &[Wire],
169 len: usize,
170) {
171 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 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 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 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 assert_concat_eq(&[1, 2], &[b"01234567", b"abcdefgh01234567"], b"01234567abcdefgh01234567");
321 }
322
323 #[test]
324 fn mutated_output_fails_constraints() {
325 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 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 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 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 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 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 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 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 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 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 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 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 #[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}