use rustls::pki_types::ServerName;
use std::{
io,
pin::Pin,
sync::Arc,
task::{Context, Poll},
time::{Duration, Instant},
};
use tokio::{
io::{AsyncRead, AsyncWrite, BufReader, ReadBuf, ReadHalf, WriteHalf},
net::TcpStream,
};
use tokio_rustls::TlsConnector;
use crate::{
error::{self, Error, Result},
frame,
pstream::{self, PObject},
};
#[derive(Debug, Clone, Copy, Default)]
pub enum TlsMode {
None,
#[default]
Verified,
Insecure,
}
struct DynStream(Box<dyn DynStreamTrait>);
trait DynStreamTrait: AsyncRead + AsyncWrite + Unpin + Send {}
impl<T: AsyncRead + AsyncWrite + Unpin + Send> DynStreamTrait for T {}
impl AsyncRead for DynStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Pin::new(&mut *self.get_mut().0).poll_read(cx, buf)
}
}
impl AsyncWrite for DynStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut *self.get_mut().0).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut *self.get_mut().0).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut *self.get_mut().0).poll_shutdown(cx)
}
}
pub struct Channel {
writer: WriteHalf<DynStream>,
pub alive_interval: Option<u64>,
last_active: std::time::Instant,
reader: BufReader<ReadHalf<DynStream>>,
}
impl Channel {
pub async fn connect(
host: &str,
port: u16,
tls: TlsMode,
tls_config: Option<&Arc<rustls::ClientConfig>>,
) -> Result<Self> {
tracing::debug!(host, port, ?tls, "tcp connect");
let tcp = TcpStream::connect((host, port)).await?;
tcp.set_nodelay(true)?;
match tls {
TlsMode::None => Ok(Self::from_stream(tcp)),
TlsMode::Verified | TlsMode::Insecure => {
let (read_half, mut write_half) = tokio::io::split(tcp);
let mut read_half = BufReader::new(read_half);
tracing::debug!("sending encrypt_channel");
let req = pmap! { "_action" => "encrypt_channel" };
frame::send_message(&mut write_half, frame::SCMD_SSL_UPGRADE, &req).await?;
let resp = pstream::decode_from(&mut read_half, None).await?;
error::check_server_error(&resp)?;
tracing::debug!("encrypt_channel accepted");
let tcp = read_half.into_inner().unsplit(write_half);
let config = match tls_config {
Some(c) => Arc::clone(c),
None => Arc::new(build_tls_config(tls)?),
};
let server_name = ServerName::try_from(host.to_string()).or_else(|_| {
host.parse::<std::net::IpAddr>()
.map(|ip| ServerName::IpAddress(ip.into()))
.map_err(|_| {
Error::InvalidConfig(format!(
"'{host}' is not a valid TLS server name or IP address"
))
})
})?;
let connector = TlsConnector::from(config);
tracing::debug!("tls handshake");
let tls_stream = connector.connect(server_name, tcp).await?;
tracing::debug!("tls complete");
Ok(Self::from_stream(tls_stream))
},
}
}
pub(crate) fn from_stream<S: AsyncRead + AsyncWrite + Unpin + Send + 'static>(
stream: S,
) -> Self {
let dyn_stream = DynStream(Box::new(stream));
let (read_half, write_half) = tokio::io::split(dyn_stream);
Self {
writer: write_half,
alive_interval: None,
last_active: Instant::now(),
reader: BufReader::new(read_half),
}
}
pub async fn send(&mut self, scmd: u8, obj: &PObject) -> Result<()> {
self.last_active = Instant::now();
frame::send_message(&mut self.writer, scmd, obj).await
}
pub async fn recv(&mut self) -> Result<PObject> {
let mut obj = self.recv_single().await?;
while has_body_continue(&obj) {
let cont = self.recv_single().await?;
let more = has_body_continue(&cont);
merge_continuation(&mut obj, cont);
if !more {
break;
}
}
self.last_active = Instant::now();
Ok(obj)
}
async fn recv_single(&mut self) -> Result<PObject> {
loop {
let obj = pstream::decode_from(&mut self.reader, None).await?;
if !pstream::is_keep_alive(&obj) {
return Ok(obj);
}
}
}
#[must_use]
pub fn is_expired(&self) -> bool {
self.alive_interval
.is_some_and(|secs| self.last_active.elapsed() >= Duration::from_secs(secs))
}
pub async fn request(&mut self, scmd: u8, obj: &PObject) -> Result<PObject> {
self.send(scmd, obj).await?;
let response = self.recv().await?;
error::check_server_error(&response)?;
Ok(response)
}
pub async fn recv_download<W: AsyncWrite + Unpin + Send>(
&mut self,
dest: &mut W,
) -> Result<PObject> {
let obj = loop {
let obj = pstream::decode_from(&mut self.reader, Some(dest)).await?;
if !pstream::is_keep_alive(&obj) {
break obj;
}
};
error::check_server_error(&obj)?;
self.last_active = Instant::now();
Ok(obj)
}
}
pub struct LazyWriter<F, W> {
make: Option<F>,
inner: Option<W>,
}
impl<F, W> LazyWriter<F, W>
where
F: FnOnce() -> io::Result<W>,
{
pub const fn new(make: F) -> Self {
Self {
inner: None,
make: Some(make),
}
}
fn ensure_inner(&mut self) -> io::Result<&mut W> {
if self.inner.is_none() {
let make = self
.make
.take()
.ok_or_else(|| io::Error::other("writer creation already failed"))?;
self.inner = Some(make()?);
}
Ok(self.inner.as_mut().unwrap())
}
}
impl<F, W> AsyncWrite for LazyWriter<F, W>
where
W: AsyncWrite + Unpin,
F: FnOnce() -> io::Result<W> + Unpin,
{
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let this = self.get_mut();
let inner = match this.ensure_inner() {
Ok(w) => w,
Err(e) => return Poll::Ready(Err(e)),
};
Pin::new(inner).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
this.inner
.as_mut()
.map_or(Poll::Ready(Ok(())), |inner| Pin::new(inner).poll_flush(cx))
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
this.inner.as_mut().map_or(Poll::Ready(Ok(())), |inner| {
Pin::new(inner).poll_shutdown(cx)
})
}
}
fn has_body_continue(obj: &PObject) -> bool {
obj.get("@proto")
.and_then(|p| p.get("body-continue"))
.and_then(PObject::as_int)
.is_some_and(|v| v != 0)
}
fn merge_continuation(base: &mut PObject, cont: PObject) {
let PObject::Map(cont_map) = cont else { return };
let Some(base_map) = base.as_map_mut() else {
return;
};
for (key, cont_val) in cont_map {
if key == "@proto" {
continue;
}
let PObject::Array(cont_items) = cont_val else {
continue;
};
match base_map.get_mut(&key) {
Some(PObject::Array(base_items)) => {
base_items.extend(cont_items);
},
_ => {
base_map.insert(key, PObject::Array(cont_items));
},
}
}
}
pub fn build_tls_config(tls: TlsMode) -> Result<rustls::ClientConfig> {
let builder = rustls::ClientConfig::builder_with_provider(Arc::new(
rustls::crypto::ring::default_provider(),
))
.with_safe_default_protocol_versions()
.map_err(Error::Tls)?;
match tls {
TlsMode::Insecure => Ok(builder
.dangerous()
.with_custom_certificate_verifier(Arc::new(NoVerifier))
.with_no_client_auth()),
TlsMode::Verified => {
let certs = rustls_native_certs::load_native_certs();
if certs.certs.is_empty() {
let err_msg = if certs.errors.is_empty() {
"no CA certificates found in system trust store".into()
} else {
format!(
"failed to load CA certificates: {}",
certs
.errors
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
.join("; ")
)
};
return Err(Error::InvalidConfig(err_msg));
}
let mut root_store = rustls::RootCertStore::empty();
for cert in certs.certs {
root_store.add(cert).ok();
}
Ok(builder
.with_root_certificates(root_store)
.with_no_client_auth())
},
TlsMode::None => unreachable!(),
}
}
#[derive(Debug)]
struct NoVerifier;
static NO_VERIFY_SCHEMES: std::sync::LazyLock<Vec<rustls::SignatureScheme>> =
std::sync::LazyLock::new(|| {
rustls::crypto::ring::default_provider()
.signature_verification_algorithms
.supported_schemes()
});
impl rustls::client::danger::ServerCertVerifier for NoVerifier {
fn verify_server_cert(
&self,
_end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> std::result::Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
NO_VERIFY_SCHEMES.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn build_download_response(file_data: &[u8]) -> Vec<u8> {
let mut buf = Vec::new();
buf.push(0x42);
pstream::encode(&PObject::Str("type".into()), &mut buf).unwrap();
pstream::encode(&PObject::Str("response".into()), &mut buf).unwrap();
pstream::encode(&PObject::Str("sync_id".into()), &mut buf).unwrap();
pstream::encode(&PObject::Integer(100), &mut buf).unwrap();
pstream::encode(&PObject::Str("file".into()), &mut buf).unwrap();
buf.push(0x42);
pstream::encode(&PObject::Str("data".into()), &mut buf).unwrap();
buf.push(0x30);
buf.extend_from_slice(&(file_data.len() as u64).to_be_bytes());
buf.extend_from_slice(file_data);
pstream::encode(&PObject::Str("hash".into()), &mut buf).unwrap();
pstream::encode(&PObject::Str("abc123".into()), &mut buf).unwrap();
pstream::encode(&PObject::Str("size".into()), &mut buf).unwrap();
pstream::encode(&PObject::Integer(file_data.len() as u64), &mut buf).unwrap();
buf.push(0x40); buf.push(0x40);
buf
}
#[tokio::test]
async fn recv_download_streams_binary_data() {
let file_data = b"hello this is file content for testing";
let response_bytes = build_download_response(file_data);
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&response_bytes).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let mut dest = Vec::new();
let obj = ch.recv_download(&mut dest).await.unwrap();
assert_eq!(dest, file_data);
assert_eq!(obj.get("sync_id").and_then(PObject::as_int), Some(100));
let file = obj.get("file").unwrap();
assert_eq!(file.get("hash").and_then(PObject::as_str), Some("abc123"));
assert_eq!(
file.get("size").and_then(PObject::as_int),
Some(file_data.len() as u64)
);
}
#[tokio::test]
async fn recv_download_skips_keepalives() {
let file_data = b"the real file data after keepalives";
let mut stream_bytes = Vec::new();
let ka = pmap! { "type" => "keep_alive" };
pstream::encode(&ka, &mut stream_bytes).unwrap();
pstream::encode(&ka, &mut stream_bytes).unwrap();
stream_bytes.extend_from_slice(&build_download_response(file_data));
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&stream_bytes).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let mut dest = Vec::new();
let obj = ch.recv_download(&mut dest).await.unwrap();
assert_eq!(dest, file_data);
assert_eq!(obj.get("sync_id").and_then(PObject::as_int), Some(100));
}
#[tokio::test]
async fn recv_download_no_binary_data() {
let response = pmap! {
"type" => "response",
"sync_id" => 100u64,
"file" => pmap! {
"refer" => true,
"hash" => "abc123",
},
};
let mut response_bytes = Vec::new();
pstream::encode(&response, &mut response_bytes).unwrap();
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&response_bytes).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let mut dest = Vec::new();
let obj = ch.recv_download(&mut dest).await.unwrap();
assert!(dest.is_empty());
assert_eq!(obj.get("sync_id").and_then(PObject::as_int), Some(100));
}
fn build_binary_ex_download_response(file_data: &[u8]) -> Vec<u8> {
let mut buf = Vec::new();
buf.push(0x42);
pstream::encode(&PObject::Str("sync_id".into()), &mut buf).unwrap();
pstream::encode(&PObject::Integer(200), &mut buf).unwrap();
pstream::encode(&PObject::Str("file".into()), &mut buf).unwrap();
buf.push(0x42);
pstream::encode(&PObject::Str("data".into()), &mut buf).unwrap();
buf.push(0x43);
pstream::encode(&PObject::Str("binary".into()), &mut buf).unwrap();
buf.push(0x30);
buf.extend_from_slice(&(file_data.len() as u64).to_be_bytes());
buf.extend_from_slice(file_data);
pstream::encode(&PObject::Str("send_hash".into()), &mut buf).unwrap();
pstream::encode(&PObject::Str("deadbeef".into()), &mut buf).unwrap();
buf.push(0x40);
pstream::encode(&PObject::Str("hash".into()), &mut buf).unwrap();
pstream::encode(&PObject::Str("filehash".into()), &mut buf).unwrap();
buf.push(0x40); buf.push(0x40);
buf
}
#[tokio::test]
async fn recv_download_binary_ex_drains_correctly() {
let file_data = b"binary_ex file content";
let mut stream = build_binary_ex_download_response(file_data);
let followup = pmap! { "type" => "followup", "value" => 42u64 };
pstream::encode(&followup, &mut stream).unwrap();
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&stream).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let mut dest = Vec::new();
let obj = ch.recv_download(&mut dest).await.unwrap();
assert_eq!(dest, file_data);
assert_eq!(obj.get("sync_id").and_then(PObject::as_int), Some(200));
let file = obj.get("file").unwrap();
assert_eq!(file.get("hash").and_then(PObject::as_str), Some("filehash"));
match file.get("data").unwrap() {
PObject::BinaryEx { send_hash, .. } => assert_eq!(send_hash, "deadbeef"),
other => panic!("expected BinaryEx, got {other:?}"),
}
let next = pstream::decode_from(&mut ch.reader, None).await.unwrap();
assert_eq!(next.get("type").and_then(PObject::as_str), Some("followup"));
assert_eq!(next.get("value").and_then(PObject::as_int), Some(42));
}
#[tokio::test]
async fn recv_download_binary_drains_correctly() {
let file_data = b"plain binary content";
let mut stream = build_download_response(file_data);
let followup = pmap! { "type" => "followup", "value" => 99u64 };
pstream::encode(&followup, &mut stream).unwrap();
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&stream).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let mut dest = Vec::new();
let obj = ch.recv_download(&mut dest).await.unwrap();
assert_eq!(dest, file_data);
assert_eq!(obj.get("sync_id").and_then(PObject::as_int), Some(100));
let next = pstream::decode_from(&mut ch.reader, None).await.unwrap();
assert_eq!(next.get("type").and_then(PObject::as_str), Some("followup"));
assert_eq!(next.get("value").and_then(PObject::as_int), Some(99));
}
#[tokio::test]
async fn recv_reassembles_body_continue_frames() {
let frame1 = pmap! {
"@proto" => pmap! {
"type" => "header",
"body-continue" => true,
},
"node_list" => PObject::Array(vec![
pmap! { "name" => "file1.txt" },
pmap! { "name" => "file2.txt" },
]),
};
let frame2 = pmap! {
"@proto" => pmap! {
"type" => "header",
"body-continue" => true,
},
"node_list" => PObject::Array(vec![
pmap! { "name" => "file3.txt" },
]),
};
let frame3 = pmap! {
"@proto" => pmap! {
"type" => "header",
"body-continue" => false,
},
"node_list" => PObject::Array(vec![
pmap! { "name" => "file4.txt" },
]),
};
let mut wire = Vec::new();
pstream::encode(&frame1, &mut wire).unwrap();
pstream::encode(&frame2, &mut wire).unwrap();
pstream::encode(&frame3, &mut wire).unwrap();
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&wire).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let result = ch.recv().await.unwrap();
let node_list = result.get("node_list").and_then(PObject::as_array).unwrap();
assert_eq!(node_list.len(), 4);
assert_eq!(
node_list[0].get("name").and_then(PObject::as_str),
Some("file1.txt")
);
assert_eq!(
node_list[1].get("name").and_then(PObject::as_str),
Some("file2.txt")
);
assert_eq!(
node_list[2].get("name").and_then(PObject::as_str),
Some("file3.txt")
);
assert_eq!(
node_list[3].get("name").and_then(PObject::as_str),
Some("file4.txt")
);
}
#[tokio::test]
async fn recv_returns_single_frame_without_body_continue() {
let frame = pmap! {
"@proto" => pmap! {
"type" => "header",
"body-continue" => false,
},
"node_list" => PObject::Array(vec![
pmap! { "name" => "only.txt" },
]),
};
let mut wire = Vec::new();
pstream::encode(&frame, &mut wire).unwrap();
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&wire).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let result = ch.recv().await.unwrap();
let node_list = result.get("node_list").and_then(PObject::as_array).unwrap();
assert_eq!(node_list.len(), 1);
assert_eq!(
node_list[0].get("name").and_then(PObject::as_str),
Some("only.txt")
);
}
#[tokio::test]
async fn recv_merges_arrays_absent_from_header_frame() {
let header = pmap! {
"@proto" => pmap! {
"type" => "header",
"body-continue" => true,
},
"action" => "list_sync_to_device",
};
let body1 = pmap! {
"@proto" => pmap! { "body-continue" => true },
"node_list" => PObject::Array(vec![
pmap! { "name" => "a.txt" },
]),
};
let body2 = pmap! {
"@proto" => pmap! { "body-continue" => false },
"node_list" => PObject::Array(vec![
pmap! { "name" => "b.txt" },
]),
};
let mut wire = Vec::new();
pstream::encode(&header, &mut wire).unwrap();
pstream::encode(&body1, &mut wire).unwrap();
pstream::encode(&body2, &mut wire).unwrap();
let (client, mut server) = tokio::io::duplex(65536);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
server.write_all(&wire).await.unwrap();
server.shutdown().await.unwrap();
});
let mut ch = Channel::from_stream(client);
let result = ch.recv().await.unwrap();
let node_list = result.get("node_list").and_then(PObject::as_array).unwrap();
assert_eq!(node_list.len(), 2);
assert_eq!(
node_list[0].get("name").and_then(PObject::as_str),
Some("a.txt")
);
assert_eq!(
node_list[1].get("name").and_then(PObject::as_str),
Some("b.txt")
);
assert_eq!(
result.get("action").and_then(PObject::as_str),
Some("list_sync_to_device")
);
}
}