use std::{cmp, io, time::Duration};
use bytes::BytesMut;
use futures::{SinkExt, StreamExt};
use prost::{DecodeError, Message};
use tokio::{
io::{AsyncRead, AsyncWrite},
time,
};
use tracing::{Instrument, Level, debug, error, span, trace, warn};
use crate::{framing::CanonicalFraming, message::MessageExt, proto, protocol::rpc::error::HandshakeRejectReason};
const LOG_TARGET: &str = "comms::rpc::handshake";
pub(super) const SUPPORTED_RPC_VERSIONS: &[u32] = &[0];
pub(super) const MAX_HANDSHAKE_FRAME_SIZE: usize = 1024;
#[derive(Debug, thiserror::Error)]
pub enum RpcHandshakeError {
#[error("Failed to decode message: {0}")]
DecodeError(#[from] DecodeError),
#[error("IO Error: {0}")]
Io(io::Error),
#[error("The client does not support any RPC protocol version supported by this node")]
ClientNoSupportedVersion,
#[error("Remote peer unexpectedly closed the RPC connection")]
ServerClosedRequest,
#[error("RPC handshake timed out")]
TimedOut,
#[error("RPC handshake was explicitly rejected: {0}")]
Rejected(#[from] HandshakeRejectReason),
#[error("The client connection is closed")]
ClientClosed,
#[error("Handshake frame was larger than the {max} byte limit")]
FrameTooLarge { max: usize },
}
impl From<io::Error> for RpcHandshakeError {
fn from(err: io::Error) -> Self {
if err
.get_ref()
.is_some_and(|inner| inner.is::<tokio_util::codec::LengthDelimitedCodecError>())
{
return RpcHandshakeError::FrameTooLarge {
max: MAX_HANDSHAKE_FRAME_SIZE,
};
}
RpcHandshakeError::Io(err)
}
}
fn send_timed_out() -> RpcHandshakeError {
RpcHandshakeError::Io(io::Error::new(
io::ErrorKind::TimedOut,
"timed out sending a handshake frame",
))
}
pub struct Handshake<'a, T> {
framed: &'a mut CanonicalFraming<T>,
timeout: Option<Duration>,
}
impl<'a, T> Handshake<'a, T>
where T: AsyncRead + AsyncWrite + Unpin
{
pub fn new(framed: &'a mut CanonicalFraming<T>) -> Self {
Self { framed, timeout: None }
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub async fn perform_server_handshake(&mut self) -> Result<u32, RpcHandshakeError> {
let version = self.receive_client_handshake().await?;
self.accept_version(version).await?;
Ok(version)
}
pub async fn receive_client_handshake(&mut self) -> Result<u32, RpcHandshakeError> {
match self.recv_next_frame().await {
Ok(Some(Ok(msg))) => {
let msg = proto::rpc::RpcSession::decode(&mut msg.freeze())?;
let version = SUPPORTED_RPC_VERSIONS
.iter()
.find(|v| msg.supported_versions.contains(v));
if let Some(version) = version {
debug!(target: LOG_TARGET, "Local server accepted version: {}", version);
return Ok(*version);
}
let span = span!(Level::INFO, "rpc::server::handshake::send_rejection");
self.reject_with_reason(HandshakeRejectReason::UnsupportedVersion)
.instrument(span)
.await?;
Err(RpcHandshakeError::ClientNoSupportedVersion)
},
Ok(Some(Err(err))) => {
trace!(target: LOG_TARGET, "Error during handshake: {err}");
Err(err.into())
},
Ok(None) => {
trace!(target: LOG_TARGET, "Error during handshake, client closed connection");
Err(RpcHandshakeError::ClientClosed)
},
Err(_) => {
trace!(target: LOG_TARGET, "Error during handshake, timed out");
Err(RpcHandshakeError::TimedOut)
},
}
}
pub async fn accept_version(&mut self, version: u32) -> Result<(), RpcHandshakeError> {
let reply = proto::rpc::RpcSessionReply {
session_result: Some(proto::rpc::rpc_session_reply::SessionResult::AcceptedVersion(version)),
..Default::default()
};
let span = span!(Level::INFO, "rpc::server::handshake::send_accept_version_reply");
self.send_bounded(reply.to_encoded_bytes().into())
.instrument(span)
.await
}
pub async fn reject_with_reason(&mut self, reject_reason: HandshakeRejectReason) -> Result<(), RpcHandshakeError> {
trace!(target: LOG_TARGET, "Rejecting handshake because {}", reject_reason);
let reply = proto::rpc::RpcSessionReply {
session_result: Some(proto::rpc::rpc_session_reply::SessionResult::Rejected(true)),
reject_reason: reject_reason.as_i32(),
};
self.send_bounded(reply.to_encoded_bytes().into()).await?;
match self.timeout {
Some(timeout) => time::timeout(timeout, self.framed.close())
.await
.map_err(|_| send_timed_out())??,
None => self.framed.close().await?,
}
Ok(())
}
async fn send_bounded(&mut self, frame: bytes::Bytes) -> Result<(), RpcHandshakeError> {
match self.timeout {
Some(timeout) => time::timeout(timeout, self.framed.send(frame))
.await
.map_err(|_| send_timed_out())??,
None => self.framed.send(frame).await?,
}
Ok(())
}
pub async fn perform_client_handshake(&mut self) -> Result<(), RpcHandshakeError> {
let msg = proto::rpc::RpcSession {
supported_versions: SUPPORTED_RPC_VERSIONS.to_vec(),
};
let payload = msg.to_encoded_bytes();
debug!(target: LOG_TARGET, "Sending client handshake ({} bytes)", payload.len());
if let Err(err) = self.framed.send(payload.into()).await {
warn!(
target: LOG_TARGET,
"IO error when sending new session handshake to peer: {}", err
);
}
self.framed.flush().await?;
match self.recv_next_frame().await {
Ok(Some(Ok(msg))) => {
let msg = proto::rpc::RpcSessionReply::decode(&mut msg.freeze())?;
let version = msg.result()?;
debug!(target: LOG_TARGET, "Remote server accepted version {}", version);
Ok(())
},
Ok(Some(Err(err))) => {
error!(target: LOG_TARGET, "Error during handshake: {}", err);
Err(err.into())
},
Ok(None) => {
warn!(target: LOG_TARGET, "Error during handshake, server closed connection");
Err(RpcHandshakeError::ServerClosedRequest)
},
Err(_) => {
error!(target: LOG_TARGET, "Error during handshake, timed out");
Err(RpcHandshakeError::TimedOut)
},
}
}
async fn recv_next_frame(&mut self) -> Result<Option<Result<BytesMut, io::Error>>, time::error::Elapsed> {
let previous_limit = self.framed.codec().max_frame_length();
self.framed
.codec_mut()
.set_max_frame_length(cmp::min(previous_limit, MAX_HANDSHAKE_FRAME_SIZE));
let result = match self.timeout {
Some(timeout) => time::timeout(timeout, self.framed.next()).await,
None => Ok(self.framed.next().await),
};
self.framed.codec_mut().set_max_frame_length(previous_limit);
result
}
}