use alloc::{string::ToString, vec::Vec};
use core::marker::PhantomData;
use miden_crypto_derive::{SilentDebug, SilentDisplay};
use num::{Complex, Float, Zero};
use num_complex::Complex64;
use rand::{CryptoRng, Rng};
use super::{
super::{
FalconVariant, LOG_N, MODULUS, N, Nonce, SIG_L2_BOUND, SIGMA, SK_LEN, ShortLatticeBasis,
Signature,
math::{
FalconFelt, FastFft, LdlTree, Polynomial, check_coefficients_bound, ffldl, ffsampling,
gram, has_acceptable_gram_schmidt_norm, normalize_tree, ntru_gen,
},
signature::SignaturePoly,
},
PublicKey,
};
use crate::{
Word,
hash::blake::Blake3_256,
utils::{
ByteReader, ByteWriter, Deserializable, DeserializationError, Serializable,
read_sensitive_array,
zeroize::{Zeroize, ZeroizeOnDrop, Zeroizing},
},
};
pub(crate) const WIDTH_BIG_POLY_COEFFICIENT: usize = 8;
pub(crate) const WIDTH_SMALL_POLY_COEFFICIENT: usize = 6;
#[derive(Clone, SilentDebug, SilentDisplay)]
pub struct SecretKey<V: FalconVariant> {
secret_key: ShortLatticeBasis,
tree: LdlTree,
variant: PhantomData<fn() -> V>,
}
impl<V: FalconVariant> Zeroize for SecretKey<V> {
fn zeroize(&mut self) {
self.secret_key.zeroize();
self.tree.zeroize();
}
}
impl<V: FalconVariant> Drop for SecretKey<V> {
fn drop(&mut self) {
self.zeroize();
}
}
impl<V: FalconVariant> ZeroizeOnDrop for SecretKey<V> {}
#[allow(clippy::new_without_default)]
impl<V: FalconVariant> SecretKey<V> {
#[cfg(feature = "std")]
pub fn new() -> Self {
let mut rng = rand::rng();
Self::with_rng(&mut rng)
}
pub fn with_rng<R: CryptoRng + Rng>(rng: &mut R) -> Self {
let basis = ntru_gen(N, rng);
Self::from_short_lattice_basis(basis)
}
pub(crate) fn from_short_lattice_basis(basis: ShortLatticeBasis) -> Self {
let basis_fft = to_complex_fft(&basis);
let gram_fft = gram(&basis_fft);
let mut tree = ffldl(&gram_fft);
normalize_tree(&mut tree, SIGMA);
Self {
secret_key: basis,
tree,
variant: PhantomData,
}
}
pub fn short_lattice_basis(&self) -> &ShortLatticeBasis {
&self.secret_key
}
pub fn public_key(&self) -> PublicKey<V> {
self.compute_pub_key_poly()
}
pub fn tree(&self) -> &LdlTree {
&self.tree
}
pub fn sign(&self, message: Word) -> Signature<V> {
use rand::SeedableRng;
use rand_chacha::ChaCha20Rng;
let seed = Zeroizing::new(self.generate_seed(&message));
let mut rng = ChaCha20Rng::from_seed(*seed);
self.sign_with_rng(message, &mut rng)
}
pub fn sign_with_rng<R: CryptoRng + Rng>(&self, message: Word, rng: &mut R) -> Signature<V> {
let nonce = Nonce::deterministic();
let h = self.compute_pub_key_poly();
let c = V::hash_message_to_point(message, &nonce);
let s2 = self.sign_helper(&c, rng);
Signature::new(nonce, h, s2)
}
#[cfg(test)]
pub(crate) fn sign_with_rng_testing<R: Rng>(
&self,
message: &[u8],
rng: &mut R,
) -> Signature<V> {
use super::super::test_utils::{ChaCha, hash_to_point_shake256};
let nonce = Nonce::random(rng);
let h = self.compute_pub_key_poly();
let c = hash_to_point_shake256(message, &nonce);
let mut chacha_prng = ChaCha::new(rng);
let s2 = self.sign_helper(&c, &mut chacha_prng);
Signature::new(nonce, h, s2)
}
fn compute_pub_key_poly(&self) -> PublicKey<V> {
let g: Polynomial<FalconFelt> = self.secret_key[0].clone().into();
let g_fft = g.fft();
let minus_f: Polynomial<FalconFelt> = self.secret_key[1].clone().into();
let f = -minus_f;
let f_fft = f.fft();
let h_fft = g_fft.hadamard_div(&f_fft);
h_fft.ifft().into()
}
fn sign_helper<R: Rng>(&self, c: &Polynomial<FalconFelt>, rng: &mut R) -> SignaturePoly {
let one_over_q = 1.0 / (MODULUS as f64);
let c_over_q_fft = c.map(|cc| Complex::new(one_over_q * cc.value() as f64, 0.0)).fft();
let [g_fft, minus_f_fft, big_g_fft, minus_big_f_fft] = to_complex_fft(&self.secret_key);
let t0 = c_over_q_fft.hadamard_mul(&minus_big_f_fft);
let t1 = -c_over_q_fft.hadamard_mul(&minus_f_fft);
loop {
let bold_s = loop {
let z = ffsampling(&(t0.clone(), t1.clone()), &self.tree, rng);
let t0_min_z0 = t0.clone() - z.0;
let t1_min_z1 = t1.clone() - z.1;
let s0 = t0_min_z0.hadamard_mul(&g_fft) + t1_min_z1.hadamard_mul(&big_g_fft);
let s1 =
t0_min_z0.hadamard_mul(&minus_f_fft) + t1_min_z1.hadamard_mul(&minus_big_f_fft);
let length_squared: f64 =
(s0.coefficients.iter().map(|a| (a * a.conj()).re).sum::<f64>()
+ s1.coefficients.iter().map(|a| (a * a.conj()).re).sum::<f64>())
/ (N as f64);
if length_squared > (SIG_L2_BOUND as f64) {
continue;
}
break [-s0, s1];
};
let s2 = bold_s[1].ifft();
let s2_coef: [i16; N] = s2
.coefficients
.iter()
.map(|a| Float::round(a.re) as i16)
.collect::<Vec<i16>>()
.try_into()
.expect("The number of coefficients should be equal to N");
if let Ok(s2) = SignaturePoly::try_from(&s2_coef) {
return s2;
}
}
}
fn generate_seed(&self, message: &Word) -> [u8; 32] {
let serialized_key = Zeroizing::new(self.to_bytes());
let mut buffer = Zeroizing::new(Vec::with_capacity(1 + SK_LEN + Word::SERIALIZED_SIZE));
buffer.push(LOG_N);
buffer.extend_from_slice(&serialized_key);
buffer.extend_from_slice(&message.to_bytes());
let digest = Blake3_256::hash(&buffer);
digest.into()
}
}
impl<V: FalconVariant> PartialEq for SecretKey<V> {
fn eq(&self, other: &Self) -> bool {
use subtle::ConstantTimeEq;
let self_bytes = Zeroizing::new(self.to_bytes());
let other_bytes = Zeroizing::new(other.to_bytes());
self_bytes.ct_eq(&other_bytes).into()
}
}
impl<V: FalconVariant> Eq for SecretKey<V> {}
impl<V: FalconVariant> Serializable for SecretKey<V> {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
let basis = &self.secret_key;
let n = basis[0].coefficients.len();
let l = n.checked_ilog2().unwrap() as u8;
let header: u8 = (5 << 4) | l;
let neg_f = &basis[1];
let g = &basis[0];
let neg_big_f = &basis[3];
let mut buffer = Zeroizing::new(Vec::with_capacity(SK_LEN));
buffer.push(header);
let f_i8 = Zeroizing::new(
neg_f
.coefficients
.iter()
.map(|&a| secret_key_coefficient_to_i8(-FalconFelt::new(a)))
.collect::<Vec<i8>>(),
);
let f_i8_encoded = Zeroizing::new(
encode_i8(&f_i8, WIDTH_SMALL_POLY_COEFFICIENT)
.expect("valid Falcon key coefficients must be encodable"),
);
buffer.extend_from_slice(&f_i8_encoded);
let g_i8 = Zeroizing::new(
g.coefficients
.iter()
.map(|&a| secret_key_coefficient_to_i8(FalconFelt::new(a)))
.collect::<Vec<i8>>(),
);
let g_i8_encoded = Zeroizing::new(
encode_i8(&g_i8, WIDTH_SMALL_POLY_COEFFICIENT)
.expect("valid Falcon key coefficients must be encodable"),
);
buffer.extend_from_slice(&g_i8_encoded);
let big_f_i8 = Zeroizing::new(
neg_big_f
.coefficients
.iter()
.map(|&a| secret_key_coefficient_to_i8(-FalconFelt::new(a)))
.collect::<Vec<i8>>(),
);
let big_f_i8_encoded = Zeroizing::new(
encode_i8(&big_f_i8, WIDTH_BIG_POLY_COEFFICIENT)
.expect("valid Falcon key coefficients must be encodable"),
);
buffer.extend_from_slice(&big_f_i8_encoded);
target.write_bytes(&buffer);
}
}
impl<V: FalconVariant> Deserializable for SecretKey<V> {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
let byte_vector = read_sensitive_array::<SK_LEN, _>(source)?;
let header = byte_vector[0];
if (header >> 4) != 5 {
return Err(DeserializationError::InvalidValue("Invalid header format".to_string()));
}
let logn = (header & 15) as usize;
let n = 1 << logn;
if n != N {
return Err(DeserializationError::InvalidValue(
"Unsupported Falcon DSA variant".to_string(),
));
}
let chunk_size_f = ((n * WIDTH_SMALL_POLY_COEFFICIENT) + 7) >> 3;
let chunk_size_g = ((n * WIDTH_SMALL_POLY_COEFFICIENT) + 7) >> 3;
let chunk_size_big_f = ((n * WIDTH_BIG_POLY_COEFFICIENT) + 7) >> 3;
let f = Zeroizing::new(
decode_i8(&byte_vector[1..chunk_size_f + 1], WIDTH_SMALL_POLY_COEFFICIENT).ok_or(
DeserializationError::InvalidValue("Failed to decode f coefficients".to_string()),
)?,
);
let g = Zeroizing::new(
decode_i8(
&byte_vector[chunk_size_f + 1..(chunk_size_f + chunk_size_g + 1)],
WIDTH_SMALL_POLY_COEFFICIENT,
)
.ok_or(DeserializationError::InvalidValue(
"Failed to decode g coefficients".to_string(),
))?,
);
let big_f = Zeroizing::new(
decode_i8(
&byte_vector[(chunk_size_f + chunk_size_g + 1)
..(chunk_size_f + chunk_size_g + chunk_size_big_f + 1)],
WIDTH_BIG_POLY_COEFFICIENT,
)
.ok_or(DeserializationError::InvalidValue(
"Failed to decode F coefficients".to_string(),
))?,
);
let mut f = Polynomial::new(f.iter().map(|&c| i16::from(c)).collect());
let g = Polynomial::new(g.iter().map(|&c| i16::from(c)).collect());
let mut big_f = Polynomial::new(big_f.iter().map(|&c| i16::from(c)).collect());
let f_fft = Polynomial::<FalconFelt>::from(&f).fft();
if f_fft.coefficients.iter().any(Zero::is_zero) {
return Err(DeserializationError::InvalidValue(
"Falcon secret key polynomial f is not invertible".to_string(),
));
}
if !has_acceptable_gram_schmidt_norm(&f, &g) {
return Err(DeserializationError::InvalidValue(
"Falcon secret key exceeds the Gram-Schmidt norm bound".to_string(),
));
}
let g_fft = Polynomial::<FalconFelt>::from(&g).fft();
let big_f_fft = Polynomial::<FalconFelt>::from(&big_f).fft();
let big_g = g_fft.hadamard_div(&f_fft).hadamard_mul(&big_f_fft).ifft();
let big_g = Polynomial::new(big_g.to_balanced_values());
let big_coefficient_bound = (1 << (WIDTH_BIG_POLY_COEFFICIENT - 1)) - 1;
if !check_coefficients_bound(&big_g, big_coefficient_bound as i16) {
return Err(DeserializationError::InvalidValue(
"Falcon secret key polynomial G exceeds its coefficient bound".to_string(),
));
}
if !satisfies_ntru_relation(&f, &g, &big_f, &big_g) {
return Err(DeserializationError::InvalidValue(
"Falcon secret key does not satisfy the NTRU equation".to_string(),
));
}
for coefficient in &mut f.coefficients {
*coefficient = -*coefficient;
}
for coefficient in &mut big_f.coefficients {
*coefficient = -*coefficient;
}
let basis = [g, f, big_g, big_f];
Ok(Self::from_short_lattice_basis(basis))
}
}
fn to_complex_fft(basis: &[Polynomial<i16>; 4]) -> [Polynomial<Complex<f64>>; 4] {
let [g, f, big_g, big_f] = basis.clone();
let g_fft = g.map(|cc| Complex64::new(*cc as f64, 0.0)).fft();
let minus_f_fft = f.map(|cc| -Complex64::new(*cc as f64, 0.0)).fft();
let big_g_fft = big_g.map(|cc| Complex64::new(*cc as f64, 0.0)).fft();
let minus_big_f_fft = big_f.map(|cc| -Complex64::new(*cc as f64, 0.0)).fft();
[g_fft, minus_f_fft, big_g_fft, minus_big_f_fft]
}
fn satisfies_ntru_relation(
f: &Polynomial<i16>,
g: &Polynomial<i16>,
big_f: &Polynomial<i16>,
big_g: &Polynomial<i16>,
) -> bool {
let f = f.map(|&coefficient| i64::from(coefficient));
let g = g.map(|&coefficient| i64::from(coefficient));
let big_f = big_f.map(|&coefficient| i64::from(coefficient));
let big_g = big_g.map(|&coefficient| i64::from(coefficient));
let determinant = (f * big_g - g * big_f).reduce_by_cyclotomic(N);
determinant == Polynomial::constant(i64::from(MODULUS))
}
fn secret_key_coefficient_to_i8(coefficient: FalconFelt) -> i8 {
i8::try_from(coefficient.balanced_value())
.expect("valid Falcon secret-key coefficients must fit in i8")
}
pub fn encode_i8(x: &[i8], bits: usize) -> Option<Vec<u8>> {
let maxv = (1 << (bits - 1)) - 1_usize;
let maxv = maxv as i8;
let minv = -maxv;
for &c in x {
if c > maxv || c < minv {
return None;
}
}
let out_len = ((N * bits) + 7) >> 3;
let mut buf = vec![0_u8; out_len];
let mut acc = 0_u32;
let mut acc_len = 0;
let mask = ((1_u16 << bits) - 1) as u8;
let mut input_pos = 0;
for &c in x {
acc = (acc << bits) | (c as u8 & mask) as u32;
acc_len += bits;
while acc_len >= 8 {
acc_len -= 8;
buf[input_pos] = (acc >> acc_len) as u8;
input_pos += 1;
}
}
if acc_len > 0 {
buf[input_pos] = (acc >> (8 - acc_len)) as u8;
}
Some(buf)
}
pub fn decode_i8(buf: &[u8], bits: usize) -> Option<Vec<i8>> {
let mut x = Zeroizing::new([0_i8; N]);
let mut i = 0;
let mut j = 0;
let mut acc = 0_u32;
let mut acc_len = 0;
let mask = (1_u32 << bits) - 1;
let a = (1 << bits) as u8;
let b = ((1 << (bits - 1)) - 1) as u8;
while i < N {
acc = (acc << 8) | (buf[j] as u32);
j += 1;
acc_len += 8;
while acc_len >= bits && i < N {
acc_len -= bits;
let w = (acc >> acc_len) & mask;
if w == 1 << (bits - 1) {
return None;
}
let w = w as u8;
let z = if w > b { w as i8 - a as i8 } else { w as i8 };
x[i] = z;
i += 1;
}
}
if (acc & ((1u32 << acc_len) - 1)) == 0 {
Some(x.to_vec())
} else {
None
}
}
#[cfg(test)]
mod tests {
use rand::SeedableRng;
use rand_chacha::ChaCha20Rng;
use super::*;
type TestSecretKey = SecretKey<super::super::super::TestVariant>;
#[test]
fn secret_key_deserialization_rejects_noninvertible_f() {
let mut encoded = vec![0u8; SK_LEN];
encoded[0] = (5 << 4) | LOG_N;
assert_invalid_key(&encoded, "Falcon secret key polynomial f is not invertible");
}
#[test]
fn secret_key_deserialization_rejects_excessive_norm() {
let mut f = Polynomial::new(vec![0i16; N]);
f.coefficients[0] = 1;
let g = Polynomial::new(vec![31i16; N]);
let big_f = Polynomial::new(vec![0i16; N]);
let encoded = encode_secret_key(&f, &g, &big_f);
assert_invalid_key(&encoded, "Falcon secret key exceeds the Gram-Schmidt norm bound");
}
#[test]
fn secret_key_deserialization_rejects_invalid_ntru_relation() {
let (f, g) = acceptable_f_and_g();
let big_f = Polynomial::new(vec![0i16; N]);
let encoded = encode_secret_key(&f, &g, &big_f);
assert_invalid_key(&encoded, "Falcon secret key does not satisfy the NTRU equation");
}
#[test]
fn secret_key_deserialization_rejects_out_of_range_reconstructed_big_g() {
let (f, g) = acceptable_f_and_g();
let big_f = constant_polynomial(127);
let encoded = encode_secret_key(&f, &g, &big_f);
assert_invalid_key(
&encoded,
"Falcon secret key polynomial G exceeds its coefficient bound",
);
}
#[test]
fn secret_key_deserialization_rejects_forbidden_minimum_coefficients() {
let f_offset = 1;
let g_offset = f_offset + N * WIDTH_SMALL_POLY_COEFFICIENT / 8;
let big_f_offset = g_offset + N * WIDTH_SMALL_POLY_COEFFICIENT / 8;
for (offset, polynomial) in [(f_offset, "f"), (g_offset, "g"), (big_f_offset, "F")] {
let mut encoded = vec![0u8; SK_LEN];
encoded[0] = (5 << 4) | LOG_N;
encoded[offset] = 0b1000_0000;
assert_invalid_key(&encoded, &format!("Failed to decode {polynomial} coefficients"));
}
}
fn constant_polynomial(value: i16) -> Polynomial<i16> {
let mut polynomial = Polynomial::new(vec![0i16; N]);
polynomial.coefficients[0] = value;
polynomial
}
fn acceptable_f_and_g() -> (Polynomial<i16>, Polynomial<i16>) {
let mut rng = ChaCha20Rng::from_seed([9u8; 32]);
let [g, minus_f, _, _] = ntru_gen(N, &mut rng);
(-minus_f, g)
}
fn encode_secret_key(
f: &Polynomial<i16>,
g: &Polynomial<i16>,
big_f: &Polynomial<i16>,
) -> Vec<u8> {
let encode = |polynomial: &Polynomial<i16>, width| {
let coefficients = polynomial
.coefficients
.iter()
.map(|&coefficient| coefficient as i8)
.collect::<Vec<_>>();
encode_i8(&coefficients, width).unwrap()
};
let mut encoded = Vec::with_capacity(SK_LEN);
encoded.push((5 << 4) | LOG_N);
encoded.extend_from_slice(&encode(f, WIDTH_SMALL_POLY_COEFFICIENT));
encoded.extend_from_slice(&encode(g, WIDTH_SMALL_POLY_COEFFICIENT));
encoded.extend_from_slice(&encode(big_f, WIDTH_BIG_POLY_COEFFICIENT));
assert_eq!(encoded.len(), SK_LEN);
encoded
}
fn assert_invalid_key(encoded: &[u8], expected_message: &str) {
let error = TestSecretKey::read_from_bytes(encoded).unwrap_err();
assert_eq!(error, DeserializationError::InvalidValue(expected_message.to_string()),);
}
}