use dusa_collection_utils::core::errors::{ErrorArrayItem, Errors};
use dusa_collection_utils::core::logger::LogLevel;
use dusa_collection_utils::{log, core::version::Version};
use tokio::io::{self, AsyncReadExt, AsyncWriteExt};
use crate::network::utils::{comms_version, get_local_ip};
use crate::protocol::{
flags::{ConnectionParams, MsgType},
handshake::{NoiseIdentity, ctx_from_handshake_result, perform_handshake_initiator, perform_handshake_responder},
message::{ConnectionCtx, ProtocolMessage},
proto::Proto,
status::ProtocolStatus,
};
pub async fn establish_connection_initiator<STREAM>(
stream: &mut STREAM,
remote_static_pubkey: &[u8; 32],
params: ConnectionParams,
) -> io::Result<ConnectionCtx>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
{
let (noise, conn_id, params) =
perform_handshake_initiator(stream, remote_static_pubkey, params).await?;
Ok(ctx_from_handshake_result(noise, conn_id, params))
}
pub async fn establish_connection_responder<STREAM>(
stream: &mut STREAM,
identity: &NoiseIdentity,
) -> io::Result<ConnectionCtx>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
{
let (noise, conn_id, params) = perform_handshake_responder(stream, identity).await?;
Ok(ctx_from_handshake_result(noise, conn_id, params))
}
fn io_err_to_item(err: io::Error) -> ErrorArrayItem {
ErrorArrayItem::new(Errors::Network, err.to_string())
}
pub async fn send_message<STREAM, DATA, RESPONSE>(
stream: &mut STREAM,
data: DATA,
proto: Proto,
conn: &mut ConnectionCtx,
) -> Result<RESPONSE, ErrorArrayItem>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
DATA: serde::de::DeserializeOwned + std::fmt::Debug + serde::Serialize + Clone + Unpin,
RESPONSE: serde::de::DeserializeOwned + std::fmt::Debug + serde::Serialize + Clone + Unpin,
{
let params = conn.params;
send_message_with_params(stream, params, data, proto, conn).await
}
pub async fn send_message_with_params<STREAM, DATA, RESPONSE>(
mut stream: &mut STREAM,
flags: ConnectionParams,
data: DATA,
proto: Proto,
conn: &mut ConnectionCtx,
) -> Result<RESPONSE, ErrorArrayItem>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
DATA: serde::de::DeserializeOwned + std::fmt::Debug + serde::Serialize + Clone + Unpin,
RESPONSE: serde::de::DeserializeOwned + std::fmt::Debug + serde::Serialize + Clone + Unpin,
{
let insecure = conn.insecure;
let mut message: ProtocolMessage<DATA> =
ProtocolMessage::new(flags, MsgType::Data, data.clone()).map_err(io_err_to_item)?;
match proto {
Proto::TCP => message.header.origin_address = get_local_ip().octets(),
Proto::UNIX => message.header.origin_address = [0, 0, 0, 0],
};
log!(LogLevel::Trace, "message serialized for sending");
message
.write_to(&mut stream, proto, Some(&mut *conn))
.await
.map_err(io_err_to_item)?;
log!(LogLevel::Trace, "Message sent over {proto}");
let response = ProtocolMessage::<RESPONSE>::read_from(&mut stream, Some(&mut *conn))
.await
.map_err(io_err_to_item)?;
let response_status: ProtocolStatus = response.status();
let response_params: ConnectionParams =
ConnectionParams::from_bits_truncate(response.header.reserved);
let response_version: Version = Version::decode(response.header.version);
let in_band = Version::compare_versions(&comms_version(), &response_version);
if !insecure && !in_band {
return Err(ProtocolStatus::NOTINBAND.to_error_item());
}
if response_status.has_flag(ProtocolStatus::SIDEGRADE) {
log!(LogLevel::Debug, "SideGrade requested");
if insecure {
return Box::pin(send_message_with_params::<STREAM, DATA, RESPONSE>(
stream,
response_params,
data,
proto,
conn,
))
.await;
} else {
log!(LogLevel::Info, "Sidegrade not allowed dropping connections");
stream.shutdown().await.map_err(io_err_to_item)?;
return Err(ProtocolStatus::REFUSED.to_error_item());
}
}
if response_status.is_error() {
return Err(response_status.to_error_item());
}
log!(LogLevel::Trace, "Received response: {:?}", response);
Ok(response.payload)
}
async fn receive_message_impl<STREAM, RESPONSE>(
stream: &mut STREAM,
auto_reply: bool,
proto: Proto,
mut conn: Option<&mut ConnectionCtx>,
required: Option<ConnectionParams>,
force: bool,
) -> io::Result<ProtocolMessage<RESPONSE>>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
RESPONSE: serde::de::DeserializeOwned + std::fmt::Debug + serde::Serialize + Clone,
{
let message: ProtocolMessage<RESPONSE> =
match ProtocolMessage::read_from(stream, conn.as_deref_mut()).await {
Ok(message) => message,
Err(err) => {
log!(LogLevel::Error, "Deserialization error: {}", err);
let _ = send_empty_err(stream, proto).await;
return Err(err);
}
};
if proto == Proto::TCP {
stream.flush().await?;
}
log!(LogLevel::Debug, "Received message: {:?}", message);
if let Some(want) = required {
if !message.flags().contains(want) {
log!(
LogLevel::Warn,
"peer used unexpected connection params: expected {:?}, got {:?}",
want,
message.flags()
);
let should_negotiate = force || conn.as_ref().map(|c| c.insecure).unwrap_or(false);
if should_negotiate {
send_sidegrade(stream, proto, want).await?;
let retried: ProtocolMessage<RESPONSE> =
ProtocolMessage::read_from(stream, conn.as_deref_mut()).await?;
if !retried.flags().contains(want) {
log!(
LogLevel::Warn,
"peer's SIDEGRADE retry still didn't satisfy required params: expected {:?}, got {:?}",
want,
retried.flags()
);
} else if let Some(ctx) = conn.as_deref_mut() {
ctx.params = want;
}
if auto_reply {
send_empty_ok(stream, proto).await?;
}
return Ok(retried);
}
}
}
if auto_reply {
send_empty_ok(stream, proto).await?;
}
Ok(message)
}
pub async fn receive_message<STREAM, RESPONSE>(
stream: &mut STREAM,
auto_reply: bool,
proto: Proto,
mut conn: Option<&mut ConnectionCtx>,
) -> io::Result<ProtocolMessage<RESPONSE>>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
RESPONSE: serde::de::DeserializeOwned + std::fmt::Debug + serde::Serialize + Clone,
{
let required = conn.as_deref().map(|c| c.params);
receive_message_impl(stream, auto_reply, proto, conn.as_deref_mut(), required, false).await
}
pub async fn receive_message_with_required_params<STREAM, RESPONSE>(
stream: &mut STREAM,
auto_reply: bool,
proto: Proto,
conn: Option<&mut ConnectionCtx>,
required: ConnectionParams,
) -> io::Result<ProtocolMessage<RESPONSE>>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
RESPONSE: serde::de::DeserializeOwned + std::fmt::Debug + serde::Serialize + Clone,
{
receive_message_impl(stream, auto_reply, proto, conn, Some(required), true).await
}
pub async fn send_sidegrade<S>(stream: &mut S, proto: Proto, desired: ConnectionParams) -> io::Result<()>
where
S: AsyncWriteExt + Unpin,
{
let mut message: ProtocolMessage<()> =
ProtocolMessage::new(ConnectionParams::NONE, MsgType::Data, ())?;
message.header.status = ProtocolStatus::SIDEGRADE.bits();
message.header.reserved = desired.bits();
message.write_to(stream, proto, None).await
}
pub async fn send_empty_err<S>(stream: &mut S, proto: Proto) -> Result<(), io::Error>
where
S: AsyncWriteExt + Unpin,
{
let mut message: ProtocolMessage<()> = ProtocolMessage::new(ConnectionParams::NONE, MsgType::Data, ())?;
message.header.status = ProtocolStatus::ERROR.bits();
message.write_to(stream, proto, None).await
}
pub async fn send_empty_ok<S>(stream: &mut S, proto: Proto) -> Result<(), io::Error>
where
S: AsyncWriteExt + Unpin,
{
let mut message: ProtocolMessage<()> = ProtocolMessage::new(ConnectionParams::NONE, MsgType::Data, ())?;
message.header.status = ProtocolStatus::OK.bits();
message.write_to(stream, proto, None).await
}
pub async fn send_data<S>(stream: &mut S, data: Vec<u8>, proto: Proto) -> Result<(), io::Error>
where
S: AsyncWriteExt + Unpin,
{
stream.write_all(&data).await?;
if proto == Proto::TCP {
stream.flush().await?
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::handshake::NoiseIdentity;
async fn establish_pair(
params: ConnectionParams,
) -> (
tokio::io::DuplexStream,
tokio::io::DuplexStream,
ConnectionCtx,
ConnectionCtx,
) {
let identity = NoiseIdentity::generate().unwrap();
let remote_pub = identity.public_key();
let (mut client, mut server) = tokio::io::duplex(8192);
let (client_ctx, server_ctx) = tokio::join!(
establish_connection_initiator(&mut client, &remote_pub, params),
establish_connection_responder(&mut server, &identity),
);
(client, server, client_ctx.unwrap(), server_ctx.unwrap())
}
#[tokio::test]
async fn transparent_sidegrade_when_insecure() {
let (mut client, mut server, mut client_ctx, mut server_ctx) =
establish_pair(ConnectionParams::ENCRYPTED | ConnectionParams::INSECURE).await;
assert!(client_ctx.insecure);
assert!(server_ctx.insecure);
let client_fut = send_message_with_params::<_, Vec<u8>, ()>(
&mut client,
ConnectionParams::NONE, b"hello".to_vec(),
Proto::TCP,
&mut client_ctx,
);
let server_fut =
receive_message::<_, Vec<u8>>(&mut server, true, Proto::TCP, Some(&mut server_ctx));
let (client_result, server_result) = tokio::join!(client_fut, server_fut);
let received = server_result.unwrap();
assert_eq!(received.payload, b"hello".to_vec());
assert!(received.flags().contains(server_ctx.params));
client_result.unwrap();
}
#[tokio::test]
async fn no_sidegrade_when_not_insecure() {
let (mut client, mut server, mut client_ctx, mut server_ctx) =
establish_pair(ConnectionParams::ENCRYPTED).await;
assert!(!client_ctx.insecure);
assert!(!server_ctx.insecure);
let client_fut = send_message_with_params::<_, Vec<u8>, ()>(
&mut client,
ConnectionParams::NONE,
b"hello".to_vec(),
Proto::TCP,
&mut client_ctx,
);
let server_fut =
receive_message::<_, Vec<u8>>(&mut server, true, Proto::TCP, Some(&mut server_ctx));
let (client_result, server_result) = tokio::join!(client_fut, server_fut);
let received = server_result.unwrap();
assert_eq!(received.flags(), ConnectionParams::NONE);
assert_ne!(received.flags(), server_ctx.params);
client_result.unwrap();
}
#[tokio::test]
async fn manual_required_params_upgrades_the_baseline() {
let (mut client, mut server, mut client_ctx, mut server_ctx) =
establish_pair(ConnectionParams::INSECURE).await;
let client_fut = send_message::<_, Vec<u8>, ()>(
&mut client,
b"sensitive".to_vec(), Proto::TCP,
&mut client_ctx,
);
let server_fut = receive_message_with_required_params::<_, Vec<u8>>(
&mut server,
true,
Proto::TCP,
Some(&mut server_ctx),
ConnectionParams::ENCRYPTED | ConnectionParams::INSECURE,
);
let (client_result, server_result) = tokio::join!(client_fut, server_fut);
let received = server_result.unwrap();
assert_eq!(received.payload, b"sensitive".to_vec());
assert!(received.flags().contains(ConnectionParams::ENCRYPTED));
assert_eq!(
server_ctx.params,
ConnectionParams::ENCRYPTED | ConnectionParams::INSECURE
);
client_result.unwrap();
}
}