use std::fmt::Write;
use crate::shortint::parameters::test_params::{
TEST_PARAM_MESSAGE_1_CARRY_1_KS_PBS_GAUSSIAN_2M128,
TEST_PARAM_MESSAGE_2_CARRY_2_KS_PBS_TUNIFORM_2M128,
TEST_PARAM_MESSAGE_3_CARRY_3_KS_PBS_GAUSSIAN_2M128,
};
use crate::shortint::prelude::*;
use crate::transciphering::ciphers::kreyvium::{
KreyviumFheState, KreyviumPlainKey, KreyviumPlainState, KreyviumState,
};
use crate::transciphering::ciphers::pack_bits_lsb_first;
use crate::transciphering::ciphers::shift_register::ShiftRegister;
use crate::transciphering::{
apply_keystream, FheKeyStream, InsufficientKeystream, StreamCipher, Transcipherer,
};
fn get_hexadecimal_string_from_bytes(bytes: &[u8]) -> String {
let mut hexadecimal = String::new();
for test in bytes {
write!(hexadecimal, "{test:02X?}").expect("writing to a String is infallible");
}
hexadecimal
}
fn hex_to_bytes_16(hex: &str) -> [u8; 16] {
let mut bytes = [0u8; 16];
for i in (0..hex.len()).step_by(2) {
bytes[i >> 1] = u8::from_str_radix(&hex[i..i + 2], 16).unwrap();
}
bytes
}
fn decrypt_keystream_to_bytes(fhe: &FheKeyStream, client_key: &ClientKey) -> Vec<u8> {
let bits: Vec<bool> = fhe.iter().map(|ct| client_key.decrypt(ct) != 0).collect();
let mut bytes = vec![0u8; bits.len().div_ceil(8)];
pack_bits_lsb_first(&bits, &mut bytes);
bytes
}
struct KreyviumTestVector {
key: &'static str,
iv: &'static str,
expected: &'static str,
}
const KREYVIUM_TV_ZERO: KreyviumTestVector = KreyviumTestVector {
key: "00000000000000000000000000000000",
iv: "00000000000000000000000000000000",
expected: "26DCF1F4BC0F1922",
};
const KREYVIUM_TV_KEY_BIT: KreyviumTestVector = KreyviumTestVector {
key: "01000000000000000000000000000000",
iv: "00000000000000000000000000000000",
expected: "4FD421D4DA3D2C8A",
};
const KREYVIUM_TV_IV_BIT: KreyviumTestVector = KreyviumTestVector {
key: "00000000000000000000000000000000",
iv: "01000000000000000000000000000000",
expected: "C9217BA0D762ACA1",
};
const KREYVIUM_TV_STANDARD: KreyviumTestVector = KreyviumTestVector {
key: "0053A6F94C9FF24598EB000000000000",
iv: "0D74DB42A91077DE45AC000000000000",
expected: "D1F0303482061111",
};
#[test]
fn kreyvium_test_plain() {
let cases = [
KREYVIUM_TV_ZERO,
KREYVIUM_TV_KEY_BIT,
KREYVIUM_TV_IV_BIT,
KREYVIUM_TV_STANDARD,
];
for KreyviumTestVector { key, iv, expected } in cases {
let key_bytes = hex_to_bytes_16(key);
let iv_bytes = hex_to_bytes_16(iv);
let mut kreyvium = KreyviumPlainState::new(key_bytes, iv_bytes);
let vec = kreyvium.next_keystream_bits(64).unwrap();
let hexadecimal = get_hexadecimal_string_from_bytes(&vec);
assert_eq!(hexadecimal, expected, "key={key} iv={iv}");
}
}
#[test]
fn kreyvium_plain_encrypt_decrypt_round_trip() {
use rand::{Rng, SeedableRng};
let seed: u64 = rand::thread_rng().gen();
println!("kreyvium_encrypt_decrypt_round_trip seed: {seed}");
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
let key: [bool; 128] = std::array::from_fn(|_| rng.gen());
let iv: [bool; 128] = std::array::from_fn(|_| rng.gen());
let message: Vec<u8> = (0..37).map(|_| rng.gen()).collect();
let mut enc_stream = KreyviumPlainState::new(key, iv);
let encrypted = enc_stream.encrypt(&message).unwrap();
assert_eq!(encrypted.bytes().len(), message.len());
let mut dec_stream = KreyviumPlainState::new(key, iv);
let decrypted = dec_stream.decrypt(&encrypted).unwrap();
assert_eq!(decrypted, message);
}
fn kreyvium_fhe_keystream_known_answer(params: ClassicPBSParameters) {
let (client_key, server_key) = gen_keys(params);
let KreyviumTestVector { key, iv, expected } = KREYVIUM_TV_STANDARD;
let key_bytes = hex_to_bytes_16(key);
let iv_bytes = hex_to_bytes_16(iv);
let cipher_key = KreyviumPlainKey::from(key_bytes).encrypt(&client_key);
let mut kreyvium = KreyviumFheState::new(cipher_key, iv_bytes, &server_key);
let cts = kreyvium.next_keystream_bits(&server_key, 64).unwrap();
let bytes = decrypt_keystream_to_bytes(&cts, &client_key);
let hexadecimal = get_hexadecimal_string_from_bytes(&bytes);
assert_eq!(expected, hexadecimal);
}
#[test]
fn kreyvium_test_fhe() {
kreyvium_fhe_keystream_known_answer(TEST_PARAM_MESSAGE_2_CARRY_2_KS_PBS_TUNIFORM_2M128);
}
#[test]
fn kreyvium_test_fhe_1_1() {
kreyvium_fhe_keystream_known_answer(TEST_PARAM_MESSAGE_1_CARRY_1_KS_PBS_GAUSSIAN_2M128);
}
#[test]
fn kreyvium_test_fhe_3_3() {
kreyvium_fhe_keystream_known_answer(TEST_PARAM_MESSAGE_3_CARRY_3_KS_PBS_GAUSSIAN_2M128);
}
#[test]
fn kreyvium_test_round_trip() {
use rand::{Rng, SeedableRng};
let seed: u64 = rand::thread_rng().gen();
println!("kreyvium_test_round_trip seed: {seed}");
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
let (client_key, server_key) = gen_keys(TEST_PARAM_MESSAGE_2_CARRY_2_KS_PBS_TUNIFORM_2M128);
const N_ITER: usize = 2;
for iter in 0..N_ITER {
let key_bits: [bool; 128] = std::array::from_fn(|_| rng.gen());
let iv_bits: [bool; 128] = std::array::from_fn(|_| rng.gen());
let input: u64 = rng.gen();
let mut sym_stream = KreyviumPlainState::new(key_bits, iv_bits);
let sym_cipher = sym_stream.encrypt(&input.to_le_bytes()).unwrap();
let cipher_key = KreyviumPlainKey::from(key_bits).encrypt(&client_key);
let mut fhe_stream = KreyviumFheState::new(cipher_key, iv_bits, &server_key);
let keystream = fhe_stream.next_keystream_bits(&server_key, 64).unwrap();
let chunks = apply_keystream(&server_key, &keystream, &sym_cipher);
let result: u64 = chunks
.iter()
.enumerate()
.map(|(i, chunk)| (client_key.decrypt(chunk) & 0b11) << (2 * i))
.sum();
assert_eq!(
result, input,
"round-trip mismatch (seed={seed}, iter={iter})"
);
}
}
#[test]
fn kreyvium_seek_plain() {
use rand::{Rng, SeedableRng};
let seed: u64 = rand::thread_rng().gen();
println!("kreyvium_seek_plain seed: {seed}");
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
let key: [bool; 128] = std::array::from_fn(|_| rng.gen());
let iv: [bool; 128] = std::array::from_fn(|_| rng.gen());
let state_at_0 = KreyviumPlainState::new(key, iv);
let mut state_at_64 = state_at_0.clone();
let head_keystream = state_at_64.next_keystream_bits(64).unwrap();
assert_eq!(state_at_64.current_counter(), 64);
let mid_keystream = state_at_64.clone().next_keystream_bits(64).unwrap();
let mut s = KreyviumPlainState::new(key, iv);
s.seek(64);
assert_eq!(s.current_counter(), 64);
assert_eq!(s, state_at_64, "forward-seek state mismatch");
s.next_keystream_bits(128).unwrap(); assert_eq!(s.current_counter(), 192);
s.seek(64);
assert_eq!(s.current_counter(), 64);
assert_eq!(s, state_at_64, "backward-seek state mismatch");
let mid_again = s.next_keystream_bits(64).unwrap();
assert_eq!(mid_again, mid_keystream);
s.seek(0);
assert_eq!(s.current_counter(), 0);
assert_eq!(s, state_at_0, "seek-to-0 state mismatch");
let head_again = s.next_keystream_bits(64).unwrap();
assert_eq!(head_again, head_keystream);
}
#[test]
fn kreyvium_seek_fhe() {
use rand::{Rng, SeedableRng};
let seed: u64 = rand::thread_rng().gen();
println!("kreyvium_seek_fhe seed: {seed}");
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
let (client_key, server_key) = gen_keys(TEST_PARAM_MESSAGE_2_CARRY_2_KS_PBS_TUNIFORM_2M128);
let key_bits: [bool; 128] = std::array::from_fn(|_| rng.gen());
let iv_bits: [bool; 128] = std::array::from_fn(|_| rng.gen());
let mut ref_stream = KreyviumPlainState::new(key_bits, iv_bits);
let ref_keystream = ref_stream.next_keystream_bits(192).unwrap();
let cipher_key = KreyviumPlainKey::from(key_bits).encrypt(&client_key);
let mut k = KreyviumFheState::new(cipher_key, iv_bits, &server_key);
assert_eq!(k.current_counter(), 0);
k.seek(&server_key, 64);
assert_eq!(k.current_counter(), 64);
let mid = k.next_keystream_bits(&server_key, 64).unwrap();
assert_eq!(k.current_counter(), 128);
assert_eq!(
decrypt_keystream_to_bytes(&mid, &client_key),
ref_keystream[8..16]
);
k.seek(&server_key, 64);
assert_eq!(k.current_counter(), 64);
let mid_again = k.next_keystream_bits(&server_key, 64).unwrap();
assert_eq!(k.current_counter(), 128);
assert_eq!(
decrypt_keystream_to_bytes(&mid_again, &client_key),
ref_keystream[8..16]
);
}
fn filled_register<const N: usize, T: Clone>(filler: &T) -> ShiftRegister<N, T> {
ShiftRegister::new(Box::new(std::array::from_fn(|_| filler.clone())))
}
fn state_at_counter<T: Clone>(filler: &T, counter: u64) -> KreyviumState<T> {
KreyviumState {
a: filled_register(filler),
b: filled_register(filler),
c: filled_register(filler),
k: filled_register(filler),
iv: filled_register(filler),
counter,
}
}
#[test]
fn kreyvium_exhaustion_at_counter_range_end() {
let (_client_key, server_key) = gen_keys(TEST_PARAM_MESSAGE_2_CARRY_2_KS_PBS_TUNIFORM_2M128);
let n_bits: usize = 8;
let start = u64::MAX - n_bits as u64;
let mut plain_stream: KreyviumPlainState = state_at_counter(&false, start);
let mut fhe_stream: KreyviumFheState = state_at_counter(&server_key.create_trivial(0), start);
assert!(matches!(
plain_stream.next_keystream_bits(n_bits + 1),
Err(InsufficientKeystream)
));
assert_eq!(plain_stream.current_counter(), start);
assert!(matches!(
fhe_stream.next_keystream_bits(&server_key, n_bits + 1),
Err(InsufficientKeystream)
));
assert_eq!(fhe_stream.current_counter(), start);
assert!(plain_stream.next_keystream_bits(n_bits).is_ok());
assert_eq!(plain_stream.current_counter(), u64::MAX);
assert!(fhe_stream.next_keystream_bits(&server_key, n_bits).is_ok());
assert_eq!(fhe_stream.current_counter(), u64::MAX);
}
#[test]
fn kreyvium_fhe_key_encrypt_decrypt_round_trip() {
let (cks, _sks) = gen_keys(TEST_PARAM_MESSAGE_2_CARRY_2_KS_PBS_TUNIFORM_2M128);
let iv = [0x5Au8; 16];
for bytes in [
[0x00u8; 16],
[0xFFu8; 16],
[
0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0xfe, 0xdc, 0xba, 0x98, 0x76, 0x54,
0x32, 0x10,
],
] {
let plain = KreyviumPlainKey::from(bytes);
let recovered = plain.encrypt(&cks).decrypt(&cks);
let value = [0u8; 8];
let from_plain = KreyviumPlainState::new(plain, iv).encrypt(&value).unwrap();
let from_recovered = KreyviumPlainState::new(recovered, iv)
.encrypt(&value)
.unwrap();
assert_eq!(
from_recovered.bytes(),
from_plain.bytes(),
"Kreyvium key did not survive the encrypt/decrypt round trip for {bytes:02x?}"
);
}
}