kcode-rust-bins 3.0.0

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