use alloc::string::{String, ToString};
use alloc::vec::Vec;
use alloc::{format, vec};
use dusk_core::BlsScalar;
use dusk_core::signatures::bls::SecretKey as BlsSecretKey;
use hkdf::Hkdf;
use sha2::{Digest, Sha256};
const SHA256_DIGEST_SIZE: usize = 32;
const HKDF_DIGESTS: usize = 255;
const HKDF_OUTPUT_SIZE: usize = SHA256_DIGEST_SIZE * HKDF_DIGESTS;
#[allow(clippy::similar_names)]
fn hkdf(salt: &[u8], ikm: &[u8], info: &[u8], okm: &mut [u8]) {
let prk = Hkdf::<Sha256>::new(Some(salt), ikm);
prk.expand(info, okm)
.expect("okm size to be a valid length HKDF-Expand");
}
#[allow(clippy::similar_names)]
fn ikm_to_lamport_sk(
ikm: &[u8],
salt: &[u8],
lamport_sk: &mut [[u8; SHA256_DIGEST_SIZE]; HKDF_DIGESTS],
) {
let mut okm = [0u8; HKDF_OUTPUT_SIZE];
hkdf(salt, ikm, b"", &mut okm);
for r in 0..HKDF_DIGESTS {
lamport_sk[r].copy_from_slice(
&okm[r * SHA256_DIGEST_SIZE..(r + 1) * SHA256_DIGEST_SIZE],
);
}
}
fn parent_sk_to_lamport_pk(parent_sk: &BlsScalar, index: u32) -> Vec<u8> {
let salt = index.to_be_bytes();
let ikm = parent_sk.to_be_bytes();
let mut lamport_0 = [[0u8; SHA256_DIGEST_SIZE]; HKDF_DIGESTS];
ikm_to_lamport_sk(ikm.as_slice(), salt.as_slice(), &mut lamport_0);
let not_ikm = ikm.map(|byte| !byte);
let mut lamport_1 = [[0u8; SHA256_DIGEST_SIZE]; HKDF_DIGESTS];
ikm_to_lamport_sk(not_ikm.as_slice(), salt.as_slice(), &mut lamport_1);
let mut lamport_combined = [[0u8; SHA256_DIGEST_SIZE]; HKDF_DIGESTS * 2];
lamport_combined[..HKDF_DIGESTS]
.clone_from_slice(&lamport_0[..HKDF_DIGESTS]);
lamport_combined[HKDF_DIGESTS..HKDF_DIGESTS * 2]
.clone_from_slice(&lamport_1[..HKDF_DIGESTS]);
let mut lamport_pk = [0u8; HKDF_OUTPUT_SIZE * 2];
for i in 0..HKDF_DIGESTS * 2 {
let sha_slice = &Sha256::digest(lamport_combined[i]);
lamport_pk[i * SHA256_DIGEST_SIZE..(i + 1) * SHA256_DIGEST_SIZE]
.clone_from_slice(sha_slice);
}
Sha256::digest(lamport_pk).to_vec()
}
#[allow(clippy::similar_names)]
fn hkdf_mod_r(ikm: &[u8], key_info: &[u8]) -> BlsScalar {
const L: usize = 48;
let ikm_combined = [ikm, &[0u8]].concat();
let key_info_combined = [
key_info,
&[0u8, u8::try_from(L).expect("L should be castable to u8")],
]
.concat();
let mut okm: [u8; L] = [0u8; L];
let mut sk = BlsScalar::zero();
let mut salt = Sha256::digest(b"BLS-SIG-KEYGEN-SALT-");
while sk.is_zero().into() {
hkdf(&salt, ikm_combined.as_ref(), &key_info_combined, &mut okm);
let mut okm_le_64 = [0u8; 64];
okm.reverse();
okm_le_64[..L].copy_from_slice(&okm);
sk = BlsScalar::from_bytes_wide(&okm_le_64);
if sk.is_zero().into() {
salt = Sha256::digest(salt);
}
}
sk
}
#[must_use]
fn derive_child_sk(parent_sk: &BlsScalar, index: u32) -> BlsSecretKey {
let lamport_pk = parent_sk_to_lamport_pk(parent_sk, index);
BlsSecretKey::from(hkdf_mod_r(lamport_pk.as_ref(), b""))
}
pub fn derive_master_sk(seed: &[u8]) -> Result<BlsSecretKey, String> {
if seed.len() < 32 {
return Err(
"seed must be greater than or equal to 32 bytes".to_string()
);
}
Ok(BlsSecretKey::from(hkdf_mod_r(seed, b"")))
}
fn get_path_indexes(path_str: &str) -> Result<Vec<u32>, String> {
let mut path: Vec<&str> = path_str.split('/').collect();
let m = path.remove(0);
if m != "m" {
return Err(format!("First node must be m, got {m}"));
}
let mut ret: Vec<u32> = vec![];
for index in path {
match index.parse::<u32>() {
Ok(v) => ret.push(v),
Err(_) => return Err("could not parse node: {index}".to_string()),
}
}
if ret.is_empty() {
return Err("Path contains no child index".to_string());
}
Ok(ret)
}
pub fn derive_bls_sk(
master_sk: &BlsSecretKey,
path: &str,
) -> Result<BlsSecretKey, String> {
let path_indexes: Vec<u32> = get_path_indexes(path)?;
let mut node_sk = master_sk.clone();
for index in &path_indexes {
node_sk = derive_child_sk(node_sk.as_ref(), *index);
}
Ok(node_sk)
}
#[cfg(test)]
mod tests {
use bip39::{Language, Mnemonic, Seed};
use dusk_bytes::Serializable;
use hex::decode;
use num_bigint::BigUint;
use super::*;
struct TestCase {
seed: &'static str,
master_sk: &'static str,
child_index: &'static str,
child_sk: &'static str,
}
#[test]
fn test_child_derivation() {
let test_cases = vec![
TestCase {
seed: "c55257c360c07c72029aebc1b53c05ed0362ada38ead3e3e9efa3708e53495531f09a6987599d18264c1e1c92f2cf141630c7a3c4ab7c81b2f001698e7463b04",
master_sk: "6083874454709270928345386274498605044986640685124978867557563392430687146096",
child_index: "0",
child_sk: "20397789859736650942317412262472558107875392172444076792671091975210932703118",
},
TestCase {
seed: "0099FF991111002299DD7744EE3355BBDD8844115566CC55663355668888CC00",
master_sk: "27580842291869792442942448775674722299803720648445448686099262467207037398656",
child_index: "4294967295",
child_sk: "29358610794459428860402234341874281240803786294062035874021252734817515685787",
},
TestCase {
seed: "3141592653589793238462643383279502884197169399375105820974944592",
master_sk: "29757020647961307431480504535336562678282505419141012933316116377660817309383",
child_index: "3141592653",
child_sk: "25457201688850691947727629385191704516744796114925897962676248250929345014287",
},
TestCase {
seed: "d4e56740f876aef8c010b86a40d5f56745a118d0906a34e69aec8c0db1cb8fa3",
master_sk: "19022158461524446591288038168518313374041767046816487870552872741050760015818",
child_index: "42",
child_sk: "31372231650479070279774297061823572166496564838472787488249775572789064611981",
},
];
for t in test_cases.iter() {
let seed = decode(t.seed).unwrap();
let master_sk = BlsSecretKey::from_bytes(
&t.master_sk
.parse::<BigUint>()
.unwrap()
.to_bytes_le()
.try_into()
.unwrap(),
)
.unwrap();
let child_index = u32::from_str_radix(t.child_index, 10).unwrap();
let child_sk = BlsSecretKey::from_bytes(
&t.child_sk
.parse::<BigUint>()
.unwrap()
.to_bytes_le()
.try_into()
.unwrap(),
)
.unwrap();
let derived_master_sk =
derive_master_sk(&seed).expect("Master SK derivation failed");
assert_eq!(derived_master_sk, master_sk);
let derived_sk = derive_child_sk(master_sk.as_ref(), child_index);
assert_eq!(derived_sk, child_sk);
}
}
#[test]
fn test_path_derivation() {
let mnemonic = Mnemonic::from_phrase(
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about",
Language::English
).unwrap();
let seed = Seed::new(&mnemonic, "TREZOR");
let seed_bytes = seed.as_bytes();
let test_cases = vec![
(
"m/0",
"20397789859736650942317412262472558107875392172444076792671091975210932703118",
),
(
"m/12381/3600/0/0/0",
"1438960529079439298020003172973761593698584351192884838483126814052706935030",
),
];
for test in test_cases {
let path = test.0;
let child_key = test.1;
let master_sk = derive_master_sk(&seed_bytes)
.expect("Master SK derivation failed");
let derived_key = derive_bls_sk(&master_sk, path).unwrap();
let expected_key = BlsSecretKey::from_bytes(
&(child_key)
.parse::<BigUint>()
.unwrap()
.to_bytes_le()
.try_into()
.unwrap(),
)
.unwrap();
assert_eq!(derived_key, expected_key);
}
}
#[test]
fn test_path_parsing() {
let seed_str = "c55257c360c07c72029aebc1b53c05ed0362ada38ead3e3e9efa3708e53495531f09a6987599d18264c1e1c92f2cf141630c7a3c4ab7c81b2f001698e7463b04";
let seed_vec = decode(seed_str).unwrap();
let seed: &[u8; 64] = seed_vec.as_slice().try_into().unwrap();
let path_test_cases = vec![
("m/12381/3600/0/0/0", true),
("x/12381/3600/0/0/0", false),
("m/qwert/3600/0/0/0", false),
("m/a/3s/1726/0", false),
("m", false),
];
for test_case in path_test_cases {
let path = test_case.0;
let expected_result = test_case.1;
let master_sk =
derive_master_sk(seed).expect("Master SK derivation failed");
let result = derive_bls_sk(&master_sk, path);
assert_eq!(result.is_ok(), expected_result);
}
}
}