Skip to main content

binius_circuits/hash_based_sig/
wots.rs

1// Copyright 2026 The Binius Developers
2// Copyright (c) 2026 leanEthereum
3//! WOTS (Winternitz one-time signature) with target-sum encoding.
4//!
5//! There are no checksum chains. Instead the signer grinds the signature randomness until the
6//! encoding's digits sum to [`TARGET_SUM`], and a verifier that checks the sum knows no digit can
7//! have been lowered without another being raised.
8
9use 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
26/// Digits carried by each of the digest's two 64-bit words.
27const DIGITS_PER_WORD: usize = V / 2;
28
29/// The encoding hashes `message | randomness`, which is one BLAKE3 block with room to spare now
30/// that the domain rides in the key rather than the payload.
31const ENCODING_PAYLOAD_LEN: usize = MESSAGE_LEN + RANDOMNESS_LEN;
32
33const _: () = assert!(ENCODING_PAYLOAD_LEN <= 64);
34const _: () = assert!(2 * DIGITS_PER_WORD == V);
35
36/// The target-sum encoding.
37///
38/// `D` is the encoding hash of `message | randomness`, truncated to 16 bytes. Each of its two
39/// little-endian 64-bit words holds 21 digits of [`W`] bits: digit `i < 21` at bits `3i` of word
40/// 0, digit `i >= 21` at bits `3(i - 21)` of word 1.
41///
42/// The encoding is valid exactly when the leftover top bit of *each* word (bits 63 and 127) is
43/// zero and the digits sum to [`TARGET_SUM`]. Grinding the top bits to zero makes each word
44/// exactly `sum(e_i * 2^{3i})` over its 21 digits, so both words decompose into digits with no
45/// slack term.
46///
47/// Returns `None` when the randomness does not produce a valid encoding, which is the signal the
48/// grinding loop retries on.
49pub 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; // the leftover top bit of each 64-bit word must be zero
62	}
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
76/// Draws randomness until it encodes validly.
77///
78/// The encoding is valid when both leftover bits are zero and the digits hit the target sum, so
79/// grinding takes fewer than `2^15` attempts on average.
80pub 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
95/// One chain step.
96///
97/// The position `chain_index * CHAIN_LENGTH + step` identifies the edge from chain value `step` to
98/// `step + 1`, so no two edges anywhere in the instance share a tweak.
99pub 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
110/// Walks chain `chain_index` for `n` steps starting at chain value `start_step`.
111pub 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
122/// Walks every chain from its tip to the public-key end.
123pub 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
135/// The Merkle leaf: the hash over the public parameter and the [`V`] concatenated chain ends.
136pub 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
148/// In-circuit form of [`wots_encode`], returning the digits and constraining them to be a valid
149/// encoding.
150///
151/// Both validity conditions are asserted rather than returned: an encoding whose leftover bits are
152/// set, or whose digits miss the target sum, has no satisfying witness.
153///
154/// # Returns
155///
156/// The [`V`] digits, each a wire holding a value below [`CHAIN_LENGTH`].
157pub 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	// `message | randomness`, which fills its wires exactly.
167	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	// With `V * W = 126` of the digest's 128 bits spent on digits, the two leftover top bits are
175	// what would otherwise let a word carry a slack term the digits do not account for.
176	for (k, &word) in digest.iter().enumerate() {
177		builder.assert_zero(format!("encoding_leftover_bit[{k}]"), builder.shr(word, 63));
178	}
179
180	// A digit is W bits from the middle of its word: lift them to the top of the word, then drop
181	// them back to the bottom, so nothing above or below them survives.
182	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	// The digits are each below CHAIN_LENGTH by construction, so the sum cannot overflow.
189	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
197/// Words the hint emits per chain hash: the input digest, then the chain, the step and the digit.
198const HINT_WORDS_PER_HASH: usize = DIGEST_WIRES + 3;
199
200/// Words the hint reads.
201const HINT_INPUTS: usize = PUBLIC_PARAM_WIRES + 1 + V + V * DIGEST_WIRES;
202
203/// Words the hint writes: one entry per chain hash, then where each chain's last entry sits.
204const HINT_OUTPUTS: usize = NUM_CHAIN_HASHES * HINT_WORDS_PER_HASH + V;
205
206/// Computes the chain hashes a verifier actually walks, and where each chain's last one sits.
207///
208/// The verifier's work is the concatenation of the chain tails: chain `i` contributes its steps
209/// from `digit_i` up to `CHAIN_LENGTH - 2`, and the tails run in chain order. How long each tail
210/// is depends on the digits, so the list cannot be laid out at circuit construction time — it is
211/// hinted here and pinned by the constraints in [`circuit_recover_public_key`].
212///
213/// The offsets are hinted for the same reason and need no separate pinning: an offset is only ever
214/// used to index the list, and what it lands on is checked.
215struct 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			// A chain walks from its digit to the last step. A digit of `CHAIN_LENGTH - 1` walks
241			// nothing, and a digit past that (an unsatisfiable witness) walks nothing either.
242			for (step, position) in (digit..CHAIN_LENGTH - 1).enumerate() {
243				// Digits that miss the target sum overrun the list. The encoding constraints
244				// reject them, so stopping short here only has to avoid a panic.
245				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(&current, &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, &current);
255				written += 1;
256				// An empty chain leaves its offset at zero; nothing reads it.
257				offsets[i] = Word::from_u64((written - 1) as u64);
258			}
259		}
260	}
261}
262
263/// Little-endian bytes from 64-bit words.
264fn 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
272/// The inverse of [`bytes_from_words`].
273fn 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
279/// In-circuit form of [`recover_public_key`], spending only the hashes a verifier walks.
280///
281/// A chain's tail is `CHAIN_LENGTH - 1 - digit` hashes, and the target sum fixes the total across
282/// all chains at [`NUM_CHAIN_HASHES`] however the digits fall. So rather than give every chain
283/// room for its longest possible tail — `V * (CHAIN_LENGTH - 1)` hashes, two thirds of them
284/// discarded — the tails are concatenated into one list of exactly that total, hinted, and pinned
285/// by constraints. Each entry carries its input digest, its chain, its digit, and the number of
286/// hashes that chain has already done; its output is the hash, not a hinted value, so nothing has
287/// to check it.
288///
289/// # What pins the list
290///
291/// Walking the list, every rule is local, which is what leaves the entries with nothing to look
292/// up:
293///
294/// - `chain` only ever increases, so a chain's entries are one contiguous run.
295/// - Within a run the step advances by one, the digit holds, and each input is the previous output.
296/// - A run opens at step zero, and so does the list.
297///
298/// Then one lookup per chain finds that chain's last entry, and checks it belongs to this chain,
299/// carries this chain's digit, and sits at the last position — `digit + step == CHAIN_LENGTH - 2`,
300/// since the last hash of a chain starts one below the chain's end. Its output is the chain's end.
301/// The offset is hinted; it needs no pinning of its own, because what it lands on is checked.
302///
303/// A run's length is therefore exactly `CHAIN_LENGTH - 1 - digit`, and the list is exactly
304/// [`NUM_CHAIN_HASHES`] long, which the target sum makes the sum of those lengths. **No chain that
305/// owes hashes can be missing a run**: if one were, the entries would not add up.
306///
307/// What a run starts *from* is left free. The tip is a hint, and verification only needs some
308/// preimage that walks a chain of the right length onto the committed public key — a prover with
309/// nothing to reveal would have to invert the hash to find one. A chain whose digit is
310/// `CHAIN_LENGTH - 1` owes no hashes at all, and its end is simply the value the signature
311/// revealed.
312pub 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	// Where a chain's last entry can sit. Chains before this one contribute at most
335	// `CHAIN_LENGTH - 1` entries each and so do the chains after, which pins the last entry of
336	// chain `i` into a window of the list — and only that window has to be muxed over.
337	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	// The hashes. Every entry's input is hinted rather than carried from the entry before it, so
345	// they are independent and pair two to a core. A chain link's tweak names the position the
346	// hash starts from, which is the chain's digit plus the hashes it has already done.
347	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	// Walking the list: every rule here is local, which is what leaves the entries with nothing
381	// to look up.
382	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			// The list opens a chain, so it opens at its first hash.
388			b.assert_eq("first_step_is_zero", step, zero);
389			continue;
390		};
391
392		// An entry either continues the one before it or opens a new chain. The chain field is
393		// what says which, and it only ever increases, so a chain's entries stay contiguous.
394		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		// Continuing: one more hash of the same chain, on the value the last one produced.
399		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		// Opening: a fresh chain starts over at its first hash. What it starts *from* is left
416		// free — the tip is a hint, and verification only needs some preimage that walks a chain
417		// of the right length onto the committed public key.
418		b.assert_eq("opens_at_zero", b.select(opens, step, zero), zero);
419	}
420
421	// One lookup per chain, into the window its last entry has to sit in. What the offset lands
422	// on is checked, so the offset itself needs no pinning.
423	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		// A chain at the last digit owes no hashes, so it has no entry to find: its end is the
446		// value the signature already revealed, and nothing about the lookup is asserted.
447		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		// The last hash of a chain starts one below the end of the chain, so its position —
454		// the digit plus the hashes done before it — is `CHAIN_LENGTH - 2`.
455		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
467/// In-circuit form of [`wots_public_key_hash`].
468pub 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	/// A signature at `epoch`: random chain preimages walked to the digits the message encodes to.
488	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			// A signature's chain tip is the secret preimage walked as far as its digit; the
508			// verifier walks the rest.
509			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		// `circuit_recovers_the_public_key` only exercises the empty-chain path if some chain is
540		// actually empty. With the target sum putting the mean digit at 4.64 that is the common
541		// case, but it is worth failing loudly if a fixture ever stops covering it.
542		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		// The one chain the verifier never advances.
559		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	/// Builds the encode-and-walk circuit, populates it from `sig`, and returns the result of
565	/// checking the constraint system.
566	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, &param_w, epoch_w, &message_w, &randomness_w);
577		let ends = circuit_recover_public_key(&b, &param_w, epoch_w, &tips_w, &digits);
578		let leaf = circuit_wots_public_key_hash(&b, &param_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(&param_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, &param_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(&param_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		// Any randomness the grinder did not settle on fails one of the two conditions, so the
648		// encoding constraints have no satisfying witness.
649		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		// Every tweak carries the epoch, so the chains and the encoding both move with it.
674		assert!(run(&sig, epoch + 1).is_err(), "an epoch it was not signed at must not verify");
675	}
676}