use crate::config::{Cipher, Secrecy};
use crate::error::{Error, Result};
use aes_gcm::aead::inout::InOutBuf;
use aes_gcm::aead::{AeadInOut, KeyInit};
use aes_gcm::Aes256Gcm;
use chacha20poly1305::ChaCha20Poly1305;
use hkdf::Hkdf;
use rand_core::RngCore;
use sha2::{Digest, Sha256};
use std::collections::HashMap;
use zeroize::Zeroize;
pub const TAG_LEN: usize = 16;
pub const NONCE_LEN: usize = 12;
const KEY_LEN: usize = 32;
pub const HANDSHAKE_MSG_LEN: usize = 2 + 1 + 1 + 32 + 32 + 32;
const PROTOCOL_LABEL: &[u8] = b"runsync-transfer/v1";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Role {
Initiator,
Responder,
}
#[derive(Clone)]
struct Key([u8; KEY_LEN]);
impl Drop for Key {
fn drop(&mut self) {
self.0.zeroize();
}
}
#[derive(Clone)]
enum Aead {
Aes(Box<Aes256Gcm>),
ChaCha(Box<ChaCha20Poly1305>),
Passthrough,
}
impl Aead {
fn new(cipher: Cipher, key: &[u8; KEY_LEN]) -> Self {
match resolve_cipher(cipher) {
Cipher::Aes256Gcm => Aead::Aes(Box::new(Aes256Gcm::new(key.into()))),
_ => Aead::ChaCha(Box::new(ChaCha20Poly1305::new(key.into()))),
}
}
fn seal(&self, nonce: &[u8; NONCE_LEN], aad: &[u8], buf: &mut [u8]) -> Result<[u8; TAG_LEN]> {
let tag = match self {
Aead::Passthrough => return Ok([0u8; TAG_LEN]),
Aead::Aes(c) => c
.encrypt_inout_detached(nonce.into(), aad, InOutBuf::from(buf))
.map_err(|_| Error::Handshake("aes-gcm seal failed".into()))?,
Aead::ChaCha(c) => c
.encrypt_inout_detached(nonce.into(), aad, InOutBuf::from(buf))
.map_err(|_| Error::Handshake("chacha20 seal failed".into()))?,
};
let mut out = [0u8; TAG_LEN];
out.copy_from_slice(&tag);
Ok(out)
}
fn open(
&self,
nonce: &[u8; NONCE_LEN],
aad: &[u8],
buf: &mut [u8],
tag: &[u8; TAG_LEN],
) -> std::result::Result<(), ()> {
match self {
Aead::Passthrough => Ok(()),
Aead::Aes(c) => c
.decrypt_inout_detached(nonce.into(), aad, InOutBuf::from(buf), tag.into())
.map_err(|_| ()),
Aead::ChaCha(c) => c
.decrypt_inout_detached(nonce.into(), aad, InOutBuf::from(buf), tag.into())
.map_err(|_| ()),
}
}
}
fn measured_preference() -> Cipher {
static PREF: std::sync::OnceLock<Cipher> = std::sync::OnceLock::new();
*PREF.get_or_init(|| {
let key = [0x42u8; KEY_LEN];
let nonce = [0u8; NONCE_LEN];
let mut buf = vec![0u8; 64 * 1024];
let bench = |a: &Aead, buf: &mut [u8]| -> u128 {
let _ = a.seal(&nonce, b"", buf);
let t = std::time::Instant::now();
for _ in 0..8 {
let _ = a.seal(&nonce, b"", buf);
}
t.elapsed().as_nanos().max(1)
};
let aes = Aead::new(Cipher::Aes256Gcm, &key);
let cha = Aead::new(Cipher::ChaCha20Poly1305, &key);
let t_aes = bench(&aes, &mut buf);
let t_cha = bench(&cha, &mut buf);
if t_aes <= t_cha {
Cipher::Aes256Gcm
} else {
Cipher::ChaCha20Poly1305
}
})
}
fn resolve_cipher(c: Cipher) -> Cipher {
match c {
Cipher::Auto => measured_preference(),
other => other,
}
}
fn cipher_id(c: Cipher) -> u8 {
match resolve_cipher(c) {
Cipher::Aes256Gcm => 1,
_ => 2,
}
}
fn cipher_from_id(id: u8) -> Result<Cipher> {
match id {
0 => Ok(Cipher::Auto),
1 => Ok(Cipher::Aes256Gcm),
2 => Ok(Cipher::ChaCha20Poly1305),
other => Err(Error::Handshake(format!("unknown cipher id {other}"))),
}
}
pub struct Handshake {
role: Role,
secrecy: Secrecy,
cipher: Cipher,
ephemeral: x25519_dalek::StaticSecret,
our_msg: [u8; HANDSHAKE_MSG_LEN],
}
impl Handshake {
pub fn new(role: Role, secrecy: &Secrecy, cipher: Cipher) -> Self {
let ephemeral = x25519_dalek::StaticSecret::random_from_rng(rand_core::OsRng);
let eph_pub = x25519_dalek::PublicKey::from(&ephemeral);
let static_pub = match secrecy {
Secrecy::Static { our_secret, .. } => {
let s = x25519_dalek::StaticSecret::from(*our_secret);
x25519_dalek::PublicKey::from(&s).to_bytes()
}
_ => [0u8; 32],
};
let mut salt = [0u8; 32];
rand_core::OsRng.fill_bytes(&mut salt);
let mut msg = [0u8; HANDSHAKE_MSG_LEN];
msg[0..2].copy_from_slice(&crate::wire::WIRE_VERSION.to_le_bytes());
msg[2] = cipher_id(cipher);
msg[3] = match secrecy {
Secrecy::TransportOnly => 0,
Secrecy::Psk(_) => 1,
Secrecy::Static { .. } => 2,
};
msg[4..36].copy_from_slice(eph_pub.as_bytes());
msg[36..68].copy_from_slice(&static_pub);
msg[68..100].copy_from_slice(&salt);
Self {
role,
secrecy: secrecy.clone(),
cipher,
ephemeral,
our_msg: msg,
}
}
pub fn message(&self) -> &[u8; HANDSHAKE_MSG_LEN] {
&self.our_msg
}
pub fn finish(self, peer_msg: &[u8]) -> Result<SessionCrypto> {
if peer_msg.len() != HANDSHAKE_MSG_LEN {
return Err(Error::Handshake(format!(
"handshake message is {} bytes, expected {HANDSHAKE_MSG_LEN}",
peer_msg.len()
)));
}
let peer_version = u16::from_le_bytes([peer_msg[0], peer_msg[1]]);
if peer_version != crate::wire::WIRE_VERSION {
return Err(Error::Version {
peer: peer_version,
ours: crate::wire::WIRE_VERSION,
});
}
let peer_mode = peer_msg[3];
let our_mode = self.our_msg[3];
if peer_mode != our_mode {
return Err(Error::Handshake(format!(
"secrecy mode mismatch: we offered {our_mode}, peer offered {peer_mode}"
)));
}
let peer_cipher = cipher_from_id(peer_msg[2])?;
let ours = resolve_cipher(self.cipher);
let negotiated = if ours == peer_cipher {
ours
} else {
Cipher::ChaCha20Poly1305
};
let mut peer_eph = [0u8; 32];
peer_eph.copy_from_slice(&peer_msg[4..36]);
let peer_eph_pub = x25519_dalek::PublicKey::from(peer_eph);
if matches!(self.secrecy, Secrecy::TransportOnly) {
return Ok(SessionCrypto::passthrough());
}
let mut ikm: Vec<u8> = Vec::with_capacity(96);
let dh_ee = self.ephemeral.diffie_hellman(&peer_eph_pub);
if !dh_ee.was_contributory() {
return Err(Error::Handshake(
"peer sent a low-order X25519 point".into(),
));
}
ikm.extend_from_slice(dh_ee.as_bytes());
if let Secrecy::Static {
our_secret,
peer_public,
} = &self.secrecy
{
let mut claimed = [0u8; 32];
claimed.copy_from_slice(&peer_msg[36..68]);
use subtle::ConstantTimeEq;
if claimed.ct_eq(peer_public).unwrap_u8() != 1 {
return Err(Error::Handshake(
"peer static public key does not match the pinned value".into(),
));
}
let our_static = x25519_dalek::StaticSecret::from(*our_secret);
let peer_static_pub = x25519_dalek::PublicKey::from(*peer_public);
let dh_es = self.ephemeral.diffie_hellman(&peer_static_pub);
let dh_se = our_static.diffie_hellman(&peer_eph_pub);
let (first, second) = match self.role {
Role::Initiator => (dh_es, dh_se),
Role::Responder => (dh_se, dh_es),
};
ikm.extend_from_slice(first.as_bytes());
ikm.extend_from_slice(second.as_bytes());
}
let salt: [u8; 32] = match &self.secrecy {
Secrecy::Psk(k) => *k,
_ => [0u8; 32],
};
let (a, b) = match self.role {
Role::Initiator => (&self.our_msg[..], peer_msg),
Role::Responder => (peer_msg, &self.our_msg[..]),
};
let mut h = Sha256::new();
h.update(PROTOCOL_LABEL);
h.update(a);
h.update(b);
let transcript = h.finalize();
let hk = Hkdf::<Sha256>::new(Some(&salt), &ikm);
ikm.zeroize();
let mut key_i2r = [0u8; KEY_LEN];
let mut key_r2i = [0u8; KEY_LEN];
expand(&hk, b"i2r", &transcript, &mut key_i2r)?;
expand(&hk, b"r2i", &transcript, &mut key_r2i)?;
let (send, recv) = match self.role {
Role::Initiator => (key_i2r, key_r2i),
Role::Responder => (key_r2i, key_i2r),
};
Ok(SessionCrypto {
cipher: negotiated,
send: Key(send),
recv: Key(recv),
passthrough: false,
})
}
}
fn expand(hk: &Hkdf<Sha256>, label: &[u8], transcript: &[u8], out: &mut [u8]) -> Result<()> {
let mut info = Vec::with_capacity(PROTOCOL_LABEL.len() + 1 + label.len() + transcript.len());
info.extend_from_slice(PROTOCOL_LABEL);
info.push(b'/');
info.extend_from_slice(label);
info.extend_from_slice(transcript);
hk.expand(&info, out)
.map_err(|e| Error::Handshake(format!("hkdf expand: {e}")))
}
pub struct SessionCrypto {
cipher: Cipher,
send: Key,
recv: Key,
passthrough: bool,
}
impl SessionCrypto {
fn passthrough() -> Self {
Self {
cipher: Cipher::ChaCha20Poly1305,
send: Key([0u8; KEY_LEN]),
recv: Key([0u8; KEY_LEN]),
passthrough: true,
}
}
pub fn is_passthrough(&self) -> bool {
self.passthrough
}
pub fn overhead(&self) -> usize {
if self.passthrough {
0
} else {
TAG_LEN
}
}
pub fn sealer(&self) -> Sealer {
Sealer::new(self.cipher, &self.send, self.passthrough)
}
pub fn opener(&self) -> Sealer {
Sealer::new(self.cipher, &self.recv, self.passthrough)
}
}
pub struct Sealer {
cipher: Cipher,
root: Key,
passthrough: bool,
files: HashMap<u32, Aead>,
}
impl Sealer {
fn new(cipher: Cipher, root: &Key, passthrough: bool) -> Self {
Self {
cipher,
root: root.clone(),
passthrough,
files: HashMap::new(),
}
}
fn for_file(&mut self, file_id: u32) -> &Aead {
if self.files.len() > 1024 {
self.files.clear();
}
let cipher = self.cipher;
let passthrough = self.passthrough;
let root = self.root.0;
self.files.entry(file_id).or_insert_with(|| {
if passthrough {
return Aead::Passthrough;
}
let hk = Hkdf::<Sha256>::from_prk(&root).expect("32-byte prk is valid");
let mut info = Vec::with_capacity(PROTOCOL_LABEL.len() + 6 + 4);
info.extend_from_slice(PROTOCOL_LABEL);
info.extend_from_slice(b"/file");
info.extend_from_slice(&file_id.to_le_bytes());
let mut sub = [0u8; KEY_LEN];
hk.expand(&info, &mut sub)
.expect("32 bytes is under the hkdf limit");
let a = Aead::new(cipher, &sub);
sub.zeroize();
a
})
}
pub fn seal(
&mut self,
file_id: u32,
chunk: u64,
epoch: u32,
aad: &[u8],
buf: &mut [u8],
) -> Result<[u8; TAG_LEN]> {
let nonce = nonce_for(chunk, epoch);
self.for_file(file_id).seal(&nonce, aad, buf)
}
pub fn open(
&mut self,
file_id: u32,
chunk: u64,
epoch: u32,
aad: &[u8],
buf: &mut [u8],
tag: &[u8; TAG_LEN],
) -> Result<()> {
let nonce = nonce_for(chunk, epoch);
self.for_file(file_id)
.open(&nonce, aad, buf, tag)
.map_err(|_| Error::Decrypt { file_id, chunk })
}
pub fn is_passthrough(&self) -> bool {
self.passthrough
}
pub fn overhead(&self) -> usize {
if self.passthrough {
0
} else {
TAG_LEN
}
}
}
impl Clone for Sealer {
fn clone(&self) -> Self {
Self {
cipher: self.cipher,
root: self.root.clone(),
passthrough: self.passthrough,
files: HashMap::new(),
}
}
}
#[inline]
fn nonce_for(chunk: u64, epoch: u32) -> [u8; NONCE_LEN] {
let mut n = [0u8; NONCE_LEN];
n[0..8].copy_from_slice(&chunk.to_le_bytes());
n[8..12].copy_from_slice(&epoch.to_le_bytes());
n
}
pub fn random_key() -> [u8; 32] {
let mut k = [0u8; 32];
rand_core::OsRng.fill_bytes(&mut k);
k
}
pub fn generate_identity() -> ([u8; 32], [u8; 32]) {
let secret = x25519_dalek::StaticSecret::random_from_rng(rand_core::OsRng);
let public = x25519_dalek::PublicKey::from(&secret);
(secret.to_bytes(), public.to_bytes())
}
#[cfg(test)]
mod tests {
use super::*;
fn exchange(a_sec: &Secrecy, b_sec: &Secrecy) -> Result<(SessionCrypto, SessionCrypto)> {
let a = Handshake::new(Role::Initiator, a_sec, Cipher::Auto);
let b = Handshake::new(Role::Responder, b_sec, Cipher::Auto);
let am = *a.message();
let bm = *b.message();
Ok((a.finish(&bm)?, b.finish(&am)?))
}
#[test]
fn psk_handshake_produces_matched_directional_keys() {
let psk = random_key();
let (a, b) = exchange(&Secrecy::Psk(psk), &Secrecy::Psk(psk)).unwrap();
assert!(!a.is_passthrough());
let mut sealer = a.sealer();
let mut opener = b.opener();
let aad = b"header-bytes";
let mut buf = b"the payload of a chunk".to_vec();
let orig = buf.clone();
let tag = sealer.seal(7, 42, 0, aad, &mut buf).unwrap();
assert_ne!(buf, orig, "ciphertext must differ from plaintext");
opener.open(7, 42, 0, aad, &mut buf, &tag).unwrap();
assert_eq!(buf, orig);
}
#[test]
fn wrong_psk_yields_keys_that_cannot_open() {
let (a, b) = exchange(&Secrecy::Psk(random_key()), &Secrecy::Psk(random_key())).unwrap();
let mut sealer = a.sealer();
let mut opener = b.opener();
let mut buf = b"secret".to_vec();
let tag = sealer.seal(1, 0, 0, b"h", &mut buf).unwrap();
assert!(opener.open(1, 0, 0, b"h", &mut buf, &tag).is_err());
}
#[test]
fn static_identity_pinning_rejects_an_impostor() {
let (a_sec, a_pub) = generate_identity();
let (b_sec, b_pub) = generate_identity();
let (impostor_sec, _) = generate_identity();
exchange(
&Secrecy::Static {
our_secret: a_sec,
peer_public: b_pub,
},
&Secrecy::Static {
our_secret: b_sec,
peer_public: a_pub,
},
)
.expect("matching pins must succeed");
let err = exchange(
&Secrecy::Static {
our_secret: a_sec,
peer_public: b_pub,
},
&Secrecy::Static {
our_secret: impostor_sec,
peer_public: a_pub,
},
);
assert!(
err.is_err(),
"an unpinned key must not complete the handshake"
);
}
#[test]
fn tampered_aad_fails_authentication() {
let psk = random_key();
let (a, b) = exchange(&Secrecy::Psk(psk), &Secrecy::Psk(psk)).unwrap();
let mut sealer = a.sealer();
let mut opener = b.opener();
let mut buf = vec![0u8; 128];
let tag = sealer.seal(3, 9, 0, b"file=3,chunk=9", &mut buf).unwrap();
assert!(opener
.open(3, 9, 0, b"file=3,chunk=8", &mut buf, &tag)
.is_err());
}
#[test]
fn wrong_chunk_index_fails() {
let psk = random_key();
let (a, b) = exchange(&Secrecy::Psk(psk), &Secrecy::Psk(psk)).unwrap();
let mut sealer = a.sealer();
let mut opener = b.opener();
let mut buf = vec![7u8; 64];
let tag = sealer.seal(1, 100, 0, b"h", &mut buf).unwrap();
assert!(opener.open(1, 101, 0, b"h", &mut buf, &tag).is_err());
assert!(opener.open(2, 100, 0, b"h", &mut buf, &tag).is_err());
}
#[test]
fn nonces_are_unique_across_chunk_and_epoch() {
let mut seen = std::collections::HashSet::new();
for chunk in 0..1000u64 {
for epoch in 0..4u32 {
assert!(seen.insert(nonce_for(chunk, epoch)), "nonce reuse");
}
}
}
#[test]
fn mode_mismatch_is_rejected() {
let r = exchange(&Secrecy::Psk(random_key()), &Secrecy::TransportOnly);
assert!(r.is_err(), "peer must not be able to downgrade us");
}
fn exchange_with(
a_cipher: Cipher,
b_cipher: Cipher,
psk: [u8; 32],
) -> Result<(SessionCrypto, SessionCrypto)> {
let a = Handshake::new(Role::Initiator, &Secrecy::Psk(psk), a_cipher);
let b = Handshake::new(Role::Responder, &Secrecy::Psk(psk), b_cipher);
let (am, bm) = (*a.message(), *b.message());
Ok((a.finish(&bm)?, b.finish(&am)?))
}
#[test]
fn peers_with_different_cipher_preferences_still_interoperate() {
let psk = random_key();
for (a, b) in [
(Cipher::Aes256Gcm, Cipher::Aes256Gcm),
(Cipher::ChaCha20Poly1305, Cipher::ChaCha20Poly1305),
(Cipher::Aes256Gcm, Cipher::ChaCha20Poly1305),
(Cipher::ChaCha20Poly1305, Cipher::Aes256Gcm),
(Cipher::Auto, Cipher::Aes256Gcm),
(Cipher::Auto, Cipher::ChaCha20Poly1305),
(Cipher::Auto, Cipher::Auto),
] {
let (sa, sb) = exchange_with(a, b, psk).unwrap();
let mut sealer = sa.sealer();
let mut opener = sb.opener();
let plain = b"a chunk of payload bytes".to_vec();
let mut buf = plain.clone();
let tag = sealer.seal(5, 11, 0, b"aad", &mut buf).unwrap();
opener
.open(5, 11, 0, b"aad", &mut buf, &tag)
.unwrap_or_else(|e| panic!("{a:?} vs {b:?} failed to interoperate: {e}"));
assert_eq!(buf, plain, "{a:?} vs {b:?}");
let mut sealer = sb.sealer();
let mut opener = sa.opener();
let mut buf = plain.clone();
let tag = sealer.seal(5, 12, 0, b"aad", &mut buf).unwrap();
opener.open(5, 12, 0, b"aad", &mut buf, &tag).unwrap();
assert_eq!(buf, plain);
}
}
#[test]
fn mismatched_preferences_fall_back_to_chacha() {
let psk = random_key();
let (sa, _sb) = exchange_with(Cipher::Aes256Gcm, Cipher::ChaCha20Poly1305, psk).unwrap();
assert_eq!(sa.cipher, Cipher::ChaCha20Poly1305);
let (sa, _sb) = exchange_with(Cipher::Aes256Gcm, Cipher::Aes256Gcm, psk).unwrap();
assert_eq!(sa.cipher, Cipher::Aes256Gcm);
}
#[test]
fn measured_preference_is_stable_and_concrete() {
let a = measured_preference();
let b = measured_preference();
assert_eq!(a, b, "calibration must be cached, not re-run");
assert_ne!(a, Cipher::Auto, "must resolve to a concrete cipher");
}
#[test]
fn transport_only_is_passthrough() {
let (a, b) = exchange(&Secrecy::TransportOnly, &Secrecy::TransportOnly).unwrap();
assert!(a.is_passthrough() && b.is_passthrough());
assert_eq!(a.overhead(), 0);
let mut sealer = a.sealer();
let mut buf = b"plain".to_vec();
let tag = sealer.seal(0, 0, 0, b"", &mut buf).unwrap();
assert_eq!(buf, b"plain");
b.opener().open(0, 0, 0, b"", &mut buf, &tag).unwrap();
}
}