use crate::hazmat::kupyna::Kupyna256;
use subtle::ConstantTimeEq;
pub const L_MAX_P: usize = 424;
const M_TILDE_BYTES: usize = L_MAX_P / 8;
const L_H_BYTES: usize = 8;
const M_PRIME_BYTES: usize = 1 + L_H_BYTES + 2 + M_TILDE_BYTES;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MessageError {
ZeroLength,
MessageTooLong,
LengthMismatch,
HashMismatch,
PaddingNotZero,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Message {
pub hash_id: u8,
pub bit_length: usize,
pub m_tilde: [u8; M_TILDE_BYTES],
}
pub fn format_m_tilde(
message: &[u8],
message_bits: usize,
) -> Result<[u8; M_TILDE_BYTES], MessageError> {
if message_bits == 0 {
return Err(MessageError::ZeroLength);
}
if message_bits > L_MAX_P {
return Err(MessageError::MessageTooLong);
}
let message_bytes = message_bits.div_ceil(8);
if message.len() != message_bytes {
return Err(MessageError::LengthMismatch);
}
let mut m_tilde = [0u8; M_TILDE_BYTES];
m_tilde[M_TILDE_BYTES - message_bytes..].copy_from_slice(message);
Ok(m_tilde)
}
#[must_use]
pub fn encode_l_m_tilde(message_bits: usize) -> [u8; 2] {
#[allow(clippy::cast_possible_truncation)] (message_bits as u16).to_be_bytes()
}
#[must_use]
pub fn build_m_prime(
hash_id: u8,
m_tilde: &[u8; M_TILDE_BYTES],
l_m_tilde: &[u8; 2],
) -> [u8; M_PRIME_BYTES] {
let mut hashed_input = [0u8; 2 + M_TILDE_BYTES];
hashed_input[..2].copy_from_slice(l_m_tilde);
hashed_input[2..].copy_from_slice(m_tilde);
let digest = Kupyna256::digest(&hashed_input);
let mut m_prime = [0u8; M_PRIME_BYTES];
m_prime[0] = hash_id;
m_prime[1..=L_H_BYTES].copy_from_slice(&digest[digest.len() - L_H_BYTES..]);
m_prime[1 + L_H_BYTES..3 + L_H_BYTES].copy_from_slice(l_m_tilde);
m_prime[3 + L_H_BYTES..].copy_from_slice(m_tilde);
m_prime
}
#[must_use]
pub fn kw_plaintext_from_m_prime(m_prime: &[u8; M_PRIME_BYTES]) -> [u8; 2 * M_PRIME_BYTES] {
let mut out = [0u8; 2 * M_PRIME_BYTES];
out[..M_PRIME_BYTES].copy_from_slice(m_prime);
out
}
pub fn parse_m_prime(m_prime: &[u8; M_PRIME_BYTES]) -> Result<Message, MessageError> {
let hash_id = m_prime[0];
let embedded_hash = &m_prime[1..=L_H_BYTES];
let mut l_m_tilde = [0u8; 2];
l_m_tilde.copy_from_slice(&m_prime[1 + L_H_BYTES..3 + L_H_BYTES]);
let mut m_tilde = [0u8; M_TILDE_BYTES];
m_tilde.copy_from_slice(&m_prime[3 + L_H_BYTES..]);
let bit_length = usize::from(u16::from_be_bytes(l_m_tilde));
if bit_length == 0 {
return Err(MessageError::ZeroLength);
}
if bit_length > L_MAX_P {
return Err(MessageError::MessageTooLong);
}
let mut hashed_input = [0u8; 2 + M_TILDE_BYTES];
hashed_input[..2].copy_from_slice(&l_m_tilde);
hashed_input[2..].copy_from_slice(&m_tilde);
let digest = Kupyna256::digest(&hashed_input);
let hash_ok: bool = digest[digest.len() - L_H_BYTES..]
.ct_eq(embedded_hash)
.into();
if !hash_ok {
return Err(MessageError::HashMismatch);
}
let message_bytes = bit_length.div_ceil(8);
let padding_len = M_TILDE_BYTES - message_bytes;
let mut bad_padding = 0u8;
for (i, &byte) in m_tilde.iter().enumerate() {
let is_padding_position = u8::from(i < padding_len);
bad_padding |= is_padding_position & u8::from(byte != 0);
}
if bad_padding != 0 {
return Err(MessageError::PaddingNotZero);
}
Ok(Message {
hash_id,
bit_length,
m_tilde,
})
}