use ic_core::Zeroize;
pub const BLOCK_LEN: usize = 16;
fn mul(a: &[u8; BLOCK_LEN], b: &[u8; BLOCK_LEN]) -> [u8; BLOCK_LEN] {
let a0 = u64::from_le_bytes(a[0..8].try_into().unwrap());
let a1 = u64::from_le_bytes(a[8..16].try_into().unwrap());
let b0 = u64::from_le_bytes(b[0..8].try_into().unwrap());
let b1 = u64::from_le_bytes(b[8..16].try_into().unwrap());
let mut z = [0u64; 4];
let mut v = [a0, a1, 0u64, 0u64];
for i in 0..128 {
let bit = if i < 64 {
(b0 >> i) & 1
} else {
(b1 >> (i - 64)) & 1
};
let mask = 0u64.wrapping_sub(bit);
for (zw, vw) in z.iter_mut().zip(v.iter()) {
*zw ^= *vw & mask;
}
v[3] = (v[3] << 1) | (v[2] >> 63);
v[2] = (v[2] << 1) | (v[1] >> 63);
v[1] = (v[1] << 1) | (v[0] >> 63);
v[0] <<= 1;
}
reduce(z)
}
const X_INVERSE_HIGH: u64 = (1 << 63) | (1 << 62) | (1 << 61) | (1 << 56);
fn reduce(mut acc: [u64; 4]) -> [u8; BLOCK_LEN] {
for _ in 0..128 {
let carry = acc[0] & 1;
acc[0] = (acc[0] >> 1) | (acc[1] << 63);
acc[1] = (acc[1] >> 1) | (acc[2] << 63);
acc[2] = (acc[2] >> 1) | (acc[3] << 63);
acc[3] >>= 1;
acc[1] ^= 0u64.wrapping_sub(carry) & X_INVERSE_HIGH;
}
debug_assert_eq!(acc[2], 0, "reduction left a high word set");
debug_assert_eq!(acc[3], 0, "reduction left a high word set");
let mut out = [0u8; BLOCK_LEN];
out[0..8].copy_from_slice(&acc[0].to_le_bytes());
out[8..16].copy_from_slice(&acc[1].to_le_bytes());
out
}
pub struct Polyval {
h: [u8; BLOCK_LEN],
acc: [u8; BLOCK_LEN],
}
impl Polyval {
pub fn new(h: [u8; BLOCK_LEN]) -> Self {
Self {
h,
acc: [0u8; BLOCK_LEN],
}
}
pub fn update_block(&mut self, block: &[u8; BLOCK_LEN]) {
for (a, b) in self.acc.iter_mut().zip(block.iter()) {
*a ^= *b;
}
self.acc = mul(&self.acc, &self.h);
}
pub fn update_padded(&mut self, data: &[u8]) {
for chunk in data.chunks(BLOCK_LEN) {
let mut block = [0u8; BLOCK_LEN];
block[..chunk.len()].copy_from_slice(chunk);
self.update_block(&block);
}
}
pub fn finish(self) -> [u8; BLOCK_LEN] {
self.acc
}
}
impl Drop for Polyval {
fn drop(&mut self) {
self.h.zeroize();
self.acc.zeroize();
}
}
#[cfg(test)]
mod tests {
use super::*;
fn byte_reverse(x: &[u8; 16]) -> [u8; 16] {
let mut out = [0u8; 16];
for (i, b) in x.iter().enumerate() {
out[15 - i] = *b;
}
out
}
fn mul_x_ghash(x: &[u8; 16]) -> [u8; 16] {
let mut out = [0u8; 16];
let mut carry = 0u8;
for i in 0..16 {
let next = x[i] & 1;
out[i] = (x[i] >> 1) | (carry << 7);
carry = next;
}
if carry != 0 {
out[0] ^= 0xe1;
}
out
}
fn polyval_via_ghash(h: &[u8; 16], blocks: &[[u8; 16]]) -> [u8; 16] {
let ghash_h = mul_x_ghash(&byte_reverse(h));
let mut acc = [0u8; 16];
for block in blocks {
let reversed = byte_reverse(block);
for (a, b) in acc.iter_mut().zip(reversed.iter()) {
*a ^= *b;
}
crate::gcm::portable_ghash_mul(&mut acc, &ghash_h);
}
byte_reverse(&acc)
}
#[test]
fn polyval_matches_the_ghash_construction() {
let mut state = 0x243f_6a88_85a3_08d3u64;
let mut next = || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for count in 0..12usize {
let mut h = [0u8; 16];
h[0..8].copy_from_slice(&next().to_le_bytes());
h[8..16].copy_from_slice(&next().to_le_bytes());
let mut blocks = Vec::new();
for _ in 0..count {
let mut b = [0u8; 16];
b[0..8].copy_from_slice(&next().to_le_bytes());
b[8..16].copy_from_slice(&next().to_le_bytes());
blocks.push(b);
}
let mut p = Polyval::new(h);
for b in &blocks {
p.update_block(b);
}
let direct = p.finish();
let via = polyval_via_ghash(&h, &blocks);
assert_eq!(
direct, via,
"POLYVAL disagreed with the GHASH construction at {count} blocks"
);
}
}
#[test]
fn polyval_handles_degenerate_inputs() {
let zero = [0u8; 16];
let mut one = [0u8; 16];
one[0] = 1;
for h in [zero, one, [0xffu8; 16]] {
for blocks in [vec![], vec![zero], vec![one, zero], vec![[0xffu8; 16]; 3]] {
let mut p = Polyval::new(h);
for b in &blocks {
p.update_block(b);
}
assert_eq!(p.finish(), polyval_via_ghash(&h, &blocks));
}
}
}
#[test]
fn padding_matches_explicit_blocks() {
let h = [0x42u8; 16];
let data = b"a partial final block";
let mut padded = Polyval::new(h);
padded.update_padded(data);
let mut explicit = Polyval::new(h);
for chunk in data.chunks(16) {
let mut block = [0u8; 16];
block[..chunk.len()].copy_from_slice(chunk);
explicit.update_block(&block);
}
assert_eq!(padded.finish(), explicit.finish());
}
}