use super::Protocol;
use super::cipher::Cipher;
use super::cipher_state::{CipherState, MAX_MESSAGE_LEN, rekey_key};
use super::error::HandshakeError;
use super::session_id::SessionId;
use super::transport::Transport;
use std::marker::PhantomData;
use std::num::NonZeroU64;
pub const MAX_EPOCH_JUMP: u64 = 2;
impl<Proto: Protocol> Transport<Proto> {
pub fn into_datagram(self) -> (DatagramSend<Proto>, DatagramRecv<Proto>) {
let (send, recv, session_id) = self.into_cipher_states();
let sender = DatagramSend {
cipher: send,
session_id: session_id.clone(),
epoch: None,
};
let receiver = DatagramRecv {
keys: RecvKeys::Plain(recv),
session_id,
};
(sender, receiver)
}
pub fn into_datagram_with_epoch(
self,
epoch_size: NonZeroU64,
) -> (DatagramSend<Proto>, DatagramRecv<Proto>) {
let (send, mut recv, session_id) = self.into_cipher_states();
let sender = DatagramSend {
cipher: send,
session_id: session_id.clone(),
epoch: Some(SendEpoch {
epoch_size,
key_epoch: 0,
}),
};
let keys = match recv.take_key() {
Some(base) => RecvKeys::Ratchet(RecvRatchet::new(epoch_size, base)),
None => RecvKeys::Plain(recv),
};
let receiver = DatagramRecv { keys, session_id };
(sender, receiver)
}
}
struct SendEpoch {
epoch_size: NonZeroU64,
key_epoch: u64,
}
pub struct DatagramSend<Proto: Protocol> {
cipher: CipherState<Proto::Cipher>,
session_id: SessionId,
epoch: Option<SendEpoch>,
}
impl<Proto: Protocol> DatagramSend<Proto> {
pub fn encrypt_next(
&mut self,
ad: &[u8],
plaintext: &[u8],
output: &mut [u8],
) -> Result<(u64, usize), HandshakeError> {
if let Some(epoch) = self.epoch.as_mut() {
if self.cipher.nonce() == u64::MAX {
return Err(HandshakeError::NonceOverflow);
}
let target = self.cipher.nonce() / epoch.epoch_size.get();
while epoch.key_epoch < target {
self.cipher.rekey()?;
epoch.key_epoch += 1;
}
}
self.cipher.encrypt_next_with_ad(ad, plaintext, output)
}
pub fn next_counter(&self) -> u64 {
self.cipher.nonce()
}
pub fn session_id(&self) -> &SessionId {
&self.session_id
}
#[cfg(test)]
pub(crate) fn set_counter_for_test(&mut self, n: u64) {
self.cipher.set_nonce_for_test(n);
}
}
pub struct DatagramRecv<Proto: Protocol> {
keys: RecvKeys<Proto::Cipher>,
session_id: SessionId,
}
enum RecvKeys<Ci: Cipher> {
Plain(CipherState<Ci>),
Ratchet(RecvRatchet<Ci>),
}
struct RecvRatchet<Ci: Cipher> {
epoch_size: NonZeroU64,
current_epoch: u64,
current_key: Ci::Key,
prev_key: Option<Ci::Key>,
_cipher: PhantomData<fn() -> Ci>,
}
impl<Ci: Cipher> RecvRatchet<Ci> {
fn new(epoch_size: NonZeroU64, base_key: Ci::Key) -> Self {
Self {
epoch_size,
current_epoch: 0,
current_key: base_key,
prev_key: None,
_cipher: PhantomData,
}
}
fn decrypt_at(
&mut self,
counter: u64,
ad: &[u8],
ciphertext: &[u8],
output: &mut [u8],
) -> Result<usize, HandshakeError> {
if ciphertext.len() > MAX_MESSAGE_LEN {
return Err(HandshakeError::MessageTooLong {
len: ciphertext.len(),
});
}
if counter == u64::MAX {
return Err(HandshakeError::NonceOverflow);
}
let msg_epoch = counter / self.epoch_size.get();
if msg_epoch == self.current_epoch {
return Ci::decrypt(&self.current_key, counter, ad, ciphertext, output);
}
if msg_epoch < self.current_epoch {
if msg_epoch + 1 == self.current_epoch
&& let Some(prev) = self.prev_key.as_ref()
{
return Ci::decrypt(prev, counter, ad, ciphertext, output);
}
return Err(HandshakeError::DecryptionFailed);
}
let steps = msg_epoch - self.current_epoch;
if steps > MAX_EPOCH_JUMP {
return Err(HandshakeError::DecryptionFailed);
}
let mut cur_cand = rekey_key::<Ci>(&self.current_key)?;
let mut prev_cand: Option<Ci::Key> = None;
for _ in 1..steps {
let next = rekey_key::<Ci>(&cur_cand)?;
prev_cand = Some(core::mem::replace(&mut cur_cand, next));
}
match Ci::decrypt(&cur_cand, counter, ad, ciphertext, output) {
Ok(len) => {
let old_current = core::mem::replace(&mut self.current_key, cur_cand);
self.prev_key = Some(prev_cand.unwrap_or(old_current));
self.current_epoch = msg_epoch;
Ok(len)
}
Err(err) => Err(err),
}
}
}
impl<Proto: Protocol> DatagramRecv<Proto> {
pub fn decrypt_at(
&mut self,
counter: u64,
ad: &[u8],
ciphertext: &[u8],
output: &mut [u8],
) -> Result<usize, HandshakeError> {
match &mut self.keys {
RecvKeys::Plain(cipher) => cipher.decrypt_at(counter, ad, ciphertext, output),
RecvKeys::Ratchet(ratchet) => ratchet.decrypt_at(counter, ad, ciphertext, output),
}
}
pub fn session_id(&self) -> &SessionId {
&self.session_id
}
}