use std::io::{Read, Write};
use crate::model::{Error, Result, RustBinInput, RustBinOutput, validate_json};
const MAGIC: [u8; 8] = *b"KRBIN\x01\0\0";
const MAX_JSON_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 {
json: 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.json, &output.objects)?;
writer
.flush()
.map_err(|error| Error::new("protocol", format!("flush stdout: {error}")))
}
pub(crate) fn write_message(
writer: &mut impl Write,
json: &str,
objects: &[Vec<u8>],
) -> Result<()> {
validate_parts(json, objects)?;
writer
.write_all(&MAGIC)
.and_then(|()| writer.write_all(&(json.len() as u64).to_le_bytes()))
.and_then(|()| writer.write_all(json.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 (json, objects) = read_message(reader)?;
Ok(RustBinOutput { json, 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 json_length = read_u64(reader, "JSON length")?;
let json_length = usize::try_from(json_length)
.map_err(|_| Error::new("protocol", "JSON length overflows usize"))?;
if json_length > MAX_JSON_BYTES {
return Err(Error::new("protocol", "JSON exceeds 16 MiB"));
}
let mut json = vec![0_u8; json_length];
read_exact(reader, &mut json, "JSON")?;
let json = String::from_utf8(json)
.map_err(|error| Error::new("protocol", format!("JSON is not UTF-8: {error}")))?;
validate_json(&json)?;
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 = json_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((json, objects)),
Ok(_) => Err(Error::new("protocol", "trailing bytes after frame")),
Err(error) => Err(Error::new(
"protocol",
format!("read frame terminator: {error}"),
)),
}
}
fn validate_parts(json: &str, objects: &[Vec<u8>]) -> Result<()> {
validate_json(json)?;
if json.len() > MAX_JSON_BYTES {
return Err(Error::new("protocol", "JSON exceeds 16 MiB"));
}
if objects.len() > MAX_OBJECTS {
return Err(Error::new("protocol", "object count exceeds 1,024"));
}
let mut total = json.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_one_json_value_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.json, "{\"answer\":42}");
assert_eq!(output.objects, [b"first".to_vec(), b"second".to_vec()]);
}
#[test]
fn rejects_invalid_json_and_trailing_bytes() {
let mut bytes = Vec::new();
assert!(write_message(&mut bytes, "not-json", &[]).is_err());
write_message(&mut bytes, "null", &[]).unwrap();
bytes.push(0);
assert!(read_output(&mut bytes.as_slice()).is_err());
}
}