1use binius_core::word::Word;
4use binius_frontend::{CircuitBuilder, Wire, WitnessFiller};
5
6pub struct Base64UrlSafe {
33 pub decoded: Vec<Wire>,
35 pub encoded: Vec<Wire>,
37 pub len_bytes: Wire,
39}
40
41impl Base64UrlSafe {
42 pub fn new(
62 builder: &CircuitBuilder,
63 decoded: Vec<Wire>,
64 encoded: Vec<Wire>,
65 len_bytes: Wire,
66 ) -> Self {
67 verify_length_bounds(builder, len_bytes, decoded.len() << 3);
69
70 let groups = (decoded.len() << 3).div_ceil(3); for group_idx in 0..groups {
74 let b = builder.subcircuit(format!("group[{group_idx}]"));
75 verify_base64_group(&b, &decoded, &encoded, len_bytes, group_idx);
76 }
77
78 Self {
79 decoded,
80 encoded,
81 len_bytes,
82 }
83 }
84
85 pub fn populate_len_bytes(&self, w: &mut WitnessFiller<'_>, len_bytes: usize) {
92 w[self.len_bytes] = Word(len_bytes as u64);
93 }
94
95 pub fn populate_decoded(&self, w: &mut WitnessFiller<'_>, data: &[u8]) {
106 w.pack_bytes_le(&self.decoded, data);
107 }
108
109 pub fn populate_encoded(&self, w: &mut WitnessFiller<'_>, data: &[u8]) {
120 w.pack_bytes_le(&self.encoded, data);
121 }
122}
123
124fn verify_length_bounds(builder: &CircuitBuilder, len_bytes: Wire, max_len_bytes: usize) {
126 let too_long = builder.icmp_ugt(len_bytes, builder.add_constant_64(max_len_bytes as u64));
128 builder.assert_false("length_check", too_long);
129}
130
131fn verify_base64_group(
142 builder: &CircuitBuilder,
143 decoded: &[Wire],
144 encoded: &[Wire],
145 len_bytes: Wire,
146 group_idx: usize,
147) {
148 let base_byte_idx = group_idx * 3;
149 let base_char_idx = group_idx * 4;
150
151 let byte0 = extract_byte(builder, decoded, base_byte_idx);
153 let byte1 = extract_byte(builder, decoded, base_byte_idx + 1);
154 let byte2 = extract_byte(builder, decoded, base_byte_idx + 2);
155
156 let has_1 = builder.icmp_ult(builder.add_constant_64(base_byte_idx as u64), len_bytes);
157 let has_2 = builder.icmp_ult(builder.add_constant_64((base_byte_idx + 1) as u64), len_bytes);
158 let has_3 = builder.icmp_ult(builder.add_constant_64((base_byte_idx + 2) as u64), len_bytes);
159
160 let zero = builder.add_constant(Word::ZERO);
161 builder.assert_eq_cond("past boundary should be empty", byte0, zero, builder.bnot(has_1));
162 builder.assert_eq_cond("past boundary should be empty", byte1, zero, builder.bnot(has_2));
163 builder.assert_eq_cond("past boundary should be empty", byte2, zero, builder.bnot(has_3));
164
165 let val0 = extract_6bit_value_0(builder, byte0);
167 let val1 = extract_6bit_value_1(builder, byte0, byte1);
168 let val2 = extract_6bit_value_2(builder, byte1, byte2);
169 let val3 = extract_6bit_value_3(builder, byte2);
170
171 let expected_char0 = compute_expected_base64_char(builder, val0);
173 let expected_char1 = compute_expected_base64_char(builder, val1);
174 let expected_char2 = compute_expected_base64_char(builder, val2);
175 let expected_char3 = compute_expected_base64_char(builder, val3);
176
177 let actual_char0 = extract_byte(builder, encoded, base_char_idx);
179 let actual_char1 = extract_byte(builder, encoded, base_char_idx + 1);
180 let actual_char2 = extract_byte(builder, encoded, base_char_idx + 2);
181 let actual_char3 = extract_byte(builder, encoded, base_char_idx + 3);
182
183 verify_base64_char(builder, expected_char0, actual_char0, has_1);
184 verify_base64_char(builder, expected_char1, actual_char1, has_1);
185 verify_base64_char(builder, expected_char2, actual_char2, has_2);
186 verify_base64_char(builder, expected_char3, actual_char3, has_3);
187}
188
189fn extract_byte(builder: &CircuitBuilder, words: &[Wire], byte_idx: usize) -> Wire {
201 let word_idx = byte_idx / 8;
202 let byte_offset = byte_idx % 8;
203
204 let zero = builder.add_constant(Word::ZERO);
205 let word = words.get(word_idx).copied().unwrap_or(zero);
206 builder.extract_byte(word, byte_offset as u32)
207}
208
209fn extract_6bit_value_0(builder: &CircuitBuilder, byte0: Wire) -> Wire {
211 builder.shr(byte0, 2)
212}
213
214fn extract_6bit_value_1(builder: &CircuitBuilder, byte0: Wire, byte1: Wire) -> Wire {
216 let byte0_low = builder.band(byte0, builder.add_constant_64(0x03));
217 builder.bxor(builder.shl(byte0_low, 4), builder.shr(byte1, 4))
218}
219
220fn extract_6bit_value_2(builder: &CircuitBuilder, byte1: Wire, byte2: Wire) -> Wire {
222 let byte1_low = builder.band(byte1, builder.add_constant_64(0x0F));
223 builder.bxor(builder.shl(byte1_low, 2), builder.shr(byte2, 6))
224}
225
226fn extract_6bit_value_3(builder: &CircuitBuilder, byte2: Wire) -> Wire {
228 builder.band(byte2, builder.add_constant_64(0x3F))
229}
230
231fn verify_base64_char(
240 builder: &CircuitBuilder,
241 expected_encoded_char: Wire,
242 actual_encoded_char: Wire,
243 is_active: Wire,
244) {
245 builder.assert_eq(
246 "base64_char",
247 actual_encoded_char,
248 builder.select(is_active, expected_encoded_char, builder.add_constant(Word::ZERO)),
249 );
250}
251
252fn compute_expected_base64_char(builder: &CircuitBuilder, six_bit_val: Wire) -> Wire {
278 let is_lowercase = builder.icmp_uge(six_bit_val, builder.add_constant_64(26));
279 let is_digit = builder.icmp_uge(six_bit_val, builder.add_constant_64(52));
280 let is_symbol = builder.icmp_uge(six_bit_val, builder.add_constant_64(62));
281
282 let offset = builder.select(
285 is_lowercase,
286 builder.add_constant_64(b'a' as u64 - 26),
287 builder.add_constant_64(b'A' as u64),
288 );
289 let offset = builder.select(is_digit, builder.add_constant_64(0xfc), offset);
290
291 let sum = builder.iadd_32(six_bit_val, offset);
293 let alphanumeric = builder.band(sum, builder.add_constant_64(0xff));
294
295 let symbol = builder.select(
298 builder.shl(six_bit_val, 63),
299 builder.add_constant_64(b'_' as u64),
300 builder.add_constant_64(b'-' as u64),
301 );
302
303 builder.select(is_symbol, symbol, alphanumeric)
304}
305
306#[cfg(test)]
307mod tests {
308 use binius_frontend::CircuitBuilder;
309
310 use super::{Base64UrlSafe, Wire};
311
312 fn encode_base64(input: &[u8]) -> Vec<u8> {
315 const BASE64_CHARS: &[u8] =
316 b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
317
318 let mut output = Vec::new();
319
320 for chunk in input.chunks(3) {
321 let b1 = chunk[0];
322 let b2 = chunk.get(1).copied().unwrap_or(0);
323 let b3 = chunk.get(2).copied().unwrap_or(0);
324
325 let n = ((b1 as u32) << 16) | ((b2 as u32) << 8) | (b3 as u32);
326
327 output.push(BASE64_CHARS[((n >> 18) & 63) as usize]);
328 output.push(BASE64_CHARS[((n >> 12) & 63) as usize]);
329
330 if chunk.len() > 1 {
331 output.push(BASE64_CHARS[((n >> 6) & 63) as usize]);
332 };
333
334 if chunk.len() > 2 {
335 output.push(BASE64_CHARS[(n & 63) as usize]);
336 };
337 }
338
339 output
340 }
341
342 fn create_base64_circuit(builder: &CircuitBuilder, max_len_decoded: usize) -> Base64UrlSafe {
344 assert!(
346 max_len_decoded.is_multiple_of(3),
347 "max_len_decoded must be a multiple of 3, got {max_len_decoded}"
348 );
349 let decoded: Vec<Wire> = (0..max_len_decoded).map(|_| builder.add_inout()).collect();
350 let max_len_encoded = (max_len_decoded / 3) * 4;
351 let encoded: Vec<Wire> = (0..max_len_encoded).map(|_| builder.add_inout()).collect();
352
353 let len_bytes = builder.add_inout();
354
355 Base64UrlSafe::new(builder, decoded, encoded, len_bytes)
356 }
357
358 fn check_base64_encoding(
360 input_bytes: &[u8],
361 encoded: &[u8],
362 max_len_decoded: usize,
363 ) -> Result<(), Box<dyn std::error::Error>> {
364 let builder = CircuitBuilder::new();
365 let circuit = create_base64_circuit(&builder, max_len_decoded);
366 let compiled = builder.build();
367
368 let mut witness = compiled.new_witness_filler();
370
371 circuit.populate_len_bytes(&mut witness, input_bytes.len());
372 circuit.populate_decoded(&mut witness, input_bytes);
373 circuit.populate_encoded(&mut witness, encoded);
374
375 compiled.populate_wire_witness(&mut witness)?;
377
378 let cs = compiled.constraint_system();
380 cs.verify(&witness.into_value_vec())?;
381
382 Ok(())
383 }
384
385 fn test_base64_encoding(input: &[u8], max_len_decoded: usize) {
387 let expected_base64 = encode_base64(input);
388 check_base64_encoding(input, &expected_base64, max_len_decoded).unwrap();
389 }
390
391 fn assert_base64_failure(input: &[u8], encoded: &[u8], max_len_decoded: usize) {
393 check_base64_encoding(input, encoded, max_len_decoded).unwrap_err();
394 }
395
396 #[test]
397 fn test_base64_hello_world() {
398 test_base64_encoding(b"Hello World!", 189);
399 }
400
401 #[test]
402 fn test_base64_empty() {
403 test_base64_encoding(b"", 189);
404 }
405
406 #[test]
407 fn test_base64_long_input() {
408 let input =
409 b"The quick brown fox jumps over the lazy dog. The quick brown fox jumps over the lazy dog.";
410 test_base64_encoding(input, 189);
411 }
412
413 #[test]
414 fn test_invalid_base64() {
415 let input = b"ABC";
416 let invalid_base64 = b"XXXX"; assert_base64_failure(input, invalid_base64, 15);
418 }
419
420 fn all_six_bit_values() -> Vec<u8> {
423 let mut bytes = vec![0u8; 48];
424 for value in 0..64usize {
425 for k in 0..6 {
426 if (value >> (5 - k)) & 1 == 1 {
427 let bit_index = value * 6 + k;
428 bytes[bit_index / 8] |= 1 << (7 - bit_index % 8);
429 }
430 }
431 }
432 bytes
433 }
434
435 #[test]
436 fn test_every_alphabet_index() {
437 const BASE64_CHARS: &[u8] =
438 b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
439
440 let input = all_six_bit_values();
441 let encoded = encode_base64(&input);
442 assert_eq!(encoded, BASE64_CHARS, "fixture must cover every alphabet index");
443
444 check_base64_encoding(&input, &encoded, 48).unwrap();
447 }
448
449 #[test]
450 fn test_every_alphabet_index_rejects_wrong_character() {
451 let input = all_six_bit_values();
452 let encoded = encode_base64(&input);
453
454 for position in 0..encoded.len() {
456 let mut corrupted = encoded.clone();
457 corrupted[position] ^= 1;
459 assert_base64_failure(&input, &corrupted, 48);
460 }
461 }
462
463 #[test]
464 fn test_url_safe_characters() {
465 let input1 = &[0b11111000]; let expected1 = encode_base64(input1);
471 assert_eq!(expected1[0], b'-', "Index 62 should map to '-' not '+'");
472
473 let input2 = &[0b11111100]; let expected2 = encode_base64(input2);
476 assert_eq!(expected2[0], b'_', "Index 63 should map to '_' not '/'");
477
478 test_base64_encoding(input1, 15);
480 test_base64_encoding(input2, 15);
481 }
482
483 #[test]
484 #[should_panic(expected = "max_len_decoded must be a multiple of 3")]
485 fn test_panic_when_max_len_not_multiple_of_3() {
486 test_base64_encoding(b"test", 13);
489 }
490
491 #[test]
492 fn test_encoding_with_padding_rejected() {
493 let input = b"A";
494 let encoding_with_padding = b"QQ==";
495 assert_base64_failure(input, encoding_with_padding, 15);
497 }
498}