use chacha20poly1305::aead::{Aead, KeyInit, Payload};
use chacha20poly1305::{ChaCha20Poly1305, Key, Nonce};
use sha2::{Digest, Sha256};
use anyhow::{Context, Result};
use x25519_dalek::{EphemeralSecret, PublicKey};
pub const HANDSHAKE_REQUEST: u8 = 0x20;
pub const HANDSHAKE_RESPONSE: u8 = 0x21;
const NONCE_LEN: usize = 12;
const TAG_LEN: usize = 16;
pub struct KeyPair {
secret: EphemeralSecret,
public: [u8; 32],
}
impl std::fmt::Debug for KeyPair {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KeyPair")
.field("public", &hex(&self.public))
.finish_non_exhaustive()
}
}
impl KeyPair {
pub fn generate() -> Self {
let secret = EphemeralSecret::random_from_rng(rand_core::OsRng);
let public = PublicKey::from(&secret).to_bytes();
Self { secret, public }
}
pub fn public(&self) -> [u8; 32] {
self.public
}
}
pub struct Direction {
cipher: ChaCha20Poly1305,
counter: u64,
}
impl std::fmt::Debug for Direction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Direction")
.field("counter", &self.counter)
.finish_non_exhaustive()
}
}
impl Direction {
pub fn new(key: [u8; 32]) -> Self {
Self {
cipher: ChaCha20Poly1305::new(Key::from_slice(&key)),
counter: 0,
}
}
fn nonce(counter: u64) -> [u8; NONCE_LEN] {
let mut n = [0u8; NONCE_LEN];
n[4..].copy_from_slice(&counter.to_be_bytes());
n
}
pub fn seal(&mut self, plaintext: &[u8]) -> Result<Vec<u8>> {
let counter = self.counter;
self.counter = self
.counter
.checked_add(1)
.context("frame counter overflowed; rekey required")?;
let raw = Self::nonce(counter);
let out = self
.cipher
.encrypt(
Nonce::from_slice(&raw),
Payload {
msg: plaintext,
aad: &aad_for(counter),
},
)
.map_err(|_| anyhow::anyhow!("frame encryption failed"))?;
let mut framed = Vec::with_capacity(4 + out.len());
framed.extend_from_slice(&((counter + 1) as u32).to_le_bytes());
framed.extend_from_slice(&out);
Ok(framed)
}
pub fn open(&mut self, framed: &[u8]) -> Result<Vec<u8>> {
anyhow::ensure!(
framed.len() > 4 + TAG_LEN,
"sealed frame is too short: {} bytes",
framed.len()
);
let wire = u32::from_le_bytes([framed[0], framed[1], framed[2], framed[3]]) as u64;
anyhow::ensure!(
wire > self.counter,
"sealed frame is a replay or out of order: counter {wire} <= {}",
self.counter
);
self.counter = wire;
let raw = Self::nonce(wire - 1);
self.cipher
.decrypt(
Nonce::from_slice(&raw),
Payload {
msg: &framed[4..],
aad: &aad_for(wire - 1),
},
)
.map_err(|_| anyhow::anyhow!("sealed frame failed authentication"))
}
pub fn rekey(&mut self, key: [u8; 32]) {
self.cipher = ChaCha20Poly1305::new(Key::from_slice(&key));
self.counter = 0;
}
}
pub struct Session {
pub host_to_viewer: Direction,
pub viewer_to_host: Direction,
}
impl std::fmt::Debug for Session {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Session")
.field("host_to_viewer", &self.host_to_viewer)
.field("viewer_to_host", &self.viewer_to_host)
.finish()
}
}
const TOKEN_DOMAIN: &[u8] = b"pcc/e2e/v1";
fn aad_for(counter: u64) -> [u8; 8] {
counter.to_le_bytes()
}
fn derive(
shared: &[u8; 32],
host_public: &[u8; 32],
viewer_public: &[u8; 32],
) -> ([u8; 32], [u8; 32]) {
let h2v = kdf(shared, host_public, viewer_public, b"pcc/h2v");
let v2h = kdf(shared, host_public, viewer_public, b"pcc/v2h");
(h2v, v2h)
}
fn kdf(shared: &[u8; 32], hp: &[u8; 32], vp: &[u8; 32], label: &[u8]) -> [u8; 32] {
let mut h = Sha256::new();
h.update(TOKEN_DOMAIN);
h.update(label);
h.update(shared);
h.update(hp);
h.update(vp);
h.finalize().into()
}
#[derive(Debug, Clone)]
pub struct Offer {
pub public: [u8; 32],
pub proof: [u8; 32],
}
#[derive(Debug, Clone)]
pub struct Reply {
pub public: [u8; 32],
pub proof: [u8; 32],
}
fn proof_of(token: &str, public: &[u8; 32]) -> [u8; 32] {
let mut h = Sha256::new();
h.update(TOKEN_DOMAIN);
h.update(b"proof");
h.update(public);
h.update(token.as_bytes());
h.finalize().into()
}
pub fn offer(mine: &KeyPair, token: &str) -> Offer {
let public = mine.public();
Offer {
public,
proof: proof_of(token, &public),
}
}
pub fn reply(mine: &KeyPair, token: &str) -> Reply {
let public = mine.public();
Reply {
public,
proof: proof_of(token, &public),
}
}
pub fn accept(offer: &Offer, mine: KeyPair, token: &str) -> Result<Session> {
anyhow::ensure!(
constant_time_eq(&proof_of(token, &offer.public), &offer.proof),
"handshake proof failed: the viewer does not hold this session token"
);
let their_public = PublicKey::from(offer.public);
let our_public = mine.public();
let shared_bytes = mine.secret.diffie_hellman(&their_public).to_bytes();
let (h2v, v2h) = derive(&shared_bytes, &offer.public, &our_public);
Ok(Session {
host_to_viewer: Direction::new(h2v),
viewer_to_host: Direction::new(v2h),
})
}
pub fn complete(
mine: KeyPair,
reply: &Reply,
offer_public: &[u8; 32],
token: &str,
) -> Result<Session> {
anyhow::ensure!(
constant_time_eq(&proof_of(token, &reply.public), &reply.proof),
"handshake proof failed: the sharer does not hold this session token"
);
let their_public = PublicKey::from(reply.public);
let shared_bytes = mine.secret.diffie_hellman(&their_public).to_bytes();
let (h2v, v2h) = derive(&shared_bytes, offer_public, &reply.public);
Ok(Session {
host_to_viewer: Direction::new(h2v),
viewer_to_host: Direction::new(v2h),
})
}
fn constant_time_eq(a: &[u8; 32], b: &[u8; 32]) -> bool {
let mut diff = 0u8;
for i in 0..32 {
diff |= a[i] ^ b[i];
}
diff == 0
}
pub fn hex(bytes: &[u8]) -> String {
bytes.iter().map(|b| format!("{b:02x}")).collect()
}
pub async fn host_handshake(
sink: &mut Box<dyn crate::network::MessageSink>,
source: &mut Box<dyn crate::network::MessageSource>,
keys: KeyPair,
token: &str,
) -> Result<Session> {
let offer = match source.recv().await? {
crate::network::Message::E2eOffer { public, proof } => Offer { public, proof },
other => anyhow::bail!(
"expected an encryption offer, got {:?}; is the viewer the same build?",
other.rev()
),
};
let reply = reply(&keys, token);
sink.send(&crate::network::Message::E2eReply {
public: reply.public,
proof: reply.proof,
})
.await?;
accept(&offer, keys, token)
}
pub async fn viewer_handshake(
sink: &mut Box<dyn crate::network::MessageSink>,
source: &mut Box<dyn crate::network::MessageSource>,
keys: KeyPair,
token: &str,
) -> Result<Session> {
let offer = offer(&keys, token);
sink.send(&crate::network::Message::E2eOffer {
public: offer.public,
proof: offer.proof,
})
.await?;
let reply = match source.recv().await? {
crate::network::Message::E2eReply { public, proof } => Reply { public, proof },
crate::network::Message::Error(text) => anyhow::bail!("{text}"),
other => anyhow::bail!(
"expected an encryption reply, got {:?}; is the sharer the same build?",
other.rev()
),
};
complete(keys, &reply, &offer.public, token)
}
#[cfg(test)]
mod tests {
use super::*;
fn handshake(token: &str) -> (Session, Session) {
let viewer_keys = KeyPair::generate();
let host_keys = KeyPair::generate();
let offer = offer(&viewer_keys, token);
let reply = reply(&host_keys, token);
let viewer = complete(viewer_keys, &reply, &offer.public, token).unwrap();
let host = accept(&offer, host_keys, token).unwrap();
(host, viewer)
}
#[test]
fn both_sides_derive_the_same_keys() {
let (mut host, mut viewer) = handshake("TOKEN12345678");
let plaintext = b"a frame of pixels";
let sealed = host.host_to_viewer.seal(plaintext).unwrap();
assert_eq!(viewer.host_to_viewer.open(&sealed).unwrap(), plaintext);
}
#[test]
fn the_viewer_can_seal_control_traffic_the_host_opens() {
let (mut host, mut viewer) = handshake("TOKEN12345678");
let sealed = viewer.viewer_to_host.seal(b"RequestKeyframe").unwrap();
assert_eq!(
host.viewer_to_host.open(&sealed).unwrap(),
b"RequestKeyframe"
);
}
#[test]
fn a_wrong_token_fails_the_handshake() {
let viewer_keys = KeyPair::generate();
let host_keys = KeyPair::generate();
let offer = offer(&viewer_keys, "TOKEN12345678");
let reply = reply(&host_keys, "TOKEN12345678");
let err = complete(viewer_keys, &reply, &offer.public, "WRONGTOKEN999")
.unwrap_err()
.to_string();
assert!(err.contains("handshake proof failed"), "unhelpful: {err}");
}
#[test]
fn a_relay_substituting_a_reply_cannot_open_any_frame() {
let viewer_keys = KeyPair::generate();
let host_keys = KeyPair::generate();
let forged_keys = KeyPair::generate();
let offer = offer(&viewer_keys, "TOKEN12345678");
let real = reply(&host_keys, "TOKEN12345678");
let forged = reply(&forged_keys, "TOKEN12345678");
assert_ne!(real.public, forged.public);
let mut viewer = complete(viewer_keys, &forged, &offer.public, "TOKEN12345678").unwrap();
let mut host = accept(&offer, host_keys, "TOKEN12345678").unwrap();
let sealed = host.host_to_viewer.seal(b"pixels").unwrap();
assert!(viewer.host_to_viewer.open(&sealed).is_err());
}
#[test]
fn the_two_directions_use_different_keys() {
let (mut s, _) = handshake("TOKEN12345678");
let sealed = s.host_to_viewer.seal(b"secret").unwrap();
assert!(s.viewer_to_host.open(&sealed).is_err());
}
#[test]
fn a_tampered_counter_fails_authentication() {
let mut d = Direction::new([7u8; 32]);
let mut sealed = d.seal(b"pixels").unwrap();
sealed[3] = sealed[3].wrapping_add(1);
let mut other = Direction::new([7u8; 32]);
let err = other.open(&sealed).unwrap_err().to_string();
assert!(
err.contains("authentication") || err.contains("replay"),
"unhelpful: {err}"
);
}
#[test]
fn a_replayed_frame_is_refused() {
let mut sender = Direction::new([1u8; 32]);
let mut receiver = Direction::new([1u8; 32]);
let sealed = sender.seal(b"pixels").unwrap();
receiver.open(&sealed).unwrap();
let err = receiver.open(&sealed).unwrap_err().to_string();
assert!(err.contains("replay"), "unhelpful: {err}");
}
#[test]
fn an_out_of_order_frame_is_refused() {
let mut sender = Direction::new([1u8; 32]);
let mut receiver = Direction::new([1u8; 32]);
let first = sender.seal(b"one").unwrap();
let second = sender.seal(b"two").unwrap();
receiver.open(&second).unwrap();
assert!(receiver.open(&first).is_err());
}
#[test]
fn a_wrong_key_cannot_open_a_frame() {
let mut sender = Direction::new([1u8; 32]);
let sealed = sender.seal(b"pixels").unwrap();
let mut wrong = Direction::new([2u8; 32]);
assert!(wrong.open(&sealed).is_err());
}
#[test]
fn a_tampered_ciphertext_is_refused() {
let mut sender = Direction::new([1u8; 32]);
let mut sealed = sender.seal(b"pixels").unwrap();
let n = sealed.len();
sealed[n - 1] ^= 0xFF;
let mut receiver = Direction::new([1u8; 32]);
assert!(receiver.open(&sealed).is_err());
}
#[test]
fn a_truncated_frame_is_refused_before_any_decryption() {
let mut sender = Direction::new([1u8; 32]);
let sealed = sender.seal(b"pixels").unwrap();
let mut receiver = Direction::new([1u8; 32]);
let err = receiver.open(&sealed[..8]).unwrap_err().to_string();
assert!(err.contains("too short"), "unhelpful: {err}");
}
#[test]
fn a_rekey_breaks_the_old_key_and_resets_the_counter() {
let mut sender = Direction::new([1u8; 32]);
sender.seal(b"before").unwrap();
let mut receiver = Direction::new([1u8; 32]);
receiver.open(&sender.seal(b"before").unwrap()).unwrap();
sender.rekey([9u8; 32]);
receiver.rekey([9u8; 32]);
let mut old = Direction::new([1u8; 32]);
let sealed = sender.seal(b"after").unwrap();
assert!(old.open(&sealed).is_err());
assert_eq!(receiver.open(&sealed).unwrap(), b"after");
}
#[test]
fn debug_output_never_contains_key_material() {
let d = Direction::new([0xAB; 32]);
let text = format!("{d:?}");
assert!(
!text.contains("ab"),
"debug output leaked key bytes: {text}"
);
}
}