use anyhow::{Context, Result, bail};
use serde::{Serialize, de::DeserializeOwned};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
pub const PROTOCOL_VERSION: u32 = 2;
const MAX_CONTROL_FRAME_BYTES: u32 = 1024 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Direction {
Download,
Upload,
}
#[derive(Debug, Serialize, serde::Deserialize)]
#[serde(tag = "type")]
pub enum ClientMessage {
Hello {
protocol_version: u32,
},
Transfer {
token: String,
direction: Direction,
remote_spec: String,
local_spec: Option<String>,
force: bool,
dry_run: bool,
use_scp: bool,
},
Status {
token: String,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Tool {
Rsync,
Scp,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum LogStream {
Stdout,
Stderr,
}
#[derive(Debug, Serialize, serde::Deserialize)]
#[serde(tag = "type")]
pub enum ServerMessage {
HelloAck {
protocol_version: u32,
},
Starting {
tool: Tool,
fell_back: bool,
note: Option<String>,
},
Log {
stream: LogStream,
chunk: String,
},
Done {
code: i32,
},
StatusResponse {
session_local_dir: String,
host: String,
},
Error {
message: String,
},
}
pub async fn write_message<W, T>(writer: &mut W, message: &T) -> Result<()>
where
W: AsyncWrite + Unpin,
T: Serialize,
{
let payload = serde_json::to_vec(message).context("failed to serialize protocol message")?;
let len = u32::try_from(payload.len()).context("protocol message too large")?;
writer
.write_all(&len.to_le_bytes())
.await
.context("failed to write protocol frame length")?;
writer
.write_all(&payload)
.await
.context("failed to write protocol frame body")?;
writer
.flush()
.await
.context("failed to flush protocol frame")?;
Ok(())
}
pub async fn read_message<R, T>(reader: &mut R) -> Result<T>
where
R: AsyncRead + Unpin,
T: DeserializeOwned,
{
let mut len_bytes = [0u8; 4];
reader
.read_exact(&mut len_bytes)
.await
.context("failed to read protocol frame length")?;
let len = u32::from_le_bytes(len_bytes);
if len > MAX_CONTROL_FRAME_BYTES {
bail!("protocol control frame of {len} bytes exceeds the maximum allowed size");
}
let mut payload = vec![0u8; len as usize];
reader
.read_exact(&mut payload)
.await
.context("failed to read protocol frame body")?;
serde_json::from_slice(&payload).context("failed to parse protocol message")
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn round_trips_a_control_message() {
let mut buf = Vec::new();
write_message(
&mut buf,
&ClientMessage::Hello {
protocol_version: PROTOCOL_VERSION,
},
)
.await
.unwrap();
let mut cursor = std::io::Cursor::new(buf);
let decoded: ClientMessage = read_message(&mut cursor).await.unwrap();
match decoded {
ClientMessage::Hello { protocol_version } => {
assert_eq!(protocol_version, PROTOCOL_VERSION);
}
other => panic!("unexpected message: {other:?}"),
}
}
#[tokio::test]
async fn round_trips_a_transfer_request() {
let mut buf = Vec::new();
write_message(
&mut buf,
&ClientMessage::Transfer {
token: "tok".to_string(),
direction: Direction::Download,
remote_spec: "/abs/logs/*.log".to_string(),
local_spec: Some("dest".to_string()),
force: true,
dry_run: false,
use_scp: false,
},
)
.await
.unwrap();
let mut cursor = std::io::Cursor::new(buf);
let decoded: ClientMessage = read_message(&mut cursor).await.unwrap();
match decoded {
ClientMessage::Transfer {
direction,
remote_spec,
force,
..
} => {
assert_eq!(direction, Direction::Download);
assert_eq!(remote_spec, "/abs/logs/*.log");
assert!(force);
}
other => panic!("unexpected message: {other:?}"),
}
}
#[tokio::test]
async fn round_trips_server_log_and_done() {
for message in [
ServerMessage::Starting {
tool: Tool::Scp,
fell_back: true,
note: Some("rsync not found; using scp".to_string()),
},
ServerMessage::Log {
stream: LogStream::Stderr,
chunk: "progress\r".to_string(),
},
ServerMessage::Done { code: 23 },
ServerMessage::StatusResponse {
session_local_dir: "/home/u".to_string(),
host: "dev".to_string(),
},
] {
let mut buf = Vec::new();
write_message(&mut buf, &message).await.unwrap();
let mut cursor = std::io::Cursor::new(buf);
let _decoded: ServerMessage = read_message(&mut cursor).await.unwrap();
}
}
#[tokio::test]
async fn rejects_oversized_control_frame() {
let mut buf = Vec::new();
buf.extend_from_slice(&(MAX_CONTROL_FRAME_BYTES + 1).to_le_bytes());
let mut cursor = std::io::Cursor::new(buf);
let result: Result<ClientMessage> = read_message(&mut cursor).await;
assert!(result.is_err());
}
}