1use binius_frontend::{CircuitBuilder, Wire};
17
18use super::{DIGEST_LEN, DIGEST_WIRES, Digest, PUBLIC_PARAM_LEN, PUBLIC_PARAM_WIRES, PublicParam};
19use crate::{
20 blake3::{KEY_BYTES, blake3_keyed_fixed, blake3_keyed_fixed_2x},
21 fixed_byte_vec::ByteVec,
22 util::{clear_high_bits, split_u32_words},
23};
24
25pub const TWEAK_TYPE_CHAIN: u8 = 0;
27pub const TWEAK_TYPE_WOTS_PK: u8 = 1;
29pub const TWEAK_TYPE_MERKLE: u8 = 2;
31pub const TWEAK_TYPE_ENCODING: u8 = 3;
33
34pub const TWEAK_LEN: usize = 16;
36
37const TWEAK_WIRES: usize = TWEAK_LEN / 8;
39
40pub type Tweak = [u8; TWEAK_LEN];
42
43const _: () = assert!(PUBLIC_PARAM_LEN + TWEAK_LEN == KEY_BYTES);
46
47pub fn make_tweak(tweak_type: u8, sub_position: u32, index: u32) -> Tweak {
52 let mut tweak = [0u8; TWEAK_LEN];
53 tweak[0] = tweak_type;
54 tweak[1..5].copy_from_slice(&sub_position.to_le_bytes());
55 tweak[5..9].copy_from_slice(&index.to_le_bytes());
56 tweak
57}
58
59pub fn make_key(
61 public_param: &PublicParam,
62 tweak_type: u8,
63 sub_position: u32,
64 index: u32,
65) -> [u8; KEY_BYTES] {
66 let mut key = [0u8; KEY_BYTES];
67 key[..PUBLIC_PARAM_LEN].copy_from_slice(public_param);
68 key[PUBLIC_PARAM_LEN..].copy_from_slice(&make_tweak(tweak_type, sub_position, index));
69 key
70}
71
72pub fn tweak_hash(
77 public_param: &PublicParam,
78 tweak_type: u8,
79 sub_position: u32,
80 index: u32,
81 payload: &[u8],
82) -> Digest {
83 let key = make_key(public_param, tweak_type, sub_position, index);
84 blake3::keyed_hash(&key, payload).as_bytes()[..DIGEST_LEN]
85 .try_into()
86 .expect("the slice is DIGEST_LEN bytes")
87}
88
89pub fn circuit_tweak_hash(
111 builder: &CircuitBuilder,
112 public_param: &[Wire; PUBLIC_PARAM_WIRES],
113 tweak_type: u8,
114 sub_position: Wire,
115 index: Wire,
116 payload: &[Wire],
117) -> [Wire; DIGEST_WIRES] {
118 let key = circuit_key(builder, public_param, tweak_type, sub_position, index);
119 let len_bytes = payload.len() * 8;
120 let message = split_u32_words(builder, payload, len_bytes / 4);
121 truncate(builder, &blake3_keyed_fixed(builder, &message, len_bytes, &key))
122}
123
124fn circuit_key(
126 builder: &CircuitBuilder,
127 public_param: &[Wire; PUBLIC_PARAM_WIRES],
128 tweak_type: u8,
129 sub_position: Wire,
130 index: Wire,
131) -> ByteVec {
132 let mut wires = Vec::with_capacity(PUBLIC_PARAM_WIRES + TWEAK_WIRES);
133 wires.extend_from_slice(public_param);
134 wires.extend_from_slice(&tweak_wires(builder, tweak_type, sub_position, index));
135 ByteVec::new_const_len(builder, wires, KEY_BYTES)
136}
137
138pub fn circuit_tweak_hash_2x(
155 builder: &CircuitBuilder,
156 public_param: &[Wire; PUBLIC_PARAM_WIRES],
157 tweak_type: u8,
158 sub_positions: [Wire; 2],
159 index: Wire,
160 payloads: [&[Wire]; 2],
161) -> [[Wire; DIGEST_WIRES]; 2] {
162 assert_eq!(
163 payloads[0].len(),
164 payloads[1].len(),
165 "both lanes must hash the same number of bytes"
166 );
167
168 let keys = sub_positions
169 .map(|sub_position| circuit_key(builder, public_param, tweak_type, sub_position, index));
170 let len_bytes = payloads[0].len() * 8;
171 let messages = payloads.map(|payload| split_u32_words(builder, payload, len_bytes / 4));
172
173 let digests = blake3_keyed_fixed_2x(
174 builder,
175 [&messages[0], &messages[1]],
176 len_bytes,
177 [&keys[0], &keys[1]],
178 );
179 digests.map(|digest| truncate(builder, &digest))
180}
181
182fn tweak_wires(
184 builder: &CircuitBuilder,
185 tweak_type: u8,
186 sub_position: Wire,
187 index: Wire,
188) -> [Wire; TWEAK_WIRES] {
189 let head = builder.add_constant_64(tweak_type as u64);
193 let sub = builder.shr(builder.shl(sub_position, 32), 24);
194 let word0 = builder.bxor(head, builder.bxor(sub, builder.shl(index, 40)));
195
196 let word1 = builder.shr(builder.shl(index, 32), 56);
201
202 [word0, word1]
203}
204
205fn truncate(builder: &CircuitBuilder, digest: &[Wire; 8]) -> [Wire; DIGEST_WIRES] {
207 std::array::from_fn(|k| {
208 let low = clear_high_bits(builder, digest[2 * k], 32);
209 builder.bxor(low, builder.shl(digest[2 * k + 1], 32))
210 })
211}
212
213#[cfg(test)]
214mod tests {
215 use binius_core::Word;
216 use proptest::prelude::*;
217
218 use super::*;
219 use crate::hash_based_sig::{PUBLIC_PARAM_LEN, V};
220
221 fn check(
223 public_param: &PublicParam,
224 tweak_type: u8,
225 sub_position: u32,
226 index: u32,
227 payload: &[u8],
228 ) {
229 assert_eq!(payload.len() % 8, 0, "payloads are a whole number of 64-bit wires");
230
231 let b = CircuitBuilder::new();
232 let param_w: [Wire; PUBLIC_PARAM_WIRES] = std::array::from_fn(|_| b.add_inout());
233 let index_w = b.add_inout();
234 let payload_w: Vec<Wire> = (0..payload.len() / 8).map(|_| b.add_inout()).collect();
235 let sub_position_w = b.add_constant_64(sub_position as u64);
236 let digest =
237 circuit_tweak_hash(&b, ¶m_w, tweak_type, sub_position_w, index_w, &payload_w);
238 let expected: [Wire; DIGEST_WIRES] = std::array::from_fn(|_| b.add_inout());
239 for k in 0..DIGEST_WIRES {
240 b.assert_eq("digest", digest[k], expected[k]);
241 }
242
243 let circuit = b.build();
244 let mut w = circuit.new_witness_filler();
245 w.pack_bytes_le(¶m_w, public_param);
246 w[index_w] = Word::from_u64(index as u64);
247 w.pack_bytes_le(&payload_w, payload);
248 w.pack_bytes_le(
249 &expected,
250 &tweak_hash(public_param, tweak_type, sub_position, index, payload),
251 );
252
253 circuit.populate_wire_witness(&mut w).unwrap();
254 circuit
255 .constraint_system()
256 .verify(&w.into_value_vec())
257 .unwrap();
258 }
259
260 #[test]
261 fn tweak_separates_everything() {
262 let pp = [7u8; PUBLIC_PARAM_LEN];
263 let x = [1u8; DIGEST_LEN];
264 let base = tweak_hash(&pp, TWEAK_TYPE_CHAIN, 3, 5, &x);
265 assert_ne!(base, tweak_hash(&pp, TWEAK_TYPE_MERKLE, 3, 5, &x));
267 assert_ne!(base, tweak_hash(&pp, TWEAK_TYPE_CHAIN, 4, 5, &x));
268 assert_ne!(base, tweak_hash(&pp, TWEAK_TYPE_CHAIN, 3, 6, &x));
269 assert_ne!(base, tweak_hash(&[8u8; PUBLIC_PARAM_LEN], TWEAK_TYPE_CHAIN, 3, 5, &x));
270 let mut extended = [0u8; 2 * DIGEST_LEN];
272 extended[..DIGEST_LEN].copy_from_slice(&x);
273 assert_ne!(base, tweak_hash(&pp, TWEAK_TYPE_CHAIN, 3, 5, &extended));
274 }
275
276 #[test]
277 fn tweak_layout_matches_the_reference() {
278 let tweak = make_tweak(TWEAK_TYPE_MERKLE, 0x0403_0201, 0x0807_0605);
280 assert_eq!(
281 tweak,
282 [
283 TWEAK_TYPE_MERKLE,
284 1,
285 2,
286 3,
287 4,
288 5,
289 6,
290 7,
291 8,
292 0,
293 0,
294 0,
295 0,
296 0,
297 0,
298 0
299 ]
300 );
301 }
302
303 #[test]
304 fn circuit_matches_reference_at_every_call_site() {
305 let pp = [3u8; PUBLIC_PARAM_LEN];
308 let payload = |len: usize| -> Vec<u8> { (0..len).map(|i| (i * 37 + 11) as u8).collect() };
309 check(&pp, TWEAK_TYPE_CHAIN, 17, 5, &payload(DIGEST_LEN));
310 check(&pp, TWEAK_TYPE_MERKLE, 4, 9, &payload(2 * DIGEST_LEN));
311 check(&pp, TWEAK_TYPE_ENCODING, 0, 5, &payload(64));
312 check(&pp, TWEAK_TYPE_WOTS_PK, 0, 5, &payload(V * DIGEST_LEN));
313 }
314
315 #[test]
316 fn circuit_matches_reference_at_the_index_extremes() {
317 let pp = [5u8; PUBLIC_PARAM_LEN];
319 for index in [0, 1, u32::MAX, u32::MAX - 1, 1 << 24, (1 << 24) - 1] {
320 check(&pp, TWEAK_TYPE_CHAIN, 0, index, &[9u8; DIGEST_LEN]);
321 }
322 }
323
324 #[test]
325 fn two_lane_hash_matches_the_reference_in_both_lanes() {
326 let pp = [6u8; PUBLIC_PARAM_LEN];
328 let payloads = [[1u8; DIGEST_LEN], [2u8; DIGEST_LEN]];
329 let sub_positions = [17u32, 25];
330 let index = 4242u32;
331
332 let b = CircuitBuilder::new();
333 let param_w: [Wire; PUBLIC_PARAM_WIRES] = std::array::from_fn(|_| b.add_inout());
334 let index_w = b.add_inout();
335 let payload_w: [[Wire; DIGEST_WIRES]; 2] =
336 std::array::from_fn(|_| std::array::from_fn(|_| b.add_inout()));
337 let digests = circuit_tweak_hash_2x(
338 &b,
339 ¶m_w,
340 TWEAK_TYPE_CHAIN,
341 sub_positions.map(|s| b.add_constant_64(s as u64)),
342 index_w,
343 [&payload_w[0], &payload_w[1]],
344 );
345 let expected: [[Wire; DIGEST_WIRES]; 2] =
346 std::array::from_fn(|_| std::array::from_fn(|_| b.add_inout()));
347 for lane in 0..2 {
348 b.assert_eq_v(format!("lane[{lane}]"), digests[lane], expected[lane]);
349 }
350
351 let circuit = b.build();
352 let mut w = circuit.new_witness_filler();
353 w.pack_bytes_le(¶m_w, &pp);
354 w[index_w] = Word::from_u64(index as u64);
355 for lane in 0..2 {
356 w.pack_bytes_le(&payload_w[lane], &payloads[lane]);
357 w.pack_bytes_le(
358 &expected[lane],
359 &tweak_hash(&pp, TWEAK_TYPE_CHAIN, sub_positions[lane], index, &payloads[lane]),
360 );
361 }
362
363 circuit.populate_wire_witness(&mut w).unwrap();
364 circuit
365 .constraint_system()
366 .verify(&w.into_value_vec())
367 .unwrap();
368 }
369
370 proptest! {
371 #[test]
372 fn circuit_matches_reference(
373 tweak_type in 0u8..=3,
374 sub_position in 0u32..=u32::MAX,
375 index in 0u32..=u32::MAX,
376 payload_wires in 1usize..=9,
377 seed in any::<u64>(),
378 ) {
379 use rand::{Rng, SeedableRng, rngs::StdRng};
380
381 let mut rng = StdRng::seed_from_u64(seed);
382 let mut public_param = [0u8; PUBLIC_PARAM_LEN];
383 rng.fill_bytes(&mut public_param);
384 let mut payload = vec![0u8; payload_wires * 8];
385 rng.fill_bytes(&mut payload);
386
387 check(&public_param, tweak_type, sub_position, index, &payload);
388 }
389 }
390}