use std::ffi::OsString;
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use super::os_bytes;
const MAX_FRAME_BYTES: u32 = 4 * 1024 * 1024;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub enum Request {
Plan(Plan),
Compiled(Compiled),
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct Plan {
pub token: String,
pub executable: Vec<u8>,
pub args: Vec<Vec<u8>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct Compiled {
pub token: String,
pub ticket: u64,
pub success: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub enum Answer {
Served,
Compile { ticket: u64 },
Recorded,
Failed { message: String },
}
impl Plan {
#[must_use]
pub fn new(token: String, executable: &std::ffi::OsStr, args: &[OsString]) -> Self {
Self {
token,
executable: os_bytes::encode(executable),
args: args.iter().map(|arg| os_bytes::encode(arg)).collect(),
}
}
pub fn executable(&self) -> Result<OsString, String> {
os_bytes::decode(&self.executable)
}
pub fn args(&self) -> Result<Vec<OsString>, String> {
self.args.iter().map(|arg| os_bytes::decode(arg)).collect()
}
}
pub async fn write_frame<W, T>(writer: &mut W, message: &T) -> Result<(), String>
where
W: AsyncWrite + Unpin + Send,
T: Serialize + Sync,
{
let body = serde_json::to_vec(message).map_err(|error| format!("encode frame: {error}"))?;
let length = u32::try_from(body.len())
.map_err(|_| format!("frame of {} bytes exceeds the wire limit", body.len()))?;
if length > MAX_FRAME_BYTES {
return Err(format!("frame of {length} bytes exceeds the wire limit"));
}
writer
.write_all(&length.to_le_bytes())
.await
.map_err(|error| format!("write frame length: {error}"))?;
writer
.write_all(&body)
.await
.map_err(|error| format!("write frame body: {error}"))?;
writer
.flush()
.await
.map_err(|error| format!("flush frame: {error}"))
}
pub async fn read_frame<R, T>(reader: &mut R) -> Result<Option<T>, String>
where
R: AsyncRead + Unpin + Send,
T: serde::de::DeserializeOwned,
{
let mut length_bytes = [0u8; 4];
match reader.read_exact(&mut length_bytes).await {
Ok(_) => {}
Err(error) if error.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
Err(error) => return Err(format!("read frame length: {error}")),
}
let length = u32::from_le_bytes(length_bytes);
if length > MAX_FRAME_BYTES {
return Err(format!("peer announced a {length}-byte frame"));
}
let mut body = vec![0u8; length as usize];
reader
.read_exact(&mut body)
.await
.map_err(|error| format!("read frame body: {error}"))?;
serde_json::from_slice(&body)
.map(Some)
.map_err(|error| format!("decode frame: {error}"))
}
#[cfg(test)]
mod tests {
use std::ffi::OsString;
use super::{Answer, Plan, Request, read_frame, write_frame};
#[tokio::test]
async fn a_plan_round_trips_through_a_frame() {
let plan = Plan::new(
"token".to_owned(),
std::ffi::OsStr::new("/usr/bin/rustc"),
&[OsString::from("--crate-name"), OsString::from("serde")],
);
let mut buffer = Vec::new();
write_frame(&mut buffer, &Request::Plan(plan.clone()))
.await
.expect("write");
let decoded: Request = read_frame(&mut buffer.as_slice())
.await
.expect("read")
.expect("a frame");
assert_eq!(decoded, Request::Plan(plan.clone()));
let Request::Plan(decoded) = decoded else {
panic!("a plan frame decodes as a plan");
};
assert_eq!(decoded.executable().expect("executable"), "/usr/bin/rustc");
assert_eq!(decoded.args().expect("args"), vec!["--crate-name", "serde"]);
}
#[tokio::test]
async fn an_empty_stream_is_a_clean_end() {
let decoded: Option<Answer> = read_frame(&mut [].as_slice()).await.expect("read");
assert!(decoded.is_none());
}
#[tokio::test]
async fn a_truncated_frame_is_an_error() {
let mut buffer = Vec::new();
write_frame(&mut buffer, &Answer::Served)
.await
.expect("write");
buffer.truncate(buffer.len() - 1);
let error = read_frame::<_, Answer>(&mut buffer.as_slice())
.await
.expect_err("a truncated frame must not read as a clean end");
assert!(error.contains("read frame body"), "{error}");
}
}