use core::{fmt, mem, ops::Range, str};
use std::borrow::Cow;
use hmac::Hmac;
use sha2::{Digest, Sha256, Sha512};
use zeroize::Zeroize;
use crate::error::Error;
use crate::language::Language;
const BITS_PER_WORD: usize = 11;
const BITS_PER_BYTE: usize = 8;
const ENTROPY_OFFSET: usize = 8;
#[derive(Copy, Clone, Debug, Ord, PartialOrd, Eq, PartialEq, Hash)]
pub enum Count {
Words12 = (128 << ENTROPY_OFFSET) | 4,
Words15 = (160 << ENTROPY_OFFSET) | 5,
Words18 = (192 << ENTROPY_OFFSET) | 6,
Words21 = (224 << ENTROPY_OFFSET) | 7,
Words24 = (256 << ENTROPY_OFFSET) | 8,
}
impl Default for Count {
fn default() -> Self {
Self::Words12
}
}
impl fmt::Display for Count {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"{} words (entropy {} bits + checksum {} bits)",
self.word_count(),
self.entropy_bits(),
self.checksum_bits()
)
}
}
impl From<Count> for usize {
fn from(count: Count) -> Self {
match count {
Count::Words12 => 12,
Count::Words15 => 15,
Count::Words18 => 18,
Count::Words21 => 21,
Count::Words24 => 24,
}
}
}
impl TryFrom<usize> for Count {
type Error = Error;
fn try_from(count: usize) -> Result<Self, Self::Error> {
Self::from_word_count(count)
}
}
impl Count {
const fn from_word_count(count: usize) -> Result<Self, Error> {
Ok(match count {
12 => Self::Words12,
15 => Self::Words15,
18 => Self::Words18,
21 => Self::Words21,
24 => Self::Words24,
others => return Err(Error::BadWordCount(others)),
})
}
const fn from_key_size(size: usize) -> Result<Self, Error> {
Ok(match size {
128 => Self::Words12,
160 => Self::Words15,
192 => Self::Words18,
224 => Self::Words21,
256 => Self::Words24,
others => return Err(Error::BadEntropyBitCount(others)),
})
}
fn from_phrase<P: AsRef<str>>(phrase: P) -> Result<Self, Error> {
let word_count = phrase.as_ref().split_whitespace().count();
Self::from_word_count(word_count)
}
pub const fn word_count(&self) -> usize {
self.total_bits() / BITS_PER_WORD
}
pub const fn total_bits(&self) -> usize {
self.entropy_bits() + self.checksum_bits()
}
pub const fn entropy_bits(&self) -> usize {
(*self as usize) >> ENTROPY_OFFSET
}
pub const fn checksum_bits(&self) -> usize {
(*self as usize) as u8 as usize
}
const fn total(&self) -> Range<usize> {
0..self.total_bits()
}
const fn entropy(&self) -> Range<usize> {
0..self.entropy_bits()
}
const fn checksum(&self) -> Range<usize> {
self.entropy_bits()..self.total_bits()
}
}
#[derive(Clone, Ord, PartialOrd, Eq, PartialEq, Hash)]
pub struct Mnemonic {
lang: Language,
phrase: String,
entropy: Vec<u8>,
}
impl fmt::Debug for Mnemonic {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.phrase())
}
}
impl fmt::Display for Mnemonic {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.phrase())
}
}
impl str::FromStr for Mnemonic {
type Err = Error;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Self::from_phrase(s)
}
}
impl AsRef<str> for Mnemonic {
fn as_ref(&self) -> &str {
self.phrase()
}
}
impl Zeroize for Mnemonic {
fn zeroize(&mut self) {
self.phrase.zeroize();
self.entropy.zeroize();
}
}
impl Drop for Mnemonic {
fn drop(&mut self) {
self.zeroize();
}
}
impl Mnemonic {
#[cfg(feature = "rand")]
pub fn generate(word_count: Count) -> Self {
Self::generate_in(Language::English, word_count)
}
#[cfg(feature = "rand")]
pub fn generate_in(lang: Language, word_count: Count) -> Self {
use rand::RngCore;
const MAX_ENTROPY_BITS: usize = Count::Words24.entropy_bits();
let mut rng = rand::thread_rng();
let mut entropy = [0u8; MAX_ENTROPY_BITS / BITS_PER_BYTE];
rng.fill_bytes(&mut entropy);
let entropy_bytes = word_count.entropy_bits() / BITS_PER_BYTE;
Self::from_entropy_in(lang, &entropy[..entropy_bytes])
.expect("valid entropy length won't fail to generate the mnemonic")
}
pub fn from_entropy<E: Into<Vec<u8>>>(entropy: E) -> Result<Self, Error> {
Self::from_entropy_in(Language::English, entropy)
}
pub fn from_entropy_in<E: Into<Vec<u8>>>(lang: Language, entropy: E) -> Result<Self, Error> {
const MAX_TOTAL_BITS: usize = Count::Words24.total_bits();
let entropy = entropy.into();
let word_count = Count::from_key_size(entropy.len() * BITS_PER_BYTE)?;
let mut bits = [false; MAX_TOTAL_BITS];
for (index, bit) in bits[word_count.entropy()].iter_mut().enumerate() {
*bit = left_index_bit(entropy[index / BITS_PER_BYTE], index % BITS_PER_BYTE);
}
let checksum_byte = Sha256::digest(&entropy)[0];
for (index, bit) in bits[word_count.checksum()].iter_mut().enumerate() {
*bit = left_index_bit(checksum_byte, index);
}
let mut words = Vec::with_capacity(word_count.word_count());
for chunk in bits[word_count.total()].chunks(BITS_PER_WORD) {
let index = bits_to_uint(chunk, BITS_PER_WORD);
words.push(lang.word_of(index));
}
let phrase = words.join(" ");
Ok(Self {
lang,
phrase,
entropy,
})
}
pub fn from_phrase<'a, P: Into<Cow<'a, str>>>(phrase: P) -> Result<Self, Error> {
Self::from_phrase_in(Language::English, phrase)
}
pub fn from_phrase_in<'a, P: Into<Cow<'a, str>>>(
lang: Language,
phrase: P,
) -> Result<Self, Error> {
let phrase = phrase.into();
let entropy = Self::phrase_to_entropy(lang, phrase.as_ref())?;
Ok(Mnemonic {
lang,
phrase: phrase.into_owned(),
entropy,
})
}
pub fn validate<'a, P: Into<Cow<'a, str>>>(phrase: P) -> Result<(), Error> {
Self::validate_in(Language::English, phrase)
}
pub fn validate_in<'a, P: Into<Cow<'a, str>>>(lang: Language, phrase: P) -> Result<(), Error> {
let _entropy = Self::phrase_to_entropy(lang, phrase)?;
Ok(())
}
fn phrase_to_entropy<'a, P: Into<Cow<'a, str>>>(
lang: Language,
phrase: P,
) -> Result<Vec<u8>, Error> {
let mut phrase = phrase.into();
normalize_utf8(&mut phrase);
let word_count = Count::from_phrase(phrase.as_ref())?;
let mut bits = vec![false; word_count.total_bits()];
for (i, word) in phrase.split_whitespace().enumerate() {
if let Some(index) = lang.index_of(word) {
index_to_bits(index, &mut bits[i * BITS_PER_WORD..], BITS_PER_WORD);
} else {
return Err(Error::UnknownWord(word.to_string()));
}
}
let mut entropy = vec![0u8; word_count.entropy_bits() / BITS_PER_BYTE];
entropy.iter_mut().enumerate().for_each(|(i, byte)| {
*byte = bits_to_uint(
&bits[i * BITS_PER_BYTE..(i + 1) * BITS_PER_BYTE],
BITS_PER_BYTE,
) as u8;
});
let checksum_bits = &bits[word_count.checksum()];
let actual_checksum = bits_to_uint(checksum_bits, word_count.checksum_bits()) as u8;
let checksum_byte = Sha256::digest(&entropy)[0];
let expected_checksum = checksum(checksum_byte, word_count.checksum_bits());
if actual_checksum != expected_checksum {
return Err(Error::InvalidChecksum);
}
Ok(entropy)
}
pub fn to_seed<P: AsRef<str>>(&self, passphrase: P) -> [u8; 64] {
const PBKDF2_ROUNDS: u32 = 2048;
const PBKDF2_BYTES: usize = 64;
let normalized_password = self.phrase();
let normalized_salt = {
let mut salt = Cow::Owned(format!("mnemonic{}", passphrase.as_ref()));
normalize_utf8(&mut salt);
salt
};
let mut seed = [0u8; PBKDF2_BYTES];
pbkdf2::pbkdf2::<Hmac<Sha512>>(
normalized_password.as_bytes(),
normalized_salt.as_bytes(),
PBKDF2_ROUNDS,
&mut seed,
);
seed
}
pub fn lang(&self) -> Language {
self.lang
}
pub fn phrase(&self) -> &str {
&self.phrase
}
pub fn into_phrase(mut self) -> String {
mem::take(&mut self.phrase)
}
pub fn entropy(&self) -> &[u8] {
&self.entropy
}
pub fn into_entropy(mut self) -> Vec<u8> {
mem::take(&mut self.entropy)
}
}
#[inline]
fn normalize_utf8(s: &mut Cow<'_, str>) {
use unicode_normalization::{is_nfkd_quick, IsNormalized, UnicodeNormalization};
if is_nfkd_quick(s.as_ref().chars()) != IsNormalized::Yes {
*s = Cow::Owned(s.as_ref().nfkd().to_string())
}
}
const fn checksum(source: u8, bits: usize) -> u8 {
source >> (BITS_PER_BYTE - bits)
}
const fn left_index_bit(source: u8, index: usize) -> bool {
let mask = 1 << (BITS_PER_BYTE - 1 - index);
source & mask > 0
}
#[inline]
fn bits_to_uint(bits: &[bool], chunk_size: usize) -> usize {
debug_assert_eq!(bits.len(), chunk_size);
bits.iter()
.take(chunk_size)
.enumerate()
.map(|(i, bit)| if *bit { 1 << (chunk_size - 1 - i) } else { 0 })
.sum::<usize>()
}
#[inline]
fn index_to_bits(index: usize, bits: &mut [bool], chunk_size: usize) {
debug_assert!(index < (2 << chunk_size));
bits.iter_mut()
.take(chunk_size)
.enumerate()
.for_each(|(i, bit)| *bit = (index >> (chunk_size - 1 - i)) & 1 == 1);
}
#[test]
fn test_left_index_bit() {
assert!(left_index_bit(0b1111_1111, 0));
assert!(left_index_bit(0b1111_1111, 3));
assert!(left_index_bit(0b1111_1111, 7));
assert!(left_index_bit(0b1111_0111, 0));
assert!(!left_index_bit(0b1111_0111, 4));
assert!(!left_index_bit(0b0100_0000, 0));
assert!(left_index_bit(0b0100_0000, 1));
}
#[test]
fn test_bits_to_uint() {
assert_eq!(bits_to_uint(&[false; 11], BITS_PER_WORD), 0b000_0000_0000); assert_eq!(bits_to_uint(&[true; 11], BITS_PER_WORD), 0b111_1111_1111); let mut bits = [false; 11];
bits[0] = true;
bits[1] = true;
bits[2] = true;
bits[3] = true;
bits[4] = true;
assert_eq!(bits_to_uint(&bits, BITS_PER_WORD), 0b111_1100_0000);
assert_eq!(bits_to_uint(&[false; 8], BITS_PER_BYTE), 0b0000_0000); assert_eq!(bits_to_uint(&[true; 8], BITS_PER_BYTE), 0b1111_1111); let mut bits = [false; 8];
bits[0] = true;
bits[1] = true;
bits[2] = true;
bits[3] = true;
bits[4] = true;
assert_eq!(bits_to_uint(&bits, BITS_PER_BYTE), 0b1111_1000); }
#[test]
fn test_index_to_bits() {
let mut bits: [bool; BITS_PER_WORD] = Default::default();
index_to_bits(0b000_0000_0000, &mut bits, BITS_PER_WORD);
assert_eq!(bits, [false; BITS_PER_WORD]);
let mut bits: [bool; BITS_PER_WORD] = Default::default();
index_to_bits(0b111_1111_1111, &mut bits, BITS_PER_WORD);
assert_eq!(bits, [true; BITS_PER_WORD]);
let mut bits: [bool; BITS_PER_WORD] = Default::default();
index_to_bits(0b111_1100_0000, &mut bits, BITS_PER_WORD);
let mut expected_bits = [false; BITS_PER_WORD];
expected_bits[0] = true;
expected_bits[1] = true;
expected_bits[2] = true;
expected_bits[3] = true;
expected_bits[4] = true;
assert_eq!(bits, expected_bits); }
#[test]
fn test_mnemonic_word_count() {
let mnemonic = Count::Words12;
assert_eq!(mnemonic.word_count(), 12);
assert_eq!(mnemonic.total_bits(), 128 + 4);
assert_eq!(mnemonic.entropy_bits(), 128);
assert_eq!(mnemonic.checksum_bits(), 4);
let mnemonic = Count::Words15;
assert_eq!(mnemonic.word_count(), 15);
assert_eq!(mnemonic.total_bits(), 160 + 5);
assert_eq!(mnemonic.entropy_bits(), 160);
assert_eq!(mnemonic.checksum_bits(), 5);
let mnemonic = Count::Words18;
assert_eq!(mnemonic.word_count(), 18);
assert_eq!(mnemonic.total_bits(), 192 + 6);
assert_eq!(mnemonic.entropy_bits(), 192);
assert_eq!(mnemonic.checksum_bits(), 6);
let mnemonic = Count::Words21;
assert_eq!(mnemonic.word_count(), 21);
assert_eq!(mnemonic.total_bits(), 224 + 7);
assert_eq!(mnemonic.entropy_bits(), 224);
assert_eq!(mnemonic.checksum_bits(), 7);
let mnemonic = Count::Words24;
assert_eq!(mnemonic.word_count(), 24);
assert_eq!(mnemonic.total_bits(), 256 + 8);
assert_eq!(mnemonic.entropy_bits(), 256);
assert_eq!(mnemonic.checksum_bits(), 8);
}
#[test]
fn test_mnemonic_zeroize_when_drop() {
let p: *const String;
let e: *const Vec<u8>;
{
let m = Mnemonic::from_entropy([1u8; 16]).unwrap();
p = &m.phrase;
e = &m.entropy;
unsafe {
println!("*p: {}", (*p));
println!("*e: {:?}", (*e));
}
}
unsafe {
assert_ne!(
(*p),
"absurd amount doctor acoustic avoid letter advice cage absurd amount doctor adjust"
);
println!("*p: {}", (*p));
assert_ne!((*e), [1u8; 16]);
println!("*e: {:?}", (*e));
}
}
#[test]
fn test_mnemonic_consume() {
let p: *const String;
{
let m = Mnemonic::from_entropy([1u8; 16]).unwrap();
p = &m.phrase;
unsafe {
println!("*p: {} ({:p})", (*p), p);
}
let phrase = m.into_phrase();
assert_ne!(p, &phrase);
println!("phrase: {} ({:p})", phrase, &phrase);
}
let e: *const Vec<u8>;
{
let m = Mnemonic::from_entropy([1u8; 16]).unwrap();
e = &m.entropy;
unsafe {
println!("*e: {:?} ({:p})", (*e), e);
}
let entropy = m.into_entropy();
assert_ne!(e, &entropy);
println!("entropy: {:?} ({:p})", entropy, &entropy);
}
}