use arrayvec::ArrayString;
use rand_core::{CryptoRng, RngCore};
use zeroize::Zeroize;
use crate::bytearray::ByteArray;
use crate::cipherstate::CipherStates;
use crate::constants::{MAX_KEY_LEN, MAX_MESSAGE_LEN, MAX_TAG_LEN};
use crate::error::{CipherResult, DhResult, HandshakeError, HandshakeResult, KemResult};
use crate::handshakepattern::HandshakePattern;
use crate::handshakestate::HandshakeStatus;
use crate::transportstate::TransportState;
use crate::KeyPair;
pub trait CryptoComponent {
fn name() -> &'static str;
}
pub trait ExtractPubKey<S, P>
where
S: ByteArray,
P: ByteArray,
{
fn pubkey(secret: &S) -> P;
}
pub trait Dh: CryptoComponent + ExtractPubKey<Self::Key, Self::PubKey> {
type Key: ByteArray;
type PubKey: ByteArray;
type Output: ByteArray;
fn genkey<R: RngCore + CryptoRng>(rng: &mut R) -> DhResult<Self::Key>;
fn dh(_: &Self::Key, _: &Self::PubKey) -> DhResult<Self::Output>;
}
pub trait Kem: CryptoComponent + ExtractPubKey<Self::SecretKey, Self::PubKey> {
type SecretKey: ByteArray;
type PubKey: ByteArray;
type Ct: ByteArray;
type Ss: ByteArray;
fn genkey<R: RngCore + CryptoRng>(
rng: &mut R,
) -> KemResult<KeyPair<Self::PubKey, Self::SecretKey>>;
fn encapsulate<R: RngCore + CryptoRng>(
pk: &[u8],
rng: &mut R,
) -> KemResult<(Self::Ct, Self::Ss)>;
fn decapsulate(ct: &[u8], sk: &[u8]) -> KemResult<Self::Ss>;
}
pub trait Hash: CryptoComponent + Default {
type Block: ByteArray;
type Output: ByteArray;
fn block_len() -> usize {
Self::Block::len()
}
fn hash_len() -> usize {
Self::Output::len()
}
fn input(&mut self, data: &[u8]);
fn result(self) -> Self::Output;
fn hash(data: &[u8]) -> Self::Output {
let mut h = Self::default();
h.input(data);
h.result()
}
fn hmac_many(key: &[u8], data: &[&[u8]]) -> Self::Output {
assert!(key.len() <= Self::block_len());
let mut ipad = Self::Block::new_with(0x36);
let mut opad = Self::Block::new_with(0x5c);
let ipad = ipad.as_mut();
let opad = opad.as_mut();
for (i, b) in key.iter().enumerate() {
ipad[i] ^= b;
opad[i] ^= b;
}
let mut hasher = Self::default();
hasher.input(ipad);
for d in data {
hasher.input(d);
}
let inner_output = hasher.result();
let mut hasher = Self::default();
hasher.input(opad);
hasher.input(inner_output.as_slice());
hasher.result()
}
fn hmac(key: &[u8], data: &[u8]) -> Self::Output {
Self::hmac_many(key, &[data])
}
fn hkdf(chaining_key: &[u8], input_key_material: &[u8]) -> (Self::Output, Self::Output) {
let temp_key = Self::hmac(chaining_key, input_key_material);
let out1 = Self::hmac(temp_key.as_slice(), &[1u8]);
let out2 = Self::hmac_many(temp_key.as_slice(), &[out1.as_slice(), &[2u8]]);
(out1, out2)
}
fn hkdf3(
chaining_key: &[u8],
input_key_material: &[u8],
) -> (Self::Output, Self::Output, Self::Output) {
let temp_key = Self::hmac(chaining_key, input_key_material);
let out1 = Self::hmac(temp_key.as_slice(), &[1u8]);
let out2 = Self::hmac_many(temp_key.as_slice(), &[out1.as_slice(), &[2u8]]);
let out3 = Self::hmac_many(temp_key.as_slice(), &[out2.as_slice(), &[3u8]]);
(out1, out2, out3)
}
}
pub trait Cipher: CryptoComponent {
type Key: ByteArray;
fn key_len() -> usize {
Self::Key::len()
}
fn tag_len() -> usize;
fn encrypt(k: &Self::Key, nonce: u64, ad: &[u8], plaintext: &[u8], out: &mut [u8]);
fn encrypt_in_place(
k: &Self::Key,
nonce: u64,
ad: &[u8],
in_out: &mut [u8],
plaintext_len: usize,
) -> usize;
fn decrypt(
k: &Self::Key,
nonce: u64,
ad: &[u8],
ciphertext: &[u8],
out: &mut [u8],
) -> CipherResult<()>;
fn decrypt_in_place(
k: &Self::Key,
nonce: u64,
ad: &[u8],
in_out: &mut [u8],
ciphertext_len: usize,
) -> CipherResult<usize>;
fn rekey(k: &Self::Key) -> Self::Key {
let mut k_new = [0u8; MAX_KEY_LEN + MAX_TAG_LEN];
let plaintext = [0u8; MAX_KEY_LEN];
Self::encrypt(
k,
u64::MAX,
&[],
&plaintext[..Self::key_len()],
&mut k_new[..Self::key_len() + Self::tag_len()],
);
let k_out = Self::Key::from_slice(&k_new[..Self::key_len()]);
k_new.zeroize();
k_out
}
}
pub(crate) trait HandshakerInternal<C, H>
where
C: Cipher,
H: Hash,
{
fn status(&self) -> HandshakeStatus;
fn set_error(&mut self);
fn write_message_impl(&mut self, payload: &[u8], out: &mut [u8]) -> HandshakeResult<usize>;
fn read_message_impl(&mut self, message: &[u8], out: &mut [u8]) -> HandshakeResult<usize>;
fn get_ciphers(&self) -> CipherStates<C>;
fn get_hash(&self) -> H::Output;
fn get_pattern(&self) -> HandshakePattern;
}
#[allow(private_bounds)] pub trait Handshaker<C, H>: HandshakerInternal<C, H>
where
C: Cipher,
H: Hash,
{
fn write_message(&mut self, payload: &[u8], out: &mut [u8]) -> HandshakeResult<usize> {
if self.status() == HandshakeStatus::Error {
return Err(HandshakeError::ErrorState);
}
if !self.is_write_turn() {
return Err(HandshakeError::InvalidState);
}
let out_len = payload.len() + self.get_next_message_overhead().unwrap();
if out_len > MAX_MESSAGE_LEN {
panic!("Maximum Noise message length exceeded");
}
if out.len() < out_len {
return Err(HandshakeError::BufferTooSmall);
}
let res = self.write_message_impl(payload, out);
if res.is_err() {
self.set_error();
}
res
}
fn read_message(&mut self, message: &[u8], out: &mut [u8]) -> HandshakeResult<usize> {
if message.len() > MAX_MESSAGE_LEN {
panic!("Maximum Noise message length exceeded");
}
if self.status() == HandshakeStatus::Error {
return Err(HandshakeError::ErrorState);
}
if self.is_write_turn() {
return Err(HandshakeError::InvalidState);
}
let out_len = message.len() - self.get_next_message_overhead().unwrap();
if !out.is_empty() && out.len() < out_len {
return Err(HandshakeError::BufferTooSmall);
}
let res = self.read_message_impl(message, out);
if res.is_err() {
self.set_error();
}
res
}
fn is_finished(&self) -> bool {
self.status() == HandshakeStatus::Ready
}
fn is_write_turn(&self) -> bool;
fn is_initiator(&self) -> bool;
fn get_next_message_overhead(&self) -> HandshakeResult<usize>;
fn build_name(pattern: &HandshakePattern) -> ArrayString<128>;
fn get_name(&self) -> ArrayString<128> {
Self::build_name(&self.get_pattern())
}
fn finalize(self) -> HandshakeResult<TransportState<C, H>>
where
Self: Sized,
{
TransportState::new(self)
}
}