#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum NafSign {
Zero = 0b00,
Pos = 0b01,
Neg = 0b11,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct NafDigit {
pub s_digit: NafSign,
pub h_digit: NafSign,
}
const MAX_NAF_DIGITS: usize = 258;
const PACKED_NAF_BYTES: usize = MAX_NAF_DIGITS.div_ceil(2);
const fn encode_digit(d: NafSign) -> u8 {
d as u8
}
const fn decode_digit(raw: u8) -> NafSign {
match raw & 0b11 {
0b00 => NafSign::Zero,
0b01 => NafSign::Pos,
0b10 => NafSign::Zero,
0b11 => NafSign::Neg,
_ => NafSign::Zero,
}
}
pub struct NafIterator {
packed: [u8; PACKED_NAF_BYTES],
len: usize,
index: usize,
}
impl NafIterator {
fn pack_set(&mut self, idx: usize, s_digit: NafSign, h_digit: NafSign) {
let byte_idx = idx / 2;
let Some(cell) = self.packed.get_mut(byte_idx) else {
return;
};
let nibble = (encode_digit(s_digit) << 2) | encode_digit(h_digit);
if idx & 1 == 0 {
*cell = (*cell & 0xF0) | nibble;
} else {
*cell = (*cell & 0x0F) | (nibble << 4);
}
}
fn pack_get(&self, idx: usize) -> NafDigit {
let byte_idx = idx / 2;
let Some(&cell) = self.packed.get(byte_idx) else {
return NafDigit {
s_digit: NafSign::Zero,
h_digit: NafSign::Zero,
};
};
let nibble = if idx & 1 == 0 { cell & 0x0F } else { cell >> 4 };
NafDigit {
s_digit: decode_digit((nibble >> 2) & 0x03),
h_digit: decode_digit(nibble & 0x03),
}
}
pub fn new<T>(s: T, h: T) -> Self
where
T: const_num_traits::Zero
+ const_num_traits::One
+ const_num_traits::WrappingAdd<Output = T>
+ Clone
+ PartialOrd
+ PartialEq
+ core::ops::ShrAssign<usize>,
for<'a> &'a T: core::ops::BitAnd<Output = T>,
{
let mut iter = NafIterator {
packed: [0u8; PACKED_NAF_BYTES],
len: 0,
index: 0,
};
let mut s_working = s;
let mut h_working = h;
let zero = T::zero();
let one = T::one();
let two = one.clone().wrapping_add(one.clone());
let three = two.wrapping_add(one.clone());
while s_working > zero || h_working > zero {
let s_bit = (&s_working & &one) == one;
let h_bit = (&h_working & &one) == one;
let (s_digit, s_carry) = if s_bit {
if (&s_working & &three) == one {
(NafSign::Pos, false)
} else {
(NafSign::Neg, true)
}
} else {
(NafSign::Zero, false)
};
let (h_digit, h_carry) = if h_bit {
if (&h_working & &three) == one {
(NafSign::Pos, false)
} else {
(NafSign::Neg, true)
}
} else {
(NafSign::Zero, false)
};
if iter.len < MAX_NAF_DIGITS {
iter.pack_set(iter.len, s_digit, h_digit);
iter.len += 1;
}
s_working >>= 1;
h_working >>= 1;
if s_carry {
s_working = s_working.wrapping_add(one.clone());
}
if h_carry {
h_working = h_working.wrapping_add(one.clone());
}
}
iter
}
pub fn digits_msb_first(&self) -> impl Iterator<Item = NafDigit> + '_ {
(0..self.len).rev().map(move |i| self.pack_get(i))
}
}
impl Iterator for NafIterator {
type Item = NafDigit;
fn next(&mut self) -> Option<Self::Item> {
if self.index < self.len {
let digit = self.pack_get(self.index);
self.index += 1;
Some(digit)
} else {
None
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_naf_basic() {
let s = 5u64;
let h = 3u64;
let naf = NafIterator::new(s, h);
let digits: Vec<_> = naf.digits_msb_first().collect();
assert_eq!(
digits,
vec![
NafDigit {
s_digit: NafSign::Pos,
h_digit: NafSign::Pos
},
NafDigit {
s_digit: NafSign::Zero,
h_digit: NafSign::Zero
},
NafDigit {
s_digit: NafSign::Pos,
h_digit: NafSign::Neg
},
]
);
}
#[test]
fn test_naf_encode_decode_roundtrip() {
for d in [NafSign::Neg, NafSign::Zero, NafSign::Pos] {
let encoded = encode_digit(d);
let decoded = decode_digit(encoded);
assert_eq!(d, decoded, "roundtrip failed for {d:?}: encoded={encoded}");
}
}
#[test]
fn test_naf_pack_roundtrip() {
let mut iter = NafIterator {
packed: [0u8; PACKED_NAF_BYTES],
len: 0,
index: 0,
};
let vals = [NafSign::Neg, NafSign::Zero, NafSign::Pos];
let mut idx = 0;
for &s in &vals {
for &h in &vals {
iter.pack_set(idx, s, h);
let got = iter.pack_get(idx);
assert_eq!(got.s_digit, s, "s mismatch at idx {idx}");
assert_eq!(got.h_digit, h, "h mismatch at idx {idx}");
idx += 1;
}
}
}
#[test]
fn test_decode_corrupted_0b10_is_zero() {
assert_eq!(decode_digit(0b10), NafSign::Zero);
}
}