1use std::iter;
12
13use binius_core::Word;
14use binius_frontend::{CircuitBuilder, Wire};
15use rand::CryptoRng;
16
17use super::{
18 DIGEST_LEN, DIGEST_WIRES, Digest, LOG_LIFETIME, MESSAGE_WIRES, Message, PUBLIC_PARAM_LEN,
19 PUBLIC_PARAM_WIRES, PublicParam, RANDOMNESS_WIRES, Randomness, V,
20 hashing::{TWEAK_TYPE_MERKLE, circuit_tweak_hash, tweak_hash},
21 wots::{
22 circuit_recover_public_key, circuit_wots_encode, circuit_wots_public_key_hash,
23 find_randomness_for_wots_encoding, iterate_hash, recover_public_key, wots_encode,
24 wots_public_key_hash,
25 },
26};
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
30pub struct XmssPublicKey {
31 pub merkle_root: Digest,
32 pub public_param: PublicParam,
33}
34
35#[derive(Debug, Clone, PartialEq, Eq)]
37pub struct XmssSignature {
38 pub randomness: Randomness,
40 pub chain_tips: [Digest; V],
42 pub merkle_path: [Digest; LOG_LIFETIME],
44}
45
46#[derive(Debug, Clone)]
48pub struct XmssSignatureWires {
49 pub randomness: [Wire; RANDOMNESS_WIRES],
50 pub chain_tips: [[Wire; DIGEST_WIRES]; V],
51 pub merkle_path: [[Wire; DIGEST_WIRES]; LOG_LIFETIME],
52}
53
54impl XmssSignatureWires {
55 pub fn new_witness(builder: &CircuitBuilder) -> Self {
57 Self {
58 randomness: std::array::from_fn(|_| builder.add_witness()),
59 chain_tips: std::array::from_fn(|_| std::array::from_fn(|_| builder.add_witness())),
60 merkle_path: std::array::from_fn(|_| std::array::from_fn(|_| builder.add_witness())),
61 }
62 }
63
64 pub fn populate(&self, w: &mut binius_frontend::WitnessFiller<'_>, signature: &XmssSignature) {
66 w.pack_bytes_le(&self.randomness, &signature.randomness);
67 for (wires, tip) in iter::zip(&self.chain_tips, &signature.chain_tips) {
68 w.pack_bytes_le(wires, tip);
69 }
70 for (wires, node) in iter::zip(&self.merkle_path, &signature.merkle_path) {
71 w.pack_bytes_le(wires, node);
72 }
73 }
74}
75
76#[derive(Debug, PartialEq, Eq, Clone, Copy)]
78pub enum XmssVerifyError {
79 InvalidWots,
81 InvalidMerklePath,
83}
84
85pub fn merkle_node(
90 public_param: &PublicParam,
91 level: usize,
92 index: u32,
93 left: &Digest,
94 right: &Digest,
95) -> Digest {
96 let mut data = [0u8; 2 * DIGEST_LEN];
97 data[..DIGEST_LEN].copy_from_slice(left);
98 data[DIGEST_LEN..].copy_from_slice(right);
99 tweak_hash(public_param, TWEAK_TYPE_MERKLE, level as u32, index, &data)
100}
101
102fn climb(
104 public_param: &PublicParam,
105 leaf: &Digest,
106 epoch: u32,
107 merkle_path: &[Digest; LOG_LIFETIME],
108) -> Digest {
109 merkle_path
110 .iter()
111 .enumerate()
112 .fold(*leaf, |current, (level, sibling)| {
113 let is_left = ((epoch >> level) & 1) == 0;
114 let (left, right) = if is_left {
115 (current, *sibling)
116 } else {
117 (*sibling, current)
118 };
119 let parent_index = ((epoch as u64) >> (level + 1)) as u32;
121 merkle_node(public_param, level + 1, parent_index, &left, &right)
122 })
123}
124
125pub fn xmss_verify(
127 public_key: &XmssPublicKey,
128 message: &Message,
129 signature: &XmssSignature,
130 epoch: u32,
131) -> Result<(), XmssVerifyError> {
132 let encoding = wots_encode(message, epoch, &public_key.public_param, &signature.randomness)
133 .ok_or(XmssVerifyError::InvalidWots)?;
134 let chain_ends =
135 recover_public_key(&signature.chain_tips, &encoding, epoch, &public_key.public_param);
136 let leaf = wots_public_key_hash(&public_key.public_param, epoch, &chain_ends);
137 if climb(&public_key.public_param, &leaf, epoch, &signature.merkle_path)
138 == public_key.merkle_root
139 {
140 Ok(())
141 } else {
142 Err(XmssVerifyError::InvalidMerklePath)
143 }
144}
145
146pub fn circuit_xmss_verify(
165 builder: &CircuitBuilder,
166 public_param: &[Wire; PUBLIC_PARAM_WIRES],
167 merkle_root: &[Wire; DIGEST_WIRES],
168 message: &[Wire; MESSAGE_WIRES],
169 epoch: Wire,
170 signature: &XmssSignatureWires,
171) {
172 builder.assert_zero("xmss_epoch_in_range", builder.shr(epoch, LOG_LIFETIME as u32));
175
176 let digits = circuit_wots_encode(builder, public_param, epoch, message, &signature.randomness);
177 let chain_ends =
178 circuit_recover_public_key(builder, public_param, epoch, &signature.chain_tips, &digits);
179 let leaf = circuit_wots_public_key_hash(builder, public_param, epoch, &chain_ends);
180
181 let root = signature
182 .merkle_path
183 .iter()
184 .enumerate()
185 .fold(leaf, |current, (level, sibling)| {
186 let is_right = builder.shl(epoch, (Word::BITS - 1 - level) as u32);
189 let left: [Wire; DIGEST_WIRES] =
190 std::array::from_fn(|k| builder.select(is_right, sibling[k], current[k]));
191 let right: [Wire; DIGEST_WIRES] =
192 std::array::from_fn(|k| builder.select(is_right, current[k], sibling[k]));
193
194 let parent_index = builder.shr(epoch, level as u32 + 1);
195 let payload = [left[0], left[1], right[0], right[1]];
196 circuit_tweak_hash(
197 builder,
198 public_param,
199 TWEAK_TYPE_MERKLE,
200 builder.add_constant_64(level as u64 + 1),
201 parent_index,
202 &payload,
203 )
204 });
205
206 builder.assert_eq_v("xmss_merkle_root", root, *merkle_root);
207}
208
209pub fn generate_signature(
219 rng: &mut impl CryptoRng,
220 message: &Message,
221 epoch: u32,
222) -> (XmssPublicKey, XmssSignature) {
223 let mut public_param = [0u8; PUBLIC_PARAM_LEN];
224 rng.fill_bytes(&mut public_param);
225
226 let (randomness, encoding) =
227 find_randomness_for_wots_encoding(message, epoch, &public_param, rng);
228
229 let chain_tips: [Digest; V] = std::array::from_fn(|i| {
231 let mut pre_image = [0u8; DIGEST_LEN];
232 rng.fill_bytes(&mut pre_image);
233 iterate_hash(&pre_image, encoding[i] as usize, &public_param, epoch, i, 0)
234 });
235
236 let chain_ends = recover_public_key(&chain_tips, &encoding, epoch, &public_param);
237 let leaf = wots_public_key_hash(&public_param, epoch, &chain_ends);
238
239 let merkle_path: [Digest; LOG_LIFETIME] = std::array::from_fn(|_| {
240 let mut node = [0u8; DIGEST_LEN];
241 rng.fill_bytes(&mut node);
242 node
243 });
244 let merkle_root = climb(&public_param, &leaf, epoch, &merkle_path);
245
246 (
247 XmssPublicKey {
248 merkle_root,
249 public_param,
250 },
251 XmssSignature {
252 randomness,
253 chain_tips,
254 merkle_path,
255 },
256 )
257}
258
259#[cfg(test)]
260mod tests {
261 use rand::{Rng, SeedableRng, rngs::StdRng};
262 use rstest::rstest;
263
264 use super::*;
265 use crate::hash_based_sig::MESSAGE_LEN;
266
267 fn run(
269 public_key: &XmssPublicKey,
270 message: &Message,
271 signature: &XmssSignature,
272 epoch: u32,
273 ) -> Result<(), String> {
274 let b = CircuitBuilder::new();
275 let param_w: [Wire; PUBLIC_PARAM_WIRES] = std::array::from_fn(|_| b.add_inout());
276 let root_w: [Wire; DIGEST_WIRES] = std::array::from_fn(|_| b.add_inout());
277 let message_w: [Wire; MESSAGE_WIRES] = std::array::from_fn(|_| b.add_inout());
278 let epoch_w = b.add_inout();
279 let sig_w = XmssSignatureWires::new_witness(&b);
280
281 circuit_xmss_verify(&b, ¶m_w, &root_w, &message_w, epoch_w, &sig_w);
282
283 let circuit = b.build();
284 let mut w = circuit.new_witness_filler();
285 w.pack_bytes_le(¶m_w, &public_key.public_param);
286 w.pack_bytes_le(&root_w, &public_key.merkle_root);
287 w.pack_bytes_le(&message_w, message);
288 w[epoch_w] = Word::from_u64(epoch as u64);
289 sig_w.populate(&mut w, signature);
290
291 circuit
292 .populate_wire_witness(&mut w)
293 .map_err(|e| format!("populate: {e:?}"))?;
294 circuit
295 .constraint_system()
296 .verify(&w.into_value_vec())
297 .map_err(|e| format!("verify: {e:?}"))
298 }
299
300 fn generate(seed: u64, epoch: u32) -> (XmssPublicKey, Message, XmssSignature) {
301 let mut rng = StdRng::seed_from_u64(seed);
302 let mut message = [0u8; MESSAGE_LEN];
303 rng.fill_bytes(&mut message);
304 let (public_key, signature) = generate_signature(&mut rng, &message, epoch);
305 (public_key, message, signature)
306 }
307
308 #[rstest]
309 #[case::first_epoch(0)]
310 #[case::odd_epoch(1)]
311 #[case::interior_epoch(0x1234_5678)]
312 #[case::last_epoch(u32::MAX)]
313 fn a_generated_signature_verifies(#[case] epoch: u32) {
314 let (public_key, message, signature) = generate(1, epoch);
315 xmss_verify(&public_key, &message, &signature, epoch).unwrap();
317 run(&public_key, &message, &signature, epoch).unwrap();
318 }
319
320 #[test]
321 fn a_tampered_path_node_is_rejected() {
322 let (public_key, message, mut signature) = generate(2, 77);
323 signature.merkle_path[0][0] ^= 0xFF;
324 assert_eq!(
325 xmss_verify(&public_key, &message, &signature, 77),
326 Err(XmssVerifyError::InvalidMerklePath)
327 );
328 assert!(run(&public_key, &message, &signature, 77).is_err());
329 }
330
331 #[test]
332 fn a_tampered_root_is_rejected() {
333 let (mut public_key, message, signature) = generate(3, 77);
334 public_key.merkle_root[0] ^= 0xFF;
335 assert!(run(&public_key, &message, &signature, 77).is_err());
336 }
337
338 #[test]
339 fn another_message_is_rejected() {
340 let (public_key, mut message, signature) = generate(4, 77);
341 message[0] ^= 0xFF;
342 assert!(run(&public_key, &message, &signature, 77).is_err());
343 }
344
345 #[test]
346 fn another_epoch_is_rejected() {
347 let (public_key, message, signature) = generate(5, 77);
350 assert!(run(&public_key, &message, &signature, 78).is_err());
351 }
352
353 #[test]
354 fn an_epoch_past_the_lifetime_is_rejected() {
355 let (public_key, message, signature) = generate(6, 5);
358 let b = CircuitBuilder::new();
359 let param_w: [Wire; PUBLIC_PARAM_WIRES] = std::array::from_fn(|_| b.add_inout());
360 let root_w: [Wire; DIGEST_WIRES] = std::array::from_fn(|_| b.add_inout());
361 let message_w: [Wire; MESSAGE_WIRES] = std::array::from_fn(|_| b.add_inout());
362 let epoch_w = b.add_inout();
363 let sig_w = XmssSignatureWires::new_witness(&b);
364 circuit_xmss_verify(&b, ¶m_w, &root_w, &message_w, epoch_w, &sig_w);
365
366 let circuit = b.build();
367 let mut w = circuit.new_witness_filler();
368 w.pack_bytes_le(¶m_w, &public_key.public_param);
369 w.pack_bytes_le(&root_w, &public_key.merkle_root);
370 w.pack_bytes_le(&message_w, &message);
371 w[epoch_w] = Word::from_u64((1u64 << LOG_LIFETIME) + 5);
372 sig_w.populate(&mut w, &signature);
373
374 assert!(
375 circuit.populate_wire_witness(&mut w).is_err(),
376 "an epoch outside the lifetime must not verify"
377 );
378 }
379
380 #[test]
381 fn the_path_climbs_the_side_its_index_says() {
382 let pp = [1u8; PUBLIC_PARAM_LEN];
385 let leaf = [2u8; DIGEST_LEN];
386 let sibling0 = [3u8; DIGEST_LEN];
387 let sibling1 = [4u8; DIGEST_LEN];
388 let parent = merkle_node(&pp, 1, 0, &sibling0, &leaf);
389 let grandparent = merkle_node(&pp, 2, 0, &parent, &sibling1);
390
391 let mut path = [[0u8; DIGEST_LEN]; LOG_LIFETIME];
392 path[0] = sibling0;
393 path[1] = sibling1;
394 let mut current = grandparent;
395 for (level, sibling) in path.iter().enumerate().skip(2) {
396 current = merkle_node(&pp, level + 1, 0, ¤t, sibling);
397 }
398 assert_eq!(climb(&pp, &leaf, 1, &path), current);
399 }
400}