use std::io::{Read, Write};
use crate::model::{Error, Result, RustBinInput, RustBinOutput};
const MAGIC: [u8; 8] = *b"KRBIN\x02\0\0";
const MAX_TEXT_BYTES: usize = 16 * 1024 * 1024;
const MAX_OBJECTS: usize = 1024;
const MAX_OBJECT_BYTES: usize = 512 * 1024 * 1024;
const MAX_TOTAL_BYTES: usize = 1024 * 1024 * 1024;
pub fn read_input() -> Result<RustBinInput> {
let stdin = std::io::stdin();
let mut input = stdin.lock();
let message = read_message(&mut input)?;
Ok(RustBinInput {
text: message.0,
objects: message.1,
})
}
pub fn write_output(output: &RustBinOutput) -> Result<()> {
let stdout = std::io::stdout();
let mut writer = stdout.lock();
write_message(&mut writer, &output.text, &output.objects)?;
writer
.flush()
.map_err(|error| Error::new("protocol", format!("flush stdout: {error}")))
}
pub(crate) fn write_message(
writer: &mut impl Write,
text: &str,
objects: &[Vec<u8>],
) -> Result<()> {
validate_parts(text, objects)?;
writer
.write_all(&MAGIC)
.and_then(|()| writer.write_all(&(text.len() as u64).to_le_bytes()))
.and_then(|()| writer.write_all(text.as_bytes()))
.and_then(|()| writer.write_all(&(objects.len() as u32).to_le_bytes()))
.map_err(|error| Error::new("protocol", format!("write frame header: {error}")))?;
for object in objects {
writer
.write_all(&(object.len() as u64).to_le_bytes())
.and_then(|()| writer.write_all(object))
.map_err(|error| Error::new("protocol", format!("write frame object: {error}")))?;
}
Ok(())
}
pub(crate) fn read_output(reader: &mut impl Read) -> Result<RustBinOutput> {
let (text, objects) = read_message(reader)?;
Ok(RustBinOutput { text, objects })
}
fn read_message(reader: &mut impl Read) -> Result<(String, Vec<Vec<u8>>)> {
let mut magic = [0_u8; 8];
read_exact(reader, &mut magic, "magic")?;
if magic != MAGIC {
return Err(Error::new("protocol", "invalid frame magic or revision"));
}
let text_length = read_u64(reader, "text length")?;
let text_length = usize::try_from(text_length)
.map_err(|_| Error::new("protocol", "text length overflows usize"))?;
if text_length > MAX_TEXT_BYTES {
return Err(Error::new("protocol", "text exceeds 16 MiB"));
}
let mut text = vec![0_u8; text_length];
read_exact(reader, &mut text, "text")?;
let text = String::from_utf8(text)
.map_err(|error| Error::new("protocol", format!("text is not UTF-8: {error}")))?;
let count = read_u32(reader, "object count")? as usize;
if count > MAX_OBJECTS {
return Err(Error::new("protocol", "object count exceeds 1,024"));
}
let mut total = text_length;
let mut objects = Vec::with_capacity(count);
for _ in 0..count {
let length = read_u64(reader, "object length")?;
let length = usize::try_from(length)
.map_err(|_| Error::new("protocol", "object length overflows usize"))?;
if length > MAX_OBJECT_BYTES {
return Err(Error::new("protocol", "object exceeds 512 MiB"));
}
total = total
.checked_add(length)
.ok_or_else(|| Error::new("protocol", "aggregate payload length overflow"))?;
if total > MAX_TOTAL_BYTES {
return Err(Error::new("protocol", "aggregate payload exceeds 1 GiB"));
}
let mut object = vec![0_u8; length];
read_exact(reader, &mut object, "object bytes")?;
objects.push(object);
}
let mut trailing = [0_u8; 1];
match reader.read(&mut trailing) {
Ok(0) => Ok((text, objects)),
Ok(_) => Err(Error::new("protocol", "trailing bytes after frame")),
Err(error) => Err(Error::new(
"protocol",
format!("read frame terminator: {error}"),
)),
}
}
fn validate_parts(text: &str, objects: &[Vec<u8>]) -> Result<()> {
if text.len() > MAX_TEXT_BYTES {
return Err(Error::new("protocol", "text exceeds 16 MiB"));
}
if objects.len() > MAX_OBJECTS {
return Err(Error::new("protocol", "object count exceeds 1,024"));
}
let mut total = text.len();
for object in objects {
if object.len() > MAX_OBJECT_BYTES {
return Err(Error::new("protocol", "object exceeds 512 MiB"));
}
total = total
.checked_add(object.len())
.ok_or_else(|| Error::new("protocol", "aggregate payload length overflow"))?;
if total > MAX_TOTAL_BYTES {
return Err(Error::new("protocol", "aggregate payload exceeds 1 GiB"));
}
}
Ok(())
}
fn read_u64(reader: &mut impl Read, label: &str) -> Result<u64> {
let mut bytes = [0_u8; 8];
read_exact(reader, &mut bytes, label)?;
Ok(u64::from_le_bytes(bytes))
}
fn read_u32(reader: &mut impl Read, label: &str) -> Result<u32> {
let mut bytes = [0_u8; 4];
read_exact(reader, &mut bytes, label)?;
Ok(u32::from_le_bytes(bytes))
}
fn read_exact(reader: &mut impl Read, bytes: &mut [u8], label: &str) -> Result<()> {
reader
.read_exact(bytes)
.map_err(|error| Error::new("protocol", format!("read {label}: {error}")))
}
#[cfg(test)]
mod tests {
use super::{read_output, write_message};
#[test]
fn round_trips_text_and_ordered_objects() {
let mut bytes = Vec::new();
write_message(
&mut bytes,
"answer: 42",
&[b"first".to_vec(), b"second".to_vec()],
)
.unwrap();
let output = read_output(&mut bytes.as_slice()).unwrap();
assert_eq!(output.text, "answer: 42");
assert_eq!(output.objects, [b"first".to_vec(), b"second".to_vec()]);
}
#[test]
fn rejects_trailing_bytes() {
let mut bytes = Vec::new();
write_message(&mut bytes, "plain text", &[]).unwrap();
bytes.push(0);
assert!(read_output(&mut bytes.as_slice()).is_err());
}
}