use super::curve512::{base_point, is_valid_scalar, point_from_x, Point};
use super::fp512::from_candidate_bytes;
use super::message512::{
build_m_prime, encode_l_m_tilde, format_m_tilde, kw_plaintext_from_m_prime, parse_m_prime,
};
use crate::hazmat::kalyna_kw::Kalyna512_512Kw;
const HASH_ID_KUPYNA256: u8 = 0x01;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EncryptError {
InvalidMessage,
InvalidEphemeralKey,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DecryptError {
InvalidCiphertext,
}
pub fn encrypt(
message: &[u8],
message_bits: usize,
q: Point,
epsilon: &[u8; 64],
) -> Result<[u8; 256], EncryptError> {
if !is_valid_scalar(epsilon) {
return Err(EncryptError::InvalidEphemeralKey);
}
let m_tilde =
format_m_tilde(message, message_bits).map_err(|_| EncryptError::InvalidMessage)?;
let l_m_tilde = encode_l_m_tilde(message_bits);
let m_prime = build_m_prime(HASH_ID_KUPYNA256, &m_tilde, &l_m_tilde);
let r_point = base_point().scalar_multiply(epsilon);
let r_bytes = r_point.x.to_be_bytes();
let t_point = q.scalar_multiply(epsilon);
let kappa = t_point.x.to_be_bytes();
let kw_plaintext = kw_plaintext_from_m_prime(&m_prime);
let mut t = [0u8; 192];
Kalyna512_512Kw::wrap(&kappa, &kw_plaintext, &mut t)
.map_err(|_| EncryptError::InvalidMessage)?;
let mut ciphertext = [0u8; 256];
ciphertext[..64].copy_from_slice(&r_bytes);
ciphertext[64..].copy_from_slice(&t);
Ok(ciphertext)
}
pub fn decrypt(ciphertext: &[u8; 256], e: &[u8; 64]) -> Result<([u8; 53], usize), DecryptError> {
if !is_valid_scalar(e) {
return Err(DecryptError::InvalidCiphertext);
}
let mut r_bytes = [0u8; 64];
r_bytes.copy_from_slice(&ciphertext[..64]);
let r_field = from_candidate_bytes(&r_bytes).ok_or(DecryptError::InvalidCiphertext)?;
let r_prime = point_from_x(r_field).ok_or(DecryptError::InvalidCiphertext)?;
let t_prime = r_prime.scalar_multiply(e);
let kappa = t_prime.x.to_be_bytes();
let mut recovered = [0u8; 128];
Kalyna512_512Kw::unwrap(&kappa, &ciphertext[64..], &mut recovered)
.map_err(|_| DecryptError::InvalidCiphertext)?;
let mut appended_block_bad = 0u8;
for &byte in &recovered[64..] {
appended_block_bad |= u8::from(byte != 0);
}
if appended_block_bad != 0 {
return Err(DecryptError::InvalidCiphertext);
}
let mut m_prime = [0u8; 64];
m_prime.copy_from_slice(&recovered[..64]);
let parsed = parse_m_prime(&m_prime).map_err(|_| DecryptError::InvalidCiphertext)?;
Ok((parsed.m_tilde, parsed.bit_length))
}