1use std::array;
21
22use anyhow::{Result, bail};
23use binius_circuits::{
24 bignum::{BigUint, select as select_biguint},
25 bitcoin::p2pkh_signature::compress_pubkey,
26 ecdsa::scalar_mul::scalar_mul,
27 hmac::hmac_sha512_fixed,
28 multiplexer::multi_wire_multiplex,
29 secp256k1::{Secp256k1, Secp256k1Affine},
30 sha256::sha256_fixed,
31 util::clear_high_bits,
32};
33use binius_core::word::Word;
34use binius_frontend::{CircuitBuilder, Wire, WitnessFiller};
35use bitcoin::{
36 NetworkKind,
37 bip32::{ChildNumber, Xpriv},
38 secp256k1::{PublicKey, Secp256k1 as BtcSecp256k1},
39};
40use clap::Args;
41use sha2::{Digest, Sha256};
42
43use crate::ExampleCircuit;
44
45const HARDENED_BIT: u32 = 0x8000_0000;
47
48pub fn bip32_derive_compressed(
67 b: &mut CircuitBuilder,
68 seed: &[Wire; 8],
69 path: &[Wire],
70 depth: Wire,
71) -> Vec<Wire> {
72 let max_depth = path.len();
73 let curve = Secp256k1::new(b);
74
75 let key_words = [
78 b.add_constant_64(u64::from_be_bytes(*b"Bitcoin ")),
79 b.add_constant_64(u64::from_be_bytes([b's', b'e', b'e', b'd', 0, 0, 0, 0])),
80 ];
81 let master = hmac_sha512_fixed(b, &key_words, seed, 64);
82 let mut k = il_scalar(&master);
83 let mut c = ir_words(&master);
84
85 let mut serp_levels: Vec<Vec<Wire>> = Vec::with_capacity(max_depth + 1);
88
89 for level in 0..max_depth {
90 let point = scalar_mul(b, &curve, &k, Secp256k1Affine::generator(b));
92 serp_levels.push(compress_pubkey(b, &point.x, &point.y));
93
94 let is_hardened = b.shl(path[level], 32);
96
97 let prefix_norm = compressed_prefix(b, &point.y);
100 let prefix = b.select(is_hardened, b.add_constant(Word::ZERO), prefix_norm);
101 let value_field = select_biguint(b, is_hardened, &k, &point.x);
104 let value = [
105 value_field.limbs[3],
106 value_field.limbs[2],
107 value_field.limbs[1],
108 value_field.limbs[0],
109 ];
110
111 let message = assemble_message(b, prefix, &value, path[level]);
112 let i = hmac_sha512_fixed(b, &c, &message, 37);
113
114 let il = il_scalar(&i);
116 k = curve.f_scalar().add(b, &il, &k);
117 c = ir_words(&i);
118 }
119
120 let point = scalar_mul(b, &curve, &k, Secp256k1Affine::generator(b));
122 serp_levels.push(compress_pubkey(b, &point.x, &point.y));
123
124 let refs: Vec<&[Wire]> = serp_levels.iter().map(Vec::as_slice).collect();
125 multi_wire_multiplex(b, &refs, depth)
126}
127
128fn il_scalar(hash: &[Wire; 8]) -> BigUint {
131 BigUint {
132 limbs: vec![hash[3], hash[2], hash[1], hash[0]],
133 }
134}
135
136const fn ir_words(hash: &[Wire; 8]) -> [Wire; 4] {
139 [hash[4], hash[5], hash[6], hash[7]]
140}
141
142fn compressed_prefix(b: &CircuitBuilder, y: &BigUint) -> Wire {
145 let y_is_odd = b.shl(y.limbs[0], 63);
146 let even = b.add_constant_64(0x02u64 << 56);
147 let odd = b.add_constant_64(0x03u64 << 56);
148 b.select(y_is_odd, odd, even)
149}
150
151fn assemble_message(b: &CircuitBuilder, prefix: Wire, value: &[Wire; 4], index: Wire) -> [Wire; 5] {
162 let m0 = b.bxor(prefix, b.shr(value[0], 8));
163 let m1 = b.bxor(b.shl(value[0], 56), b.shr(value[1], 8));
164 let m2 = b.bxor(b.shl(value[1], 56), b.shr(value[2], 8));
165 let m3 = b.bxor(b.shl(value[2], 56), b.shr(value[3], 8));
166 let index_bytes = clear_high_bits(b, index, 32);
167 let m4 = b.bxor(b.shl(value[3], 56), b.shl(index_bytes, 24));
168 [m0, m1, m2, m3, m4]
169}
170
171pub struct Bip32Example {
177 seed: [Wire; 8],
178 path: Vec<Wire>,
179 depth: Wire,
180 expected_hash: [Wire; 8],
181 max_depth: usize,
182}
183
184#[derive(Args, Debug, Clone)]
185pub struct Params {
186 #[arg(long, default_value_t = 5)]
188 pub max_depth: usize,
189}
190
191#[derive(Args, Debug, Clone)]
192pub struct Instance {
193 #[arg(long)]
195 pub seed: Option<String>,
196
197 #[arg(long, value_delimiter = ',', value_parser = parse_child, default_value = "0'")]
200 pub path: Vec<u32>,
201}
202
203impl ExampleCircuit for Bip32Example {
204 type Params = Params;
205 type Instance = Instance;
206
207 fn build(params: Params, builder: &mut CircuitBuilder) -> Result<Self> {
208 let max_depth = params.max_depth;
209 let seed: [Wire; 8] = array::from_fn(|_| builder.add_witness());
210 let path: Vec<Wire> = (0..max_depth).map(|_| builder.add_witness()).collect();
211 let depth = builder.add_witness();
212
213 let pubkey = bip32_derive_compressed(builder, &seed, &path, depth);
214
215 let digest = sha256_fixed(builder, &pubkey, 33);
217 let expected_hash: [Wire; 8] = array::from_fn(|_| builder.add_inout());
218 for (idx, (&computed, &expected)) in digest.iter().zip(&expected_hash).enumerate() {
219 builder.assert_eq(format!("pubkey_hash[{idx}]"), computed, expected);
220 }
221
222 Ok(Self {
223 seed,
224 path,
225 depth,
226 expected_hash,
227 max_depth,
228 })
229 }
230
231 fn populate_witness(&self, instance: Instance, w: &mut WitnessFiller<'_>) -> Result<()> {
232 if instance.path.len() > self.max_depth {
233 bail!("path depth {} exceeds max_depth {}", instance.path.len(), self.max_depth);
234 }
235
236 let seed_bytes = match &instance.seed {
237 Some(hex_str) => {
238 let bytes = hex::decode(hex_str.trim_start_matches("0x"))
239 .map_err(|e| anyhow::anyhow!("invalid seed hex: {e}"))?;
240 let bytes: [u8; 64] = bytes
241 .try_into()
242 .map_err(|_| anyhow::anyhow!("seed must be exactly 64 bytes (512 bits)"))?;
243 bytes
244 }
245 None => array::from_fn(|i| i as u8),
246 };
247
248 for i in 0..8 {
250 let word = u64::from_be_bytes(seed_bytes[8 * i..8 * i + 8].try_into().unwrap());
251 w[self.seed[i]] = Word::from_u64(word);
252 }
253
254 for i in 0..self.max_depth {
256 let idx = instance.path.get(i).copied().unwrap_or(0);
257 w[self.path[i]] = Word::from_u64(idx as u64);
258 }
259 w[self.depth] = Word::from_u64(instance.path.len() as u64);
260
261 let pubkey = derive_compressed_pubkey(&seed_bytes, &instance.path)?;
263 let hash: [u8; 32] = Sha256::digest(pubkey).into();
264 let words = sha256_digest_words(&hash);
265 for i in 0..8 {
266 w[self.expected_hash[i]] = Word::from_u64(words[i]);
267 }
268
269 tracing::info!(
270 "BIP32 compressed pubkey {} -> SHA-256 {} (depth {})",
271 hex::encode(pubkey),
272 hex::encode(hash),
273 instance.path.len()
274 );
275 Ok(())
276 }
277}
278
279fn parse_child(s: &str) -> Result<u32, String> {
282 let (digits, hardened) = s
283 .strip_suffix(['\'', 'h', 'H'])
284 .map_or((s, false), |rest| (rest, true));
285 let idx: u32 = digits
286 .parse()
287 .map_err(|e| format!("invalid child index '{s}': {e}"))?;
288 if idx >= HARDENED_BIT {
289 return Err(format!("child index {idx} out of range (must be < 2^31)"));
290 }
291 Ok(if hardened { idx | HARDENED_BIT } else { idx })
292}
293
294fn derive_compressed_pubkey(seed: &[u8; 64], path: &[u32]) -> Result<[u8; 33]> {
296 let secp = BtcSecp256k1::new();
297 let master = Xpriv::new_master(NetworkKind::Main, seed)
298 .map_err(|e| anyhow::anyhow!("invalid master key: {e}"))?;
299 let children: Vec<ChildNumber> = path
300 .iter()
301 .map(|&idx| child_number(idx))
302 .collect::<Result<_>>()?;
303 let derived = master
304 .derive_priv(&secp, &children)
305 .map_err(|e| anyhow::anyhow!("derivation failed: {e}"))?;
306 let pubkey = PublicKey::from_secret_key(&secp, &derived.private_key);
307 Ok(pubkey.serialize())
308}
309
310fn child_number(idx: u32) -> Result<ChildNumber> {
312 let child = if idx & HARDENED_BIT != 0 {
313 ChildNumber::from_hardened_idx(idx & !HARDENED_BIT)
314 } else {
315 ChildNumber::from_normal_idx(idx)
316 };
317 child.map_err(|e| anyhow::anyhow!("invalid child index {idx}: {e}"))
318}
319
320fn sha256_digest_words(digest: &[u8; 32]) -> [u64; 8] {
323 array::from_fn(|i| u32::from_be_bytes(digest[4 * i..4 * i + 4].try_into().unwrap()) as u64)
324}
325
326#[cfg(test)]
327mod tests {
328 use binius_frontend::CircuitBuilder;
329
330 use super::*;
331
332 const COMPRESSED_PUBKEY_WORDS: usize = 9;
334
335 fn compressed_pubkey_words(compressed: &[u8; 33]) -> [u64; COMPRESSED_PUBKEY_WORDS] {
338 let mut padded = [0u8; 36];
339 padded[..33].copy_from_slice(compressed);
340 array::from_fn(|i| u32::from_be_bytes(padded[4 * i..4 * i + 4].try_into().unwrap()) as u64)
341 }
342
343 fn check_derivation(seed: &[u8; 64], path: &[u32], max_depth: usize) {
346 assert!(path.len() <= max_depth);
347 let builder = CircuitBuilder::new();
348
349 let seed_wires: [Wire; 8] = array::from_fn(|_| builder.add_witness());
350 let path_wires: Vec<Wire> = (0..max_depth).map(|_| builder.add_witness()).collect();
351 let depth_wire = builder.add_witness();
352
353 let mut b = builder;
355 let derived = bip32_derive_compressed(&mut b, &seed_wires, &path_wires, depth_wire);
356 let builder = b;
357
358 let expected_wires: Vec<Wire> = (0..derived.len()).map(|_| builder.add_inout()).collect();
359 for (idx, (&computed, &expected)) in derived.iter().zip(&expected_wires).enumerate() {
360 builder.assert_eq(format!("pubkey[{idx}]"), computed, expected);
361 }
362
363 let circuit = builder.build();
364 let mut w = circuit.new_witness_filler();
365
366 for i in 0..8 {
367 let word = u64::from_be_bytes(seed[8 * i..8 * i + 8].try_into().unwrap());
368 w[seed_wires[i]] = Word::from_u64(word);
369 }
370 for i in 0..max_depth {
371 let idx = path.get(i).copied().unwrap_or(0);
372 w[path_wires[i]] = Word::from_u64(idx as u64);
373 }
374 w[depth_wire] = Word::from_u64(path.len() as u64);
375
376 let expected = derive_compressed_pubkey(seed, path).expect("oracle derivation");
377 let words = compressed_pubkey_words(&expected);
378 for i in 0..COMPRESSED_PUBKEY_WORDS {
379 w[expected_wires[i]] = Word::from_u64(words[i]);
380 }
381
382 circuit
383 .populate_wire_witness(&mut w)
384 .expect("witness population");
385 circuit
386 .constraint_system()
387 .verify(&w.into_value_vec())
388 .expect("constraints satisfied");
389 }
390
391 fn test_seed() -> [u8; 64] {
392 array::from_fn(|i| (i as u8).wrapping_mul(7).wrapping_add(1))
393 }
394
395 #[test]
396 fn master_pubkey_depth_zero() {
397 check_derivation(&test_seed(), &[], 3);
398 }
399
400 #[test]
401 fn single_hardened_step() {
402 check_derivation(&test_seed(), &[HARDENED_BIT], 3);
403 }
404
405 #[test]
406 fn single_normal_step() {
407 check_derivation(&test_seed(), &[7], 3);
408 }
409
410 #[test]
411 fn mixed_path_full_depth() {
412 let path = [HARDENED_BIT, 1, 2 | HARDENED_BIT, 2, 1_000_000_000];
414 check_derivation(&test_seed(), &path, 5);
415 }
416
417 #[test]
418 fn short_path_with_padding() {
419 check_derivation(&test_seed(), &[5 | HARDENED_BIT, 9], 5);
421 }
422
423 #[test]
424 fn hardened_boundary_indices() {
425 check_derivation(&test_seed(), &[HARDENED_BIT - 1, HARDENED_BIT], 2);
427 }
428
429 #[test]
430 fn parse_child_accepts_hardened_and_normal() {
431 assert_eq!(parse_child("0").unwrap(), 0);
432 assert_eq!(parse_child("44'").unwrap(), 44 | HARDENED_BIT);
433 assert_eq!(parse_child("5h").unwrap(), 5 | HARDENED_BIT);
434 assert!(parse_child("2147483648").is_err());
435 }
436
437 #[test]
440 fn example_proves_pubkey_hash() {
441 let mut builder = CircuitBuilder::new();
442 let example = Bip32Example::build(Params { max_depth: 3 }, &mut builder)
443 .expect("build example circuit");
444 let circuit = builder.build();
445
446 let mut w = circuit.new_witness_filler();
447 let instance = Instance {
448 seed: None,
449 path: vec![HARDENED_BIT, 1],
450 };
451 example
452 .populate_witness(instance, &mut w)
453 .expect("populate witness");
454
455 circuit
456 .populate_wire_witness(&mut w)
457 .expect("witness population");
458 circuit
459 .constraint_system()
460 .verify(&w.into_value_vec())
461 .expect("constraints satisfied");
462 }
463}