use std::time::Instant;
use ts_keys::{NodeKeyPair, NodePublicKey};
use ts_noise::ikpsk2;
use ts_packet::PacketMut;
use ts_time::Handle;
use zerocopy::IntoBytes;
use crate::{
config::Psk,
endpoint::Event,
macs::{MACReceiver, MACSender, Mac},
messages::*,
session::{ReceiveSession, TransmitSession},
time::TAI64N,
};
const PROLOGUE: &[u8] = b"WireGuard v1 zx2c4 Jason@zx2c4.com";
pub struct ReceivedHandshake {
send_id: SessionId,
noise: ikpsk2::ReceivedHandshake,
pub timestamp: TAI64N,
}
impl ReceivedHandshake {
pub fn new(
pkt: &mut HandshakeInitiation,
my_static: &NodeKeyPair,
macs: &MACReceiver,
) -> Option<ReceivedHandshake> {
if !macs.verify_macs(pkt.as_bytes()) {
return None;
};
let (noise, timestamp) =
ikpsk2::ReceivedHandshake::new(&mut pkt.noise, PROLOGUE, my_static.into())?;
Some(ReceivedHandshake {
send_id: pkt.sender_id,
noise,
timestamp: *timestamp,
})
}
pub fn respond(
self,
session_id: SessionId,
psk: &Psk,
macs: &MACSender,
now: Instant,
) -> (SessionPair, PacketMut) {
let mut response = HandshakeResponse {
sender_id: session_id,
receiver_id: self.send_id,
..Default::default()
};
let session_keys = self.noise.finish(psk, response.noise.as_mut_bytes());
let send = TransmitSession::new(session_keys.send, self.send_id, now);
let recv = ReceiveSession::new(session_keys.recv, session_id, now);
let mut pkt = PacketMut::new(size_of::<HandshakeResponse>());
response.write_to(pkt.as_mut()).unwrap();
macs.write_macs(pkt.as_mut());
(SessionPair { send, recv }, pkt)
}
pub fn peer_static(&self) -> NodePublicKey {
self.noise.peer_static_pub.to_bytes().into()
}
}
pub fn initiate_handshake(
endpoint_static: &NodeKeyPair,
peer_static: &NodePublicKey,
session_id: SessionId,
timestamp: TAI64N,
) -> (SentHandshake, HandshakeInitiation) {
let mut pkt = HandshakeInitiation {
sender_id: session_id,
..Default::default()
};
let noise = ikpsk2::SentHandshake::new(
endpoint_static.into(),
peer_static.into(),
PROLOGUE,
timestamp,
pkt.noise.as_mut_bytes(),
);
let ret = SentHandshake {
id: session_id,
noise,
};
(ret, pkt)
}
pub struct SentHandshake {
pub id: SessionId,
noise: ikpsk2::SentHandshake<TAI64N>,
}
pub struct SessionPair {
pub send: TransmitSession,
pub recv: ReceiveSession,
}
pub(crate) enum Handshake {
None,
Initiated(SentHandshake, Handle<Event>, Mac),
Responded(Box<SessionPair>),
}
impl Handshake {
pub(crate) fn is_active(&self) -> bool {
!matches!(self, Handshake::None)
}
pub(crate) fn session_id(&self) -> Option<SessionId> {
match self {
Handshake::Initiated(handshake, ..) => Some(handshake.id),
Handshake::Responded(tentative) => Some(tentative.recv.id()),
Handshake::None => None,
}
}
pub(crate) fn take_initiated(&mut self) -> Option<(SentHandshake, Handle<Event>, Mac)> {
match std::mem::replace(self, Handshake::None) {
Handshake::Initiated(sent, timeout, mac) => Some((sent, timeout, mac)),
other => {
*self = other;
None
}
}
}
pub(crate) fn respond(
&mut self,
session_id: SessionId,
handshake: ReceivedHandshake,
psk: &Psk,
cookie_sender: &MACSender,
now: Instant,
) -> PacketMut {
let (session, packet) = handshake.respond(session_id, psk, cookie_sender, now);
*self = Handshake::Responded(Box::new(session));
packet
}
pub(crate) fn finish(
&mut self,
packet: &mut HandshakeResponse,
endpoint_static: &NodeKeyPair,
psk: &Psk,
cookies: &MACReceiver,
now: Instant,
) -> Option<SessionPair> {
let (mut sent_handshake, timeout, _) = self.take_initiated()?;
if !cookies.verify_macs(packet.as_bytes()) {
return None;
};
let session_keys =
match sent_handshake
.noise
.try_finish(&mut packet.noise, endpoint_static.into(), psk)
{
Ok(session_keys) => session_keys,
Err(handshake) => {
sent_handshake.noise = handshake;
return None;
}
};
let send = TransmitSession::new(session_keys.send, packet.sender_id, now);
let recv = ReceiveSession::new(session_keys.recv, sent_handshake.id, now);
timeout.cancel();
Some(SessionPair { send, recv })
}
pub(crate) fn confirm(
&mut self,
session_id: SessionId,
mut packets: Vec<PacketMut>,
) -> Option<(SessionPair, Vec<PacketMut>)> {
let Handshake::Responded(tentative) = self else {
return None;
};
if tentative.recv.id() != session_id {
return None;
};
packets = tentative.recv.decrypt(packets);
if packets.is_empty() {
return None;
}
let Handshake::Responded(tentative) = std::mem::replace(self, Handshake::None) else {
unreachable!();
};
Some((*tentative, packets))
}
}
#[cfg(test)]
mod tests {
use ts_keys::NodeKeyPair;
use ts_time::Scheduler;
use zerocopy::TryFromBytes;
use super::*;
#[test]
fn test_handshake() {
let (a_static, b_static) = (NodeKeyPair::new(), NodeKeyPair::new());
let psk = rand::random();
let a_mac_send = MACSender::new(&b_static.public);
let a_mac_recv = MACReceiver::new(&a_static.public);
let a_session = SessionId::random(); let a_init_time = TAI64N::now();
let (a_handshake, init_pkt) =
initiate_handshake(&a_static, &b_static.public, a_session, a_init_time);
let mut init_pkt = PacketMut::from(init_pkt.as_bytes());
let handshake_mac = a_mac_send.write_macs(init_pkt.as_mut());
let mut scheduler = Scheduler::default();
let timeout = scheduler.add(
ts_time::TimeRange::new_around(Instant::now(), std::time::Duration::from_secs(1000)),
crate::Event::HandshakeTimeout(crate::config::PeerId(0)),
);
let mut a_handshake = Handshake::Initiated(a_handshake, timeout, handshake_mac);
let init_pkt = HandshakeInitiation::try_mut_from_bytes(init_pkt.as_mut())
.expect("init_pkt should be a valid handshake initiation message");
let b_mac_send = MACSender::new(&a_static.public);
let b_mac_recv = MACReceiver::new(&b_static.public);
let b_handshake = ReceivedHandshake::new(init_pkt, &b_static, &b_mac_recv)
.expect("peer B should successfully process A's handshake initiation");
assert_eq!(b_handshake.peer_static(), a_static.public);
assert_eq!(b_handshake.timestamp, a_init_time);
let b_session = SessionId::random(); let (mut b_session, mut response_pkt) =
b_handshake.respond(b_session, &psk, &b_mac_send, Instant::now());
let response_pkt = HandshakeResponse::try_mut_from_bytes(response_pkt.as_mut())
.expect("response_pkt should be a valid handshake response message");
let Some(mut a_session) =
a_handshake.finish(response_pkt, &a_static, &psk, &a_mac_recv, Instant::now())
else {
panic!("failed to process handshake response from peer B");
};
let a_plaintext = vec![PacketMut::from("xyzzy".as_bytes())];
let mut packets = a_plaintext.clone();
a_session.send.encrypt(packets.iter_mut());
let b_received = b_session.recv.decrypt(packets);
assert_eq!(b_received, a_plaintext);
let b_plaintext = vec![PacketMut::from("plover".as_bytes())];
packets = b_plaintext.clone();
b_session.send.encrypt(&mut packets);
let a_received = a_session.recv.decrypt(packets);
assert_eq!(a_received, b_plaintext);
}
}