use sha3::{Sha3_256, Digest};
use crate::{Result, QsshError};
use zeroize::Zeroize;
pub struct DoubleRatchet {
root_key: [u8; 32],
chain_key_send: [u8; 32],
chain_key_recv: [u8; 32],
message_number_send: u64,
message_number_recv: u64,
prev_chain_key: Option<[u8; 32]>,
}
impl DoubleRatchet {
pub fn new(root_key: &[u8]) -> Self {
let mut hasher = Sha3_256::new();
hasher.update(b"QSSH-RATCHET-ROOT");
hasher.update(root_key);
let mut root = [0u8; 32];
root.copy_from_slice(&hasher.finalize());
let chain_key_send = Self::kdf(&root, b"QSSH-CHAIN-SEND");
let chain_key_recv = Self::kdf(&root, b"QSSH-CHAIN-RECV");
Self {
root_key: root,
chain_key_send,
chain_key_recv,
message_number_send: 0,
message_number_recv: 0,
prev_chain_key: None,
}
}
pub fn dh_ratchet(&mut self, shared_secret: &[u8]) {
self.prev_chain_key = Some(self.chain_key_recv);
let (new_root, new_recv_chain) = self.kdf_chain(&self.root_key, shared_secret);
self.root_key = new_root;
self.chain_key_recv = new_recv_chain;
self.message_number_recv = 0;
let (new_root2, new_send_chain) = self.kdf_chain(&self.root_key, b"QSSH-SEND-RATCHET");
self.root_key = new_root2;
self.chain_key_send = new_send_chain;
self.message_number_send = 0;
}
pub fn next_send_key(&mut self) -> [u8; 32] {
let (new_chain, message_key) = self.chain_step(&self.chain_key_send);
self.chain_key_send = new_chain;
self.message_number_send += 1;
message_key
}
pub fn next_recv_key(&mut self) -> [u8; 32] {
let (new_chain, message_key) = self.chain_step(&self.chain_key_recv);
self.chain_key_recv = new_chain;
self.message_number_recv += 1;
message_key
}
pub fn next_key(&mut self) -> [u8; 32] {
let send_key = self.next_send_key();
let recv_key = self.next_recv_key();
let mut combined = [0u8; 32];
for i in 0..32 {
combined[i] = send_key[i] ^ recv_key[i];
}
combined
}
pub fn skip_keys(&mut self, until: u64) -> Result<()> {
if until > self.message_number_recv + 1000 {
return Err(QsshError::Crypto("Too many keys to skip".into()));
}
while self.message_number_recv < until {
let _ = self.next_recv_key();
}
Ok(())
}
fn chain_step(&self, chain_key: &[u8; 32]) -> ([u8; 32], [u8; 32]) {
let message_key = Self::kdf(chain_key, b"QSSH-MESSAGE-KEY");
let next_chain = Self::kdf(chain_key, b"QSSH-CHAIN-KEY");
(next_chain, message_key)
}
fn kdf_chain(&self, root: &[u8; 32], input: &[u8]) -> ([u8; 32], [u8; 32]) {
let mut hasher = Sha3_256::new();
hasher.update(b"QSSH-KDF-CHAIN");
hasher.update(root);
hasher.update(input);
let output = hasher.finalize();
let mut key1 = [0u8; 32];
let mut key2 = [0u8; 32];
let mut hasher1 = Sha3_256::new();
hasher1.update(&output);
hasher1.update(&[0x01]);
key1.copy_from_slice(&hasher1.finalize());
let mut hasher2 = Sha3_256::new();
hasher2.update(&output);
hasher2.update(&[0x02]);
key2.copy_from_slice(&hasher2.finalize());
(key1, key2)
}
fn kdf(key: &[u8], label: &[u8]) -> [u8; 32] {
let mut hasher = Sha3_256::new();
hasher.update(label);
hasher.update(key);
let mut output = [0u8; 32];
output.copy_from_slice(&hasher.finalize());
output
}
}
impl Drop for DoubleRatchet {
fn drop(&mut self) {
self.root_key.zeroize();
self.chain_key_send.zeroize();
self.chain_key_recv.zeroize();
if let Some(ref mut key) = self.prev_chain_key {
key.zeroize();
}
}
}
#[derive(Clone, Debug)]
pub struct RatchetHeader {
pub ephemeral_key: Vec<u8>,
pub prev_chain_len: u32,
pub message_number: u32,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_ratchet_basic() {
let mut ratchet = DoubleRatchet::new(b"test-root-key");
let key1 = ratchet.next_key();
let key2 = ratchet.next_key();
let key3 = ratchet.next_key();
assert_ne!(key1, key2);
assert_ne!(key2, key3);
assert_ne!(key1, key3);
}
#[test]
fn test_dh_ratchet() {
let mut ratchet = DoubleRatchet::new(b"test-root-key");
let key_before = ratchet.next_send_key();
ratchet.dh_ratchet(b"new-shared-secret");
let key_after = ratchet.next_send_key();
assert_ne!(key_before, key_after);
}
}