unwound context scheme

This commit is contained in:
integritychain 2024-01-18 12:55:20 -06:00
parent 464a7d4e7c
commit 44667557df
2 changed files with 161 additions and 190 deletions

View file

@ -12,8 +12,6 @@ use crate::types::{
Adrs, ForsPk, ForsSig, HtSig, SlhDsaSig, SlhPrivateKey, SlhPublicKey, WotsPk, WotsSig, XmssSig, Adrs, ForsPk, ForsSig, HtSig, SlhDsaSig, SlhPrivateKey, SlhPublicKey, WotsPk, WotsSig, XmssSig,
}; };
use crate::types::{FORS_PRF, FORS_ROOTS, FORS_TREE, TREE, WOTS_HASH, WOTS_PK, WOTS_PRF}; use crate::types::{FORS_PRF, FORS_ROOTS, FORS_TREE, TREE, WOTS_HASH, WOTS_PK, WOTS_PRF};
use crate::Context;
/// Algorithm 1: `toInt(X, n)` on page 14. /// Algorithm 1: `toInt(X, n)` on page 14.
/// Convert a byte string to an integer. /// Convert a byte string to an integer.
@ -22,6 +20,7 @@ use crate::Context;
/// Output: Integer value of `X`. /// Output: Integer value of `X`.
pub(crate) fn to_int(x: &[u8], n: usize) -> u64 { pub(crate) fn to_int(x: &[u8], n: usize) -> u64 {
assert_eq!(x.len(), n); assert_eq!(x.len(), n);
println!("byte count {}", x.len());
// 1: total ← 0 // 1: total ← 0
let mut total = 0_u64; let mut total = 0_u64;
@ -150,12 +149,12 @@ pub(crate) fn f<N: ArrayLength>(
/// Input: Input string `X`, start index `i`, number of steps `s`, public seed `PK.seed`, address `ADRS`. <br> /// Input: Input string `X`, start index `i`, number of steps `s`, public seed `PK.seed`, address `ADRS`. <br>
/// Output: Value of `F` iterated `s` times on `X`. /// Output: Value of `F` iterated `s` times on `X`.
pub(crate) fn chain<N: ArrayLength>( pub(crate) fn chain<N: ArrayLength>(
context: &Context, cap_x: GenericArray<u8, N>, i: usize, s: usize, pk_seed: &[u8], adrs: &Adrs, cap_x: GenericArray<u8, N>, i: usize, s: usize, pk_seed: &[u8], adrs: &Adrs,
) -> Option<GenericArray<u8, N>> { ) -> Option<GenericArray<u8, N>> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
// 1: if (i + s) ≥ w then // 1: if (i + s) ≥ w then
if (i + s) >= context.w { if (i + s) >= crate::W as usize {
// //
// 2: return NULL // 2: return NULL
return None; return None;
@ -194,7 +193,7 @@ pub(crate) fn prf<N: ArrayLength>(
pub(crate) fn tlen<LEN: ArrayLength, N: ArrayLength>( pub(crate) fn tlen<LEN: ArrayLength, N: ArrayLength>(
_context: &Context, pk_seed: &[u8], adrs: &Adrs, ml: &GenericArray<GenericArray<u8, N>, LEN>, pk_seed: &[u8], adrs: &Adrs, ml: &GenericArray<GenericArray<u8, N>, LEN>,
) -> GenericArray<u8, N> { ) -> GenericArray<u8, N> {
let mut hasher = Shake256::default(); let mut hasher = Shake256::default();
hasher.update(pk_seed); hasher.update(pk_seed);
@ -217,7 +216,7 @@ pub(crate) fn tlen<LEN: ArrayLength, N: ArrayLength>(
/// Output: WOTS+ public key `pk`. /// Output: WOTS+ public key `pk`.
#[allow(clippy::similar_names)] #[allow(clippy::similar_names)]
pub(crate) fn wots_pkgen<LEN: ArrayLength, N: ArrayLength>( pub(crate) fn wots_pkgen<LEN: ArrayLength, N: ArrayLength>(
context: &Context, sk_seed: &[u8], pk_seed: &[u8], adrs: &Adrs, sk_seed: &[u8], pk_seed: &[u8], adrs: &Adrs,
) -> WotsPk<N> { ) -> WotsPk<N> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
let mut tmp: GenericArray<GenericArray<u8, N>, LEN> = GenericArray::default(); let mut tmp: GenericArray<GenericArray<u8, N>, LEN> = GenericArray::default();
@ -233,7 +232,8 @@ pub(crate) fn wots_pkgen<LEN: ArrayLength, N: ArrayLength>(
// 4: for i from 0 to len 1 do // 4: for i from 0 to len 1 do
//#[allow(clippy::cast_possible_truncation)] // steps 5 and 7 //#[allow(clippy::cast_possible_truncation)] // steps 5 and 7
for i in 0..context.len1 { let len = 2 * N::to_u32() + 3;
for i in 0..len {
// //
// 5: skADRS.setChainAddress(i) // 5: skADRS.setChainAddress(i)
sk_adrs.set_chain_address(i); sk_adrs.set_chain_address(i);
@ -246,7 +246,7 @@ pub(crate) fn wots_pkgen<LEN: ArrayLength, N: ArrayLength>(
// 8: tmp[i] ← chain(sk, 0, w 1, PK.seed, ADRS) ▷ Compute public value for chain i // 8: tmp[i] ← chain(sk, 0, w 1, PK.seed, ADRS) ▷ Compute public value for chain i
tmp[i as usize] = tmp[i as usize] =
chain(context, sk, 0, context.w - 1, pk_seed, &adrs).expect("chain broek!"); chain(sk, 0, crate::W as usize - 1, pk_seed, &adrs).expect("chain broek!");
// 9: end for // 9: end for
} }
@ -261,7 +261,7 @@ pub(crate) fn wots_pkgen<LEN: ArrayLength, N: ArrayLength>(
wotspk_adrs.set_key_pair_address(adrs.get_key_pair_address()); wotspk_adrs.set_key_pair_address(adrs.get_key_pair_address());
// 13: pk ← Tlen (PK.seed, wotspkADRS,tmp) ▷ Compress public key // 13: pk ← Tlen (PK.seed, wotspkADRS,tmp) ▷ Compress public key
let pk = tlen(context, pk_seed, &wotspk_adrs, &tmp); let pk = tlen(pk_seed, &wotspk_adrs, &tmp);
// 14: return pk // 14: return pk
WotsPk(pk) WotsPk(pk)
@ -275,7 +275,7 @@ pub(crate) fn wots_pkgen<LEN: ArrayLength, N: ArrayLength>(
/// Output: WOTS+ signature sig. /// Output: WOTS+ signature sig.
#[allow(clippy::similar_names)] #[allow(clippy::similar_names)]
pub(crate) fn wots_sign<N: ArrayLength, LEN: ArrayLength>( pub(crate) fn wots_sign<N: ArrayLength, LEN: ArrayLength>(
context: &Context, m: &[u8], sk_seed: &[u8], pk_seed: &[u8], adrs: &Adrs, m: &[u8], sk_seed: &[u8], pk_seed: &[u8], adrs: &Adrs,
) -> WotsSig<N, LEN> { ) -> WotsSig<N, LEN> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
let mut sig: WotsSig<N, LEN> = WotsSig::default(); let mut sig: WotsSig<N, LEN> = WotsSig::default();
@ -285,27 +285,28 @@ pub(crate) fn wots_sign<N: ArrayLength, LEN: ArrayLength>(
// 2: // 2:
// 3: msg ← base_2b(M, lgw, len1) ▷ Convert message to base w // 3: msg ← base_2b(M, lgw, len1) ▷ Convert message to base w
let mut msg = base_2b(m, context.lgw, context.len1 as usize); let mut msg = base_2b(m, crate::LGW, 2 * N::to_usize());
// 4: // 4:
// 5: for i from 0 to len1 1 do ▷ Compute checksum // 5: for i from 0 to len1 1 do ▷ Compute checksum
for item in msg.iter().take(context.len1 as usize) { for item in msg.iter().take(2 * N::to_usize()) {
// //
// 6: csum ← csum + w 1 msg[i] // 6: csum ← csum + w 1 msg[i]
csum += context.w as u64 - 1 - item; csum += u64::from(crate::W) - 1 - item;
// 7: end for // 7: end for
} }
// 8: // 8:
// 9: csum ← csum ≪ ((8 ((len2·lgw) mod 8)) mod 8) ▷ For lgw = 4 left shift by 4 // 9: csum ← csum ≪ ((8 ((len2·lgw) mod 8)) mod 8) ▷ For lgw = 4 left shift by 4
csum <<= (8 - ((context.len2 * context.lgw as usize) % 8)) % 8; let len2 = 3_usize; //
csum <<= (8 - ((len2 as u64 * u64::from(crate::LGW)) % 8)) % 8;
// 10: msg ← msg ∥ base_2^b(toByte(csum, ceil(len2·lgw/8)), lgw, len2) ▷ Convert csum to base w // 10: msg ← msg ∥ base_2^b(toByte(csum, ceil(len2·lgw/8)), lgw, len2) ▷ Convert csum to base w
msg.extend(&base_2b( msg.extend(&base_2b(
&to_byte(csum, (context.len2 * context.lgw as usize).div_ceil(8)), &to_byte(csum, (len2 * crate::LGW as usize).div_ceil(8)),
context.lgw, crate::LGW,
context.len2, len2,
)); ));
// 11: // 11:
@ -320,7 +321,8 @@ pub(crate) fn wots_sign<N: ArrayLength, LEN: ArrayLength>(
// 15: for i from 0 to len 1 do // 15: for i from 0 to len 1 do
#[allow(clippy::cast_possible_truncation)] // step 18/19 #[allow(clippy::cast_possible_truncation)] // step 18/19
for (i, item) in msg.iter().enumerate().take(context.len) { let len = 2 * N::to_usize() + 3;
for (i, item) in msg.iter().enumerate().take(len) {
// //
// 16: skADRS.setChainAddress(i) // 16: skADRS.setChainAddress(i)
sk_addrs.set_chain_address(i as u32); sk_addrs.set_chain_address(i as u32);
@ -332,7 +334,7 @@ pub(crate) fn wots_sign<N: ArrayLength, LEN: ArrayLength>(
adrs.set_chain_address(i as u32); adrs.set_chain_address(i as u32);
// 19: sig[i] ← chain(sk, 0, msg[i], PK.seed, ADRS) ▷ Compute signature value for chain i // 19: sig[i] ← chain(sk, 0, msg[i], PK.seed, ADRS) ▷ Compute signature value for chain i
sig.data[i] = chain(context, sk, 0, *item as usize, pk_seed, &adrs).unwrap(); sig.data[i] = chain(sk, 0, *item as usize, pk_seed, &adrs).unwrap();
// 20: end for // 20: end for
} }
@ -348,7 +350,7 @@ pub(crate) fn wots_sign<N: ArrayLength, LEN: ArrayLength>(
/// Input: WOTS+ signature `sig`, message `M`, public seed `PK.seed`, address `ADRS`. <br> /// Input: WOTS+ signature `sig`, message `M`, public seed `PK.seed`, address `ADRS`. <br>
/// Output: WOTS+ public key `pksig` derived from `sig`. /// Output: WOTS+ public key `pksig` derived from `sig`.
pub(crate) fn wots_pk_from_sig<LEN: ArrayLength, N: ArrayLength>( pub(crate) fn wots_pk_from_sig<LEN: ArrayLength, N: ArrayLength>(
context: &Context, sig: &WotsSig<LEN, N>, m: &[u8], pk_seed: &[u8], adrs: &Adrs, sig: &WotsSig<LEN, N>, m: &[u8], pk_seed: &[u8], adrs: &Adrs,
) -> WotsPk<N> { ) -> WotsPk<N> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
let mut tmp: GenericArray<GenericArray<u8, N>, LEN> = GenericArray::default(); let mut tmp: GenericArray<GenericArray<u8, N>, LEN> = GenericArray::default();
@ -358,42 +360,43 @@ pub(crate) fn wots_pk_from_sig<LEN: ArrayLength, N: ArrayLength>(
// 2: // 2:
// 3: msg ← base_2b (M, lgw , len1 ) ▷ Convert message to base w // 3: msg ← base_2b (M, lgw , len1 ) ▷ Convert message to base w
let mut msg = base_2b(m, context.lgw, context.len1 as usize); let len1 = 2 * N::to_usize();
let mut msg = base_2b(m, crate::LGW, len1);
// 4: // 4:
// 5: for i from 0 to len1 1 do ▷ Compute checksum // 5: for i from 0 to len1 1 do ▷ Compute checksum
for item in msg.iter().take(context.len1 as usize) { for item in msg.iter().take(len1) {
// //
// 6: csum ← csum + w 1 msg[i] // 6: csum ← csum + w 1 msg[i]
csum += context.w as u64 - 1 - item; csum += u64::from(crate::W) - 1 - item;
// 7: end for // 7: end for
} }
// 8: // 8:
// 9: csum ← csum ≪ ((8 ((len2·lgw) mod 8)) mod 8) ▷ For lgw = 4 left shift by 4 // 9: csum ← csum ≪ ((8 ((len2·lgw) mod 8)) mod 8) ▷ For lgw = 4 left shift by 4
csum <<= (8 - ((context.len2 * context.lgw as usize) % 8)) % 8; let len2 = 3;
csum <<= (8 - ((len2 * crate::LGW as usize) % 8)) % 8;
// 10: msg ← msg ∥ base_2^b(toByte(csum, ceil(len2·lgw/8)), lgw, len2) ▷ Convert csum to base w // 10: msg ← msg ∥ base_2^b(toByte(csum, ceil(len2·lgw/8)), lgw, len2) ▷ Convert csum to base w
msg.extend(&base_2b( msg.extend(&base_2b(
&to_byte(csum, (context.len2 * context.lgw as usize).div_ceil(8)), &to_byte(csum, (len2 * crate::LGW as usize).div_ceil(8)),
context.lgw, crate::LGW,
context.len2, len2,
)); ));
// 11: for i from 0 to len 1 do // 11: for i from 0 to len 1 do
#[allow(clippy::cast_possible_truncation)] // steps 12 and 13 #[allow(clippy::cast_possible_truncation)] // steps 12 and 13
for i in 0..context.len { for i in 0..LEN::to_usize() {
// //
// 12: ADRS.setChainAddress(i) // 12: ADRS.setChainAddress(i)
adrs.set_chain_address(i as u32); adrs.set_chain_address(i as u32);
// 13: tmp[i] ← chain(sig[i], msg[i], w 1 msg[i], PK.seed, ADRS) // 13: tmp[i] ← chain(sig[i], msg[i], w 1 msg[i], PK.seed, ADRS)
tmp[i] = chain( tmp[i] = chain(
context,
sig.data[i].clone(), sig.data[i].clone(),
usize::try_from(msg[i]).unwrap(), usize::try_from(msg[i]).unwrap(),
context.w - 1 - msg[i] as usize, crate::W as usize - 1 - msg[i] as usize,
pk_seed, pk_seed,
&adrs, &adrs,
) )
@ -412,7 +415,7 @@ pub(crate) fn wots_pk_from_sig<LEN: ArrayLength, N: ArrayLength>(
wotspk_adrs.set_key_pair_address(adrs.get_key_pair_address()); wotspk_adrs.set_key_pair_address(adrs.get_key_pair_address());
// 18: pksig ← Tlen (PK.seed, wotspkADRS, tmp) // 18: pksig ← Tlen (PK.seed, wotspkADRS, tmp)
let pk = tlen(context, pk_seed, &wotspk_adrs, &tmp); let pk = tlen(pk_seed, &wotspk_adrs, &tmp);
// 19: return pksig // 19: return pksig
WotsPk(pk) WotsPk(pk)
@ -441,13 +444,13 @@ pub(crate) fn h<N: ArrayLength>(
/// `address ADRS`. <br> /// `address ADRS`. <br>
/// Output: n-byte root `node`. /// Output: n-byte root `node`.
#[allow(clippy::similar_names)] // sk_seed and pk_seed #[allow(clippy::similar_names)] // sk_seed and pk_seed
pub(crate) fn xmss_node<LEN: ArrayLength, N: ArrayLength>( pub(crate) fn xmss_node<H: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>(
context: &Context, sk_seed: &[u8], i: u32, z: u32, pk_seed: &[u8], adrs: &Adrs, sk_seed: &[u8], i: u32, z: u32, pk_seed: &[u8], adrs: &Adrs,
) -> Result<GenericArray<u8, N>, &'static str> { ) -> Result<GenericArray<u8, N>, &'static str> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
// 1: if z > h or i ≥ 2^{h z} then // 1: if z > h or i ≥ 2^{h z} then
if (z > context.h_prime) | (i >= 2u32.pow(context.h - z)) { if (z > HP::to_u32()) | (i as u64 >= 2u64.pow(H::to_u32() - z)) {
// //
// 2: return NULL // 2: return NULL
return Err("Alg8: fail"); return Err("Alg8: fail");
@ -465,18 +468,16 @@ pub(crate) fn xmss_node<LEN: ArrayLength, N: ArrayLength>(
adrs.set_key_pair_address(i); adrs.set_key_pair_address(i);
// 7: node ← wots_PKgen(SK.seed, PK.seed, ADRS) // 7: node ← wots_PKgen(SK.seed, PK.seed, ADRS)
wots_pkgen::<LEN, N>(context, sk_seed, pk_seed, &adrs) wots_pkgen::<LEN, N>(sk_seed, pk_seed, &adrs).0.clone() // TODO revisit (remove clone?)
.0
.clone() // TODO revisit
// 8: else // 8: else
} else { } else {
// //
// 9: lnode ← xmss_node(SK.seed, 2 * i, z 1, PK.seed, ADRS) // 9: lnode ← xmss_node(SK.seed, 2 * i, z 1, PK.seed, ADRS)
let lnode = xmss_node::<LEN, N>(context, sk_seed, 2 * i, z - 1, pk_seed, &adrs)?; let lnode = xmss_node::<H, HP, LEN, N>(sk_seed, 2 * i, z - 1, pk_seed, &adrs)?;
// 10: rnode ← xmss_node(SK.seed, 2 * i + 1, z 1, PK.seed, ADRS) // 10: rnode ← xmss_node(SK.seed, 2 * i + 1, z 1, PK.seed, ADRS)
let rnode = xmss_node::<LEN, N>(context, sk_seed, 2 * i + 1, z - 1, pk_seed, &adrs)?; let rnode = xmss_node::<H, HP, LEN, N>(sk_seed, 2 * i + 1, z - 1, pk_seed, &adrs)?;
// 11: ADRS.setTypeAndClear(TREE) // 11: ADRS.setTypeAndClear(TREE)
adrs.set_type_and_clear(TREE); adrs.set_type_and_clear(TREE);
@ -504,8 +505,8 @@ pub(crate) fn xmss_node<LEN: ArrayLength, N: ArrayLength>(
/// Input: n-byte message `M`, secret seed `SK.seed`, index `idx`, public seed `PK.seed`, address `ADRS`. <br> /// Input: n-byte message `M`, secret seed `SK.seed`, index `idx`, public seed `PK.seed`, address `ADRS`. <br>
/// Output: XMSS signature SIGXMSS = (sig ∥ AUTH). /// Output: XMSS signature SIGXMSS = (sig ∥ AUTH).
#[allow(clippy::similar_names)] // sk_seed and pk_seed #[allow(clippy::similar_names)] // sk_seed and pk_seed
pub(crate) fn xmss_sign<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>( pub(crate) fn xmss_sign<H: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>(
context: &Context, m: &[u8], sk_seed: &[u8], idx: u32, pk_seed: &[u8], adrs: &Adrs, m: &[u8], sk_seed: &[u8], idx: u32, pk_seed: &[u8], adrs: &Adrs,
) -> Result<XmssSig<HP, LEN, N>, &'static str> { ) -> Result<XmssSig<HP, LEN, N>, &'static str> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
let mut sig_xmss = XmssSig::default(); let mut sig_xmss = XmssSig::default();
@ -517,7 +518,7 @@ pub(crate) fn xmss_sign<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>(
let k = (idx / 2) ^ 1; let k = (idx / 2) ^ 1;
// 3: AUTH[j] ← xmss_node(SK.seed, k, j, PK.seed, ADRS) // 3: AUTH[j] ← xmss_node(SK.seed, k, j, PK.seed, ADRS)
sig_xmss.auth[j as usize] = xmss_node::<LEN, N>(context, sk_seed, k, j, pk_seed, &adrs)?; sig_xmss.auth[j as usize] = xmss_node::<H, HP, LEN, N>(sk_seed, k, j, pk_seed, &adrs)?;
// 4: end for // 4: end for
} }
@ -530,7 +531,7 @@ pub(crate) fn xmss_sign<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>(
adrs.set_key_pair_address(idx); adrs.set_key_pair_address(idx);
// 8: sig ← wots_sign(M, SK.seed, PK.seed, ADRS) // 8: sig ← wots_sign(M, SK.seed, PK.seed, ADRS)
sig_xmss.sig_wots = wots_sign(context, m, sk_seed, pk_seed, &adrs); sig_xmss.sig_wots = wots_sign(m, sk_seed, pk_seed, &adrs);
// 9: SIG_XMSS ← sig ∥ AUTH // 9: SIG_XMSS ← sig ∥ AUTH
// struct constructed above // struct constructed above
@ -547,8 +548,7 @@ pub(crate) fn xmss_sign<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>(
/// address `ADRS`. <br> /// address `ADRS`. <br>
/// Output: n-byte root value `node[0]`. /// Output: n-byte root value `node[0]`.
pub(crate) fn xmss_pk_from_sig<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>( pub(crate) fn xmss_pk_from_sig<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>(
context: &Context, idx: u32, sig_xmss: &XmssSig<HP, LEN, N>, m: &[u8], pk_seed: &[u8], idx: u32, sig_xmss: &XmssSig<HP, LEN, N>, m: &[u8], pk_seed: &[u8], adrs: &Adrs,
adrs: &Adrs,
) -> GenericArray<u8, N> { ) -> GenericArray<u8, N> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
@ -565,9 +565,7 @@ pub(crate) fn xmss_pk_from_sig<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength
let auth = sig_xmss.get_xmss_auth(); let auth = sig_xmss.get_xmss_auth();
// 5: node[0] ← wots_PKFromSig(sig, M, PK.seed, ADRS) // 5: node[0] ← wots_PKFromSig(sig, M, PK.seed, ADRS)
let mut node_0 = wots_pk_from_sig::<LEN, N>(context, sig, m, pk_seed, &adrs) let mut node_0 = wots_pk_from_sig::<LEN, N>(sig, m, pk_seed, &adrs).0.clone();
.0
.clone();
// 6: // 6:
// 7: ADRS.setTypeAndClear(TREE) ▷ Compute root from WOTS+ pk and AUTH // 7: ADRS.setTypeAndClear(TREE) ▷ Compute root from WOTS+ pk and AUTH
@ -584,7 +582,7 @@ pub(crate) fn xmss_pk_from_sig<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength
// 11: if idx/2^k is even then // 11: if idx/2^k is even then
#[allow(clippy::if_not_else)] // Follows the algorithm as written #[allow(clippy::if_not_else)] // Follows the algorithm as written
let node_1 = if (idx & 2_u32.pow(k)) != 0 { let node_1 = if (idx >> k) % 2 == 0 {
// 12: ADRS.setTreeIndex(ADRS.getTreeIndex()/2) // 12: ADRS.setTreeIndex(ADRS.getTreeIndex()/2)
let tmp = adrs.get_tree_index() / 2; let tmp = adrs.get_tree_index() / 2;
adrs.set_tree_index(tmp); adrs.set_tree_index(tmp);
@ -623,8 +621,14 @@ pub(crate) fn xmss_pk_from_sig<HP: ArrayLength, LEN: ArrayLength, N: ArrayLength
/// index `idx_leaf`. <br> /// index `idx_leaf`. <br>
/// Output: HT signature `SIG_HT`. /// Output: HT signature `SIG_HT`.
#[allow(clippy::similar_names)] // sk_seed and pk_seed #[allow(clippy::similar_names)] // sk_seed and pk_seed
pub(crate) fn ht_sign<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>( pub(crate) fn ht_sign<
context: &Context, m: &[u8], sk_seed: &[u8], pk_seed: &[u8], idx_tree: u32, idx_leaf: u32, D: ArrayLength,
H: ArrayLength,
HP: ArrayLength,
LEN: ArrayLength,
N: ArrayLength,
>(
m: &[u8], sk_seed: &[u8], pk_seed: &[u8], idx_tree: u32, idx_leaf: u32,
) -> Result<HtSig<D, HP, LEN, N>, &'static str> { ) -> Result<HtSig<D, HP, LEN, N>, &'static str> {
// //
// 1: ADRS ← toByte(0, 32) // 1: ADRS ← toByte(0, 32)
@ -635,17 +639,17 @@ pub(crate) fn ht_sign<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: Arra
adrs.set_tree_address(idx_tree); adrs.set_tree_address(idx_tree);
// 4: SIG_tmp ← xmss_sign(M, SK.seed, idxleaf, PK.seed, ADRS) // 4: SIG_tmp ← xmss_sign(M, SK.seed, idxleaf, PK.seed, ADRS)
let mut sig_tmp = xmss_sign::<HP, LEN, N>(context, m, sk_seed, idx_leaf, pk_seed, &adrs)?; let mut sig_tmp = xmss_sign::<H, HP, LEN, N>(m, sk_seed, idx_leaf, pk_seed, &adrs)?;
// 5: SIG_HT ← SIG_tmp // 5: SIG_HT ← SIG_tmp
let mut sig_ht = HtSig::default(); let mut sig_ht = HtSig::default();
sig_ht.xmss_sigs[0] = sig_tmp.clone(); sig_ht.xmss_sigs[0] = sig_tmp.clone();
// 6: root ← xmss_PKFromSig(idx_leaf, SIG_tmp, M, PK.seed, ADRS) // 6: root ← xmss_PKFromSig(idx_leaf, SIG_tmp, M, PK.seed, ADRS)
let mut root = xmss_pk_from_sig::<HP, LEN, N>(context, idx_leaf, &sig_tmp, m, pk_seed, &adrs); let mut root = xmss_pk_from_sig::<HP, LEN, N>(idx_leaf, &sig_tmp, m, pk_seed, &adrs);
// 7: for j from 1 to d 1 do // 7: for j from 1 to d 1 do
for j in 1..context.d { for j in 1..D::to_u32() {
// //
// 8: idx_leaf ← idx_tree mod 2^{h} ▷ h least significant bits of idx_tree // 8: idx_leaf ← idx_tree mod 2^{h} ▷ h least significant bits of idx_tree
let idx_leaf = idx_tree % 2u32.pow(HP::to_u32()); let idx_leaf = idx_tree % 2u32.pow(HP::to_u32());
@ -660,17 +664,16 @@ pub(crate) fn ht_sign<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: Arra
adrs.set_tree_address(idx_tree); adrs.set_tree_address(idx_tree);
// 12: SIG_tmp ← xmss_sign(root, SK.seed, idx_leaf, PK.seed, ADRS) // 12: SIG_tmp ← xmss_sign(root, SK.seed, idx_leaf, PK.seed, ADRS)
sig_tmp = xmss_sign::<HP, LEN, N>(context, &root, sk_seed, idx_leaf, pk_seed, &adrs)?; sig_tmp = xmss_sign::<H, HP, LEN, N>(&root, sk_seed, idx_leaf, pk_seed, &adrs)?;
// 13: SIG_HT ← SIG_HT ∥ SIG_tmp // 13: SIG_HT ← SIG_HT ∥ SIG_tmp
sig_ht.xmss_sigs[j as usize] = sig_tmp.clone(); sig_ht.xmss_sigs[j as usize] = sig_tmp.clone();
// 14: if j < d 1 then // 14: if j < d 1 then
if j < (context.d - 1) { if j < (D::to_u32() - 1) {
// //
// 15: root ← xmss_PKFromSig(idx_leaf, SIG_tmp, root, PK.seed, ADRS) // 15: root ← xmss_PKFromSig(idx_leaf, SIG_tmp, root, PK.seed, ADRS)
root = root = xmss_pk_from_sig::<HP, LEN, N>(idx_leaf, &sig_tmp, &root, pk_seed, &adrs);
xmss_pk_from_sig::<HP, LEN, N>(context, idx_leaf, &sig_tmp, &root, pk_seed, &adrs);
// 16: end if // 16: end if
} }
@ -689,8 +692,8 @@ pub(crate) fn ht_sign<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: Arra
/// HT public key `PK.root`. <br> /// HT public key `PK.root`. <br>
/// Output: Boolean. /// Output: Boolean.
pub(crate) fn ht_verify<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>( pub(crate) fn ht_verify<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: ArrayLength>(
context: &Context, m: &[u8], sig_ht: &HtSig<D, HP, LEN, N>, pk_seed: &[u8], idx_tree: u32, m: &[u8], sig_ht: &HtSig<D, HP, LEN, N>, pk_seed: &[u8], idx_tree: u32, idx_leaf: u32,
idx_leaf: u32, pk_root: &GenericArray<u8, N>, pk_root: &GenericArray<u8, N>,
) -> bool { ) -> bool {
// //
// 1: ADRS ← toByte(0, 32) // 1: ADRS ← toByte(0, 32)
@ -704,10 +707,10 @@ pub(crate) fn ht_verify<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: Ar
let sig_tmp = sig_ht.xmss_sigs[0].clone(); let sig_tmp = sig_ht.xmss_sigs[0].clone();
// 5: node ← xmss_PKFromSig(idx_leaf, SIG_tmp, M, PK.seed, ADRS) // 5: node ← xmss_PKFromSig(idx_leaf, SIG_tmp, M, PK.seed, ADRS)
let mut node = xmss_pk_from_sig(context, idx_leaf, &sig_tmp, m, pk_seed, &adrs); let mut node = xmss_pk_from_sig(idx_leaf, &sig_tmp, m, pk_seed, &adrs);
// 6: for j from 1 to d 1 do // 6: for j from 1 to d 1 do
for j in 1..(context.d) { for j in 1..D::to_u32() {
// //
// 7: idx_leaf ← idx_tree mod 2^{h} ▷ h least significant bits of idx_tree // 7: idx_leaf ← idx_tree mod 2^{h} ▷ h least significant bits of idx_tree
let idx_leaf = idx_tree % 2u32.pow(HP::to_u32()); let idx_leaf = idx_tree % 2u32.pow(HP::to_u32());
@ -725,7 +728,7 @@ pub(crate) fn ht_verify<D: ArrayLength, HP: ArrayLength, LEN: ArrayLength, N: Ar
let sig_tmp = sig_ht.xmss_sigs[j as usize].clone(); let sig_tmp = sig_ht.xmss_sigs[j as usize].clone();
// 12: node ← xmss_PKFromSig(idx_leaf, SIG_tmp, node, PK.seed, ADRS) // 12: node ← xmss_PKFromSig(idx_leaf, SIG_tmp, node, PK.seed, ADRS)
node = xmss_pk_from_sig(context, idx_leaf, &sig_tmp, &node, pk_seed, &adrs); node = xmss_pk_from_sig(idx_leaf, &sig_tmp, &node, pk_seed, &adrs);
// 13: end for // 13: end for
} }
@ -772,13 +775,13 @@ pub(crate) fn fors_sk_gen<N: ArrayLength>(
/// address `ADRS`. <br> /// address `ADRS`. <br>
/// Output: n-byte root node. /// Output: n-byte root node.
#[allow(clippy::similar_names)] // sk_seed and pk_seed #[allow(clippy::similar_names)] // sk_seed and pk_seed
pub(crate) fn fors_node<N: ArrayLength>( pub(crate) fn fors_node<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
context: &Context, sk_seed: &[u8], i: u32, z: u32, pk_seed: &[u8], adrs: &Adrs, sk_seed: &[u8], i: u32, z: u32, pk_seed: &[u8], adrs: &Adrs,
) -> Result<GenericArray<u8, N>, &'static str> { ) -> Result<GenericArray<u8, N>, &'static str> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
// 1: if z > a or i ≥ k · 2(az) then // 1: if z > a or i ≥ k · 2(az) then
if (z > context.a) | (i > context.k * 2 * (context.a - z)) { if (z > A::to_u32()) | (i > K::to_u32() * 2 * (A::to_u32() - z)) {
// //
// 2: return NULL // 2: return NULL
return Err("Alg14 fails"); return Err("Alg14 fails");
@ -804,10 +807,10 @@ pub(crate) fn fors_node<N: ArrayLength>(
// 9: else // 9: else
} else { } else {
// 10: lnode ← fors_node(SK.seed, 2i, z 1, PK.seed, ADRS) // 10: lnode ← fors_node(SK.seed, 2i, z 1, PK.seed, ADRS)
let lnode = fors_node::<N>(context, sk_seed, 2 * i, z - 1, pk_seed, &adrs)?; let lnode = fors_node::<A, K, N>(sk_seed, 2 * i, z - 1, pk_seed, &adrs)?;
// 11: rnode ← fors_node(SK.seed, 2i + 1, z 1, PK.seed, ADRS) // 11: rnode ← fors_node(SK.seed, 2i + 1, z 1, PK.seed, ADRS)
let rnode = fors_node::<N>(context, sk_seed, 2 * i + 1, z - 1, pk_seed, &adrs)?; let rnode = fors_node::<A, K, N>(sk_seed, 2 * i + 1, z - 1, pk_seed, &adrs)?;
// 12: ADRS.setTreeHeight(z) // 12: ADRS.setTreeHeight(z)
adrs.set_tree_height(z); adrs.set_tree_height(z);
@ -832,36 +835,36 @@ pub(crate) fn fors_node<N: ArrayLength>(
/// Output: FORS signature `SIG_FORS`. /// Output: FORS signature `SIG_FORS`.
#[allow(clippy::similar_names)] // sk_seed and pk_seed #[allow(clippy::similar_names)] // sk_seed and pk_seed
pub(crate) fn fors_sign<A: ArrayLength, K: ArrayLength, N: ArrayLength>( pub(crate) fn fors_sign<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
context: &Context, md: &[u8], sk_seed: &[u8], adrs: &Adrs, pk_seed: &[u8], md: &[u8], sk_seed: &[u8], adrs: &Adrs, pk_seed: &[u8],
) -> Result<ForsSig<A, K, N>, &'static str> { ) -> Result<ForsSig<A, K, N>, &'static str> {
// 1: SIG_FORS = NULL ▷ Initialize SIG_FORS as a zero-length byte string // 1: SIG_FORS = NULL ▷ Initialize SIG_FORS as a zero-length byte string
let mut sig_fors = ForsSig::default(); let mut sig_fors = ForsSig::default();
// 2: indices ← base_2^b(md, a, k) // 2: indices ← base_2^b(md, a, k)
let indices = base_2b(md, context.a, context.k as usize); let indices = base_2b(md, A::to_u32(), K::to_usize());
// 3: for i from 0 to k 1 do ▷ Compute signature elements // 3: for i from 0 to k 1 do ▷ Compute signature elements
#[allow(clippy::cast_possible_truncation)] #[allow(clippy::cast_possible_truncation)]
for i in 0..context.k { for i in 0..K::to_u32() {
// //
// 4: SIG_FORS ← SIG_FORS ∥ fors_SKgen(SK.seed, PK.seed, ADRS, i · 2a + indices[i]) // 4: SIG_FORS ← SIG_FORS ∥ fors_SKgen(SK.seed, PK.seed, ADRS, i · 2a + indices[i])
sig_fors.private_key_value[i as usize] = fors_sk_gen::<N>( sig_fors.private_key_value[i as usize] = fors_sk_gen::<N>(
sk_seed, sk_seed,
pk_seed, pk_seed,
adrs, adrs,
i * 2 * context.a + indices[i as usize] as u32, i * 2 * A::to_u32() + indices[i as usize] as u32,
); );
// 5: // 5:
// 6: for j from 0 to a 1 do ▷ Compute auth path // 6: for j from 0 to a 1 do ▷ Compute auth path
for j in 0..context.a { for j in 0..A::to_u32() {
// //
// 7: s ← indices[i]/2^j xor 1 // 7: s ← indices[i]/2^j xor 1
let s = (indices[i as usize] >> j) ^ 1; let s = (indices[i as usize] >> j) ^ 1;
// 8: AUTH[j] ← fors_node(SK.seed, i · 2^{aj} + s, j, PK.seed, ADRS) // 8: AUTH[j] ← fors_node(SK.seed, i · 2^{aj} + s, j, PK.seed, ADRS)
sig_fors.auth[j as usize].tree[i as usize] = // TODO: check order of j and i sig_fors.auth[j as usize].tree[i as usize] = // TODO: check order of j and i
fors_node::<N>(context, sk_seed, i * 2u32.pow(context.a - j) + s as u32, j, pk_seed, adrs)?; fors_node::<A, K, N>(sk_seed, i * 2u32.pow(A::to_u32() - j) + s as u32, j, pk_seed, adrs)?;
// 9: end for // 9: end for
} }
@ -883,17 +886,17 @@ pub(crate) fn fors_sign<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
/// Input: FORS signature `SIG_FORS`, message digest `md`, public seed `PK.seed`, address `ADRS`. <br> /// Input: FORS signature `SIG_FORS`, message digest `md`, public seed `PK.seed`, address `ADRS`. <br>
/// Output: FORS public key. /// Output: FORS public key.
pub(crate) fn fors_pk_from_sig<A: ArrayLength, K: ArrayLength, N: ArrayLength>( pub(crate) fn fors_pk_from_sig<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
context: &Context, sig_fors: &ForsSig<A, K, N>, md: &[u8], pk_seed: &[u8], adrs: &Adrs, sig_fors: &ForsSig<A, K, N>, md: &[u8], pk_seed: &[u8], adrs: &Adrs,
) -> ForsPk<N> { ) -> ForsPk<N> {
let mut adrs = adrs.clone(); let mut adrs = adrs.clone();
// 1: indices ← base_2^b(md, a, k) // 1: indices ← base_2^b(md, a, k)
let indices = base_2b(md, context.a, context.k as usize); let indices = base_2b(md, A::to_u32(), K::to_usize());
// 2: for i from 0 to k 1 do // 2: for i from 0 to k 1 do
let mut root: GenericArray<GenericArray<u8, N>, K> = GenericArray::default(); let mut root: GenericArray<GenericArray<u8, N>, K> = GenericArray::default();
#[allow(clippy::cast_possible_truncation)] // Step 5 #[allow(clippy::cast_possible_truncation)] // Step 5
for i in 0..context.k { for i in 0..K::to_u32() {
// //
// 3: sk ← SIG_FORS.getSK(i) ▷ SIG_FORS [i · (a + 1) · n : (i · (a + 1) + 1) · n] // 3: sk ← SIG_FORS.getSK(i) ▷ SIG_FORS [i · (a + 1) · n : (i · (a + 1) + 1) · n]
let sk = sig_fors.private_key_value[i as usize].clone(); let sk = sig_fors.private_key_value[i as usize].clone();
@ -902,7 +905,7 @@ pub(crate) fn fors_pk_from_sig<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
adrs.set_tree_height(0); adrs.set_tree_height(0);
// 5: ADRS.setTreeIndex(i · 2^a + indices[i]) // 5: ADRS.setTreeIndex(i · 2^a + indices[i])
adrs.set_tree_index(i * 2u32.pow(context.a) + indices[i as usize] as u32); adrs.set_tree_index(i * 2u32.pow(A::to_u32()) + indices[i as usize] as u32);
// 6: node[0] ← F(PK.seed, ADRS, sk) // 6: node[0] ← F(PK.seed, ADRS, sk)
let mut node_0 = f(pk_seed, &adrs, &sk); let mut node_0 = f(pk_seed, &adrs, &sk);
@ -912,7 +915,7 @@ pub(crate) fn fors_pk_from_sig<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
let auth = sig_fors.auth[i as usize].clone(); let auth = sig_fors.auth[i as usize].clone();
// 9: for j from 0 to a 1 do ▷ Compute root from leaf and AUTH // 9: for j from 0 to a 1 do ▷ Compute root from leaf and AUTH
for j in 0..context.a { for j in 0..A::to_u32() {
// //
// 10: ADRS.setTreeHeight(j + 1) // 10: ADRS.setTreeHeight(j + 1)
adrs.set_tree_height(j + 1); adrs.set_tree_height(j + 1);
@ -962,7 +965,7 @@ pub(crate) fn fors_pk_from_sig<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
fors_pk_adrs.set_key_pair_address(adrs.get_key_pair_address()); fors_pk_adrs.set_key_pair_address(adrs.get_key_pair_address());
// 25: pk ← Tk(PK.seed, forspkADRS, root) // 25: pk ← Tk(PK.seed, forspkADRS, root)
let pk = tlen(context, pk_seed, &fors_pk_adrs, &root); let pk = tlen(pk_seed, &fors_pk_adrs, &root);
// 26: return pk; // 26: return pk;
ForsPk { key: pk } ForsPk { key: pk }
@ -975,8 +978,14 @@ pub(crate) fn fors_pk_from_sig<A: ArrayLength, K: ArrayLength, N: ArrayLength>(
/// Input: (none) <br> /// Input: (none) <br>
/// Output: SLH-DSA key pair `(SK, PK)`. /// Output: SLH-DSA key pair `(SK, PK)`.
#[allow(clippy::similar_names)] // sk_seed and pk_seed #[allow(clippy::similar_names)] // sk_seed and pk_seed
pub(crate) fn slh_keygen_with_rng<LEN: ArrayLength, N: ArrayLength>( pub(crate) fn slh_keygen_with_rng<
context: &Context, rng: &mut impl CryptoRngCore, D: ArrayLength,
H: ArrayLength,
HP: ArrayLength,
LEN: ArrayLength,
N: ArrayLength,
>(
rng: &mut impl CryptoRngCore,
) -> Result<(SlhPrivateKey<N>, SlhPublicKey<N>), &'static str> { ) -> Result<(SlhPrivateKey<N>, SlhPublicKey<N>), &'static str> {
// 1: SK.seed ←$ B^n ▷ Set SK.seed, SK.prf, and PK.seed to random n-byte // 1: SK.seed ←$ B^n ▷ Set SK.seed, SK.prf, and PK.seed to random n-byte
let mut sk_seed = GenericArray::default(); let mut sk_seed = GenericArray::default();
@ -998,10 +1007,10 @@ pub(crate) fn slh_keygen_with_rng<LEN: ArrayLength, N: ArrayLength>(
let mut adrs = Adrs::default(); let mut adrs = Adrs::default();
// 6: ADRS.setLayerAddress(d 1) // 6: ADRS.setLayerAddress(d 1)
adrs.set_layer_address(context.d - 1); adrs.set_layer_address(D::to_u32() - 1);
// 7: PK.root ← xmss_node(SK.seed, 0, h, PK.seed, ADRS) // 7: PK.root ← xmss_node(SK.seed, 0, h, PK.seed, ADRS)
let pk_root = xmss_node::<LEN, N>(context, &sk_seed, 0, context.h_prime, &pk_seed, &adrs)?; let pk_root = xmss_node::<H, HP, LEN, N>(&sk_seed, 0, HP::to_u32(), &pk_seed, &adrs)?;
// 8: // 8:
// 9: return ( (SK.seed, SK.prf, PK.seed, PK.root), (PK.seed, PK.root) ) // 9: return ( (SK.seed, SK.prf, PK.seed, PK.root), (PK.seed, PK.root) )
@ -1020,13 +1029,14 @@ pub(crate) fn slh_keygen_with_rng<LEN: ArrayLength, N: ArrayLength>(
pub(crate) fn slh_sign_with_rng< pub(crate) fn slh_sign_with_rng<
A: ArrayLength, A: ArrayLength,
D: ArrayLength, D: ArrayLength,
H: ArrayLength,
HP: ArrayLength, HP: ArrayLength,
K: ArrayLength, K: ArrayLength,
LEN: ArrayLength, LEN: ArrayLength,
M: ArrayLength,
N: ArrayLength, N: ArrayLength,
>( >(
context: &Context, rng: &mut impl CryptoRngCore, m: &[u8], sk: &SlhPrivateKey<N>, rng: &mut impl CryptoRngCore, m: &[u8], sk: &SlhPrivateKey<N>, randomize: bool,
randomize: bool,
) -> Result<SlhDsaSig<A, D, HP, K, LEN, N>, &'static str> { ) -> Result<SlhDsaSig<A, D, HP, K, LEN, N>, &'static str> {
// 1: ADRS ← toByte(0, 32) // 1: ADRS ← toByte(0, 32)
let mut adrs = Adrs::default(); let mut adrs = Adrs::default();
@ -1054,28 +1064,29 @@ pub(crate) fn slh_sign_with_rng<
// 9: // 9:
// 10: digest ← H_msg(R, PK.seed, PK.root, M) ▷ Compute message digest // 10: digest ← H_msg(R, PK.seed, PK.root, M) ▷ Compute message digest
let digest = h::<N>(&r, &sk.pk_seed, &sk.pk_root, m); let digest = h::<M>(&r, &sk.pk_seed, &sk.pk_root, m);
// 11: md ← digest[0 : ceil(k·a/8)] ▷ first ceil(k·a/8) bytes // 11: md ← digest[0 : ceil(k·a/8)] ▷ first ceil(k·a/8) bytes
let index1 = (context.k * context.a).div_ceil(8) as usize; let index1 = (K::to_usize() * A::to_usize()).div_ceil(8);
let md = &digest[0..index1]; let md = &digest[0..index1];
// 12: tmp_idx_tree ← digest[ceil(k·a/8) : ceil(k·a/8) + ceil((h-h/d)/8)] ▷ next ceil((h-h/d)/8) bytes // 12: tmp_idx_tree ← digest[ceil(k·a/8) : ceil(k·a/8) + ceil((h-h/d)/8)] ▷ next ceil((h-h/d)/8) bytes
let index2 = index1 + (context.h - context.h / context.d).div_ceil(8) as usize; let index2 = index1 + (H::to_usize() - H::to_usize() / D::to_usize()).div_ceil(8);
let tmp_idx_tree = &digest[index1..index2]; let tmp_idx_tree = &digest[index1..index2];
// 13: tmp_idx_leaf ← digest[ceil(k·a/8) + ceil((h-h/d)/8) : ceil(k·a/8) + ceil((h-h/d)/8) + ceil(h/8d)] ▷ next ceil(h/8d) bytes // 13: tmp_idx_leaf ← digest[ceil(k·a/8) + ceil((h-h/d)/8) : ceil(k·a/8) + ceil((h-h/d)/8) + ceil(h/8d)] ▷ next ceil(h/8d) bytes
let index3 = index2 + context.h.div_ceil(8 * context.d) as usize; let index3 = index2 + H::to_usize().div_ceil(8 * D::to_usize());
let tmp_idx_leaf = &digest[index2..index3]; let tmp_idx_leaf = &digest[index2..index3];
// 14: // 14:
// 15: idx_tree ← toInt(tmp_idx_tree, ceil((h-h/d)/8)) mod 2^{hh/d} // 15: idx_tree ← toInt(tmp_idx_tree, ceil((h-h/d)/8)) mod 2^{hh/d}
let idx_tree = to_int(tmp_idx_tree, context.h.div_ceil(8 * context.d) as usize) let idx_tree =
% 2u64.pow(context.h - context.h / context.d); to_int(tmp_idx_tree, (H::to_usize() - H::to_usize() / D::to_usize()).div_ceil(8))
% 2u64.pow(H::to_u32() - H::to_u32() / D::to_u32());
// 16: idx_leaf ← toInt(tmp_idx_leaf, ceil(h/8d) mod 2^{h/d} // 16: idx_leaf ← toInt(tmp_idx_leaf, ceil(h/8d) mod 2^{h/d}
let idx_leaf = to_int(tmp_idx_leaf, context.h.div_ceil(8 * context.d) as usize) let idx_leaf = to_int(tmp_idx_leaf, H::to_usize().div_ceil(8 * D::to_usize()))
% 2u64.pow(context.h / context.d); // TODO: indicates size of int!! % 2u64.pow(H::to_u32() / D::to_u32());
// 17: // 17:
// 18: ADRS.setTreeAddress(idx_tree) // 18: ADRS.setTreeAddress(idx_tree)
@ -1089,17 +1100,16 @@ pub(crate) fn slh_sign_with_rng<
// 21: SIG_FORS ← fors_sign(md, SK.seed, PK.seed, ADRS) // 21: SIG_FORS ← fors_sign(md, SK.seed, PK.seed, ADRS)
// 22: SIG ← SIG ∥ SIG_FORS // 22: SIG ← SIG ∥ SIG_FORS
sig.fors_sig = fors_sign(context, md, &sk.sk_seed, &adrs, &sk.pk_seed)?; // TODO: adrs swapped position? sig.fors_sig = fors_sign(md, &sk.sk_seed, &adrs, &sk.pk_seed)?; // TODO: adrs swapped position?
// 23: // 23:
// 24: PK_FORS ← fors_pkFromSig(SIG_FORS , md, PK.seed, ADRS) ▷ Get FORS key // 24: PK_FORS ← fors_pkFromSig(SIG_FORS , md, PK.seed, ADRS) ▷ Get FORS key
let pk_fors = fors_pk_from_sig::<A, K, N>(context, &sig.fors_sig, md, &sk.pk_seed, &adrs); let pk_fors = fors_pk_from_sig::<A, K, N>(&sig.fors_sig, md, &sk.pk_seed, &adrs);
// 25: // 25:
// 26: SIG_HT ← ht_sign(PK_FORS , SK.seed, PK.seed, idx_tree, idx_leaf) // 26: SIG_HT ← ht_sign(PK_FORS , SK.seed, PK.seed, idx_tree, idx_leaf)
// 27: SIG ← SIG ∥ SIG_HT // 27: SIG ← SIG ∥ SIG_HT
sig.ht_sig = ht_sign::<D, HP, LEN, N>( sig.ht_sig = ht_sign::<D, H, HP, LEN, N>(
context,
&pk_fors.key, &pk_fors.key,
&sk.sk_seed, &sk.sk_seed,
&sk.pk_seed, &sk.pk_seed,
@ -1120,17 +1130,19 @@ pub(crate) fn slh_sign_with_rng<
pub(crate) fn slh_verify< pub(crate) fn slh_verify<
A: ArrayLength, A: ArrayLength,
D: ArrayLength, D: ArrayLength,
H: ArrayLength,
HP: ArrayLength, HP: ArrayLength,
K: ArrayLength, K: ArrayLength,
LEN: ArrayLength, LEN: ArrayLength,
M: ArrayLength,
N: ArrayLength, N: ArrayLength,
>( >(
context: &Context, m: &[u8], sig: &SlhDsaSig<A, D, HP, K, LEN, N>, pk: &SlhPublicKey<N>, m: &[u8], sig: &SlhDsaSig<A, D, HP, K, LEN, N>, pk: &SlhPublicKey<N>,
) -> bool { ) -> bool {
// 1: if |SIG| != (1 + k(1 + a) + h + d · len) · n then // 1: if |SIG| != (1 + k(1 + a) + h + d · len) · n then
// 2: return false // 2: return false
// 3: end if // 3: end if
// TODO: THIS FUNCTION PROBABLY WANTS A BYTE ARRAY, THEN DESERIALIZE // TODO: THIS FUNCTION PROBABLY WANTS A BYTE ARRAY SIGNATURE, THEN DESERIALIZE (??)
// 4: ADRS ← toByte(0, 32) // 4: ADRS ← toByte(0, 32)
let mut adrs = Adrs::default(); let mut adrs = Adrs::default();
@ -1146,29 +1158,31 @@ pub(crate) fn slh_verify<
// 8: // 8:
// 9: digest ← Hmsg(R, PK.seed, PK.root, M) ▷ Compute message digest // 9: digest ← Hmsg(R, PK.seed, PK.root, M) ▷ Compute message digest
let digest = h::<N>(r, &pk.pk_seed, &pk.pk_root, m); let digest = h::<M>(r, &pk.pk_seed, &pk.pk_root, m);
// 10: md ← digest[0 : ceil(k·a/8)] ▷ first ceil(k·a/8) bytes // 10: md ← digest[0 : ceil(k·a/8)] ▷ first ceil(k·a/8) bytes
let index1 = (context.k * context.a).div_ceil(8) as usize; let index1 = (K::to_usize() * A::to_usize()).div_ceil(8);
let md = &digest[0..index1]; let md = &digest[0..index1];
// 11: tmp_idx_tree ← digest[ceil(k·a/8) : ceil(k·a/8) + ceil((h - h/d)/8)] ▷ next ceil((h - h/d)/8) bytes // 11: tmp_idx_tree ← digest[ceil(k·a/8) : ceil(k·a/8) + ceil((h - h/d)/8)] ▷ next ceil((h - h/d)/8) bytes
let index2 = index1 + (context.h - context.h / context.d).div_ceil(8) as usize; let index2 = index1 + (H::to_usize() - H::to_usize() / D::to_usize()).div_ceil(8);
let tmp_idx_tree = &digest[index1..index2]; let tmp_idx_tree = &digest[index1..index2];
// 12: tmp_idx_leaf ← digest[ceil(k·a/8) + ceil((h - h/d)/8) : ceil(k·a/8) + ceil((h - h/d)/8) + ceil(h/8d)] ▷ next ceil(h/8d) bytes // 12: tmp_idx_leaf ← digest[ceil(k·a/8) + ceil((h - h/d)/8) : ceil(k·a/8) + ceil((h - h/d)/8) + ceil(h/8d)] ▷ next ceil(h/8d) bytes
let index3 = index2 + context.h.div_ceil(8 * context.d) as usize; let index3 = index2 + H::to_usize().div_ceil(8 * D::to_usize());
let tmp_idx_leaf = &digest[index2..index3]; let tmp_idx_leaf = &digest[index2..index3];
// 13: // 13:
// 14: idx_tree ← toInt(tmp_idx_tree, ceil((h - h/d)/8)) mod 2^{hh/d} // 14: idx_tree ← toInt(tmp_idx_tree, ceil((h - h/d)/8)) mod 2^{hh/d}
let idx_tree = to_int(tmp_idx_tree, context.h.div_ceil(8 * context.d) as usize) let idx_tree = to_int(tmp_idx_tree, H::to_usize() - H::to_usize() / D::to_usize()).div_ceil(8)
% 2u64.pow(context.h - context.h / context.d); % 2u64.pow(H::to_u32() - H::to_u32() / D::to_u32());
// 15: idx_leaf ← toInt(tmp_idx_leaf, ceil(h/8d) mod 2^{h/d} // 15: idx_leaf ← toInt(tmp_idx_leaf, ceil(h/8d) mod 2^{h/d}
// 16: idx_leaf ← toInt(tmp_idx_leaf, ceil(h/8d) mod 2^{h/d} // 16: idx_leaf ← toInt(tmp_idx_leaf, ceil(h/8d) mod 2^{h/d}
let idx_leaf = to_int(tmp_idx_leaf, context.h.div_ceil(8 * context.d) as usize) let idx_leaf = to_int(tmp_idx_leaf, H::to_usize().div_ceil(8 * D::to_usize()))
% 2u64.pow(context.h / context.d); // TODO: indicates size of int!! % 2u64.pow(H::to_u32() / D::to_u32());
println!("mod h/d ={}", H::to_u64() as f64 / D::to_u64() as f64);
// 16: // 16:
// 17: ADRS.setTreeAddress(idx_tree) ▷ Compute FORS public key // 17: ADRS.setTreeAddress(idx_tree) ▷ Compute FORS public key
@ -1182,12 +1196,11 @@ pub(crate) fn slh_verify<
// 20: // 20:
// 21: PK_FORS ← fors_pkFromSig(SIG_FORS, md, PK.seed, ADRS) // 21: PK_FORS ← fors_pkFromSig(SIG_FORS, md, PK.seed, ADRS)
let pk_fors = fors_pk_from_sig::<A, K, N>(context, sig_fors, md, &pk.pk_seed, &adrs); let pk_fors = fors_pk_from_sig::<A, K, N>(sig_fors, md, &pk.pk_seed, &adrs);
// 22: // 22:
// 23: return ht_verify(PK_FORS, SIG_HT, PK.seed, idx_tree , idx_leaf, PK.root) // 23: return ht_verify(PK_FORS, SIG_HT, PK.seed, idx_tree , idx_leaf, PK.root)
ht_verify::<D, HP, LEN, N>( ht_verify::<D, HP, LEN, N>(
context,
&pk_fors.key, &pk_fors.key,
sig_ht, sig_ht,
&pk.pk_seed, &pk.pk_seed,

View file

@ -1,82 +1,50 @@
#![no_std] //#![no_std]
#![deny(clippy::pedantic)] #![deny(clippy::pedantic)]
#![deny(warnings)] #![deny(warnings)]
#![deny(missing_docs)] #![deny(missing_docs)]
#![allow(dead_code)] #![allow(dead_code)]
// TODO // TODO
// 1. check 12-byte adrs fields // 1. Get one instance working (or at least not erroring)
// 2. revisit/clean hash functions // 2. check 12-byte adrs fields -- how big is the integer really?
// 3. revisit/clean hash functions
// 4. adrs - store in be or le; how to account for sha2/shake?? (different size)
//! TKTK crate doc //! TKTK crate doc
extern crate alloc; extern crate alloc; // TODO: remove (with vecs)
mod algs; mod algs;
mod traits; mod traits;
mod types; mod types;
/// to be deleted // Per eqns 5.1-4 on page 16, LGW=4, W=16 and LEN2=3 are constant across all parameter sets.
#[must_use] const LGW: u32 = 4;
pub fn add(left: usize, right: usize) -> usize { left + right } const W: u32 = 16;
struct Context {
lgw: u32,
w: usize,
len1: u32,
len2: usize,
len: usize,
h: u32,
h_prime: u32,
d: u32,
a: u32,
k: u32,
}
macro_rules! functionality { macro_rules! functionality {
() => { () => {
use crate::traits::PK; #[cfg(test)]
use crate::types::SlhDsaSig; mod tests {
use crate::Context; use super::*;
use generic_array::typenum::{U16, U5}; use crate::algs::{slh_keygen_with_rng, slh_sign_with_rng, slh_verify};
use zeroize::{Zeroize, ZeroizeOnDrop}; use generic_array::typenum::{Prod, Sum, U2, U3};
// ----- 'EXTERNAL' DATA TYPES ----- use rand_core::OsRng;
const W: usize = 2_usize.pow(LGW); #[test]
const LEN1: usize = (8 * N).div_ceil(LGW as usize); fn it_works1111() {
const LEN2: usize = ((LEN1 * (W - 1)).ilog2() / LGW) as usize + 1; let m = [0u8, 1, 2, 3];
const LEN: usize = LEN1 + LEN2; // TODO: can we push LEN SUM<PROD> calculation downwards? (and remove a generic arg)
let (sk, pk) =
static CONTEXT: Context = Context { slh_keygen_with_rng::<D, H, HP, Sum<Prod<U2, N>, U3>, N>(&mut OsRng).unwrap();
lgw: LGW, let sig = slh_sign_with_rng::<A, D, H, HP, K, Sum<Prod<U2, N>, U3>, M, N>(
w: W, &mut OsRng, &m, &sk, false,
len1: LEN1 as u32, )
len2: LEN2, .unwrap();
len: LEN, let result =
h: H, slh_verify::<A, D, H, HP, K, Sum<Prod<U2, N>, U3>, M, N>(&m, &sig, &pk);
h_prime: H_PRIME, assert_eq!(result, false);
d: D, }
a: A,
k: K,
};
// Dummy placeholder TODO: fix
fn sign() -> SlhDsaSig<U5, U5, U16, U5, U16, U5> {
SlhDsaSig::<U5, U5, U16, U5, U16, U5>::default()
}
/// Correctly sized private key specific to the target security parameter set. <br>
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct PrivateKey {
pub(crate) sk_seed: [u8; N],
sk_prf: [u8; N],
pk_seed: [u8; N],
pk_root: [u8; N],
}
impl PK for PrivateKey {
type Seed = [u8; N];
fn seed(&self) -> [u8; N] { self.sk_seed }
} }
}; };
} }
@ -84,14 +52,15 @@ macro_rules! functionality {
/// TKTK /// TKTK
#[cfg(feature = "slh_dsa_sha2_128s")] #[cfg(feature = "slh_dsa_sha2_128s")]
pub mod slh_dsa_sha2_128s { pub mod slh_dsa_sha2_128s {
const N: usize = 16; use generic_array::typenum::{U12, U14, U16, U30, U63, U7, U9};
const H: u32 = 63;
const D: u32 = 7; type N = U16;
const H_PRIME: u32 = 9; type H = U63;
const A: u32 = 12; type D = U7;
const K: u32 = 14; type HP = U9;
const LGW: u32 = 4; type A = U12;
const M: u32 = 30; type K = U14;
type M = U30;
const PK_LEN: usize = 32; const PK_LEN: usize = 32;
const SIG_LEN: usize = 7856; const SIG_LEN: usize = 7856;
const SK_LEN: usize = 0000; const SK_LEN: usize = 0000;
@ -296,14 +265,3 @@ pub mod slh_dsa_shake_256f {
functionality!(); functionality!();
} }
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn it_works() {
let result = add(2, 2);
assert_eq!(result, 4);
}
}