kcode-rust-bins 1.0.0

Author, validate, object-publish, and run small Rust binaries
Documentation
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());
    }
}