use crate::io::Connection;
use crate::node::{Creation, LocalNode, NodeName, NodeNameError, PeerNode};
use crate::{DistributionFlags, LOWEST_DISTRIBUTION_PROTOCOL_VERSION};
use futures::io::{AsyncRead, AsyncWrite};
const PROTOCOL_VERSION: u16 = LOWEST_DISTRIBUTION_PROTOCOL_VERSION;
const NODE_NAME_VERSION: u16 = 5;
#[derive(Debug)]
pub struct ClientSideHandshake<T> {
local_node: LocalNode,
local_challenge: Challenge,
cookie: String,
connection: Connection<T>,
send_name_status: Option<HandshakeStatus>,
may_need_complement: bool,
}
impl<T> ClientSideHandshake<T>
where
T: AsyncRead + AsyncWrite + Unpin,
{
pub fn new(connection: T, local_node: LocalNode, cookie: &str) -> Self {
Self {
local_node,
local_challenge: Challenge::new(),
cookie: cookie.to_owned(),
connection: Connection::new(connection),
send_name_status: None,
may_need_complement: false,
}
}
pub async fn execute_send_name(
&mut self,
protocol_version: u16,
) -> Result<HandshakeStatus, HandshakeError> {
self.send_name(protocol_version).await?;
let status = self.recv_status().await?;
self.send_name_status = Some(status.clone());
Ok(status)
}
pub async fn execute_rest(
mut self,
do_continue: bool,
) -> Result<(T, PeerNode), HandshakeError> {
match self.send_name_status {
None => {
return Err(HandshakeError::PhaseError {
current: "ClientSideHandshake::execute_rest()",
depends_on: "ClientSideHandshake::execute_send_name()",
});
}
Some(HandshakeStatus::Nok) => return Err(HandshakeError::OngoingHandshake),
Some(HandshakeStatus::NotAllowed) => return Err(HandshakeError::NotAllowed),
Some(HandshakeStatus::Alive) => {
self.send_status(if do_continue { "true" } else { "false" })
.await?;
if !do_continue {
return Err(HandshakeError::AlreadyActive);
}
}
_ => {}
}
let (peer_node, peer_challenge) = self.recv_challenge().await?;
if self.may_need_complement && peer_node.creation.is_some() {
self.send_complement().await?;
}
self.send_challenge_reply(peer_challenge).await?;
self.recv_challenge_ack().await?;
let connection = self.connection.into_inner();
Ok((connection, peer_node))
}
async fn send_name(&mut self, protocol_version: u16) -> Result<(), HandshakeError> {
let mut writer = self.connection.handshake_message_writer();
match protocol_version {
PROTOCOL_VERSION => {
writer.write_u8(b'N')?;
writer.write_u64(self.local_node.flags.bits())?;
writer.write_u32(self.local_node.creation.get())?;
if self.local_node.flags.contains(DistributionFlags::NAME_ME) {
writer.write_u16(self.local_node.name.host().len() as u16)?;
writer.write_all(self.local_node.name.host().as_bytes())?;
} else {
writer.write_u16(self.local_node.name.len() as u16)?;
writer.write_all(self.local_node.name.to_string().as_bytes())?;
}
}
value => {
return Err(HandshakeError::UnknownProtocolVersion { value });
}
}
writer.finish().await?;
Ok(())
}
async fn recv_status(&mut self) -> Result<HandshakeStatus, HandshakeError> {
let mut reader = self.connection.handshake_message_reader().await?;
let tag = reader.read_u8().await?;
if tag != b's' {
return Err(HandshakeError::UnexpectedTag {
message: "STATUS",
tag,
});
}
let status = reader.read_bytes().await?;
let status = match status.as_slice() {
b"ok" => HandshakeStatus::Ok,
b"ok_simultaneous" => HandshakeStatus::OkSimultaneous,
b"nok" => HandshakeStatus::Nok,
b"not_allowed" => HandshakeStatus::NotAllowed,
b"alive" => HandshakeStatus::Alive,
_ => {
if status.starts_with(b"named:") {
let bytes = &status["named:".len()..];
if bytes.len() < 2 {
return Err(HandshakeError::Io(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"unexpected eof",
)));
}
let n = u16::from_be_bytes([bytes[0], bytes[1]]) as usize;
let bytes = &bytes[2..];
if bytes.len() < n + 4 {
return Err(HandshakeError::Io(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"unexpected eof",
)));
}
let name = std::str::from_utf8(&bytes[..n])
.map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
"invalid UTF-8 in node name",
)
})?
.to_owned();
let bytes = &bytes[n..];
let node_name: NodeName = name.parse()?;
let name = node_name.name().to_owned();
let creation =
Creation::new(u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]));
HandshakeStatus::Named { name, creation }
} else {
let status = String::from_utf8_lossy(&status).to_string();
return Err(HandshakeError::UnknownStatus { status });
}
}
};
reader.finish().await?;
Ok(status)
}
async fn send_status(&mut self, status: &str) -> Result<(), HandshakeError> {
let mut writer = self.connection.handshake_message_writer();
writer.write_u8(b's')?;
writer.write_all(status.as_bytes())?;
Ok(())
}
async fn recv_challenge(&mut self) -> Result<(PeerNode, Challenge), HandshakeError> {
let mut reader = self.connection.handshake_message_reader().await?;
let (node, challenge) = match reader.read_u8().await? {
b'n' => {
let version = reader.read_u16().await?;
if version != NODE_NAME_VERSION {
return Err(HandshakeError::InvalidVersionValue { value: version });
}
let flags =
DistributionFlags::from_bits_truncate(u64::from(reader.read_u32().await?));
let challenge = Challenge(reader.read_u32().await?);
let name = reader.read_string().await?.parse()?;
let node = PeerNode {
name,
flags,
creation: None,
};
(node, challenge)
}
b'N' => {
let flags = DistributionFlags::from_bits_truncate(reader.read_u64().await?);
let challenge = Challenge(reader.read_u32().await?);
let creation = Creation::new(reader.read_u32().await?);
let name = reader.read_u16_string().await?.parse()?;
let node = PeerNode {
name,
flags,
creation: Some(creation),
};
(node, challenge)
}
tag => {
return Err(HandshakeError::UnexpectedTag {
message: "CHALLENGE",
tag,
});
}
};
reader.finish().await?;
Ok((node, challenge))
}
async fn send_complement(&mut self) -> Result<(), HandshakeError> {
let mut writer = self.connection.handshake_message_writer();
writer.write_u8(b'c')?;
writer.write_u32((self.local_node.flags.bits() >> 32) as u32)?;
writer.write_u32(self.local_node.creation.get())?;
writer.finish().await?;
Ok(())
}
async fn send_challenge_reply(
&mut self,
peer_challenge: Challenge,
) -> Result<(), HandshakeError> {
let mut writer = self.connection.handshake_message_writer();
writer.write_u8(b'r')?;
writer.write_u32(self.local_challenge.0)?;
writer.write_all(&peer_challenge.digest(&self.cookie).0)?;
writer.finish().await?;
Ok(())
}
async fn recv_challenge_ack(&mut self) -> Result<(), HandshakeError> {
let mut reader = self.connection.handshake_message_reader().await?;
let tag = reader.read_u8().await?;
if tag != b'a' {
return Err(HandshakeError::UnexpectedTag {
message: "CHALLENGE_ACK",
tag,
});
}
let mut digest = [0; 16];
reader.read_exact(&mut digest).await?;
if digest != self.local_challenge.digest(&self.cookie).0 {
return Err(HandshakeError::CookieMismatch);
}
reader.finish().await?;
Ok(())
}
}
#[derive(Debug)]
pub struct ServerSideHandshake<T> {
local_node: LocalNode,
local_challenge: Challenge,
cookie: String,
connection: Connection<T>,
peer_node: Option<PeerNode>,
}
impl<T> ServerSideHandshake<T>
where
T: AsyncRead + AsyncWrite + Unpin,
{
pub fn new(connection: T, local_node: LocalNode, cookie: &str) -> Self {
Self {
local_node,
local_challenge: Challenge::new(),
cookie: cookie.to_owned(),
connection: Connection::new(connection),
peer_node: None,
}
}
pub async fn execute_recv_name(&mut self) -> Result<Option<NodeName>, HandshakeError> {
let mut reader = self.connection.handshake_message_reader().await?;
let tag = reader.read_u8().await?;
let node = match tag {
b'n' => {
let version = reader.read_u16().await?;
if version != NODE_NAME_VERSION {
return Err(HandshakeError::InvalidVersionValue { value: version });
}
let flags =
DistributionFlags::from_bits_truncate(u64::from(reader.read_u32().await?));
let name = reader.read_string().await?.parse()?;
PeerNode {
name,
flags,
creation: None,
}
}
b'N' => {
let flags = DistributionFlags::from_bits_truncate(reader.read_u64().await?);
let creation = Creation::new(reader.read_u32().await?);
let name = if flags.contains(DistributionFlags::NAME_ME) {
let host = reader.read_u16_string().await?;
NodeName::new("nonode", &host)?
} else {
reader.read_u16_string().await?.parse()?
};
PeerNode {
name,
flags,
creation: Some(creation),
}
}
_ => {
return Err(HandshakeError::UnexpectedTag {
message: "NAME",
tag,
});
}
};
reader.finish().await?;
let name = node.name.clone();
let is_dynamic = node.flags.contains(DistributionFlags::NAME_ME);
self.peer_node = Some(node);
if is_dynamic { Ok(None) } else { Ok(Some(name)) }
}
pub async fn execute_rest(
mut self,
status: HandshakeStatus,
) -> Result<(T, PeerNode), HandshakeError> {
let (peer_flags, peer_creation) = if let Some(peer) = &self.peer_node {
(peer.flags, peer.creation)
} else {
return Err(HandshakeError::PhaseError {
current: "ServerSideHandshake::execute_rest()",
depends_on: "ServerSideHandshake::execute_recv_name()",
});
};
self.send_status(status).await?;
self.send_challenge(peer_flags).await?;
if peer_flags.contains(DistributionFlags::HANDSHAKE_23) && peer_creation.is_none() {
self.recv_complement().await?;
}
let peer_challenge = self.recv_challenge_reply().await?;
self.send_challenge_ack(peer_challenge).await?;
let peer_node = self.peer_node.take().expect("unreachable");
let connection = self.connection.into_inner();
Ok((connection, peer_node))
}
async fn send_status(&mut self, status: HandshakeStatus) -> Result<(), HandshakeError> {
let mut writer = self.connection.handshake_message_writer();
writer.write_u8(b's')?;
match &status {
HandshakeStatus::Ok => writer.write_all(b"ok")?,
HandshakeStatus::OkSimultaneous => writer.write_all(b"ok_simultaneous")?,
HandshakeStatus::Nok => writer.write_all(b"nok")?,
HandshakeStatus::NotAllowed => writer.write_all(b"not_allowed")?,
HandshakeStatus::Alive => writer.write_all(b"alive")?,
HandshakeStatus::Named { name, creation } => {
let peer_node = self.peer_node.as_mut().expect("unreachable");
let node_name = NodeName::new(name, peer_node.name.host())?;
writer.write_all(b"named:")?;
writer.write_u16(node_name.len() as u16)?;
writer.write_all(node_name.to_string().as_bytes())?;
writer.write_u32(creation.get())?;
peer_node.name = node_name;
peer_node.creation = Some(*creation);
self.local_node.flags |= DistributionFlags::NAME_ME;
}
}
writer.finish().await?;
match status {
HandshakeStatus::Nok => Err(HandshakeError::OngoingHandshake),
HandshakeStatus::NotAllowed => Err(HandshakeError::NotAllowed),
_ => Ok(()),
}
}
async fn send_challenge(
&mut self,
peer_flags: DistributionFlags,
) -> Result<(), HandshakeError> {
let mut writer = self.connection.handshake_message_writer();
if peer_flags.contains(DistributionFlags::HANDSHAKE_23) {
writer.write_u8(b'N')?;
writer.write_u64(self.local_node.flags.bits())?;
writer.write_u32(self.local_challenge.0)?;
writer.write_u32(self.local_node.creation.get())?;
writer.write_u16(self.local_node.name.len() as u16)?;
writer.write_all(self.local_node.name.to_string().as_bytes())?;
} else {
writer.write_u8(b'n')?;
writer.write_u16(5)?;
writer.write_u32(self.local_node.flags.bits() as u32)?;
writer.write_u32(self.local_challenge.0)?;
writer.write_all(self.local_node.name.to_string().as_bytes())?;
}
writer.finish().await?;
Ok(())
}
async fn recv_complement(&mut self) -> Result<(), HandshakeError> {
let mut reader = self.connection.handshake_message_reader().await?;
let tag = reader.read_u8().await?;
if tag != b'c' {
return Err(HandshakeError::UnexpectedTag {
message: "send_complement",
tag,
});
}
let flags_high =
DistributionFlags::from_bits_truncate(u64::from(reader.read_u32().await?) << 32);
let creation = Creation::new(reader.read_u32().await?);
reader.finish().await?;
let peer = self.peer_node.as_mut().expect("unreachable");
peer.flags |= flags_high;
peer.creation = Some(creation);
Ok(())
}
async fn recv_challenge_reply(&mut self) -> Result<Challenge, HandshakeError> {
let mut reader = self.connection.handshake_message_reader().await?;
let tag = reader.read_u8().await?;
if tag != b'r' {
return Err(HandshakeError::UnexpectedTag {
message: "challenge_reply",
tag,
});
}
let peer_challenge = Challenge(reader.read_u32().await?);
let mut digest = Digest([0; 16]);
reader.read_exact(&mut digest.0).await?;
reader.finish().await?;
if self.local_challenge.digest(&self.cookie) != digest {
return Err(HandshakeError::CookieMismatch);
}
Ok(peer_challenge)
}
async fn send_challenge_ack(
&mut self,
peer_challenge: Challenge,
) -> Result<(), HandshakeError> {
let mut writer = self.connection.handshake_message_writer();
writer.write_u8(b'a')?;
writer.write_all(&peer_challenge.digest(&self.cookie).0)?;
writer.finish().await?;
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum HandshakeStatus {
Ok,
OkSimultaneous,
Nok,
NotAllowed,
Alive,
Named {
name: String,
creation: Creation,
},
}
#[derive(Debug, Clone, Copy)]
struct Challenge(u32);
impl Challenge {
fn new() -> Self {
Self(rand::random())
}
fn digest(self, cookie: &str) -> Digest {
Digest(md5::compute(format!("{}{}", cookie, self.0)).0)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct Digest([u8; 16]);
#[derive(Debug)]
#[non_exhaustive]
#[allow(missing_docs)]
pub enum HandshakeError {
OngoingHandshake,
NotAllowed,
AlreadyActive,
UnknownProtocolVersion { value: u16 },
UnknownStatus { status: String },
UnexpectedTag { message: &'static str, tag: u8 },
CookieMismatch,
InvalidVersionValue { value: u16 },
PhaseError {
current: &'static str,
depends_on: &'static str,
},
NodeNameError(NodeNameError),
Io(std::io::Error),
}
impl std::fmt::Display for HandshakeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::OngoingHandshake => {
write!(f, "peer already has an ongoing handshake with this node")
}
Self::NotAllowed => write!(
f,
"the connection is disallowed for some (unspecified) security reason"
),
Self::AlreadyActive => write!(f, "a connection to the node is already active"),
Self::UnknownProtocolVersion { value } => {
write!(f, "unknown distribution protocol version {value:?}")
}
Self::UnknownStatus { status } => {
write!(f, "received an unknown status {status:?}")
}
Self::UnexpectedTag { message, tag } => {
write!(f, "received an unexpected tag {tag} for {message:?}")
}
Self::CookieMismatch => write!(f, "cookie mismatch"),
Self::InvalidVersionValue { value } => {
write!(
f,
"the 'version' value of an old 'send_name' message must be {NODE_NAME_VERSION}, but got {value}"
)
}
Self::PhaseError {
current,
depends_on,
} => {
write!(
f,
"{current:?} was unexpectedly executed before {depends_on:?}"
)
}
Self::NodeNameError(error) => write!(f, "{error}"),
Self::Io(error) => write!(f, "{error}"),
}
}
}
impl std::error::Error for HandshakeError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::NodeNameError(error) => Some(error),
Self::Io(error) => Some(error),
_ => None,
}
}
}
impl From<std::io::Error> for HandshakeError {
fn from(value: std::io::Error) -> Self {
Self::Io(value)
}
}
impl From<NodeNameError> for HandshakeError {
fn from(value: NodeNameError) -> Self {
Self::NodeNameError(value)
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::StreamExt;
#[test]
fn client_side_handshake_works() {
let peer_name = "client_side_handshake_works";
smol::block_on(async {
let erl_node = crate::tests::TestErlangNode::new(peer_name)
.await
.expect("failed to run a test erlang node");
let peer_entry = crate::tests::epmd_client()
.await
.get_node(peer_name)
.await
.expect("failed to get node");
let peer_entry = peer_entry.expect("no such node");
let connection = smol::net::TcpStream::connect(("localhost", peer_entry.port))
.await
.expect("failed to connect");
let local_node = LocalNode::new("foo@localhost".parse().unwrap(), Creation::random());
let mut handshake =
ClientSideHandshake::new(connection, local_node, crate::tests::COOKIE);
let status = handshake
.execute_send_name(crate::LOWEST_DISTRIBUTION_PROTOCOL_VERSION)
.await
.expect("failed to execute send name");
assert_eq!(status, HandshakeStatus::Ok);
let (_, peer_node) = handshake
.execute_rest(true)
.await
.expect("failed to execute handshake");
assert_eq!(peer_entry.name, peer_node.name.name());
std::mem::drop(erl_node);
});
}
#[test]
fn server_side_handshake_works() {
smol::block_on(async {
let listener = smol::net::TcpListener::bind("0.0.0.0:0").await.unwrap();
let listening_port = listener.local_addr().unwrap().port();
let (tx, rx) = futures::channel::oneshot::channel();
let connection = smol::net::TcpStream::connect(("localhost", listening_port))
.await
.unwrap();
smol::spawn(async move {
let local_node =
LocalNode::new("foo@localhost".parse().unwrap(), Creation::random());
let mut handshake =
ClientSideHandshake::new(connection, local_node, crate::tests::COOKIE);
let _status = handshake
.execute_send_name(crate::LOWEST_DISTRIBUTION_PROTOCOL_VERSION)
.await
.unwrap();
let (connection, _) = handshake.execute_rest(true).await.unwrap();
let _ = tx.send(connection);
})
.detach();
let mut incoming = listener.incoming();
if let Some(connection) = incoming.next().await {
let local_node =
LocalNode::new("bar@localhost".parse().unwrap(), Creation::random());
let mut handshake =
ServerSideHandshake::new(connection.unwrap(), local_node, crate::tests::COOKIE);
let peer_name = handshake.execute_recv_name().await.unwrap();
assert!(peer_name.is_some());
handshake.execute_rest(HandshakeStatus::Ok).await.unwrap();
}
let _ = rx.await;
})
}
}