1use std::iter;
10
11use binius_core::Word;
12use binius_frontend::{CircuitBuilder, Hint, Wire};
13use rand::CryptoRng;
14
15use super::{
16 CHAIN_LENGTH, DIGEST_LEN, DIGEST_WIRES, Digest, MESSAGE_LEN, MESSAGE_WIRES, Message,
17 NUM_CHAIN_HASHES, PUBLIC_PARAM_LEN, PUBLIC_PARAM_WIRES, PublicParam, RANDOMNESS_LEN,
18 RANDOMNESS_WIRES, Randomness, TARGET_SUM, V, W,
19 hashing::{
20 TWEAK_TYPE_CHAIN, TWEAK_TYPE_ENCODING, TWEAK_TYPE_WOTS_PK, circuit_tweak_hash,
21 circuit_tweak_hash_2x, tweak_hash,
22 },
23};
24use crate::multiplexer::multi_wire_multiplex;
25
26const DIGITS_PER_WORD: usize = V / 2;
28
29const ENCODING_PAYLOAD_LEN: usize = MESSAGE_LEN + RANDOMNESS_LEN;
32
33const _: () = assert!(ENCODING_PAYLOAD_LEN <= 64);
34const _: () = assert!(2 * DIGITS_PER_WORD == V);
35
36pub fn wots_encode(
50 message: &Message,
51 epoch: u32,
52 public_param: &PublicParam,
53 randomness: &Randomness,
54) -> Option<[u8; V]> {
55 let mut data = [0u8; ENCODING_PAYLOAD_LEN];
56 data[..MESSAGE_LEN].copy_from_slice(message);
57 data[MESSAGE_LEN..][..RANDOMNESS_LEN].copy_from_slice(randomness);
58 let digest = tweak_hash(public_param, TWEAK_TYPE_ENCODING, 0, epoch, &data);
59
60 if digest[7] >> 7 != 0 || digest[DIGEST_LEN - 1] >> 7 != 0 {
61 return None; }
63 let bit = |j: usize| (digest[j / 8] >> (j % 8)) & 1;
64 let pos = |i: usize| {
65 if i < DIGITS_PER_WORD {
66 W * i
67 } else {
68 64 + W * (i - DIGITS_PER_WORD)
69 }
70 };
71 let encoding: [u8; V] =
72 std::array::from_fn(|i| (0..W).fold(0, |acc, k| acc | (bit(pos(i) + k) << k)));
73 (encoding.iter().map(|&x| x as usize).sum::<usize>() == TARGET_SUM).then_some(encoding)
74}
75
76pub fn find_randomness_for_wots_encoding(
81 message: &Message,
82 epoch: u32,
83 public_param: &PublicParam,
84 rng: &mut impl CryptoRng,
85) -> (Randomness, [u8; V]) {
86 loop {
87 let mut randomness = [0u8; RANDOMNESS_LEN];
88 rng.fill_bytes(&mut randomness);
89 if let Some(encoding) = wots_encode(message, epoch, public_param, &randomness) {
90 return (randomness, encoding);
91 }
92 }
93}
94
95pub fn chain_step(
100 public_param: &PublicParam,
101 epoch: u32,
102 chain_index: usize,
103 step: usize,
104 x: &Digest,
105) -> Digest {
106 let position = (chain_index * CHAIN_LENGTH + step) as u32;
107 tweak_hash(public_param, TWEAK_TYPE_CHAIN, position, epoch, x)
108}
109
110pub fn iterate_hash(
112 a: &Digest,
113 n: usize,
114 public_param: &PublicParam,
115 epoch: u32,
116 chain_index: usize,
117 start_step: usize,
118) -> Digest {
119 (0..n).fold(*a, |acc, j| chain_step(public_param, epoch, chain_index, start_step + j, &acc))
120}
121
122pub fn recover_public_key(
124 chain_tips: &[Digest; V],
125 encoding: &[u8; V],
126 epoch: u32,
127 public_param: &PublicParam,
128) -> [Digest; V] {
129 std::array::from_fn(|i| {
130 let digit = encoding[i] as usize;
131 iterate_hash(&chain_tips[i], CHAIN_LENGTH - 1 - digit, public_param, epoch, i, digit)
132 })
133}
134
135pub fn wots_public_key_hash(
137 public_param: &PublicParam,
138 epoch: u32,
139 chain_ends: &[Digest; V],
140) -> Digest {
141 let mut data = [0u8; V * DIGEST_LEN];
142 for (chunk, end) in iter::zip(data.chunks_exact_mut(DIGEST_LEN), chain_ends) {
143 chunk.copy_from_slice(end);
144 }
145 tweak_hash(public_param, TWEAK_TYPE_WOTS_PK, 0, epoch, &data)
146}
147
148pub fn circuit_wots_encode(
158 builder: &CircuitBuilder,
159 public_param: &[Wire; PUBLIC_PARAM_WIRES],
160 epoch: Wire,
161 message: &[Wire; MESSAGE_WIRES],
162 randomness: &[Wire; RANDOMNESS_WIRES],
163) -> [Wire; V] {
164 let zero = builder.add_constant_64(0);
165
166 let mut payload = Vec::with_capacity(ENCODING_PAYLOAD_LEN / 8);
168 payload.extend_from_slice(message);
169 payload.extend_from_slice(randomness);
170
171 let digest =
172 circuit_tweak_hash(builder, public_param, TWEAK_TYPE_ENCODING, zero, epoch, &payload);
173
174 for (k, &word) in digest.iter().enumerate() {
177 builder.assert_zero(format!("encoding_leftover_bit[{k}]"), builder.shr(word, 63));
178 }
179
180 let digits: [Wire; V] = std::array::from_fn(|i| {
183 let word = digest[i / DIGITS_PER_WORD];
184 let shift = (W * (i % DIGITS_PER_WORD)) as u32;
185 builder.shr(builder.shl(word, u64::BITS - shift - W as u32), u64::BITS - W as u32)
186 });
187
188 let sum = digits
190 .iter()
191 .fold(zero, |acc, &digit| builder.iadd(acc, digit).0);
192 builder.assert_eq("encoding_target_sum", sum, builder.add_constant_64(TARGET_SUM as u64));
193
194 digits
195}
196
197const HINT_WORDS_PER_HASH: usize = DIGEST_WIRES + 3;
199
200const HINT_INPUTS: usize = PUBLIC_PARAM_WIRES + 1 + V + V * DIGEST_WIRES;
202
203const HINT_OUTPUTS: usize = NUM_CHAIN_HASHES * HINT_WORDS_PER_HASH + V;
205
206struct ChainHashesHint;
216
217impl Hint for ChainHashesHint {
218 const NAME: &'static str = "binius.xmss_wots_chain_hashes";
219
220 fn shape(&self, _dimensions: &[usize]) -> (usize, usize) {
221 (HINT_INPUTS, HINT_OUTPUTS)
222 }
223
224 fn execute(&self, _dimensions: &[usize], inputs: &[Word], outputs: &mut [Word]) {
225 let public_param = bytes_from_words::<PUBLIC_PARAM_LEN>(&inputs[..PUBLIC_PARAM_WIRES]);
226 let epoch = inputs[PUBLIC_PARAM_WIRES].as_u64() as u32;
227 let digits = &inputs[PUBLIC_PARAM_WIRES + 1..][..V];
228 let tips = &inputs[PUBLIC_PARAM_WIRES + 1 + V..];
229
230 let (hashes, offsets) = outputs.split_at_mut(NUM_CHAIN_HASHES * HINT_WORDS_PER_HASH);
231 hashes.fill(Word::ZERO);
232 offsets.fill(Word::ZERO);
233
234 let mut written = 0;
235 for i in 0..V {
236 let digit = digits[i].as_u64() as usize;
237 let mut current =
238 bytes_from_words::<DIGEST_LEN>(&tips[i * DIGEST_WIRES..][..DIGEST_WIRES]);
239
240 for (step, position) in (digit..CHAIN_LENGTH - 1).enumerate() {
243 if written == NUM_CHAIN_HASHES {
246 break;
247 }
248 let slot = &mut hashes[written * HINT_WORDS_PER_HASH..][..HINT_WORDS_PER_HASH];
249 bytes_to_words(¤t, &mut slot[..DIGEST_WIRES]);
250 slot[DIGEST_WIRES] = Word::from_u64(i as u64);
251 slot[DIGEST_WIRES + 1] = Word::from_u64(step as u64);
252 slot[DIGEST_WIRES + 2] = Word::from_u64(digit as u64);
253
254 current = chain_step(&public_param, epoch, i, position, ¤t);
255 written += 1;
256 offsets[i] = Word::from_u64((written - 1) as u64);
258 }
259 }
260 }
261}
262
263fn bytes_from_words<const N: usize>(words: &[Word]) -> [u8; N] {
265 let mut bytes = [0u8; N];
266 for (chunk, word) in iter::zip(bytes.chunks_exact_mut(8), words) {
267 chunk.copy_from_slice(&word.as_u64().to_le_bytes());
268 }
269 bytes
270}
271
272fn bytes_to_words(bytes: &[u8], words: &mut [Word]) {
274 for (word, chunk) in iter::zip(words, bytes.chunks_exact(8)) {
275 *word = Word::from_u64(u64::from_le_bytes(chunk.try_into().expect("eight bytes")));
276 }
277}
278
279pub fn circuit_recover_public_key(
313 builder: &CircuitBuilder,
314 public_param: &[Wire; PUBLIC_PARAM_WIRES],
315 epoch: Wire,
316 chain_tips: &[[Wire; DIGEST_WIRES]; V],
317 digits: &[Wire; V],
318) -> [[Wire; DIGEST_WIRES]; V] {
319 let mut hint_inputs = Vec::with_capacity(HINT_INPUTS);
320 hint_inputs.extend_from_slice(public_param);
321 hint_inputs.push(epoch);
322 hint_inputs.extend_from_slice(digits);
323 hint_inputs.extend(chain_tips.iter().flatten().copied());
324 let hinted = builder.call_hint(ChainHashesHint, &[], &hint_inputs);
325
326 let input_of = |k: usize| -> [Wire; DIGEST_WIRES] {
327 std::array::from_fn(|w| hinted[k * HINT_WORDS_PER_HASH + w])
328 };
329 let chain_of = |k: usize| hinted[k * HINT_WORDS_PER_HASH + DIGEST_WIRES];
330 let step_of = |k: usize| hinted[k * HINT_WORDS_PER_HASH + DIGEST_WIRES + 1];
331 let digit_of = |k: usize| hinted[k * HINT_WORDS_PER_HASH + DIGEST_WIRES + 2];
332 let offset_of = |i: usize| hinted[NUM_CHAIN_HASHES * HINT_WORDS_PER_HASH + i];
333
334 let window = |i: usize| -> (usize, usize) {
338 let per_chain = CHAIN_LENGTH - 1;
339 let earliest = (NUM_CHAIN_HASHES).saturating_sub((V - i) * per_chain);
340 let latest = ((i + 1) * per_chain - 1).min(NUM_CHAIN_HASHES - 1);
341 (earliest, latest)
342 };
343
344 let sub_position = |k: usize| {
348 let position = builder.iadd(digit_of(k), step_of(k)).0;
349 builder.bxor(builder.shl(chain_of(k), W as u32), position)
350 };
351 let mut outputs = Vec::with_capacity(NUM_CHAIN_HASHES);
352 for pair in 0..NUM_CHAIN_HASHES / 2 {
353 let (a, b) = (2 * pair, 2 * pair + 1);
354 let (in_a, in_b) = (input_of(a), input_of(b));
355 let digests = circuit_tweak_hash_2x(
356 builder,
357 public_param,
358 TWEAK_TYPE_CHAIN,
359 [sub_position(a), sub_position(b)],
360 epoch,
361 [&in_a, &in_b],
362 );
363 outputs.extend_from_slice(&digests);
364 }
365 if NUM_CHAIN_HASHES % 2 == 1 {
366 let k = NUM_CHAIN_HASHES - 1;
367 outputs.push(circuit_tweak_hash(
368 builder,
369 public_param,
370 TWEAK_TYPE_CHAIN,
371 sub_position(k),
372 epoch,
373 &input_of(k),
374 ));
375 }
376
377 let zero = builder.add_constant_64(0);
378 let one = builder.add_constant_64(1);
379
380 for k in 0..NUM_CHAIN_HASHES {
383 let b = builder.subcircuit(format!("chain_hash[{k}]"));
384 let (chain, step) = (chain_of(k), step_of(k));
385
386 let Some(previous) = k.checked_sub(1) else {
387 b.assert_eq("first_step_is_zero", step, zero);
389 continue;
390 };
391
392 let continues = b.icmp_eq(chain, chain_of(previous));
395 let opens = b.bnot(continues);
396 b.assert_true("chain_non_decreasing", b.icmp_ule(chain_of(previous), chain));
397
398 let next_step = b.iadd(step_of(previous), one).0;
400 b.assert_eq("step_advances", b.select(continues, step, next_step), next_step);
401 b.assert_eq(
402 "digit_holds",
403 b.select(continues, digit_of(k), digit_of(previous)),
404 digit_of(previous),
405 );
406 for w in 0..DIGEST_WIRES {
407 let carried = outputs[previous][w];
408 b.assert_eq(
409 format!("input_continues[{w}]"),
410 b.select(continues, input_of(k)[w], carried),
411 carried,
412 );
413 }
414
415 b.assert_eq("opens_at_zero", b.select(opens, step, zero), zero);
419 }
420
421 std::array::from_fn(|i| {
424 let b = builder.subcircuit(format!("chain_end[{i}]"));
425 let (earliest, latest) = window(i);
426 let entries = (earliest..=latest)
427 .map(|k| {
428 vec![
429 chain_of(k),
430 step_of(k),
431 digit_of(k),
432 outputs[k][0],
433 outputs[k][1],
434 ]
435 })
436 .collect::<Vec<_>>();
437 let rows = entries.iter().map(|e| e.as_slice()).collect::<Vec<_>>();
438
439 let earliest_wire = b.add_constant_64(earliest as u64);
440 let index = b.isub_bin_bout(offset_of(i), earliest_wire, zero).0;
441 let found = multi_wire_multiplex(&b, &rows, index);
442 let (chain, step, digit) = (found[0], found[1], found[2]);
443 let end: [Wire; DIGEST_WIRES] = std::array::from_fn(|w| found[DIGEST_WIRES + 1 + w]);
444
445 let walks = b.bnot(b.icmp_eq(digits[i], b.add_constant_64((CHAIN_LENGTH - 1) as u64)));
448
449 let expected_chain = b.add_constant_64(i as u64);
450 b.assert_eq("chain_is_this_one", b.select(walks, chain, expected_chain), expected_chain);
451 b.assert_eq("digit_is_the_encoding", b.select(walks, digit, digits[i]), digits[i]);
452
453 let last_position = b.add_constant_64((CHAIN_LENGTH - 2) as u64);
456 let position = b.iadd(digit, step).0;
457 b.assert_eq(
458 "ends_at_the_last_position",
459 b.select(walks, position, last_position),
460 last_position,
461 );
462
463 std::array::from_fn(|w| b.select(walks, end[w], chain_tips[i][w]))
464 })
465}
466
467pub fn circuit_wots_public_key_hash(
469 builder: &CircuitBuilder,
470 public_param: &[Wire; PUBLIC_PARAM_WIRES],
471 epoch: Wire,
472 chain_ends: &[[Wire; DIGEST_WIRES]; V],
473) -> [Wire; DIGEST_WIRES] {
474 let payload = chain_ends.iter().flatten().copied().collect::<Vec<_>>();
475 let zero = builder.add_constant_64(0);
476 circuit_tweak_hash(builder, public_param, TWEAK_TYPE_WOTS_PK, zero, epoch, &payload)
477}
478
479#[cfg(test)]
480mod tests {
481 use binius_core::Word;
482 use rand::{Rng, SeedableRng, rngs::StdRng};
483
484 use super::*;
485 use crate::hash_based_sig::PUBLIC_PARAM_LEN;
486
487 struct TestSignature {
489 public_param: PublicParam,
490 message: Message,
491 randomness: Randomness,
492 encoding: [u8; V],
493 chain_tips: [Digest; V],
494 chain_ends: [Digest; V],
495 }
496
497 impl TestSignature {
498 fn generate(rng: &mut StdRng, epoch: u32) -> Self {
499 let mut public_param = [0u8; PUBLIC_PARAM_LEN];
500 rng.fill_bytes(&mut public_param);
501 let mut message = [0u8; MESSAGE_LEN];
502 rng.fill_bytes(&mut message);
503
504 let (randomness, encoding) =
505 find_randomness_for_wots_encoding(&message, epoch, &public_param, rng);
506
507 let mut pre_images = [[0u8; DIGEST_LEN]; V];
510 for pre_image in pre_images.iter_mut() {
511 rng.fill_bytes(pre_image);
512 }
513 let chain_tips: [Digest; V] = std::array::from_fn(|i| {
514 iterate_hash(&pre_images[i], encoding[i] as usize, &public_param, epoch, i, 0)
515 });
516 let chain_ends = recover_public_key(&chain_tips, &encoding, epoch, &public_param);
517
518 Self {
519 public_param,
520 message,
521 randomness,
522 encoding,
523 chain_tips,
524 chain_ends,
525 }
526 }
527 }
528
529 #[test]
530 fn encoding_is_valid_by_construction() {
531 let mut rng = StdRng::seed_from_u64(0);
532 let sig = TestSignature::generate(&mut rng, 7);
533 assert_eq!(sig.encoding.iter().map(|&e| e as usize).sum::<usize>(), TARGET_SUM);
534 assert!(sig.encoding.iter().all(|&e| (e as usize) < CHAIN_LENGTH));
535 }
536
537 #[test]
538 fn the_fixture_covers_empty_and_walked_chains() {
539 let mut rng = StdRng::seed_from_u64(1);
543 let sig = TestSignature::generate(&mut rng, 12345);
544 assert!(
545 sig.encoding.iter().any(|&e| e as usize == CHAIN_LENGTH - 1),
546 "no chain is empty, so the zero-hash path goes unchecked"
547 );
548 assert!(
549 sig.encoding
550 .iter()
551 .any(|&e| (e as usize) < CHAIN_LENGTH - 1),
552 "every chain is empty, so no chain hash is walked"
553 );
554 }
555
556 #[test]
557 fn a_chain_end_is_its_tip_at_the_last_digit() {
558 let pp = [4u8; PUBLIC_PARAM_LEN];
560 let tip = [9u8; DIGEST_LEN];
561 assert_eq!(iterate_hash(&tip, CHAIN_LENGTH - 1 - (CHAIN_LENGTH - 1), &pp, 3, 0, 7), tip);
562 }
563
564 fn run(sig: &TestSignature, epoch: u32) -> Result<(), String> {
567 let b = CircuitBuilder::new();
568 let param_w: [Wire; PUBLIC_PARAM_WIRES] = std::array::from_fn(|_| b.add_inout());
569 let epoch_w = b.add_inout();
570 let message_w: [Wire; MESSAGE_WIRES] = std::array::from_fn(|_| b.add_inout());
571 let randomness_w: [Wire; RANDOMNESS_WIRES] = std::array::from_fn(|_| b.add_witness());
572 let tips_w: [[Wire; DIGEST_WIRES]; V] =
573 std::array::from_fn(|_| std::array::from_fn(|_| b.add_witness()));
574 let leaf_w: [Wire; DIGEST_WIRES] = std::array::from_fn(|_| b.add_inout());
575
576 let digits = circuit_wots_encode(&b, ¶m_w, epoch_w, &message_w, &randomness_w);
577 let ends = circuit_recover_public_key(&b, ¶m_w, epoch_w, &tips_w, &digits);
578 let leaf = circuit_wots_public_key_hash(&b, ¶m_w, epoch_w, &ends);
579 b.assert_eq_v("leaf", leaf, leaf_w);
580
581 let circuit = b.build();
582 let mut w = circuit.new_witness_filler();
583 w.pack_bytes_le(¶m_w, &sig.public_param);
584 w[epoch_w] = Word::from_u64(epoch as u64);
585 w.pack_bytes_le(&message_w, &sig.message);
586 w.pack_bytes_le(&randomness_w, &sig.randomness);
587 for (wires, tip) in tips_w.iter().zip(&sig.chain_tips) {
588 w.pack_bytes_le(wires, tip);
589 }
590 w.pack_bytes_le(&leaf_w, &wots_public_key_hash(&sig.public_param, epoch, &sig.chain_ends));
591
592 circuit
593 .populate_wire_witness(&mut w)
594 .map_err(|e| format!("populate: {e:?}"))?;
595 circuit
596 .constraint_system()
597 .verify(&w.into_value_vec())
598 .map_err(|e| format!("verify: {e:?}"))
599 }
600
601 #[test]
602 fn circuit_recovers_the_public_key() {
603 let mut rng = StdRng::seed_from_u64(1);
604 let epoch = 12345;
605 let sig = TestSignature::generate(&mut rng, epoch);
606 run(&sig, epoch).unwrap();
607 }
608
609 #[test]
610 fn circuit_digits_match_the_reference_encoding() {
611 let mut rng = StdRng::seed_from_u64(2);
612 let epoch = 9;
613 let sig = TestSignature::generate(&mut rng, epoch);
614
615 let b = CircuitBuilder::new();
616 let param_w: [Wire; PUBLIC_PARAM_WIRES] = std::array::from_fn(|_| b.add_inout());
617 let epoch_w = b.add_inout();
618 let message_w: [Wire; MESSAGE_WIRES] = std::array::from_fn(|_| b.add_inout());
619 let randomness_w: [Wire; RANDOMNESS_WIRES] = std::array::from_fn(|_| b.add_inout());
620 let digits = circuit_wots_encode(&b, ¶m_w, epoch_w, &message_w, &randomness_w);
621 let expected: [Wire; V] = std::array::from_fn(|_| b.add_inout());
622 b.assert_eq_v("digits", digits, expected);
623
624 let circuit = b.build();
625 let mut w = circuit.new_witness_filler();
626 w.pack_bytes_le(¶m_w, &sig.public_param);
627 w[epoch_w] = Word::from_u64(epoch as u64);
628 w.pack_bytes_le(&message_w, &sig.message);
629 w.pack_bytes_le(&randomness_w, &sig.randomness);
630 for (wire, &digit) in expected.iter().zip(&sig.encoding) {
631 w[*wire] = Word::from_u64(digit as u64);
632 }
633
634 circuit.populate_wire_witness(&mut w).unwrap();
635 circuit
636 .constraint_system()
637 .verify(&w.into_value_vec())
638 .unwrap();
639 }
640
641 #[test]
642 fn circuit_rejects_randomness_that_does_not_encode() {
643 let mut rng = StdRng::seed_from_u64(3);
644 let epoch = 4;
645 let mut sig = TestSignature::generate(&mut rng, epoch);
646
647 let mut bad = sig.randomness;
650 bad[0] ^= 0xFF;
651 assert!(
652 wots_encode(&sig.message, epoch, &sig.public_param, &bad).is_none(),
653 "the tampered randomness happened to encode validly; pick another"
654 );
655 sig.randomness = bad;
656 assert!(run(&sig, epoch).is_err(), "an invalid encoding must not verify");
657 }
658
659 #[test]
660 fn circuit_rejects_a_tampered_chain_tip() {
661 let mut rng = StdRng::seed_from_u64(4);
662 let epoch = 4;
663 let mut sig = TestSignature::generate(&mut rng, epoch);
664 sig.chain_tips[0][0] ^= 0xFF;
665 assert!(run(&sig, epoch).is_err(), "a tampered tip must not reach the public key");
666 }
667
668 #[test]
669 fn circuit_rejects_a_signature_from_another_epoch() {
670 let mut rng = StdRng::seed_from_u64(5);
671 let epoch = 4;
672 let sig = TestSignature::generate(&mut rng, epoch);
673 assert!(run(&sig, epoch + 1).is_err(), "an epoch it was not signed at must not verify");
675 }
676}