use super::tables::{apply_forward_matrix, forward_sbox_mds, ROWS, SBOXES};
const MAX_COLUMNS: usize = 16;
const MAX_BLOCK_BYTES: usize = MAX_COLUMNS * ROWS;
#[allow(dead_code)]
fn sub_bytes(state: &mut [[u8; ROWS]]) {
for column in state.iter_mut() {
for (row, byte) in column.iter_mut().enumerate() {
*byte = SBOXES[row % 4][*byte as usize];
}
}
}
#[allow(dead_code)]
fn shift_bytes(state: &mut [[u8; ROWS]], last_row_shift: usize) {
let columns = state.len();
let mut shifted = [[0u8; ROWS]; MAX_COLUMNS];
for row in 0..ROWS {
let shift = if row == ROWS - 1 { last_row_shift } else { row };
for col in 0..columns {
shifted[(col + shift) % columns][row] = state[col][row];
}
}
state[..columns].copy_from_slice(&shifted[..columns]);
}
#[allow(dead_code)]
fn mix_columns(state: &mut [[u8; ROWS]]) {
apply_forward_matrix(state);
}
#[allow(dead_code)]
fn sub_shift_mix(state: &mut [[u8; ROWS]], last_row_shift: usize) {
let columns = state.len();
debug_assert!(columns.is_power_of_two());
let columns_mask = columns - 1;
let mut result = [[0u8; ROWS]; MAX_COLUMNS];
for (out_col, out_word) in result[..columns].iter_mut().enumerate() {
let mut acc = 0u64;
#[allow(clippy::needless_range_loop)]
for row in 0..ROWS {
let shift = if row == ROWS - 1 { last_row_shift } else { row };
let src_col = (out_col + columns - shift) & columns_mask;
let byte = state[src_col][row];
acc ^= forward_sbox_mds(row, byte);
}
*out_word = acc.to_le_bytes();
}
state[..columns].copy_from_slice(&result[..columns]);
}
fn sub_shift_mix_n<const COLUMNS: usize>(state: &mut [[u8; ROWS]; COLUMNS], last_row_shift: usize) {
debug_assert!(COLUMNS.is_power_of_two());
let columns_mask = COLUMNS - 1;
let mut result = [[0u8; ROWS]; COLUMNS];
for (out_col, out_word) in result.iter_mut().enumerate() {
let mut acc = 0u64;
#[allow(clippy::needless_range_loop)]
for row in 0..ROWS {
let shift = if row == ROWS - 1 { last_row_shift } else { row };
let src_col = (out_col + COLUMNS - shift) & columns_mask;
let byte = state[src_col][row];
acc ^= forward_sbox_mds(row, byte);
}
*out_word = acc.to_le_bytes();
}
*state = result;
}
#[allow(dead_code)]
#[allow(clippy::cast_possible_truncation)] fn add_round_constant_xor(state: &mut [[u8; ROWS]], round: u8) {
for (col, column) in state.iter_mut().enumerate() {
column[0] ^= (col as u8).wrapping_mul(0x10) ^ round;
}
}
#[allow(clippy::cast_possible_truncation)] fn add_round_constant_xor_n<const COLUMNS: usize>(state: &mut [[u8; ROWS]; COLUMNS], round: u8) {
for (col, column) in state.iter_mut().enumerate() {
column[0] ^= (col as u8).wrapping_mul(0x10) ^ round;
}
}
#[allow(dead_code)]
#[allow(clippy::cast_possible_truncation)] fn add_round_constant_add(state: &mut [[u8; ROWS]], round: u8) {
let columns = state.len();
for (col, column) in state.iter_mut().enumerate() {
let top_byte = ((columns - 1 - col) as u8).wrapping_mul(0x10) ^ round;
let addend = u64::from_le_bytes([0xF3, 0xF0, 0xF0, 0xF0, 0xF0, 0xF0, 0xF0, top_byte]);
let word = u64::from_le_bytes(*column).wrapping_add(addend);
*column = word.to_le_bytes();
}
}
#[allow(clippy::cast_possible_truncation)] fn add_round_constant_add_n<const COLUMNS: usize>(state: &mut [[u8; ROWS]; COLUMNS], round: u8) {
for (col, column) in state.iter_mut().enumerate() {
let top_byte = ((COLUMNS - 1 - col) as u8).wrapping_mul(0x10) ^ round;
let addend = u64::from_le_bytes([0xF3, 0xF0, 0xF0, 0xF0, 0xF0, 0xF0, 0xF0, top_byte]);
let word = u64::from_le_bytes(*column).wrapping_add(addend);
*column = word.to_le_bytes();
}
}
#[allow(dead_code)]
#[allow(clippy::cast_possible_truncation)] fn t_transform(state: &mut [[u8; ROWS]], rounds: usize, last_row_shift: usize) {
for round in 0..rounds {
add_round_constant_xor(state, round as u8);
sub_shift_mix(state, last_row_shift);
}
}
#[allow(dead_code)]
#[allow(clippy::cast_possible_truncation)] fn t_plus_transform(state: &mut [[u8; ROWS]], rounds: usize, last_row_shift: usize) {
for round in 0..rounds {
add_round_constant_add(state, round as u8);
sub_shift_mix(state, last_row_shift);
}
}
#[allow(clippy::cast_possible_truncation)] fn t_transform_n<const COLUMNS: usize, const ROUNDS: usize>(
state: &mut [[u8; ROWS]; COLUMNS],
last_row_shift: usize,
) {
for round in 0..ROUNDS {
add_round_constant_xor_n(state, round as u8);
sub_shift_mix_n(state, last_row_shift);
}
}
#[allow(clippy::cast_possible_truncation)] fn t_plus_transform_n<const COLUMNS: usize, const ROUNDS: usize>(
state: &mut [[u8; ROWS]; COLUMNS],
last_row_shift: usize,
) {
for round in 0..ROUNDS {
add_round_constant_add_n(state, round as u8);
sub_shift_mix_n(state, last_row_shift);
}
}
#[allow(dead_code)]
fn compress(h: &mut [[u8; ROWS]], block: &[[u8; ROWS]], rounds: usize, last_row_shift: usize) {
let columns = h.len();
let mut t_input = [[0u8; ROWS]; MAX_COLUMNS];
let mut q_input = [[0u8; ROWS]; MAX_COLUMNS];
for col in 0..columns {
for row in 0..ROWS {
t_input[col][row] = h[col][row] ^ block[col][row];
q_input[col][row] = block[col][row];
}
}
t_transform(&mut t_input[..columns], rounds, last_row_shift);
t_plus_transform(&mut q_input[..columns], rounds, last_row_shift);
for col in 0..columns {
for row in 0..ROWS {
h[col][row] ^= t_input[col][row] ^ q_input[col][row];
}
}
}
fn compress_n<const COLUMNS: usize, const ROUNDS: usize>(
h: &mut [[u8; ROWS]; COLUMNS],
block: &[[u8; ROWS]; COLUMNS],
last_row_shift: usize,
) {
let mut t_input = [[0u8; ROWS]; COLUMNS];
let mut q_input = [[0u8; ROWS]; COLUMNS];
for col in 0..COLUMNS {
for row in 0..ROWS {
t_input[col][row] = h[col][row] ^ block[col][row];
q_input[col][row] = block[col][row];
}
}
t_transform_n::<COLUMNS, ROUNDS>(&mut t_input, last_row_shift);
t_plus_transform_n::<COLUMNS, ROUNDS>(&mut q_input, last_row_shift);
for col in 0..COLUMNS {
for row in 0..ROWS {
h[col][row] ^= t_input[col][row] ^ q_input[col][row];
}
}
}
#[allow(dead_code)]
fn bytes_to_columns(bytes: &[u8], columns: usize) -> [[u8; ROWS]; MAX_COLUMNS] {
let mut out = [[0u8; ROWS]; MAX_COLUMNS];
for col in 0..columns {
out[col].copy_from_slice(&bytes[col * ROWS..col * ROWS + ROWS]);
}
out
}
fn bytes_to_columns_n<const COLUMNS: usize>(bytes: &[u8]) -> [[u8; ROWS]; COLUMNS] {
let mut out = [[0u8; ROWS]; COLUMNS];
for (col, word) in out.iter_mut().enumerate() {
word.copy_from_slice(&bytes[col * ROWS..col * ROWS + ROWS]);
}
out
}
fn state_array_mut_kupyna<const COLUMNS: usize>(
full: &mut [[u8; ROWS]; MAX_COLUMNS],
) -> &mut [[u8; ROWS]; COLUMNS] {
match (&mut full[..COLUMNS]).try_into() {
Ok(array) => array,
Err(_) => unreachable!("KupynaCore only ever constructs COLUMNS <= MAX_COLUMNS"),
}
}
fn h_to_array<const COLUMNS: usize>(h: &[[u8; ROWS]; MAX_COLUMNS]) -> [[u8; ROWS]; COLUMNS] {
match h[..COLUMNS].try_into() {
Ok(array) => array,
Err(_) => unreachable!("KupynaCore only ever constructs columns=8 or 16"),
}
}
pub(crate) fn kupyna_padding(
prefix: &[u8],
msg_bits: u64,
block_bytes: usize,
) -> ([u8; 2 * MAX_BLOCK_BYTES], usize) {
let mut tail = [0u8; 2 * MAX_BLOCK_BYTES];
let mut pos = prefix.len();
tail[..pos].copy_from_slice(prefix);
tail[pos] = 0x80;
pos += 1;
let used = pos + 12;
let zero_bytes = (block_bytes - (used % block_bytes)) % block_bytes;
pos += zero_bytes;
tail[pos..pos + 8].copy_from_slice(&msg_bits.to_le_bytes());
pos += 12; (tail, pos)
}
pub(crate) struct KupynaCore {
h: [[u8; ROWS]; MAX_COLUMNS],
buffer: [u8; MAX_BLOCK_BYTES],
buffer_len: usize,
total_len: u64,
columns: usize,
#[allow(dead_code)]
rounds: usize,
last_row_shift: usize,
}
impl KupynaCore {
#[allow(clippy::cast_possible_truncation)] pub(crate) fn new(columns: usize, rounds: usize, last_row_shift: usize) -> Self {
let block_bytes = columns * ROWS;
let mut h = [[0u8; ROWS]; MAX_COLUMNS];
h[0][0] = block_bytes as u8; Self {
h,
buffer: [0u8; MAX_BLOCK_BYTES],
buffer_len: 0,
total_len: 0,
columns,
rounds,
last_row_shift,
}
}
pub(crate) fn block_bytes(&self) -> usize {
self.columns * ROWS
}
pub(crate) fn buffered(&self) -> &[u8] {
&self.buffer[..self.buffer_len]
}
fn compress_block(&mut self, block: &[u8]) {
match self.columns {
8 => {
let columns_buf = bytes_to_columns_n::<8>(block);
let h = state_array_mut_kupyna::<8>(&mut self.h);
compress_n::<8, 10>(h, &columns_buf, self.last_row_shift);
}
16 => {
let columns_buf = bytes_to_columns_n::<16>(block);
let h = state_array_mut_kupyna::<16>(&mut self.h);
compress_n::<16, 14>(h, &columns_buf, self.last_row_shift);
}
_ => unreachable!("KupynaCore::new only ever constructed with columns=8 or 16"),
}
}
#[allow(clippy::cast_possible_truncation)] pub(crate) fn update(&mut self, mut data: &[u8]) {
self.total_len += data.len() as u64;
let block_bytes = self.block_bytes();
if self.buffer_len > 0 {
let take = (block_bytes - self.buffer_len).min(data.len());
self.buffer[self.buffer_len..self.buffer_len + take].copy_from_slice(&data[..take]);
self.buffer_len += take;
data = &data[take..];
if self.buffer_len < block_bytes {
debug_assert!(data.is_empty());
return;
}
let block = self.buffer;
self.compress_block(&block[..block_bytes]);
self.buffer_len = 0;
}
let mut full_blocks = data.chunks_exact(block_bytes);
for block in &mut full_blocks {
self.compress_block(block);
}
let remainder = full_blocks.remainder();
self.buffer[..remainder.len()].copy_from_slice(remainder);
self.buffer_len = remainder.len();
}
pub(crate) fn finalize(mut self, output_bytes: usize) -> [u8; 64] {
let block_bytes = self.block_bytes();
let msg_bits: u64 = self.total_len * 8;
let (tail, pos) = kupyna_padding(&self.buffer[..self.buffer_len], msg_bits, block_bytes);
let tail_blocks = pos / block_bytes;
for i in 0..tail_blocks {
self.compress_block(&tail[i * block_bytes..(i + 1) * block_bytes]);
}
#[allow(clippy::needless_range_loop)]
match self.columns {
8 => {
let mut t_final = h_to_array::<8>(&self.h);
t_transform_n::<8, 10>(&mut t_final, self.last_row_shift);
for col in 0..8 {
for row in 0..ROWS {
self.h[col][row] ^= t_final[col][row];
}
}
}
16 => {
let mut t_final = h_to_array::<16>(&self.h);
t_transform_n::<16, 14>(&mut t_final, self.last_row_shift);
for col in 0..16 {
for row in 0..ROWS {
self.h[col][row] ^= t_final[col][row];
}
}
}
_ => unreachable!("KupynaCore only ever constructs columns=8 or 16"),
}
let mut flat = [0u8; MAX_BLOCK_BYTES];
for col in 0..self.columns {
flat[col * ROWS..(col + 1) * ROWS].copy_from_slice(&self.h[col]);
}
let mut out = [0u8; 64];
out[..output_bytes].copy_from_slice(&flat[block_bytes - output_bytes..block_bytes]);
out
}
}
fn digest_generic(
message: &[u8],
columns: usize,
rounds: usize,
last_row_shift: usize,
output_bytes: usize,
) -> [u8; 64] {
let mut core = KupynaCore::new(columns, rounds, last_row_shift);
core.update(message);
core.finalize(output_bytes)
}
pub struct Kupyna256;
impl Kupyna256 {
#[must_use]
pub fn digest(message: &[u8]) -> [u8; 32] {
let full = digest_generic(message, 8, 10, 7, 32);
let mut out = [0u8; 32];
out.copy_from_slice(&full[..32]);
out
}
}
pub struct Kupyna512;
impl Kupyna512 {
#[must_use]
pub fn digest(message: &[u8]) -> [u8; 64] {
digest_generic(message, 16, 14, 11, 64)
}
}
pub struct Kupyna256Hasher(KupynaCore);
impl Kupyna256Hasher {
#[must_use]
pub fn new() -> Self {
Self(KupynaCore::new(8, 10, 7))
}
pub fn update(&mut self, data: &[u8]) {
self.0.update(data);
}
#[must_use]
pub fn finalize(self) -> [u8; 32] {
let full = self.0.finalize(32);
let mut out = [0u8; 32];
out.copy_from_slice(&full[..32]);
out
}
}
impl Default for Kupyna256Hasher {
fn default() -> Self {
Self::new()
}
}
pub struct Kupyna512Hasher(KupynaCore);
impl Kupyna512Hasher {
#[must_use]
pub fn new() -> Self {
Self(KupynaCore::new(16, 14, 11))
}
pub fn update(&mut self, data: &[u8]) {
self.0.update(data);
}
#[must_use]
pub fn finalize(self) -> [u8; 64] {
self.0.finalize(64)
}
}
impl Default for Kupyna512Hasher {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod fused_round_tests {
use super::{mix_columns, shift_bytes, sub_bytes, sub_shift_mix, MAX_COLUMNS, ROWS};
use proptest::prelude::*;
fn naive_sub_shift_mix(state: &mut [[u8; ROWS]], last_row_shift: usize) {
sub_bytes(state);
shift_bytes(state, last_row_shift);
mix_columns(state);
}
fn arb_state(columns: usize) -> impl Strategy<Value = Vec<[u8; ROWS]>> {
proptest::collection::vec(proptest::array::uniform8(any::<u8>()), columns)
}
proptest! {
#[test]
fn fused_sub_shift_mix_matches_naive_256(state in arb_state(8)) {
let mut fused = [[0u8; ROWS]; MAX_COLUMNS];
fused[..8].copy_from_slice(&state);
let mut naive = fused;
sub_shift_mix(&mut fused[..8], 7);
naive_sub_shift_mix(&mut naive[..8], 7);
prop_assert_eq!(fused, naive);
}
#[test]
fn fused_sub_shift_mix_matches_naive_512(state in arb_state(16)) {
let mut fused = [[0u8; ROWS]; MAX_COLUMNS];
fused[..16].copy_from_slice(&state);
let mut naive = fused;
sub_shift_mix(&mut fused[..16], 11);
naive_sub_shift_mix(&mut naive[..16], 11);
prop_assert_eq!(fused, naive);
}
}
}
#[cfg(test)]
mod const_shift_mix_tests {
use super::{
bytes_to_columns, bytes_to_columns_n, compress, compress_n, sub_shift_mix, sub_shift_mix_n,
ROWS,
};
use proptest::prelude::*;
fn arb_state<const COLUMNS: usize>() -> impl Strategy<Value = [[u8; ROWS]; COLUMNS]> {
proptest::collection::vec(proptest::array::uniform8(any::<u8>()), COLUMNS).prop_map(|v| {
let mut out = [[0u8; ROWS]; COLUMNS];
out.copy_from_slice(&v);
out
})
}
macro_rules! const_matches_dyn_test {
($shift_mix_test:ident, $compress_test:ident, $bytes_test:ident, $columns:literal, $rounds:literal, $last_row_shift:literal) => {
proptest! {
#[test]
fn $shift_mix_test(state in arb_state::<$columns>()) {
let mut dynamic = state;
let mut constant = state;
sub_shift_mix(&mut dynamic[..], $last_row_shift);
sub_shift_mix_n::<$columns>(&mut constant, $last_row_shift);
prop_assert_eq!(dynamic, constant);
}
#[test]
fn $compress_test(h in arb_state::<$columns>(), block in arb_state::<$columns>()) {
let mut dynamic = h;
let mut constant = h;
compress(&mut dynamic[..], &block[..], $rounds, $last_row_shift);
compress_n::<$columns, $rounds>(&mut constant, &block, $last_row_shift);
prop_assert_eq!(dynamic, constant);
}
#[test]
fn $bytes_test(bytes in proptest::collection::vec(any::<u8>(), $columns * ROWS)) {
let dynamic = bytes_to_columns(&bytes, $columns);
let constant = bytes_to_columns_n::<$columns>(&bytes);
prop_assert_eq!(&dynamic[..$columns], &constant[..]);
}
}
};
}
const_matches_dyn_test!(
sub_shift_mix_n_matches_dyn_columns8,
compress_n_matches_dyn_columns8,
bytes_to_columns_n_matches_dyn_columns8,
8,
10,
7
);
const_matches_dyn_test!(
sub_shift_mix_n_matches_dyn_columns16,
compress_n_matches_dyn_columns16,
bytes_to_columns_n_matches_dyn_columns16,
16,
14,
11
);
}