use arrayvec::ArrayString;
pub 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::symmetricstate::SymmetricState;
use crate::transportstate::TransportState;
use crate::KeyPair;
pub trait CryptoComponent: Clone {
fn name() -> &'static str;
}
pub trait Rng: RngCore + CryptoRng + Default + Clone {}
impl<T: RngCore + CryptoRng + Default + Clone> Rng for T {}
pub trait Dh: CryptoComponent {
type PrivateKey: ByteArray;
type PubKey: ByteArray;
type Output: ByteArray;
fn genkey_rng<R: Rng>(rng: &mut R) -> DhResult<KeyPair<Self::PubKey, Self::PrivateKey>>;
#[cfg(feature = "getrandom")]
fn genkey() -> DhResult<KeyPair<Self::PubKey, Self::PrivateKey>> {
Self::genkey_rng(&mut crate::crypto::rng::DefaultRng)
}
fn pubkey(k: &Self::PrivateKey) -> Self::PubKey;
fn dh(_: &Self::PrivateKey, _: &Self::PubKey) -> DhResult<Self::Output>;
}
pub trait Kem: CryptoComponent {
type SecretKey: ByteArray;
type PubKey: ByteArray;
type Ct: ByteArray;
type Ss: ByteArray;
fn genkey_rng<R: Rng>(rng: &mut R) -> KemResult<KeyPair<Self::PubKey, Self::SecretKey>>;
#[cfg(feature = "getrandom")]
fn genkey() -> KemResult<KeyPair<Self::PubKey, Self::SecretKey>> {
Self::genkey_rng(&mut crate::crypto::rng::DefaultRng)
}
fn encapsulate<R: Rng>(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],
) -> CipherResult<()>;
fn encrypt_in_place(
k: &Self::Key,
nonce: u64,
ad: &[u8],
in_out: &mut [u8],
plaintext_len: usize,
) -> CipherResult<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) -> CipherResult<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();
Ok(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) -> CipherResult<CipherStates<C>>;
fn get_hash(&self) -> H::Output;
fn mix_hash(&mut self, data: &[u8]);
fn mix_key_and_hash(&mut self, data: &[u8]);
fn get_pattern(&self) -> HandshakePattern;
}
#[allow(private_bounds)] pub trait Handshaker<C, H>: HandshakerInternal<C, H>
where
C: Cipher,
H: Hash,
{
type E;
type S;
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 overhead = self.get_next_message_overhead().unwrap();
if message.len() < overhead {
return Err(HandshakeError::InvalidMessage);
}
let out_len = message.len() - overhead;
if 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 push_psk(&mut self, psk: &[u8]);
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 get_remote_static(&self) -> Option<Self::S>;
fn get_remote_ephemeral(&self) -> Option<Self::E>;
fn finalize(self) -> HandshakeResult<TransportState<C, H>>
where
Self: Sized,
{
TransportState::new(self)
}
fn get_state(&self) -> SymmetricState<C, H>;
fn get_state_mut(&mut self) -> &mut SymmetricState<C, H>;
}